fix: defer model lifecycle to Core
This commit is contained in:
@@ -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())
|
||||
@@ -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)"),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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):
|
||||
"""
|
||||
|
||||
@@ -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
|
||||
@@ -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",)
|
||||
Reference in New Issue
Block a user