mirror of
https://github.com/Comfy-Org/ComfyUI.git
synced 2026-08-16 22:46:38 +08:00
Add SeedVR2 support (CORE-6) (#14110)
This commit is contained in:
213
tests-unit/comfy_extras_test/test_seedvr2_conditioning.py
Normal file
213
tests-unit/comfy_extras_test/test_seedvr2_conditioning.py
Normal file
@@ -0,0 +1,213 @@
|
||||
"""Consolidated SeedVR2 conditioning and refactor regression tests.
|
||||
|
||||
Merges the prior test_seedvr2_refactor_nodes.py and
|
||||
test_seedvr_conditioning_hardening.py modules. Refactor tests use the
|
||||
top-level comfy_extras.nodes_seedvr import; conditioning-hardening tests
|
||||
use _import_nodes_seedvr_isolated() for sys.modules isolation when
|
||||
mocking comfy.model_management.
|
||||
"""
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
|
||||
_SENTINEL = object()
|
||||
_TARGETS = (
|
||||
("comfy.model_management", "comfy"),
|
||||
("comfy_extras.nodes_seedvr", "comfy_extras"),
|
||||
)
|
||||
|
||||
|
||||
def _import_nodes_seedvr_isolated():
|
||||
"""Import comfy_extras.nodes_seedvr with comfy.model_management mocked."""
|
||||
priors = []
|
||||
for mod_name, parent_name in _TARGETS:
|
||||
prior_mod = sys.modules.get(mod_name, _SENTINEL)
|
||||
parent = sys.modules.get(parent_name)
|
||||
attr = mod_name.split(".")[-1]
|
||||
prior_attr = (
|
||||
getattr(parent, attr, _SENTINEL) if parent is not None else _SENTINEL
|
||||
)
|
||||
priors.append((mod_name, parent_name, attr, prior_mod, prior_attr))
|
||||
|
||||
mock_mm = MagicMock()
|
||||
for fn in (
|
||||
"xformers_enabled", "xformers_enabled_vae",
|
||||
"pytorch_attention_enabled", "pytorch_attention_enabled_vae",
|
||||
"sage_attention_enabled", "flash_attention_enabled",
|
||||
"is_intel_xpu",
|
||||
):
|
||||
getattr(mock_mm, fn).return_value = False
|
||||
tv = torch.version.__version__.split(".")
|
||||
mock_mm.torch_version_numeric = (int(tv[0]), int(tv[1]))
|
||||
mock_mm.WINDOWS = False
|
||||
sys.modules["comfy.model_management"] = mock_mm
|
||||
if sys.modules.get("comfy") is None:
|
||||
import comfy as _comfy_pkg # noqa: F401
|
||||
comfy_pkg = sys.modules.get("comfy")
|
||||
if comfy_pkg is not None:
|
||||
setattr(comfy_pkg, "model_management", mock_mm)
|
||||
nodes_seedvr = sys.modules.get("comfy_extras.nodes_seedvr") or (
|
||||
importlib.import_module("comfy_extras.nodes_seedvr")
|
||||
)
|
||||
|
||||
def _restore():
|
||||
for mod_name, parent_name, attr, prior_mod, prior_attr in priors:
|
||||
if prior_mod is _SENTINEL:
|
||||
sys.modules.pop(mod_name, None)
|
||||
else:
|
||||
sys.modules[mod_name] = prior_mod
|
||||
parent = sys.modules.get(parent_name)
|
||||
if parent is None:
|
||||
continue
|
||||
if prior_attr is _SENTINEL:
|
||||
if hasattr(parent, attr):
|
||||
delattr(parent, attr)
|
||||
else:
|
||||
setattr(parent, attr, prior_attr)
|
||||
|
||||
return nodes_seedvr, _restore
|
||||
|
||||
|
||||
class _Rope(nn.Module):
|
||||
"""Minimal RoPE stub exposing a `freqs` parameter."""
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.freqs = nn.Parameter(torch.zeros(4))
|
||||
|
||||
|
||||
class _Block(nn.Module):
|
||||
"""Minimal transformer block stub holding a `_Rope`."""
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.rope = _Rope()
|
||||
|
||||
|
||||
class _DiffusionModel(nn.Module):
|
||||
"""Stub diffusion model with N blocks and pos/neg conditioning buffers."""
|
||||
def __init__(self, n_blocks=3, zero_conditioning=False, conditioning_dtype=torch.float32):
|
||||
super().__init__()
|
||||
self.blocks = nn.ModuleList([_Block() for _ in range(n_blocks)])
|
||||
pos = torch.zeros if zero_conditioning else torch.ones
|
||||
self.register_buffer("positive_conditioning", pos((2, 4), dtype=conditioning_dtype))
|
||||
self.register_buffer("negative_conditioning", torch.zeros((3, 4), dtype=conditioning_dtype))
|
||||
|
||||
|
||||
class _ModelInner:
|
||||
"""Inner model wrapper exposing `.diffusion_model`."""
|
||||
def __init__(self, diffusion_model):
|
||||
self.diffusion_model = diffusion_model
|
||||
|
||||
|
||||
class _ModelPatcher:
|
||||
"""ModelPatcher stub exposing `.model._ModelInner`."""
|
||||
def __init__(self, diffusion_model):
|
||||
self.model = _ModelInner(diffusion_model)
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_schema_exposes_model_passthrough_output():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
schema = nodes_seedvr.SeedVR2Conditioning.define_schema()
|
||||
assert [input_item.id for input_item in schema.inputs] == [
|
||||
"model",
|
||||
"vae_conditioning",
|
||||
]
|
||||
assert schema.inputs[1].display_name == "latent"
|
||||
assert [output.display_name for output in schema.outputs] == [
|
||||
"model",
|
||||
"positive",
|
||||
"negative",
|
||||
"latent",
|
||||
]
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_returns_packed_input_latent_deterministically():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
diffusion_model = _DiffusionModel()
|
||||
patcher = _ModelPatcher(diffusion_model)
|
||||
samples = torch.arange(1, 25, dtype=torch.float32).reshape(1, 2, 3, 2, 2)
|
||||
vae_conditioning = {"samples": samples}
|
||||
|
||||
_, first_positive, first_negative, first_latent = (
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher,
|
||||
vae_conditioning,
|
||||
)
|
||||
)
|
||||
_, second_positive, second_negative, second_latent = (
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher,
|
||||
vae_conditioning,
|
||||
)
|
||||
)
|
||||
|
||||
expected_latent = samples.reshape(1, 6, 2, 2)
|
||||
channel_last = samples.movedim(1, -1).contiguous()
|
||||
expected_condition = torch.cat(
|
||||
[
|
||||
channel_last,
|
||||
torch.ones((*channel_last.shape[:-1], 1)),
|
||||
],
|
||||
dim=-1,
|
||||
).movedim(-1, 1).reshape(1, 9, 2, 2)
|
||||
|
||||
assert torch.equal(first_latent["samples"], expected_latent)
|
||||
assert torch.equal(second_latent["samples"], expected_latent)
|
||||
assert torch.equal(
|
||||
first_positive[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
second_positive[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
first_negative[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
assert torch.equal(
|
||||
second_negative[0][1]["condition"],
|
||||
expected_condition,
|
||||
)
|
||||
finally:
|
||||
restore()
|
||||
|
||||
|
||||
def test_seedvr2_conditioning_fails_loud_on_zero_buffers():
|
||||
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
|
||||
try:
|
||||
diffusion_model = _DiffusionModel(zero_conditioning=True)
|
||||
patcher = _ModelPatcher(diffusion_model)
|
||||
vae_conditioning = {"samples": torch.zeros((1, 2, 1, 1, 1))}
|
||||
|
||||
with pytest.raises(RuntimeError) as excinfo:
|
||||
nodes_seedvr.SeedVR2Conditioning.execute(
|
||||
patcher, vae_conditioning,
|
||||
)
|
||||
|
||||
message = str(excinfo.value)
|
||||
assert message.startswith(
|
||||
nodes_seedvr._SEEDVR2_INVALID_MODEL_MSG_PREFIX
|
||||
), (
|
||||
"Fail-loud message must use the standard "
|
||||
"_SEEDVR2_INVALID_MODEL_MSG_PREFIX so callers/log scrapers "
|
||||
f"can match it. Got: {message!r}"
|
||||
)
|
||||
assert "positive_conditioning" in message
|
||||
assert "negative_conditioning" in message
|
||||
finally:
|
||||
restore()
|
||||
55
tests-unit/comfy_extras_test/test_seedvr2_nodes.py
Normal file
55
tests-unit/comfy_extras_test/test_seedvr2_nodes.py
Normal file
@@ -0,0 +1,55 @@
|
||||
import importlib
|
||||
import inspect
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
|
||||
def test_seedvr_node_signature_matches_schema():
|
||||
mock_mm = MagicMock()
|
||||
mock_mm.xformers_enabled.return_value = False
|
||||
mock_mm.xformers_enabled_vae.return_value = False
|
||||
mock_mm.sage_attention_enabled.return_value = False
|
||||
mock_mm.flash_attention_enabled.return_value = False
|
||||
|
||||
sentinel = object()
|
||||
prior_cpu = cli_args.cpu
|
||||
cli_args.cpu = True
|
||||
prior_module = sys.modules.get("comfy_extras.nodes_seedvr", sentinel)
|
||||
comfy_pkg = sys.modules.get("comfy")
|
||||
prior_mm_attr = getattr(comfy_pkg, "model_management", sentinel) if comfy_pkg else sentinel
|
||||
|
||||
with patch.dict(sys.modules, {"comfy.model_management": mock_mm}):
|
||||
if comfy_pkg is not None:
|
||||
setattr(comfy_pkg, "model_management", mock_mm)
|
||||
sys.modules.pop("comfy_extras.nodes_seedvr", None)
|
||||
try:
|
||||
nodes_seedvr = importlib.import_module("comfy_extras.nodes_seedvr")
|
||||
for node_cls in (nodes_seedvr.SeedVR2Preprocess, nodes_seedvr.SeedVR2PostProcessing, nodes_seedvr.SeedVR2Conditioning, nodes_seedvr.SeedVR2ProgressiveSampler):
|
||||
schema_ids = [i.id for i in node_cls.define_schema().inputs]
|
||||
exec_params = [
|
||||
p for p in inspect.signature(node_cls.execute).parameters.keys()
|
||||
if p != "cls"
|
||||
]
|
||||
assert schema_ids == exec_params, (
|
||||
f"{node_cls.__name__} schema/execute drift: "
|
||||
f"schema_ids={schema_ids}, exec_params={exec_params}"
|
||||
)
|
||||
finally:
|
||||
cli_args.cpu = prior_cpu
|
||||
if prior_module is sentinel:
|
||||
sys.modules.pop("comfy_extras.nodes_seedvr", None)
|
||||
else:
|
||||
sys.modules["comfy_extras.nodes_seedvr"] = prior_module
|
||||
if comfy_pkg is not None:
|
||||
if prior_mm_attr is sentinel:
|
||||
if hasattr(comfy_pkg, "model_management"):
|
||||
delattr(comfy_pkg, "model_management")
|
||||
else:
|
||||
setattr(comfy_pkg, "model_management", prior_mm_attr)
|
||||
57
tests-unit/comfy_extras_test/test_seedvr2_post_processing.py
Normal file
57
tests-unit/comfy_extras_test/test_seedvr2_post_processing.py
Normal file
@@ -0,0 +1,57 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
from comfy_extras import nodes_seedvr # noqa: E402
|
||||
|
||||
|
||||
def _schema_ids(items):
|
||||
return [item.id for item in items]
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_schema():
|
||||
schema = nodes_seedvr.SeedVR2PostProcessing.define_schema()
|
||||
|
||||
assert _schema_ids(schema.inputs) == ["images", "original_resized_images", "color_correction_method"]
|
||||
assert schema.inputs[2].options == ["lab", "wavelet", "adain", "none"]
|
||||
assert schema.inputs[2].default == "lab"
|
||||
assert schema.outputs[0].get_io_type() == "IMAGE"
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_oom_error_uses_color_correction_method(monkeypatch):
|
||||
decoded = torch.full((1, 3, 4, 4), 0.25)
|
||||
reference = torch.full((1, 3, 4, 4), 0.75)
|
||||
|
||||
def _lab(content, style):
|
||||
raise torch.cuda.OutOfMemoryError("CUDA out of memory")
|
||||
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "vae_device", lambda: torch.device("cpu"))
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "get_free_memory", lambda device: 1_000_000)
|
||||
monkeypatch.setattr(nodes_seedvr.comfy.model_management, "soft_empty_cache", lambda: None)
|
||||
|
||||
with patch.object(nodes_seedvr, "lab_color_transfer", _lab):
|
||||
try:
|
||||
nodes_seedvr.SeedVR2PostProcessing._color_transfer_chunked(
|
||||
decoded, reference, torch.device("cpu"), "lab",
|
||||
)
|
||||
except RuntimeError as exc:
|
||||
assert "color_correction_method=lab" in str(exc)
|
||||
assert " method=lab" not in str(exc)
|
||||
else:
|
||||
raise AssertionError("expected RuntimeError for one-frame LAB OOM")
|
||||
|
||||
|
||||
def test_seedvr2_post_processing_unknown_color_correction_method_raises():
|
||||
decoded = torch.zeros(1, 2, 4, 4, 3)
|
||||
original = torch.zeros(1, 2, 4, 4, 3)
|
||||
try:
|
||||
nodes_seedvr.SeedVR2PostProcessing.execute(decoded, original, "bogus")
|
||||
except ValueError as exc:
|
||||
assert "color_correction_method" in str(exc)
|
||||
else:
|
||||
raise AssertionError("expected ValueError for unknown color_correction_method")
|
||||
@@ -73,6 +73,24 @@ def _make_flux_schnell_comfyui_sd():
|
||||
return sd
|
||||
|
||||
|
||||
def _make_seedvr2_7b_separate_mm_sd():
|
||||
return {
|
||||
"blocks.35.mlp.vid.proj_in.weight": torch.empty(1, 3072),
|
||||
}
|
||||
|
||||
|
||||
def _make_seedvr2_7b_shared_mm_sd():
|
||||
return {
|
||||
"blocks.35.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
|
||||
}
|
||||
|
||||
|
||||
def _make_seedvr2_3b_shared_mm_sd():
|
||||
return {
|
||||
"blocks.31.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
|
||||
}
|
||||
|
||||
|
||||
class TestModelDetection:
|
||||
"""Verify that first-match model detection selects the correct model
|
||||
based on list ordering and unet_config specificity."""
|
||||
@@ -125,6 +143,48 @@ class TestModelDetection:
|
||||
assert model_config is not None
|
||||
assert type(model_config).__name__ == "FluxSchnell"
|
||||
|
||||
def test_seedvr2_7b_separate_mm_detection_config(self):
|
||||
sd = _make_seedvr2_7b_separate_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 3072
|
||||
assert unet_config["heads"] == 24
|
||||
assert unet_config["num_layers"] == 36
|
||||
assert unet_config["mm_layers"] == 36
|
||||
assert unet_config["mlp_type"] == "normal"
|
||||
assert unet_config["qk_rope"] is True
|
||||
assert unet_config["rope_type"] == "rope3d"
|
||||
assert unet_config["rope_dim"] == 64
|
||||
|
||||
def test_seedvr2_7b_shared_mm_detection_config(self):
|
||||
sd = _make_seedvr2_7b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 3072
|
||||
assert unet_config["heads"] == 24
|
||||
assert unet_config["num_layers"] == 36
|
||||
assert unet_config["mm_layers"] == 10
|
||||
assert unet_config["mlp_type"] == "swiglu"
|
||||
assert unet_config["qk_rope"] is True
|
||||
assert unet_config["rope_type"] == "rope3d"
|
||||
assert unet_config["rope_dim"] == 64
|
||||
|
||||
def test_seedvr2_3b_shared_mm_detection_config(self):
|
||||
sd = _make_seedvr2_3b_shared_mm_sd()
|
||||
unet_config = detect_unet_config(sd, "")
|
||||
|
||||
assert unet_config is not None
|
||||
assert unet_config["image_model"] == "seedvr2"
|
||||
assert unet_config["vid_dim"] == 2560
|
||||
assert unet_config["heads"] == 20
|
||||
assert unet_config["num_layers"] == 32
|
||||
assert unet_config["mlp_type"] == "swiglu"
|
||||
assert unet_config["qk_rope"] is None
|
||||
|
||||
def test_unet_config_and_required_keys_combination_is_unique(self):
|
||||
"""Each model in the registry must have a unique combination of
|
||||
``unet_config`` and ``required_keys``. If two models share the same
|
||||
|
||||
90
tests-unit/comfy_test/seedvr_vae_forward_test.py
Normal file
90
tests-unit/comfy_test/seedvr_vae_forward_test.py
Normal file
@@ -0,0 +1,90 @@
|
||||
"""Regression: ``comfy.ldm.seedvr.vae.VideoAutoencoderKL.forward`` must
|
||||
honor the actual tensor/tuple return contract of ``encode()`` and
|
||||
``decode_()`` and must NOT dereference diffusers-style ``.latent_dist``
|
||||
or ``.sample`` attributes on those returns.
|
||||
|
||||
The pre-fix body raised ``AttributeError: 'Tensor' object has no
|
||||
attribute 'latent_dist'`` for ``mode in {"encode", "all"}`` and
|
||||
``AttributeError: 'VideoAutoencoderKL' object has no attribute 'decode'``
|
||||
for ``mode == "decode"`` (the class only defines ``decode_`` with a
|
||||
trailing underscore). The post-fix body unwraps the optional one-element
|
||||
tuple shape that ``return_dict=False`` produces and returns the tensor
|
||||
directly.
|
||||
|
||||
Tests construct a stub subclass of ``VideoAutoencoderKL`` that bypasses
|
||||
the heavy ``__init__`` via ``torch.nn.Module.__init__(self)`` and
|
||||
overrides ``encode``/``decode_`` with known tensors so the contract can
|
||||
be probed without loading any real VAE weights.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
from comfy.ldm.seedvr.vae import VideoAutoencoderKL # noqa: E402
|
||||
|
||||
|
||||
_LATENT_SHAPE = (1, 16, 2, 2, 2)
|
||||
_DECODED_SHAPE = (1, 3, 5, 16, 16)
|
||||
_INPUT_ENCODE_SHAPE = (1, 3, 5, 16, 16)
|
||||
_INPUT_DECODE_SHAPE = (1, 16, 2, 2, 2)
|
||||
|
||||
|
||||
class _StubVAE(VideoAutoencoderKL):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self._encode_out = torch.zeros(*_LATENT_SHAPE)
|
||||
self._decode_out = torch.zeros(*_DECODED_SHAPE)
|
||||
|
||||
def encode(self, x, return_dict=True):
|
||||
return self._encode_out
|
||||
|
||||
def decode_(self, z, return_dict=True):
|
||||
return self._decode_out
|
||||
|
||||
|
||||
def test_forward_encode_returns_tensor():
|
||||
vae = _StubVAE()
|
||||
x = torch.zeros(*_INPUT_ENCODE_SHAPE)
|
||||
result = vae.forward(x, mode="encode")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_LATENT_SHAPE)
|
||||
|
||||
|
||||
def test_forward_decode_returns_tensor():
|
||||
vae = _StubVAE()
|
||||
z = torch.zeros(*_INPUT_DECODE_SHAPE)
|
||||
result = vae.forward(z, mode="decode")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_DECODED_SHAPE)
|
||||
|
||||
|
||||
class _TupleReturningStubVAE(VideoAutoencoderKL):
|
||||
"""Stub variant whose ``encode``/``decode_`` return the
|
||||
``(tensor,)`` one-element tuple shape ``return_dict=False`` produces
|
||||
in the parent class. Exercises the unwrap branch of
|
||||
``VideoAutoencoderKL.forward``.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self._encode_tensor = torch.zeros(*_LATENT_SHAPE)
|
||||
self._decode_tensor = torch.zeros(*_DECODED_SHAPE)
|
||||
|
||||
def encode(self, x, return_dict=True):
|
||||
return (self._encode_tensor,)
|
||||
|
||||
def decode_(self, z, return_dict=True):
|
||||
return (self._decode_tensor,)
|
||||
|
||||
|
||||
def test_forward_all_unwraps_one_tuple_at_each_step():
|
||||
vae = _TupleReturningStubVAE()
|
||||
x = torch.zeros(*_INPUT_ENCODE_SHAPE)
|
||||
result = vae.forward(x, mode="all")
|
||||
assert type(result) is torch.Tensor
|
||||
assert result.shape == torch.Size(_DECODED_SHAPE)
|
||||
47
tests-unit/comfy_test/test_seedvr2_dtype.py
Normal file
47
tests-unit/comfy_test/test_seedvr2_dtype.py
Normal file
@@ -0,0 +1,47 @@
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.sd
|
||||
import comfy.supported_models
|
||||
import comfy.ldm.seedvr.model as seedvr_model
|
||||
|
||||
|
||||
def test_seedvr2_fp16_manual_cast_only_for_bf16_device(monkeypatch):
|
||||
bf16_device = object()
|
||||
fp16_device = object()
|
||||
|
||||
monkeypatch.setattr(
|
||||
comfy.supported_models.comfy.model_management,
|
||||
"should_use_bf16",
|
||||
lambda device=None: device is bf16_device,
|
||||
)
|
||||
|
||||
bf16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"})
|
||||
bf16_config.set_inference_dtype(torch.float16, None, device=bf16_device)
|
||||
assert bf16_config.manual_cast_dtype is torch.bfloat16
|
||||
|
||||
fp16_config = comfy.supported_models.SeedVR2({"image_model": "seedvr2"})
|
||||
fp16_config.set_inference_dtype(torch.float16, None, device=fp16_device)
|
||||
assert fp16_config.manual_cast_dtype is None
|
||||
|
||||
|
||||
def test_seedvr2_text_conditioning_accepts_cfg1_single_branch():
|
||||
context = torch.arange(6, dtype=torch.float32).reshape(1, 3, 2)
|
||||
|
||||
txt, txt_shape = seedvr_model.NaDiT._resolve_text_conditioning(object(), context, [0])
|
||||
|
||||
torch.testing.assert_close(txt, context.squeeze(0))
|
||||
torch.testing.assert_close(txt_shape, torch.tensor([[3]], device=context.device))
|
||||
|
||||
|
||||
def test_seedvr2_vae_decode_memory_covers_full_frame_lab_transfer():
|
||||
estimate = comfy.sd._seedvr2_vae_decode_memory_used((1, 16, 26, 120, 160))
|
||||
old_estimate = 16 * 120 * 160 * (4 * 8 * 8) * 2
|
||||
|
||||
assert estimate == 101 * 960 * 1280 * 160
|
||||
assert estimate > 15 * 1024 ** 3
|
||||
assert estimate > old_estimate * 100
|
||||
341
tests-unit/comfy_test/test_seedvr2_internals.py
Normal file
341
tests-unit/comfy_test/test_seedvr2_internals.py
Normal file
@@ -0,0 +1,341 @@
|
||||
"""Consolidated SeedVR2 internals regression tests.
|
||||
|
||||
Sources (all merged verbatim, helper names disambiguated where colliding):
|
||||
|
||||
* RoPE rewrite — NaMMRotaryEmbedding3d.forward must match the legacy
|
||||
apply_rotary_emb wrapper oracle at fp32.
|
||||
* GroupNorm limit gate — causal_norm_wrapper at vae.py:509 must compare
|
||||
memory_occupy against get_norm_limit(), not float('inf').
|
||||
* SeedVR2 variable-length attention split-loop contract.
|
||||
|
||||
Pre-import CPU-only guard is required because comfy.ldm.seedvr.model and
|
||||
comfy.ldm.modules.attention transitively pull in comfy.model_management,
|
||||
which probes torch.cuda.current_device() at import time unless args.cpu is
|
||||
set first.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.model as seedvr_model # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
import comfy.ldm.modules.attention as attention # noqa: E402
|
||||
import comfy.ops as comfy_ops # noqa: E402
|
||||
from comfy.ldm.seedvr.model import ( # noqa: E402
|
||||
Cache,
|
||||
NaMMRotaryEmbedding3d,
|
||||
)
|
||||
from comfy.ldm.seedvr.vae import ( # noqa: E402
|
||||
causal_norm_wrapper,
|
||||
set_norm_limit,
|
||||
)
|
||||
from comfy.ldm.modules.attention import var_attention_optimized_split # noqa: E402
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# RoPE rewrite tests (test_seedvr_rope_rewrite.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# Test rig dimensions. dim=192 → per-axis rope dim = 64 (even, lucidrains
|
||||
# requirement). vid_shape=(2,4,4) → L_vid = 32. txt_shape=(8,) → L_txt = 8.
|
||||
_DIM = 192
|
||||
_HEADS = 4
|
||||
_VID_T, _VID_H, _VID_W = 2, 4, 4
|
||||
_TXT_L = 8
|
||||
_L_VID = _VID_T * _VID_H * _VID_W
|
||||
_SEED = 0
|
||||
|
||||
|
||||
def _make_inputs(dtype=torch.float32, device="cpu"):
|
||||
"""Construct the 6 forward inputs + cache. Deterministic via local
|
||||
Generator so global RNG state is not mutated.
|
||||
"""
|
||||
g = torch.Generator(device=device).manual_seed(_SEED)
|
||||
vid_q = torch.randn(_L_VID, _HEADS, _DIM, dtype=dtype, device=device, generator=g)
|
||||
vid_k = torch.randn(_L_VID, _HEADS, _DIM, dtype=dtype, device=device, generator=g)
|
||||
txt_q = torch.randn(_TXT_L, _HEADS, _DIM, dtype=dtype, device=device, generator=g)
|
||||
txt_k = torch.randn(_TXT_L, _HEADS, _DIM, dtype=dtype, device=device, generator=g)
|
||||
vid_shape = torch.tensor([[_VID_T, _VID_H, _VID_W]], dtype=torch.long, device=device)
|
||||
txt_shape = torch.tensor([[_TXT_L]], dtype=torch.long, device=device)
|
||||
cache = Cache(disable=True)
|
||||
return vid_q, vid_k, vid_shape, txt_q, txt_k, txt_shape, cache
|
||||
|
||||
|
||||
def _legacy_get_freqs(rope: NaMMRotaryEmbedding3d, vid_shape, txt_shape):
|
||||
"""Reproduce the pre-rewrite ``get_freqs`` body verbatim against
|
||||
``self.get_axial_freqs`` (parent ``RotaryEmbeddingBase`` method,
|
||||
unchanged by the rewrite).
|
||||
"""
|
||||
max_temporal = 0
|
||||
max_height = 0
|
||||
max_width = 0
|
||||
max_txt_len = 0
|
||||
for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()):
|
||||
max_temporal = max(max_temporal, l + f)
|
||||
max_height = max(max_height, h)
|
||||
max_width = max(max_width, w)
|
||||
max_txt_len = max(max_txt_len, l)
|
||||
with torch.amp.autocast(device_type="cuda", enabled=False):
|
||||
vid_freqs_full = rope.get_axial_freqs(
|
||||
min(max_temporal + 16, 1024),
|
||||
min(max_height + 4, 128),
|
||||
min(max_width + 4, 128),
|
||||
).float()
|
||||
txt_freqs_full = rope.get_axial_freqs(min(max_txt_len + 16, 1024))
|
||||
vid_freq_list, txt_freq_list = [], []
|
||||
for (f, h, w), l in zip(vid_shape.tolist(), txt_shape[:, 0].tolist()):
|
||||
vid_freq = vid_freqs_full[l : l + f, :h, :w].reshape(-1, vid_freqs_full.size(-1))
|
||||
txt_freq = txt_freqs_full[:l].repeat(1, 3).reshape(-1, vid_freqs_full.size(-1))
|
||||
vid_freq_list.append(vid_freq)
|
||||
txt_freq_list.append(txt_freq)
|
||||
return torch.cat(vid_freq_list, dim=0), torch.cat(txt_freq_list, dim=0)
|
||||
|
||||
|
||||
def _legacy_forward(rope: NaMMRotaryEmbedding3d, vid_q, vid_k, vid_shape,
|
||||
txt_q, txt_k, txt_shape):
|
||||
"""Compute expected forward output via the unchanged
|
||||
``apply_rotary_emb`` wrapper fed with legacy-shape freqs. This is the
|
||||
oracle. The wrapper itself is out of scope for the rewrite (Shape B).
|
||||
"""
|
||||
vid_freqs, txt_freqs = _legacy_get_freqs(rope, vid_shape, txt_shape)
|
||||
vid_freqs = vid_freqs.to(vid_q.device)
|
||||
txt_freqs = txt_freqs.to(txt_q.device)
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
vid_q = rearrange(vid_q, "L h d -> h L d")
|
||||
vid_k = rearrange(vid_k, "L h d -> h L d")
|
||||
vid_q_out = seedvr_model.apply_rotary_emb(vid_freqs, vid_q.float()).to(vid_q.dtype)
|
||||
vid_k_out = seedvr_model.apply_rotary_emb(vid_freqs, vid_k.float()).to(vid_k.dtype)
|
||||
vid_q_out = rearrange(vid_q_out, "h L d -> L h d")
|
||||
vid_k_out = rearrange(vid_k_out, "h L d -> L h d")
|
||||
|
||||
txt_q = rearrange(txt_q, "L h d -> h L d")
|
||||
txt_k = rearrange(txt_k, "L h d -> h L d")
|
||||
txt_q_out = seedvr_model.apply_rotary_emb(txt_freqs, txt_q.float()).to(txt_q.dtype)
|
||||
txt_k_out = seedvr_model.apply_rotary_emb(txt_freqs, txt_k.float()).to(txt_k.dtype)
|
||||
txt_q_out = rearrange(txt_q_out, "h L d -> L h d")
|
||||
txt_k_out = rearrange(txt_k_out, "h L d -> L h d")
|
||||
return vid_q_out, vid_k_out, txt_q_out, txt_k_out
|
||||
|
||||
|
||||
def test_namm_forward_output_tensor_equal_against_legacy_oracle():
|
||||
rope = NaMMRotaryEmbedding3d(dim=_DIM)
|
||||
vid_q, vid_k, vid_shape, txt_q, txt_k, txt_shape, cache = _make_inputs()
|
||||
|
||||
expected_vid_q, expected_vid_k, expected_txt_q, expected_txt_k = _legacy_forward(
|
||||
rope,
|
||||
vid_q.clone(), vid_k.clone(), vid_shape,
|
||||
txt_q.clone(), txt_k.clone(), txt_shape,
|
||||
)
|
||||
|
||||
actual_vid_q, actual_vid_k, actual_txt_q, actual_txt_k = rope.forward(
|
||||
vid_q.clone(), vid_k.clone(), vid_shape,
|
||||
txt_q.clone(), txt_k.clone(), txt_shape, cache,
|
||||
)
|
||||
|
||||
torch.testing.assert_close(actual_vid_q, expected_vid_q, rtol=0, atol=0,
|
||||
msg="vid_q output diverges from wrapper oracle")
|
||||
torch.testing.assert_close(actual_vid_k, expected_vid_k, rtol=0, atol=0,
|
||||
msg="vid_k output diverges from wrapper oracle")
|
||||
torch.testing.assert_close(actual_txt_q, expected_txt_q, rtol=0, atol=0,
|
||||
msg="txt_q output diverges from wrapper oracle")
|
||||
torch.testing.assert_close(actual_txt_k, expected_txt_k, rtol=0, atol=0,
|
||||
msg="txt_k output diverges from wrapper oracle")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# GroupNorm limit tests (test_seedvr_groupnorm_limit.py)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_NUM_CHANNELS = 8
|
||||
_NUM_GROUPS = 4
|
||||
_TENSOR_SHAPE = (1, 8, 2, 4, 4)
|
||||
|
||||
_GROUPNORM_SUBCLASSES = [
|
||||
pytest.param(comfy_ops.disable_weight_init.GroupNorm, id="disable_weight_init"),
|
||||
pytest.param(comfy_ops.manual_cast.GroupNorm, id="manual_cast"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("groupnorm_cls", _GROUPNORM_SUBCLASSES)
|
||||
def test_seedvr_groupnorm_low_limit_uses_chunked_groupnorm_path(groupnorm_cls):
|
||||
real_group_norm = vae_mod.F.group_norm
|
||||
set_norm_limit(1e-9)
|
||||
try:
|
||||
gn = groupnorm_cls(num_channels=_NUM_CHANNELS, num_groups=_NUM_GROUPS)
|
||||
gn.eval()
|
||||
|
||||
forward_hook_calls = []
|
||||
|
||||
def _hook(module, inputs, output):
|
||||
forward_hook_calls.append(tuple(inputs[0].shape))
|
||||
|
||||
spy_calls = []
|
||||
|
||||
def _group_norm_spy(input_tensor, num_groups_arg, *args, **kwargs):
|
||||
spy_calls.append({"num_groups": int(num_groups_arg)})
|
||||
return real_group_norm(input_tensor, num_groups_arg, *args, **kwargs)
|
||||
|
||||
handle = gn.register_forward_hook(_hook)
|
||||
try:
|
||||
with patch.object(vae_mod.F, "group_norm", side_effect=_group_norm_spy):
|
||||
out_tensor = causal_norm_wrapper(gn, torch.randn(*_TENSOR_SHAPE))
|
||||
finally:
|
||||
handle.remove()
|
||||
|
||||
full_calls = len(forward_hook_calls)
|
||||
chunked_calls = sum(1 for entry in spy_calls if entry["num_groups"] < _NUM_GROUPS)
|
||||
|
||||
assert tuple(int(s) for s in out_tensor.shape) == _TENSOR_SHAPE
|
||||
assert full_calls == 0, (
|
||||
f"low-limit GroupNorm gate must NOT take the full-forward path; got full_calls={full_calls}"
|
||||
)
|
||||
assert chunked_calls > 0, (
|
||||
f"low-limit GroupNorm gate must take the chunked path; got chunked_calls={chunked_calls}"
|
||||
)
|
||||
finally:
|
||||
set_norm_limit(None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# SeedVR2 var_attention split-loop tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def test_var_attention_registry_contains_always_available_entries():
|
||||
assert (
|
||||
attention.REGISTERED_ATTENTION_FUNCTIONS["var_attention_optimized_split"]
|
||||
is attention.var_attention_optimized_split
|
||||
)
|
||||
|
||||
|
||||
def test_seedvr2_7b_swin_attention_forward_uses_optimized_var_attention(monkeypatch):
|
||||
dim = 8
|
||||
heads = 2
|
||||
head_dim = 4
|
||||
attn = seedvr_model.NaSwinAttention(
|
||||
vid_dim=dim,
|
||||
txt_dim=dim,
|
||||
heads=heads,
|
||||
head_dim=head_dim,
|
||||
qk_bias=False,
|
||||
qk_norm=seedvr_model.CustomRMSNorm,
|
||||
qk_norm_eps=1e-6,
|
||||
rope_type=None,
|
||||
rope_dim=head_dim,
|
||||
shared_weights=False,
|
||||
window=(2, 1, 1),
|
||||
window_method="720pwin_by_size_bysize",
|
||||
version=True,
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
operations=comfy_ops.disable_weight_init,
|
||||
)
|
||||
generator = torch.Generator(device="cpu").manual_seed(11)
|
||||
vid = torch.randn(8, dim, generator=generator)
|
||||
txt = torch.randn(3, dim, generator=generator)
|
||||
vid_shape = torch.tensor([[2, 2, 2]], dtype=torch.long)
|
||||
txt_shape = torch.tensor([[3]], dtype=torch.long)
|
||||
calls = []
|
||||
|
||||
def fake_optimized_var_attention(**kwargs):
|
||||
calls.append(kwargs)
|
||||
return kwargs["q"]
|
||||
|
||||
monkeypatch.setattr(seedvr_model, "optimized_var_attention", fake_optimized_var_attention)
|
||||
|
||||
vid_out, txt_out = attn(vid, txt, vid_shape, txt_shape, seedvr_model.Cache(disable=True))
|
||||
|
||||
assert tuple(vid_out.shape) == (8, dim)
|
||||
assert tuple(txt_out.shape) == (3, dim)
|
||||
assert len(calls) == 1
|
||||
call = calls[0]
|
||||
assert tuple(call["q"].shape) == (14, heads, head_dim)
|
||||
assert tuple(call["k"].shape) == (14, heads, head_dim)
|
||||
assert tuple(call["v"].shape) == (14, heads, head_dim)
|
||||
assert call["heads"] == heads
|
||||
assert call["skip_reshape"] is True
|
||||
assert call["skip_output_reshape"] is True
|
||||
torch.testing.assert_close(
|
||||
call["cu_seqlens_q"],
|
||||
torch.tensor([0, 7, 14], dtype=torch.int32),
|
||||
rtol=0,
|
||||
atol=0,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
call["cu_seqlens_k"],
|
||||
torch.tensor([0, 7, 14], dtype=torch.int32),
|
||||
rtol=0,
|
||||
atol=0,
|
||||
)
|
||||
|
||||
|
||||
def test_var_attention_optimized_split_calls_dense_backend_per_window(monkeypatch):
|
||||
heads = 2
|
||||
head_dim = 3
|
||||
q = torch.arange(30, dtype=torch.float32).reshape(5, heads, head_dim)
|
||||
k = q + 100
|
||||
v = q + 200
|
||||
cu = torch.tensor([0, 2, 5], dtype=torch.int32)
|
||||
calls = []
|
||||
|
||||
def fake_optimized_attention(q_arg, k_arg, v_arg, heads_arg, **kwargs):
|
||||
calls.append(
|
||||
{
|
||||
"q_shape": tuple(q_arg.shape),
|
||||
"k_shape": tuple(k_arg.shape),
|
||||
"v_shape": tuple(v_arg.shape),
|
||||
"heads": heads_arg,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
)
|
||||
return q_arg + v_arg
|
||||
|
||||
monkeypatch.setattr(attention, "optimized_attention", fake_optimized_attention)
|
||||
|
||||
out = var_attention_optimized_split(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
heads,
|
||||
cu,
|
||||
cu,
|
||||
skip_reshape=True,
|
||||
skip_output_reshape=True,
|
||||
)
|
||||
|
||||
assert tuple(out.shape) == (5, heads, head_dim)
|
||||
assert len(calls) == 2
|
||||
assert calls[0]["q_shape"] == (1, heads, 2, head_dim)
|
||||
assert calls[1]["q_shape"] == (1, heads, 3, head_dim)
|
||||
assert all(call["heads"] == heads for call in calls)
|
||||
assert all(call["kwargs"]["skip_reshape"] is True for call in calls)
|
||||
assert all(call["kwargs"]["skip_output_reshape"] is True for call in calls)
|
||||
torch.testing.assert_close(out, q + v, rtol=0, atol=0)
|
||||
|
||||
|
||||
def test_var_attention_optimized_split_rejects_bad_offsets():
|
||||
q = torch.randn(5, 2, 3)
|
||||
cu_bad = torch.tensor([0, 2, 6], dtype=torch.int32)
|
||||
cu_ok = torch.tensor([0, 2, 5], dtype=torch.int32)
|
||||
|
||||
with pytest.raises(ValueError, match="cu_seqlens_q does not match token count"):
|
||||
var_attention_optimized_split(
|
||||
q,
|
||||
q,
|
||||
q,
|
||||
2,
|
||||
cu_bad,
|
||||
cu_ok,
|
||||
skip_reshape=True,
|
||||
skip_output_reshape=True,
|
||||
)
|
||||
308
tests-unit/comfy_test/test_seedvr2_model.py
Normal file
308
tests-unit/comfy_test/test_seedvr2_model.py
Normal file
@@ -0,0 +1,308 @@
|
||||
"""Consolidated SeedVR2 model/graph/forward regression tests.
|
||||
|
||||
Merged from:
|
||||
- seedvr_model_test.py
|
||||
- test_seedvr_7b_final_block_text_path.py
|
||||
- test_seedvr_forward_no_device_cast.py
|
||||
- test_seedvr_latent_format.py
|
||||
- test_seedvr2_vae_graph_boundaries.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from comfy.cli_args import args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
args.cpu = True
|
||||
|
||||
import comfy # noqa: E402
|
||||
import comfy.latent_formats # noqa: E402
|
||||
import comfy.ldm.seedvr.model # noqa: E402
|
||||
import comfy.ldm.seedvr.model as seedvr_model # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402
|
||||
import comfy.model_management # noqa: E402
|
||||
import comfy.sample # noqa: E402
|
||||
import comfy.sd as sd_mod # noqa: E402
|
||||
import nodes as nodes_mod # noqa: E402
|
||||
from comfy.ldm.seedvr.model import NaDiT # noqa: E402
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers from seedvr_model_test.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_standin(positive_conditioning):
|
||||
class _StandIn(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.register_buffer(
|
||||
"positive_conditioning", positive_conditioning
|
||||
)
|
||||
|
||||
_resolve_text_conditioning = NaDiT._resolve_text_conditioning
|
||||
|
||||
return _StandIn()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers from test_seedvr_7b_final_block_text_path.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _StubModule(nn.Module):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
|
||||
def _capture_last_layer_flags(monkeypatch, vid_dim: int, txt_in_dim: int) -> list[bool]:
|
||||
flags = []
|
||||
|
||||
class _Block(_StubModule):
|
||||
def __init__(self, *args, **kwargs):
|
||||
flags.append(kwargs["is_last_layer"])
|
||||
super().__init__()
|
||||
|
||||
monkeypatch.setattr(seedvr_model, "NaPatchIn", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "NaPatchOut", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "TimeEmbedding", _StubModule)
|
||||
monkeypatch.setattr(seedvr_model, "NaMMSRTransformerBlock", _Block)
|
||||
|
||||
seedvr_model.NaDiT(
|
||||
norm_eps=1e-5,
|
||||
qk_rope=None,
|
||||
num_layers=4,
|
||||
mlp_type="normal",
|
||||
vid_dim=vid_dim,
|
||||
txt_in_dim=txt_in_dim,
|
||||
heads=24,
|
||||
mm_layers=3,
|
||||
)
|
||||
|
||||
return flags
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers from test_seedvr_latent_format.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _Model:
|
||||
def __init__(self, latent_format):
|
||||
self._latent_format = latent_format
|
||||
|
||||
def get_model_object(self, name):
|
||||
assert name == "latent_format"
|
||||
return self._latent_format
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers from test_seedvr2_vae_graph_boundaries.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _Patcher:
|
||||
def get_free_memory(self, device):
|
||||
return 1024 * 1024 * 1024
|
||||
|
||||
|
||||
class _EncodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self, encoded):
|
||||
nn.Module.__init__(self)
|
||||
self.encoded = encoded
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.seen = []
|
||||
|
||||
def encode(self, x):
|
||||
self.seen.append(tuple(x.shape))
|
||||
return self.encoded.to(device=x.device, dtype=x.dtype)
|
||||
|
||||
|
||||
class _DecodeWrapper(seedvr_vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.calls = []
|
||||
|
||||
def decode(self, z, seedvr2_tiling=None):
|
||||
self.calls.append({"shape": tuple(z.shape), "seedvr2_tiling": seedvr2_tiling})
|
||||
if z.ndim == 4:
|
||||
b, tc, h, w = z.shape
|
||||
t = tc // 16
|
||||
else:
|
||||
b, _, t, h, w = z.shape
|
||||
return torch.zeros(b, 3, t, h * 8, w * 8, dtype=z.dtype, device=z.device)
|
||||
|
||||
|
||||
def _make_vae(wrapper):
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = wrapper
|
||||
vae.device = torch.device("cpu")
|
||||
vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.latent_channels = 16
|
||||
vae.latent_dim = 3
|
||||
vae.downscale_ratio = (lambda a: max(0, (a + 3) // 4), 8, 8)
|
||||
vae.upscale_ratio = (lambda a: max(0, a * 4 - 3), 8, 8)
|
||||
vae.output_channels = 3
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.crop_input = False
|
||||
vae.not_video = False
|
||||
vae.patcher = _Patcher()
|
||||
vae.process_input = lambda image: image
|
||||
vae.process_output = lambda image: image.add(1.0).div(2.0).clamp(0.0, 1.0)
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.memory_used_encode = lambda shape, dtype: 1
|
||||
vae.memory_used_decode = lambda shape, dtype: 1
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.vae_encode_crop_pixels = lambda pixels: pixels
|
||||
vae.spacial_compression_decode = lambda: 8
|
||||
vae.temporal_compression_decode = lambda: 4
|
||||
return vae
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests from seedvr_model_test.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_missing_context_falls_back_to_positive_buffer():
|
||||
"""AC: ``context is None`` falls back to the registered
|
||||
``positive_conditioning`` buffer and runs to completion — no
|
||||
silent zero substitution, no raised exception.
|
||||
"""
|
||||
pos_buffer = torch.full((58, 5120), 7.0)
|
||||
standin = _make_standin(pos_buffer)
|
||||
txt, txt_shape = standin._resolve_text_conditioning(None)
|
||||
assert txt.shape == (58, 5120)
|
||||
assert (txt == 7.0).all(), (
|
||||
"fallback path must use the positive_conditioning buffer "
|
||||
"verbatim, not a zero tensor"
|
||||
)
|
||||
assert txt_shape.shape == (1, 1)
|
||||
assert txt_shape[0, 0].item() == 58
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests from test_seedvr_7b_final_block_text_path.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_seedvr2_7b_keeps_final_block_text_path(monkeypatch):
|
||||
assert _capture_last_layer_flags(monkeypatch, vid_dim=3072, txt_in_dim=3072) == [
|
||||
False,
|
||||
False,
|
||||
False,
|
||||
False,
|
||||
]
|
||||
|
||||
|
||||
def test_seedvr2_7b_rope3d_matches_wrapper_oracle():
|
||||
rope = seedvr_model.get_na_rope("rope3d", dim=64)
|
||||
generator = torch.Generator(device="cpu").manual_seed(0)
|
||||
q = torch.randn(4, 2, 128, generator=generator)
|
||||
k = torch.randn(4, 2, 128, generator=generator)
|
||||
shape = torch.tensor([[1, 2, 2]], dtype=torch.long)
|
||||
freqs = rope.get_axial_freqs(1, 2, 2).reshape(4, -1)
|
||||
|
||||
expected_q = seedvr_model._apply_seedvr2_rotary_emb(
|
||||
freqs,
|
||||
q.permute(1, 0, 2).float(),
|
||||
).to(q.dtype).permute(1, 0, 2)
|
||||
expected_k = seedvr_model._apply_seedvr2_rotary_emb(
|
||||
freqs,
|
||||
k.permute(1, 0, 2).float(),
|
||||
).to(k.dtype).permute(1, 0, 2)
|
||||
|
||||
actual_q, actual_k = rope(q.clone(), k.clone(), shape, seedvr_model.Cache(disable=True))
|
||||
|
||||
torch.testing.assert_close(actual_q, expected_q, rtol=0, atol=0)
|
||||
torch.testing.assert_close(actual_k, expected_k, rtol=0, atol=0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests from test_seedvr_latent_format.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_seedvr2_latent_format_uses_16_channels_without_3d_empty_latent_expansion():
|
||||
latent_format = comfy.latent_formats.SeedVR2()
|
||||
latent_image = torch.zeros(1, 1, 4, 5)
|
||||
|
||||
fixed = comfy.sample.fix_empty_latent_channels(_Model(latent_format), latent_image)
|
||||
|
||||
assert latent_format.latent_channels == 16
|
||||
assert latent_format.latent_dimensions == 2
|
||||
assert fixed.shape == (1, 16, 4, 5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests from test_seedvr2_vae_graph_boundaries.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_seedvr2_encode_and_encode_tiled_preserve_native_latent_contract(monkeypatch):
|
||||
monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None)
|
||||
|
||||
encoded = torch.full((1, 16, 2, 4, 5), 2.0)
|
||||
vae = _make_vae(_EncodeWrapper(encoded))
|
||||
pixels = torch.zeros(1, 5, 32, 40, 3)
|
||||
|
||||
node_output = nodes_mod.VAEEncode().encode(vae, pixels)[0]
|
||||
node_latent = node_output["samples"]
|
||||
assert set(node_output) == {"samples"}
|
||||
assert tuple(node_latent.shape) == (1, 16, 2, 4, 5)
|
||||
assert node_latent.dtype == torch.float32
|
||||
assert node_latent.stride()[-1] == 1
|
||||
assert torch.equal(node_latent, torch.full_like(node_latent, 2.0 * 0.9152))
|
||||
|
||||
tiled = torch.full((1, 16, 2, 4, 5), 3.0)
|
||||
monkeypatch.setattr(seedvr_vae_mod, "tiled_vae", MagicMock(return_value=tiled))
|
||||
tiled_output = nodes_mod.VAEEncodeTiled().encode(
|
||||
vae,
|
||||
pixels,
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=16,
|
||||
temporal_overlap=4,
|
||||
)[0]
|
||||
tiled_latent = tiled_output["samples"]
|
||||
assert set(tiled_output) == {"samples"}
|
||||
assert tuple(tiled_latent.shape) == (1, 16, 2, 4, 5)
|
||||
assert tiled_latent.dtype == torch.float32
|
||||
assert torch.equal(tiled_latent, torch.full_like(tiled_latent, 3.0 * 0.9152))
|
||||
|
||||
|
||||
def test_vaedecode_tiled_visible_inputs_are_seedvr2_decode_tiling_authority(monkeypatch):
|
||||
monkeypatch.setattr(sd_mod.model_management, "load_models_gpu", lambda *a, **k: None)
|
||||
vae = _make_vae(_DecodeWrapper())
|
||||
|
||||
nodes_mod.VAEDecodeTiled().decode(
|
||||
vae,
|
||||
{"samples": torch.zeros(1, 16, 2, 4, 5)},
|
||||
tile_size=512,
|
||||
overlap=64,
|
||||
temporal_size=16,
|
||||
temporal_overlap=4,
|
||||
)
|
||||
|
||||
assert vae.first_stage_model.calls == [
|
||||
{
|
||||
"shape": (1, 16, 2, 4, 5),
|
||||
"seedvr2_tiling": {
|
||||
"enable_tiling": True,
|
||||
"tile_size": (512, 512),
|
||||
"tile_overlap": (64, 64),
|
||||
"temporal_size": 16,
|
||||
"temporal_overlap": 4,
|
||||
},
|
||||
}
|
||||
]
|
||||
91
tests-unit/comfy_test/test_seedvr2_vae_decode.py
Normal file
91
tests-unit/comfy_test/test_seedvr2_vae_decode.py
Normal file
@@ -0,0 +1,91 @@
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
from comfy_extras import nodes_seedvr # noqa: E402
|
||||
|
||||
|
||||
def _make_wrapper() -> vae_mod.VideoAutoencoderKLWrapper:
|
||||
wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
nn.Module.__init__(wrapper)
|
||||
return wrapper
|
||||
|
||||
|
||||
def _fingerprint_decode_(self, z, return_dict=True):
|
||||
b = int(z.shape[0])
|
||||
t = int(z.shape[2])
|
||||
h = int(z.shape[3])
|
||||
w = int(z.shape[4])
|
||||
out = torch.empty(b, 3, t, h * 8, w * 8)
|
||||
for batch_idx in range(b):
|
||||
out[batch_idx].fill_(float(batch_idx + 1))
|
||||
return out
|
||||
|
||||
|
||||
def _decode_with_patches(wrapper, z):
|
||||
with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _fingerprint_decode_):
|
||||
return wrapper.decode(z)
|
||||
|
||||
|
||||
def test_decode_b2_t3_multi_frame_batch_unchanged():
|
||||
wrapper = _make_wrapper()
|
||||
|
||||
out = _decode_with_patches(wrapper, torch.zeros(2, 16 * 3, 2, 2))
|
||||
|
||||
assert tuple(out.shape) == (2, 3, 3, 16, 16)
|
||||
|
||||
|
||||
class _Wrapper(vae_mod.VideoAutoencoderKLWrapper):
|
||||
def __init__(self):
|
||||
nn.Module.__init__(self)
|
||||
self.calls = []
|
||||
|
||||
def parameters(self):
|
||||
return iter([torch.nn.Parameter(torch.zeros(()))])
|
||||
|
||||
def _decode_stub(self, latent):
|
||||
self.calls.append(tuple(latent.shape))
|
||||
return torch.zeros(latent.shape[0], 3, latent.shape[2], latent.shape[3] * 8, latent.shape[4] * 8)
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_decode_accepts_5d_channel_first_latents_without_preprocessor_state():
|
||||
wrapper = _Wrapper()
|
||||
|
||||
with patch.object(vae_mod.VideoAutoencoderKL, "decode_", _decode_stub):
|
||||
out = wrapper.decode(torch.zeros(1, 16, 2, 4, 5))
|
||||
|
||||
assert tuple(out.shape) == (1, 3, 2, 32, 40)
|
||||
assert wrapper.calls == [(1, 16, 2, 4, 5)]
|
||||
|
||||
|
||||
def test_seedvr2_wrapper_decode_rejects_wrong_rank_latents():
|
||||
wrapper = _Wrapper()
|
||||
|
||||
with pytest.raises(RuntimeError, match=r"latent input must be 4-D collapsed .* or 5-D"):
|
||||
wrapper.decode(torch.zeros(1, 16, 4))
|
||||
|
||||
|
||||
def _t_padded(t_in: int) -> int:
|
||||
if t_in == 1:
|
||||
return 1
|
||||
if t_in <= 4:
|
||||
return 5
|
||||
if (t_in - 1) % 4 == 0:
|
||||
return t_in
|
||||
return t_in + (4 - ((t_in - 1) % 4))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("t_in", [1, 5, 9])
|
||||
def test_t_padded_matches_cut_videos(t_in):
|
||||
dummy = torch.zeros(1, t_in, 1, 1, 1)
|
||||
assert nodes_seedvr.cut_videos(dummy).shape[1] == _t_padded(t_in)
|
||||
347
tests-unit/comfy_test/test_seedvr2_vae_tiled.py
Normal file
347
tests-unit/comfy_test/test_seedvr2_vae_tiled.py
Normal file
@@ -0,0 +1,347 @@
|
||||
from contextlib import ExitStack
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.ldm.seedvr.vae as vae_mod # noqa: E402
|
||||
import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402
|
||||
import comfy.sd as sd_mod # noqa: E402
|
||||
from comfy.ldm.seedvr.vae import MemoryState, tiled_vae # noqa: E402
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# From test_seedvr_vae_tiled_decode_latent_min_size_override.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_runtime_decode_zero_temporal_size_disables_slicing_for_call():
|
||||
from comfy.ldm.seedvr.vae import MemoryState, VideoAutoencoderKL, tiled_vae
|
||||
|
||||
class StubVAEModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.slicing_latent_min_size = 2
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self.use_slicing = True
|
||||
self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
self.decode_min_sizes = []
|
||||
self.memory_states = []
|
||||
|
||||
def decode_(self, t_chunk):
|
||||
self.decode_min_sizes.append(self.slicing_latent_min_size)
|
||||
return VideoAutoencoderKL.slicing_decode(self, t_chunk)
|
||||
|
||||
def _decode(self, z, memory_state=MemoryState.DISABLED):
|
||||
self.memory_states.append(memory_state)
|
||||
b, c, d, h, w = z.shape
|
||||
return torch.zeros((b, 3, d, h * 8, w * 8), dtype=z.dtype)
|
||||
|
||||
vae = StubVAEModel()
|
||||
z = torch.zeros((1, 16, 5, 8, 8), dtype=torch.float32)
|
||||
|
||||
tiled_vae(
|
||||
z,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=0,
|
||||
temporal_overlap=0,
|
||||
encode=False,
|
||||
)
|
||||
|
||||
assert vae.decode_min_sizes == [5]
|
||||
assert vae.memory_states == [MemoryState.DISABLED]
|
||||
assert vae.slicing_latent_min_size == 2
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# From test_seedvr_vae_tiled_encode_runt_slice_override.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_zero_temporal_size_preserves_min_size_when_encode_raises():
|
||||
from comfy.ldm.seedvr.vae import tiled_vae
|
||||
|
||||
class RaisingVAEModel(torch.nn.Module):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.slicing_sample_min_size = 4
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self._dummy = torch.nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
|
||||
def encode(self, t_chunk):
|
||||
raise RuntimeError("simulated encode failure")
|
||||
|
||||
vae = RaisingVAEModel()
|
||||
x = torch.zeros((1, 3, 12, 64, 64), dtype=torch.float32)
|
||||
|
||||
raised = False
|
||||
try:
|
||||
tiled_vae(
|
||||
x,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=0,
|
||||
temporal_overlap=0,
|
||||
encode=True,
|
||||
)
|
||||
except RuntimeError as exc:
|
||||
if "simulated encode failure" not in str(exc):
|
||||
raise
|
||||
raised = True
|
||||
|
||||
assert raised
|
||||
assert vae.slicing_sample_min_size == 4
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# From test_seedvr_vae_tiled_temporal_slicing.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _SlicingDecodeVAE(nn.Module):
|
||||
def __init__(self, slicing_latent_min_size):
|
||||
super().__init__()
|
||||
self.slicing_latent_min_size = slicing_latent_min_size
|
||||
self.spatial_downsample_factor = 8
|
||||
self.temporal_downsample_factor = 4
|
||||
self.device = torch.device("cpu")
|
||||
self.use_slicing = True
|
||||
self._dummy = nn.Parameter(torch.zeros(1, dtype=torch.float32))
|
||||
self.decode_min_sizes = []
|
||||
self.memory_states = []
|
||||
|
||||
def decode_(self, z):
|
||||
self.decode_min_sizes.append(self.slicing_latent_min_size)
|
||||
return vae_mod.VideoAutoencoderKL.slicing_decode(self, z)
|
||||
|
||||
def _decode(self, z, memory_state=MemoryState.DISABLED):
|
||||
self.memory_states.append(memory_state)
|
||||
x = z[:, :1].repeat(
|
||||
1,
|
||||
3,
|
||||
1,
|
||||
self.spatial_downsample_factor,
|
||||
self.spatial_downsample_factor,
|
||||
)
|
||||
return x
|
||||
|
||||
|
||||
def test_decode_tiled_vae_maps_temporal_args_to_latent_slicing_min_size():
|
||||
vae = _SlicingDecodeVAE(slicing_latent_min_size=2)
|
||||
z = torch.arange(1 * 16 * 5 * 8 * 8, dtype=torch.float32).reshape(1, 16, 5, 8, 8)
|
||||
|
||||
tiled_vae(
|
||||
z,
|
||||
vae,
|
||||
tile_size=(64, 64),
|
||||
tile_overlap=(0, 0),
|
||||
temporal_size=12,
|
||||
temporal_overlap=4,
|
||||
encode=False,
|
||||
)
|
||||
|
||||
assert vae.decode_min_sizes == [2]
|
||||
assert vae.memory_states == [MemoryState.INITIALIZING, MemoryState.ACTIVE]
|
||||
assert vae.slicing_latent_min_size == 2
|
||||
|
||||
wrapper = vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
nn.Module.__init__(wrapper)
|
||||
seedvr2_tiling = {
|
||||
"enable_tiling": True,
|
||||
"tile_size": (64, 64),
|
||||
"tile_overlap": (0, 0),
|
||||
"temporal_size": 8,
|
||||
"temporal_overlap": 7,
|
||||
}
|
||||
|
||||
captured = {}
|
||||
|
||||
def _fake_tiled_vae(latent, model, **kwargs):
|
||||
captured.update(kwargs)
|
||||
return torch.zeros(1, 3, 1, 16, 16)
|
||||
|
||||
with patch.object(vae_mod, "tiled_vae", side_effect=_fake_tiled_vae):
|
||||
wrapper.decode(torch.zeros(1, 16, 2, 2), seedvr2_tiling=seedvr2_tiling)
|
||||
|
||||
assert captured["temporal_overlap"] == 7
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# From test_vae_decode_tiled_dispatcher_seedvr2_4d.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _force_oom(*a, **k):
|
||||
raise torch.cuda.OutOfMemoryError("forced OOM for dispatcher test")
|
||||
|
||||
|
||||
def _make_vae(first_stage_model, latent_channels, latent_dim):
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = first_stage_model
|
||||
vae.patcher = MagicMock()
|
||||
vae.patcher.get_free_memory = MagicMock(return_value=8 * 1024 * 1024 * 1024)
|
||||
vae.device = vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.upscale_ratio = vae.downscale_ratio = 8
|
||||
vae.upscale_index_formula = vae.downscale_index_formula = None
|
||||
vae.output_channels = 3
|
||||
vae.latent_channels = latent_channels
|
||||
vae.latent_dim = latent_dim
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.spacial_compression_decode = lambda: 8
|
||||
vae.process_input = lambda x: x
|
||||
vae.process_output = lambda x: x
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.memory_used_decode = lambda *a, **k: 1
|
||||
return vae
|
||||
|
||||
|
||||
def _dispatch(vae, samples, seedvr2_call, generic_call, patch_wrapper_decode):
|
||||
mm = sd_mod.model_management
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(patch.object(mm, "raise_non_oom", lambda e: None))
|
||||
stack.enter_context(patch.object(mm, "load_models_gpu", lambda *a, **k: None))
|
||||
stack.enter_context(patch.object(mm, "soft_empty_cache", lambda: None))
|
||||
stack.enter_context(patch.object(sd_mod.VAE, "decode_tiled_seedvr2", seedvr2_call))
|
||||
stack.enter_context(patch.object(sd_mod.VAE, "decode_tiled_", generic_call))
|
||||
if patch_wrapper_decode:
|
||||
stack.enter_context(patch.object(
|
||||
seedvr_vae_mod.VideoAutoencoderKLWrapper, "decode",
|
||||
side_effect=_force_oom))
|
||||
vae.decode(samples)
|
||||
|
||||
|
||||
def test_4d_seedvr2_latent_routes_to_decode_tiled_seedvr2():
|
||||
wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
seedvr_vae_mod.VideoAutoencoderKLWrapper)
|
||||
vae = _make_vae(wrapper, latent_channels=16, latent_dim=3)
|
||||
seedvr2_call = MagicMock(return_value=torch.zeros(1, 3, 9, 64, 64))
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, 3, 64, 64))
|
||||
_dispatch(vae, torch.zeros(1, 16 * 3, 8, 8), seedvr2_call, generic_call, True)
|
||||
assert seedvr2_call.call_count == 1
|
||||
assert generic_call.call_count == 0
|
||||
|
||||
|
||||
def test_4d_non_seedvr2_latent_still_routes_to_generic_decode_tiled():
|
||||
first_stage = MagicMock()
|
||||
first_stage.decode = MagicMock(side_effect=_force_oom)
|
||||
vae = _make_vae(first_stage, latent_channels=4, latent_dim=2)
|
||||
seedvr2_call = MagicMock(return_value=torch.zeros(1, 3, 9, 64, 64))
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, 3, 64, 64))
|
||||
_dispatch(vae, torch.zeros(1, 4, 8, 8), seedvr2_call, generic_call, False)
|
||||
assert generic_call.call_count == 1
|
||||
assert seedvr2_call.call_count == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# From test_vae_encode_tiled_fallback_dispatcher_seedvr2.py
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _populate_common_vae_attrs_fallback(vae):
|
||||
vae.patcher = MagicMock()
|
||||
vae.patcher.get_free_memory = MagicMock(return_value=8 * 1024 * 1024 * 1024)
|
||||
vae.device = torch.device("cpu")
|
||||
vae.output_device = torch.device("cpu")
|
||||
vae.vae_dtype = torch.float32
|
||||
vae.disable_offload = True
|
||||
vae.extra_1d_channel = None
|
||||
vae.upscale_ratio = 8
|
||||
vae.upscale_index_formula = None
|
||||
vae.output_channels = 3
|
||||
vae.latent_channels = 16
|
||||
vae.latent_dim = 3
|
||||
vae.downscale_ratio = 8
|
||||
vae.downscale_index_formula = None
|
||||
vae.not_video = False
|
||||
vae.crop_input = False
|
||||
vae.pad_channel_value = None
|
||||
|
||||
vae.vae_output_dtype = lambda: torch.float32
|
||||
vae.spacial_compression_encode = lambda: 8
|
||||
vae.process_input = lambda x: x
|
||||
vae.process_output = lambda x: x
|
||||
vae.throw_exception_if_invalid = lambda: None
|
||||
vae.memory_used_encode = lambda *a, **k: 1
|
||||
|
||||
|
||||
def _make_seedvr2_vae_fallback():
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(
|
||||
seedvr_vae_mod.VideoAutoencoderKLWrapper
|
||||
)
|
||||
vae.first_stage_model = wrapper
|
||||
_populate_common_vae_attrs_fallback(vae)
|
||||
return vae
|
||||
|
||||
|
||||
def _make_non_seedvr2_vae_fallback():
|
||||
vae = sd_mod.VAE.__new__(sd_mod.VAE)
|
||||
vae.first_stage_model = MagicMock()
|
||||
_populate_common_vae_attrs_fallback(vae)
|
||||
return vae
|
||||
|
||||
|
||||
def _force_regular_encode_oom(*args, **kwargs):
|
||||
raise torch.cuda.OutOfMemoryError("forced OOM for dispatcher test")
|
||||
|
||||
|
||||
def test_seedvr2_3d_routes_to_encode_tiled_seedvr2_on_oom():
|
||||
vae = _make_seedvr2_vae_fallback()
|
||||
pixel_samples = torch.zeros((1, 8, 64, 64, 3))
|
||||
|
||||
seedvr2_call = MagicMock(return_value=torch.zeros(1, 16, 2, 8, 8))
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, 16, 2, 8, 8))
|
||||
|
||||
with patch.object(sd_mod.model_management, "raise_non_oom",
|
||||
lambda e: None), \
|
||||
patch.object(sd_mod.model_management, "load_models_gpu",
|
||||
lambda *a, **k: None), \
|
||||
patch.object(sd_mod.model_management, "soft_empty_cache",
|
||||
lambda: None), \
|
||||
patch.object(seedvr_vae_mod.VideoAutoencoderKLWrapper, "encode",
|
||||
side_effect=_force_regular_encode_oom), \
|
||||
patch.object(sd_mod.VAE, "encode_tiled_seedvr2", seedvr2_call,
|
||||
create=True), \
|
||||
patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call):
|
||||
vae.encode(pixel_samples)
|
||||
|
||||
assert seedvr2_call.call_count == 1, (
|
||||
f"Expected encode_tiled_seedvr2 to be called once for a SeedVR2 3D "
|
||||
f"input under OOM fallback; got {seedvr2_call.call_count} calls."
|
||||
)
|
||||
assert generic_call.call_count == 0, (
|
||||
f"encode_tiled_3d must NOT be called for a SeedVR2 input; got "
|
||||
f"{generic_call.call_count} calls."
|
||||
)
|
||||
|
||||
|
||||
def test_non_seedvr2_encode_tiled_3d_default_overlap_is_concrete():
|
||||
vae = _make_non_seedvr2_vae_fallback()
|
||||
vae.downscale_ratio = (lambda a: max(1, a // 4), 8, 8)
|
||||
vae.upscale_ratio = (lambda a: a * 4, 8, 8)
|
||||
generic_call = MagicMock(return_value=torch.zeros(1, 16, 2, 8, 8))
|
||||
pixel_samples = torch.zeros((1, 8, 64, 64, 3))
|
||||
|
||||
with patch.object(sd_mod.model_management, "load_models_gpu",
|
||||
lambda *a, **k: None), \
|
||||
patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call):
|
||||
vae.encode_tiled(pixel_samples)
|
||||
|
||||
assert generic_call.call_args.kwargs["overlap"] == (1, 64, 64)
|
||||
126
tests-unit/comfy_test/test_seedvr_progressive_sampler.py
Normal file
126
tests-unit/comfy_test/test_seedvr_progressive_sampler.py
Normal file
@@ -0,0 +1,126 @@
|
||||
"""Unit tests for ``comfy_extras.nodes_seedvr.SeedVR2ProgressiveSampler``."""
|
||||
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from comfy.cli_args import args as cli_args
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
cli_args.cpu = True
|
||||
|
||||
import comfy.sample # noqa: E402
|
||||
import comfy_extras.nodes_seedvr as nodes_seedvr_mod # noqa: E402
|
||||
from comfy_extras.nodes_seedvr import SeedVR2ProgressiveSampler # noqa: E402
|
||||
|
||||
_LAT_C = 16
|
||||
_COND_C = 17
|
||||
|
||||
|
||||
def _make_inputs(B: int = 1, T: int = 5, H: int = 8, W: int = 8):
|
||||
"""Build minimal SeedVR2-shaped sampling inputs."""
|
||||
samples_5d = torch.arange(
|
||||
B * _LAT_C * T * H * W, dtype=torch.float32
|
||||
).reshape(B, _LAT_C, T, H, W)
|
||||
samples = samples_5d.reshape(B, _LAT_C * T, H, W).contiguous()
|
||||
|
||||
cond_5d = torch.arange(
|
||||
B * _COND_C * T * H * W, dtype=torch.float32
|
||||
).reshape(B, _COND_C, T, H, W) + 10000.0
|
||||
cond = cond_5d.reshape(B, _COND_C * T, H, W).contiguous()
|
||||
|
||||
text_pos = torch.zeros(1, 4, 32)
|
||||
text_neg = torch.zeros(1, 4, 32)
|
||||
positive = [[text_pos, {"condition": cond.clone()}]]
|
||||
negative = [[text_neg, {"condition": cond.clone()}]]
|
||||
latent_image = {"samples": samples}
|
||||
return latent_image, positive, negative, samples_5d, cond_5d
|
||||
|
||||
|
||||
def _identity_fix_empty(model, latent_image, downscale_ratio_spacial=None):
|
||||
return latent_image
|
||||
|
||||
|
||||
def _fingerprinted_prepare_noise(latent_image, seed, batch_inds=None):
|
||||
"""Return a tensor whose values encode ``(seed, position)``."""
|
||||
base = torch.arange(
|
||||
latent_image.numel(), dtype=torch.float32
|
||||
).reshape(latent_image.shape)
|
||||
return base + float(seed) * 1e6
|
||||
|
||||
|
||||
def test_progressive_sampler_schema_exposes_manual_default_auto_chunking():
|
||||
schema = SeedVR2ProgressiveSampler.define_schema()
|
||||
inputs = {item.id: item for item in schema.inputs}
|
||||
|
||||
assert inputs["chunking_mode"].options == ["manual", "auto"]
|
||||
assert inputs["chunking_mode"].default == "manual"
|
||||
|
||||
|
||||
def test_auto_chunking_walks_two_three_four_chunk_ladder():
|
||||
"""Auto mode must walk 2-, 3-, then 4-chunk geometries on OOM."""
|
||||
latent, pos, neg, _, _ = _make_inputs(T=17)
|
||||
calls = []
|
||||
|
||||
def _oom_until_four_chunks(model, noise, steps, cfg, sampler_name,
|
||||
scheduler, positive, negative,
|
||||
latent_image, denoise=1.0,
|
||||
noise_mask=None, seed=None):
|
||||
calls.append(tuple(latent_image.shape))
|
||||
if latent_image.shape[1] > _LAT_C * 5:
|
||||
raise torch.cuda.OutOfMemoryError("chunk too large")
|
||||
return latent_image.clone()
|
||||
|
||||
with patch.object(comfy.sample, "sample",
|
||||
side_effect=_oom_until_four_chunks), \
|
||||
patch.object(comfy.sample, "fix_empty_latent_channels",
|
||||
side_effect=_identity_fix_empty), \
|
||||
patch.object(comfy.sample, "prepare_noise",
|
||||
side_effect=_fingerprinted_prepare_noise), \
|
||||
patch.object(nodes_seedvr_mod.comfy.model_management,
|
||||
"soft_empty_cache") as soft_empty:
|
||||
out = SeedVR2ProgressiveSampler.execute(
|
||||
model=None, seed=0, steps=2, cfg=1.0,
|
||||
sampler_name="euler", scheduler="simple",
|
||||
positive=pos, negative=neg, latent=latent,
|
||||
denoise=1.0, frames_per_chunk=65, temporal_overlap=0,
|
||||
chunking_mode="auto",
|
||||
)
|
||||
|
||||
assert calls[:4] == [
|
||||
(1, _LAT_C * 17, 8, 8),
|
||||
(1, _LAT_C * 9, 8, 8),
|
||||
(1, _LAT_C * 6, 8, 8),
|
||||
(1, _LAT_C * 5, 8, 8),
|
||||
]
|
||||
assert torch.equal(out.result[0]["samples"], latent["samples"])
|
||||
assert soft_empty.call_count == 3
|
||||
|
||||
|
||||
@pytest.mark.parametrize("bad_chunk", [0, -1, 2])
|
||||
def test_t3_invalid_frames_per_chunk_raises_value_error(bad_chunk):
|
||||
"""``frames_per_chunk`` violating 4n+1 (or <1) must raise ``ValueError`` before any model invocation."""
|
||||
latent, pos, neg, _, _ = _make_inputs(T=5)
|
||||
|
||||
sampler_called = {"n": 0}
|
||||
|
||||
def _should_not_be_called(*args, **kwargs):
|
||||
sampler_called["n"] += 1
|
||||
return torch.zeros(1)
|
||||
|
||||
with patch.object(comfy.sample, "sample",
|
||||
side_effect=_should_not_be_called), \
|
||||
patch.object(comfy.sample, "fix_empty_latent_channels",
|
||||
side_effect=_identity_fix_empty), \
|
||||
patch.object(comfy.sample, "prepare_noise",
|
||||
side_effect=_fingerprinted_prepare_noise):
|
||||
with pytest.raises(ValueError) as excinfo:
|
||||
SeedVR2ProgressiveSampler.execute(
|
||||
model=None, seed=0, steps=2, cfg=1.0,
|
||||
sampler_name="euler", scheduler="simple",
|
||||
positive=pos, negative=neg, latent=latent,
|
||||
denoise=1.0, frames_per_chunk=bad_chunk, temporal_overlap=0,
|
||||
)
|
||||
assert str(bad_chunk) in str(excinfo.value)
|
||||
assert sampler_called["n"] == 0
|
||||
Reference in New Issue
Block a user