mirror of
https://github.com/comfyanonymous/ComfyUI.git
synced 2026-07-21 23:41:28 +08:00
Improvements to execution
- Validation errors that occur early in the lifecycle of prompt execution now get propagated to their callers in the EmbeddedComfyClient. This includes error messages about missing node classes. - The execution context now includes the node_id and the prompt_id - Latent previews are now sent with a node_id. This is not backwards compatible with old frontends. - Dependency execution errors are now modeled correctly. - Distributed progress encodes image previews with node and prompt IDs. - Typing for models - The frontend was updated to use node IDs with previews - Improvements to torch.compile experiments - Some controlnet_aux nodes were upstreamed
This commit is contained in:
@@ -1,10 +1,8 @@
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy import model_management
|
||||
from comfy.model_base import Flux
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.nodes.base_nodes import UNETLoader
|
||||
from comfy.nodes.base_nodes import UNETLoader, CheckpointLoaderSimple
|
||||
from comfy_extras.nodes.nodes_torch_compile import QuantizeModel
|
||||
|
||||
has_torchao = True
|
||||
@@ -20,45 +18,42 @@ except (ImportError, ModuleNotFoundError):
|
||||
has_tensorrt = False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpoint_name", ["flux1-dev.safetensors"])
|
||||
@pytest.fixture(scope="function", params=["flux1-dev.safetensors"])
|
||||
def model_patcher_obj(request) -> ModelPatcher:
|
||||
checkpoint_name = request.param
|
||||
model_obj = None
|
||||
try:
|
||||
if "flux" in checkpoint_name:
|
||||
model_obj, = UNETLoader().load_unet(checkpoint_name, weight_dtype="default")
|
||||
yield model_obj
|
||||
else:
|
||||
objs = CheckpointLoaderSimple().load_checkpoint(checkpoint_name)
|
||||
model_obj = objs[0]
|
||||
yield model_obj
|
||||
finally:
|
||||
model_management.unload_all_models()
|
||||
if model_obj is not None:
|
||||
model_obj.unpatch_model()
|
||||
del model_obj
|
||||
|
||||
model_management.soft_empty_cache(force=True)
|
||||
|
||||
|
||||
@pytest.mark.forked
|
||||
@pytest.mark.skipif(not has_torchao, reason="torchao not installed")
|
||||
async def test_unit_torchao(checkpoint_name):
|
||||
# Downloads FLUX.1-dev and loads it using ComfyUI's models
|
||||
model, = UNETLoader().load_unet(checkpoint_name, weight_dtype="default")
|
||||
model: ModelPatcher = model.clone()
|
||||
|
||||
transformer: Flux = model.get_model_object("diffusion_model")
|
||||
quantize_(transformer, int8_dynamic_activation_int8_weight(), device=model_management.get_torch_device())
|
||||
assert transformer is not None
|
||||
del transformer
|
||||
model_management.unload_all_models()
|
||||
@pytest.mark.skipif(True, reason="wip")
|
||||
async def test_unit_torchao(model_patcher_obj):
|
||||
quantize_(model_patcher_obj.diffusion_model, int8_dynamic_activation_int8_weight(), device=model_management.get_torch_device())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpoint_name", ["flux1-dev.safetensors"])
|
||||
@pytest.mark.forked
|
||||
@pytest.mark.parametrize("strategy", ["torchao", "torchao-autoquant"])
|
||||
@pytest.mark.skipif(not has_torchao, reason="torchao not installed")
|
||||
async def test_torchao_node(checkpoint_name, strategy):
|
||||
model, = UNETLoader().load_unet(checkpoint_name, weight_dtype="default")
|
||||
model: ModelPatcher = model.clone()
|
||||
|
||||
quantized_model, = QuantizeModel().execute(model, strategy=strategy)
|
||||
|
||||
transformer = quantized_model.get_model_object("diffusion_model")
|
||||
del transformer
|
||||
model_management.unload_all_models()
|
||||
@pytest.mark.skipif(True, reason="wip")
|
||||
async def test_torchao_node(model_patcher_obj, strategy):
|
||||
QuantizeModel().execute(model_patcher_obj, strategy=strategy)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("checkpoint_name", ["flux1-dev.safetensors"])
|
||||
@pytest.mark.parametrize("strategy", ["torchao", "torchao-autoquant"])
|
||||
@pytest.mark.skipif(True, reason="not yet supported")
|
||||
async def test_torchao_into_tensorrt(checkpoint_name, strategy):
|
||||
model, = UNETLoader().load_unet(checkpoint_name, weight_dtype="default")
|
||||
model: ModelPatcher = model.clone()
|
||||
model_management.load_models_gpu([model], force_full_load=True)
|
||||
model.diffusion_model = model.diffusion_model.to(memory_format=torch.channels_last)
|
||||
model.diffusion_model = torch.compile(model.diffusion_model, mode="max-autotune", fullgraph=True)
|
||||
|
||||
quantized_model, = QuantizeModel().execute(model, strategy=strategy)
|
||||
|
||||
STATIC_TRT_MODEL_CONVERSION().convert(quantized_model, "test", 1, 1024, 1024, 1, 14)
|
||||
model_management.unload_all_models()
|
||||
@pytest.mark.forked
|
||||
@pytest.mark.skipif(True, reason="wip")
|
||||
async def test_tensorrt(model_patcher_obj):
|
||||
STATIC_TRT_MODEL_CONVERSION().convert(model_patcher_obj, "test", 1, 1024, 1024, 1, 14)
|
||||
|
||||
Reference in New Issue
Block a user