diff --git a/managed_models.py b/managed_models.py new file mode 100644 index 0000000..86f62de --- /dev/null +++ b/managed_models.py @@ -0,0 +1,44 @@ +from collections.abc import Callable +from typing import TypeVar + +import torch + +import comfy.model_management +import comfy.model_patcher + + +T = TypeVar("T", bound=torch.Tensor) + + +class ManagedAuxiliaryModel: + """Keep a small auxiliary model under ComfyUI's load/offload lifecycle.""" + + def __init__(self, factory: Callable[[], torch.nn.Module]): + model = factory().eval().requires_grad_(False) + load_device = comfy.model_management.vae_device() + offload_device = comfy.model_management.vae_offload_device() + model.to(offload_device) + + self.model = model + self.patcher = comfy.model_patcher.CoreModelPatcher( + model, + load_device=load_device, + offload_device=offload_device, + ) + + def _dtype(self) -> torch.dtype: + parameter = next(self.model.parameters(), None) + return parameter.dtype if parameter is not None else torch.float32 + + def run( + self, + value: torch.Tensor, + postprocess: Callable[[torch.Tensor], T] | None = None, + ) -> T: + comfy.model_management.load_models_gpu([self.patcher]) + device = self.patcher.load_device + with comfy.model_management.cuda_device_context(device): + output = self.model(value.to(device=device, dtype=self._dtype())) + if postprocess is not None: + output = postprocess(output) + return output.to(comfy.model_management.intermediate_device()) diff --git a/nodes.py b/nodes.py index 9eccfc1..9b3a2d2 100644 --- a/nodes.py +++ b/nodes.py @@ -12,6 +12,7 @@ from nodes import VAELoader from .src.sd import CustomVAE from .latent_upscale.model import latent_upscale_models from .latent_upscale.latent_projector import Wan21_latent_projector +from .managed_models import ManagedAuxiliaryModel class VAEUtils_CustomVAELoader(VAELoader): @@ -41,9 +42,8 @@ class VAEUtils_CustomVAELoader(VAELoader): else: vae_path = folder_paths.get_full_path_or_raise("vae", vae_name) sd = comfy.utils.load_torch_file(vae_path) - vae = CustomVAE(sd=sd) + vae = CustomVAE(sd=sd, disable_offload=disable_offload) vae.throw_exception_if_invalid() - vae.disable_offload = disable_offload return (vae, ) @@ -63,6 +63,13 @@ class VAEUtils_DisableVAEOffload: def set_offload(self, vae, disable_offload): vae = copy.copy(vae) + if hasattr(vae, "patcher"): + vae.patcher = vae.patcher.clone() + vae.patcher.offload_device = ( + vae.patcher.load_device + if disable_offload + else comfy.model_management.vae_offload_device() + ) vae.disable_offload = disable_offload return (vae, ) @@ -136,6 +143,10 @@ class VAEUtils_VAEDecodeTiled: class VAEUtils_LatentUpscale: + def __init__(self): + self._managed_model = None + self._managed_model_name = None + @classmethod def INPUT_TYPES(s): return { @@ -150,19 +161,22 @@ class VAEUtils_LatentUpscale: CATEGORY = "VAE-Utils" def upscale(self, samples, model): - device = comfy.model_management.get_torch_device() - model = latent_upscale_models[model]().to(device) - - latents = samples["samples"].to(dtype=torch.float32, device=device) - upscaled_latents = model(latents).to(comfy.model_management.intermediate_device()) - - samples = copy.deepcopy(samples) + if self._managed_model is None or self._managed_model_name != model: + self._managed_model = ManagedAuxiliaryModel(latent_upscale_models[model]) + self._managed_model_name = model + + upscaled_latents = self._managed_model.run(samples["samples"]) + + samples = samples.copy() samples["samples"] = upscaled_latents return (samples, ) class VAEUtils_WanLatentPreview: + def __init__(self): + self._managed_projector = None + @classmethod def INPUT_TYPES(s): return { @@ -176,18 +190,21 @@ class VAEUtils_WanLatentPreview: CATEGORY = "VAE-Utils" def upscale(self, samples): - device = comfy.model_management.intermediate_device() - projector = Wan21_latent_projector().to(device) - - latents = samples["samples"].to(dtype=torch.float32, device=device) - pixels = projector(latents).to(comfy.model_management.intermediate_device()) - pixels = pixels * 0.5 + 0.5 - - f, h, w = pixels.shape[-3:] - pixels = F.interpolate(pixels, size=(f, h//8, w//8), mode="area") - - pixels = [b.movedim(0, -1) for b in pixels] # CFHW -> FHWC - pixels = torch.cat(pixels, dim=0) # (BF)HWC + if self._managed_projector is None: + self._managed_projector = ManagedAuxiliaryModel(Wan21_latent_projector) + + def postprocess(pixels): + pixels = pixels * 0.5 + 0.5 + frames, height, width = pixels.shape[-3:] + pixels = F.interpolate( + pixels, + size=(frames, height // 8, width // 8), + mode="area", + ) + pixels = [batch.movedim(0, -1) for batch in pixels] + return torch.cat(pixels, dim=0) + + pixels = self._managed_projector.run(samples["samples"], postprocess) return (pixels, ) @@ -447,4 +464,4 @@ COMBINED_MAPPINGS = { "VAEUtils_TileModelPatch": (VAEUtils_TileModelPatch, "Tile Model Patch (VAE Utils)"), "VAEUtils_VisualizeTiles": (VAEUtils_VisualizeTiles, "Visualize Tiles (VAE Utils)"), "VAEUtils_ScaleLatents": (VAEUtils_ScaleLatents, "Scale/Unscale Latents (VAE Utils)"), -} \ No newline at end of file +} diff --git a/src/sd.py b/src/sd.py index e364c6a..90184d2 100644 --- a/src/sd.py +++ b/src/sd.py @@ -27,7 +27,7 @@ import comfy.taesd.taesd from comfy.sd import VAE class CustomVAE(VAE): - def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None): + def __init__(self, sd=None, device=None, config=None, dtype=None, metadata=None, disable_offload=False): if model_management.is_amd(): VAE_KL_MEM_RATIO = 2.73 else: @@ -46,7 +46,7 @@ class CustomVAE(VAE): self.process_input = lambda image: image * 2.0 - 1.0 self.process_output = lambda image: torch.clamp((image + 1.0) / 2.0, min=0.0, max=1.0) self.working_dtypes = [torch.bfloat16, torch.float32] - self.disable_offload = False + self.disable_offload = disable_offload self.not_video = False self.size = None @@ -361,7 +361,9 @@ class CustomVAE(VAE): self.output_device = model_management.intermediate_device() self.patcher = comfy.model_patcher.ModelPatcher(self.first_stage_model, load_device=self.device, offload_device=offload_device) - logging.info("VAE load device: {}, offload device: {}, dtype: {}".format(self.device, offload_device, self.vae_dtype)) + if self.disable_offload: + self.patcher.offload_device = self.device + logging.info("VAE load device: {}, offload device: {}, dtype: {}".format(self.device, self.patcher.offload_device, self.vae_dtype)) def decode(self, samples_in): """ diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..43c492f --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,31 @@ +import importlib.util +import sys +from pathlib import Path + +import pytest + + +REPOSITORY_ROOT = Path(__file__).resolve().parents[1] + + +def _comfyui_root() -> Path: + for parent in REPOSITORY_ROOT.parents: + if (parent / "comfy").is_dir() and (parent / "nodes.py").is_file(): + return parent + raise RuntimeError("Could not locate the containing ComfyUI installation.") + + +@pytest.fixture(scope="session") +def vae_utils_package(): + comfyui_root = _comfyui_root() + sys.path.insert(0, str(comfyui_root)) + package_name = "comfyui_vae_utils_under_test" + spec = importlib.util.spec_from_file_location( + package_name, + REPOSITORY_ROOT / "__init__.py", + submodule_search_locations=[str(REPOSITORY_ROOT)], + ) + module = importlib.util.module_from_spec(spec) + sys.modules[package_name] = module + spec.loader.exec_module(module) + return module diff --git a/tests/test_managed_lifecycle.py b/tests/test_managed_lifecycle.py new file mode 100644 index 0000000..8474d6c --- /dev/null +++ b/tests/test_managed_lifecycle.py @@ -0,0 +1,131 @@ +import contextlib +import copy + +import torch + + +class DummyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor(2.0)) + + def forward(self, value): + return value * self.weight + + +class DummyPatcher: + def __init__(self, model=None, load_device="cpu", offload_device="cpu"): + self.model = model + self.load_device = load_device + self.offload_device = offload_device + + def clone(self): + return copy.copy(self) + + +def test_managed_model_uses_core_lifecycle(monkeypatch, vae_utils_package): + managed_models = vae_utils_package.nodes.ManagedAuxiliaryModel.__module__ + managed_models = __import__(managed_models, fromlist=["ManagedAuxiliaryModel"]) + loads = [] + + monkeypatch.setattr(managed_models.comfy.model_patcher, "CoreModelPatcher", DummyPatcher) + monkeypatch.setattr(managed_models.comfy.model_management, "vae_device", lambda: torch.device("cpu")) + monkeypatch.setattr(managed_models.comfy.model_management, "vae_offload_device", lambda: torch.device("cpu")) + monkeypatch.setattr(managed_models.comfy.model_management, "intermediate_device", lambda: torch.device("cpu")) + monkeypatch.setattr(managed_models.comfy.model_management, "load_models_gpu", lambda patchers: loads.append(patchers)) + monkeypatch.setattr(managed_models.comfy.model_management, "cuda_device_context", lambda _device: contextlib.nullcontext()) + + managed = managed_models.ManagedAuxiliaryModel(DummyModel) + output = managed.run(torch.ones(2, dtype=torch.float64)) + + assert loads == [[managed.patcher]] + assert managed.model.training is False + assert all(parameter.requires_grad is False for parameter in managed.model.parameters()) + assert output.dtype == torch.float32 + assert output.requires_grad is False + + +def test_latent_upscaler_reuses_managed_model_and_shallow_copies(monkeypatch, vae_utils_package): + nodes = vae_utils_package.nodes + created = [] + + class FakeManagedModel: + def __init__(self, factory): + created.append(factory) + + def run(self, value, postprocess=None): + output = value + 1 + return postprocess(output) if postprocess is not None else output + + factory = object() + monkeypatch.setattr(nodes, "ManagedAuxiliaryModel", FakeManagedModel) + monkeypatch.setitem(nodes.latent_upscale_models, "test", factory) + node = nodes.VAEUtils_LatentUpscale() + metadata = {"seed": 1} + source = {"samples": torch.zeros(1), "metadata": metadata} + + first = node.upscale(source, "test")[0] + second = node.upscale(source, "test")[0] + + assert created == [factory] + assert first is not source and second is not source + assert first["metadata"] is metadata + assert torch.equal(source["samples"], torch.zeros(1)) + assert torch.equal(first["samples"], torch.ones(1)) + + +def test_preview_reuses_projector_and_preserves_output_layout(monkeypatch, vae_utils_package): + nodes = vae_utils_package.nodes + created = [] + + class FakeManagedModel: + def __init__(self, factory): + created.append(factory) + + def run(self, _value, postprocess=None): + pixels = torch.zeros(2, 3, 1, 16, 16) + return postprocess(pixels) + + monkeypatch.setattr(nodes, "ManagedAuxiliaryModel", FakeManagedModel) + node = nodes.VAEUtils_WanLatentPreview() + + first = node.upscale({"samples": torch.zeros(1)})[0] + second = node.upscale({"samples": torch.zeros(1)})[0] + + assert len(created) == 1 + assert first.shape == (2, 2, 2, 3) + assert second.shape == first.shape + + +def test_disable_offload_clones_patcher(monkeypatch, vae_utils_package): + nodes = vae_utils_package.nodes + monkeypatch.setattr(nodes.comfy.model_management, "vae_offload_device", lambda: "cpu") + source = type("VAE", (), {})() + source.patcher = DummyPatcher(load_device="cuda", offload_device="cpu") + source.disable_offload = False + + result = nodes.VAEUtils_DisableVAEOffload().set_offload(source, True)[0] + + assert result is not source + assert result.patcher is not source.patcher + assert result.patcher.offload_device == "cuda" + assert source.patcher.offload_device == "cpu" + assert result.disable_offload is True + + +def test_public_node_contract_is_unchanged(vae_utils_package): + nodes = vae_utils_package.nodes + + assert set(nodes.COMBINED_MAPPINGS) == { + "VAEUtils_CustomVAELoader", + "VAEUtils_DisableVAEOffload", + "VAEUtils_VAEDecodeTiled", + "VAEUtils_LatentUpscale", + "VAEUtils_WanLatentPreview", + "VAEUtils_TileModelPatch", + "VAEUtils_VisualizeTiles", + "VAEUtils_ScaleLatents", + } + assert nodes.VAEUtils_LatentUpscale.RETURN_TYPES == ("LATENT",) + assert nodes.VAEUtils_WanLatentPreview.RETURN_TYPES == ("IMAGE",) + assert nodes.VAEUtils_DisableVAEOffload.RETURN_TYPES == ("VAE",)