Merge branch 'master' into matt/be-1899-listassets-contract-fields

This commit is contained in:
Matt Miller
2026-07-29 10:52:35 -07:00
committed by GitHub
131 changed files with 13920 additions and 865 deletions

View File

@@ -24,6 +24,28 @@ def app(model_manager):
app.add_routes(routes)
return app
async def test_get_model_folders_includes_registered_extensions(aiohttp_client, app, tmp_path):
"""Folders expose their registered extension set verbatim; an empty list
means match-all (filter_files_extensions semantics)."""
with patch('folder_paths.folder_names_and_paths', {
'test_checkpoints': ([str(tmp_path)], {'.safetensors', '.ckpt'}),
'test_configs': ([str(tmp_path)], ['.yaml']),
'test_match_all': ([str(tmp_path)], set()),
'configs': ([str(tmp_path)], ['.yaml']),
}):
client = await aiohttp_client(app)
response = await client.get('/experiment/models')
assert response.status == 200
folders = {f['name']: f for f in await response.json()}
assert 'configs' not in folders # blocklisted
assert folders['test_checkpoints']['folders'] == [str(tmp_path)]
assert folders['test_checkpoints']['extensions'] == ['.ckpt', '.safetensors']
assert folders['test_configs']['extensions'] == ['.yaml']
# Match-all registrations are exposed honestly, not substituted.
assert folders['test_match_all']['extensions'] == []
async def test_get_model_preview_safetensors(aiohttp_client, app, tmp_path):
img = Image.new('RGB', (100, 100), 'white')
img_byte_arr = BytesIO()

View File

@@ -2,11 +2,12 @@ import pytest
import torch
import tempfile
import os
import sys
import av
import io
from fractions import Fraction
from comfy_api.input_impl.video_types import VideoFromFile, VideoFromComponents
from comfy_api.util.video_types import VideoComponents
from comfy_api.util.video_types import VideoComponents, VideoContainer, VideoCodec
from comfy_api.input.basic_types import AudioInput
from av.error import InvalidDataError
@@ -237,3 +238,526 @@ def test_duration_consistency(video_components):
manual_duration = float(components.images.shape[0] / components.frame_rate)
assert duration == pytest.approx(manual_duration)
def create_transcode_source(
width=64, height=64, frames=30, fps=30, audio_streams=1, undecodable_audio=0, rotation=False,
container_format="mov", audio_codec="pcm_s16le",
):
"""Create a temp video that save_to must transcode (mpeg4 video, so codec != h264).
``undecodable_audio`` trailing PCM streams get their fourcc corrupted so no decoder exists
(``codec_context is None``), like the APAC track in iPhone spatial-audio recordings.
``rotation`` patches a 90-degree display matrix into the video track header.
"""
buffer = io.BytesIO()
with av.open(buffer, mode="w", format=container_format) as container:
video_stream = container.add_stream("mpeg4", rate=fps)
video_stream.width = width
video_stream.height = height
video_stream.pix_fmt = "yuv420p"
audio = []
for _ in range(audio_streams + undecodable_audio):
stream = container.add_stream(audio_codec, rate=44100)
stream.sample_rate = 44100
audio.append(stream)
for i in range(frames):
frame = av.VideoFrame.from_ndarray(
torch.full((height, width, 3), (i * 7) % 256, dtype=torch.uint8).numpy(),
format="rgb24",
)
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
# write audio in 1024-sample frames, like real decoders produce, so the
# per-frame skip/cap logic in the transcode path actually runs
for stream in audio:
for offset in range(0, 44100 * frames // fps, 1024):
n = min(1024, 44100 * frames // fps - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, n, dtype=torch.int16).numpy(), format="s16", layout="mono"
)
audio_frame.sample_rate = 44100
audio_frame.pts = offset
container.mux(stream.encode(audio_frame))
for stream in [video_stream, *audio]:
container.mux(stream.encode(None))
data = bytearray(buffer.getvalue())
end = len(data)
for _ in range(undecodable_audio):
end = data.rindex(b"sowt", 0, end)
data[end:end + 4] = b"Xpac"
if rotation:
# the 3x3 display matrix sits 40 bytes into the version-0 tkhd payload; first tkhd
# inside moov = video track (search from moov so mdat bytes can't false-match)
matrix_offset = data.index(b"tkhd", data.rindex(b"moov")) + 4 + 40
values = [0, 1 << 16, 0, -(1 << 16), 0, 0, 0, 0, 1 << 30]
data[matrix_offset:matrix_offset + 36] = b"".join(v.to_bytes(4, "big", signed=True) for v in values)
tmp = tempfile.NamedTemporaryFile(suffix=f".{container_format}", delete=False)
tmp.write(bytes(data))
tmp.close()
return tmp.name
def transcode_and_probe(video):
buffer = io.BytesIO()
video.save_to(buffer, format=VideoContainer.MP4, codec=VideoCodec.H264)
buffer.seek(0)
with av.open(buffer) as container:
video_stream = container.streams.video[0]
audio_stream = container.streams.audio[0] if container.streams.audio else None
frames = 0
first_pts = None
for packet in container.demux(video_stream):
for frame in packet.decode():
if first_pts is None:
first_pts = frame.pts
frames += 1
return {
"codec": video_stream.codec_context.name,
"width": video_stream.codec_context.width,
"height": video_stream.codec_context.height,
"frames": frames,
"first_pts": first_pts,
"video_seconds": float(video_stream.duration * video_stream.time_base) if video_stream.duration else None,
"audio_seconds": float(audio_stream.duration * audio_stream.time_base)
if audio_stream and audio_stream.duration else None,
"audio_codecs": [s.codec_context.name for s in container.streams.audio],
}
def test_save_to_transcode_streams_without_buffering_frames():
"""Transcoding must not decode the whole video into memory first (~2 GiB for this source)"""
resource = pytest.importorskip("resource") # no getrusage on Windows
rss_scale = 1 if sys.platform == "darwin" else 1024 # ru_maxrss: bytes on macOS, KiB elsewhere
# ru_maxrss is a lifetime peak: a heavier test running earlier would shrink the measured
# delta and quietly defang this canary, so keep this source the biggest thing in the suite
file_path = create_transcode_source(width=640, height=480, frames=300)
try:
rss_before = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale
result = transcode_and_probe(VideoFromFile(file_path))
rss_delta = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * rss_scale - rss_before
assert result["codec"] == "h264"
assert result["frames"] == 300
assert rss_delta < 500 * 2**20, f"transcode buffered frames in RAM (peak grew {rss_delta / 2**20:.0f} MiB)"
finally:
os.unlink(file_path)
def test_save_to_transcode_honors_trim_window():
"""start_time/duration trim applies to both video and audio on the streaming path"""
file_path = create_transcode_source(frames=90) # 3s @ 30fps
try:
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
assert result["frames"] == pytest.approx(30, abs=2)
assert result["first_pts"] == 0 # trimmed output is rebased to start at zero
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
finally:
os.unlink(file_path)
def test_save_to_transcode_keeps_audio_of_sparse_video():
"""Audio that runs ahead of a sparse video track (slideshows, timelapses) must be
kept in full — it is only clamped to the video's end, never to the video cursor."""
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
for t in (0, 30, 60): # 3 frames spread over 60 seconds
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), t * 4, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = t * 15360
frame.time_base = Fraction(1, 15360)
container.mux(video_stream.encode(frame))
container.mux(video_stream.encode(None))
for offset in range(0, 48000 * 60, 1024):
n = min(1024, 48000 * 60 - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
)
audio_frame.sample_rate = 48000
audio_frame.pts = offset
audio_frame.time_base = Fraction(1, 48000)
container.mux(audio_stream.encode(audio_frame))
container.mux(audio_stream.encode(None))
buffer.seek(0)
result = transcode_and_probe(VideoFromFile(buffer))
assert result["audio_seconds"] == pytest.approx(60.0, abs=1.0)
def test_save_to_transcode_vfr_audio_covers_video_span():
"""A trim window in the sparse region of a VFR file keeps audio for the true pts span
of the kept frames. Deriving the span as frames/average_rate undercuts it badly: the
average is dominated by the dense region (and can be plain wrong on MediaRecorder files)."""
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=48000, layout="stereo")
# 10 frames inside the first second, then one every 1.25 s
for i, t in enumerate([x / 10 for x in range(10)] + [1.0, 2.25, 3.5, 4.75]):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 16) % 256, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = int(t * 15360)
frame.time_base = Fraction(1, 15360)
container.mux(video_stream.encode(frame))
container.mux(video_stream.encode(None))
for offset in range(0, 48000 * 6, 1024):
n = min(1024, 48000 * 6 - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(2, n, dtype=torch.float32).numpy(), format="fltp", layout="stereo"
)
audio_frame.sample_rate = 48000
audio_frame.pts = offset
audio_frame.time_base = Fraction(1, 48000)
container.mux(audio_stream.encode(audio_frame))
container.mux(audio_stream.encode(None))
buffer.seek(0)
result = transcode_and_probe(VideoFromFile(buffer, start_time=1, duration=5))
# kept frames: 1.0/2.25/3.5/4.75 s -> rebased span 3.75 s + one nominal interval
assert result["frames"] == 4
assert result["audio_seconds"] == pytest.approx(4.0, abs=0.45)
def test_save_to_transcode_trims_audio_in_stream_time_base_units():
"""Matroska audio timestamps tick in 1/1000, not 1/sample_rate; trim and audio timing
must convert through the frame's time base instead of assuming sample units. AAC audio,
because it decodes straight to the encoder's format and hits the resampler passthrough
that keeps the source time base on the frames."""
file_path = create_transcode_source(frames=90, container_format="matroska", audio_codec="aac")
try:
result = transcode_and_probe(VideoFromFile(file_path, start_time=1, duration=1))
assert result["audio_codecs"] == ["aac"]
assert result["video_seconds"] == pytest.approx(1.0, abs=0.1)
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
finally:
os.unlink(file_path)
def test_save_to_transcode_learns_unprobed_audio_params():
"""mpegts is only probed a few seconds deep at open, so an audio stream whose first
packet comes later (live captures where audio kicks in late) still has sample_rate 0
when the transcode starts; the parameters must be learned from the stream itself."""
sample_rate, fps, video_seconds, audio_start = 48000, 30, 13, 12
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mpegts") as container:
video_stream = container.add_stream("mpeg4", rate=fps)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
for i in range(video_seconds * fps):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
)
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
for offset in range(0, (video_seconds - audio_start) * sample_rate, 1024):
n = min(1024, (video_seconds - audio_start) * sample_rate - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
)
audio_frame.sample_rate = sample_rate
audio_frame.pts = audio_start * sample_rate + offset
container.mux(audio_stream.encode(audio_frame))
for stream in (video_stream, audio_stream):
container.mux(stream.encode(None))
buffer.seek(0)
with av.open(buffer) as container:
# the scenario requires unprobed parameters; if a future FFmpeg probes deeper,
# push audio_start/video_seconds further out to restore it
assert container.streams.audio[0].codec_context.sample_rate == 0
result = transcode_and_probe(VideoFromFile(buffer))
assert result["frames"] == video_seconds * fps
assert result["audio_codecs"] == ["aac"]
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
buffer.seek(0)
trimmed_before_audio = transcode_and_probe(VideoFromFile(buffer, duration=1))
assert trimmed_before_audio["frames"] == fps
assert trimmed_before_audio["audio_codecs"] == []
assert trimmed_before_audio["audio_seconds"] is None
buffer.seek(0)
trimmed_crossing_audio = transcode_and_probe(VideoFromFile(buffer, start_time=11.5, duration=1))
assert trimmed_crossing_audio["frames"] == fps
assert trimmed_crossing_audio["audio_codecs"] == ["aac"]
assert trimmed_crossing_audio["video_seconds"] == pytest.approx(1.0, abs=0.05)
assert trimmed_crossing_audio["audio_seconds"] == pytest.approx(0.5, abs=0.1)
def test_save_to_transcode_trimmed_fragmented_mp4_keeps_audio():
"""Fragmented mp4 (MediaRecorder, DASH/HLS-derived files) delivers audio well behind
video, so when the trim window's last video frame arrives the audio demuxed so far
does not cover the window yet; the transcode must keep demuxing audio until it does
instead of finalizing on the first audio frame it sees afterwards."""
sample_rate, fps, seconds = 48000, 30, 6
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4", options={"movflags": "frag_keyframe+empty_moov"}) as container:
video_stream = container.add_stream("h264", rate=fps)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
next_audio_pts = 0
for i in range(seconds * fps):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
)
container.mux(video_stream.encode(frame.reformat(format="yuv420p")))
while next_audio_pts / sample_rate <= i / fps: # feed audio alongside, like a live pipeline
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, 1024, dtype=torch.float32).numpy(), format="fltp", layout="mono"
)
audio_frame.sample_rate = sample_rate
audio_frame.pts = next_audio_pts
container.mux(audio_stream.encode(audio_frame))
next_audio_pts += 1024
for stream in (video_stream, audio_stream):
container.mux(stream.encode(None))
result = transcode_and_probe(VideoFromFile(buffer, start_time=0.5, duration=1.0))
assert result["video_seconds"] == pytest.approx(1.0, abs=0.05)
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.05)
def test_save_to_transcode_sparse_video_keeps_true_duration():
"""average_rate is not a frame duration: a 3-frame video spanning 60 s averages
0.05 fps, and padding the last frame with 1/average_rate used to extend the
output — and the audio kept with it — about 20 s past the source span."""
sample_rate = 48000
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
audio_stream = container.add_stream("aac", rate=sample_rate, layout="mono")
for i, second in enumerate((0, 30, 60)):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), i * 80, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = second * 30
frame.time_base = Fraction(1, 30)
container.mux(video_stream.encode(frame))
for offset in range(0, 90 * sample_rate, 1024):
n = min(1024, 90 * sample_rate - offset)
audio_frame = av.AudioFrame.from_ndarray(
torch.zeros(1, n, dtype=torch.float32).numpy(), format="fltp", layout="mono"
)
audio_frame.sample_rate = sample_rate
audio_frame.pts = offset
container.mux(audio_stream.encode(audio_frame))
for stream in (video_stream, audio_stream):
container.mux(stream.encode(None))
result = transcode_and_probe(VideoFromFile(buffer))
assert result["frames"] == 3
# the last frame keeps its true stts duration (1/30 s), not 1/average_rate (~20 s)
assert result["video_seconds"] == pytest.approx(60.03, abs=0.05)
assert result["audio_seconds"] == pytest.approx(60.03, abs=0.1)
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=45))
assert trimmed["frames"] == 2
# a kept frame whose source duration crosses the window end is clamped to it
assert trimmed["video_seconds"] == pytest.approx(45.0, abs=0.05)
assert trimmed["audio_seconds"] == pytest.approx(45.0, abs=0.1)
def test_save_to_transcode_clamps_final_pts_to_declared_stream_duration():
"""Some iPhone MOVs report a video stream duration that ends before the final
decoded frame's nominal duration. A transcode must not turn that trailing
timestamp quirk into an extra frame interval compared to the source/remux path."""
fps = 30
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=fps)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
for i, pts in enumerate([*range(31), 32]):
frame = av.VideoFrame.from_ndarray(
torch.full((64, 64, 3), (i * 7) % 256, dtype=torch.uint8).numpy(), format="rgb24"
).reformat(format="yuv420p")
frame.pts = pts
frame.time_base = Fraction(1, fps)
container.mux(video_stream.encode(frame))
container.mux(video_stream.encode(None))
class _StreamProxy:
def __init__(self, stream, duration):
self._stream = stream
self.duration = duration
def __getattr__(self, name):
return getattr(self._stream, name)
class _StreamsProxy:
def __init__(self, video_stream):
self.video = [video_stream]
self.audio = []
class _PacketProxy:
def __init__(self, packet, stream):
self._packet = packet
self.stream = stream
def __getattr__(self, name):
return getattr(self._packet, name)
class _ContainerProxy:
def __init__(self, container, stream):
self._container = container
self._stream = stream
self.streams = _StreamsProxy(stream)
def __getattr__(self, name):
return getattr(self._container, name)
def demux(self, *streams):
for packet in self._container.demux(self._stream._stream):
yield _PacketProxy(packet, self._stream)
buffer.seek(0)
output = io.BytesIO()
with av.open(buffer) as container:
real_stream = container.streams.video[0]
declared_duration = 32 * int(round((1 / fps) / real_stream.time_base))
stream = _StreamProxy(real_stream, declared_duration)
VideoFromFile(buffer)._save_transcoded(
_ContainerProxy(container, stream), output, VideoContainer.MP4, VideoCodec.H264, None, 8
)
output.seek(0)
with av.open(output) as container:
video_stream = container.streams.video[0]
frames = [f for p in container.demux(video_stream) for f in p.decode()]
assert len(frames) == 32
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(32 / fps, abs=0.01)
assert float(frames[-1].pts * frames[-1].time_base) == pytest.approx(31 / fps, abs=0.01)
def test_save_to_transcode_irregular_vfr_keeps_span():
"""B-frames reorder packets, and mp4 sample durations follow decode order: the dts
timeline ends before the pts timeline, so an irregular-VFR source's tail holds fell
out of the container (this 20.23 s span used to come out as 15.27 s, and the 10 s
trim as 6.03 s). The transcode encodes without B-frames so every sample keeps its
true display duration."""
durations = [1, 1, 60, 1, 1, 120, 1, 180, 1, 1, 150, 90] # 1/30 s ticks, span 20.2333 s
generator = torch.Generator().manual_seed(7)
buffer = io.BytesIO()
with av.open(buffer, mode="w", format="mp4") as container:
video_stream = container.add_stream("mpeg4", rate=30)
video_stream.width = video_stream.height = 64
video_stream.pix_fmt = "yuv420p"
pts = 0
for duration in durations:
# textured frames, so an encoder with default settings has B-frames to gain from
frame = av.VideoFrame.from_ndarray(
torch.randint(0, 255, (64, 64, 3), generator=generator, dtype=torch.uint8).numpy(),
format="rgb24",
).reformat(format="yuv420p")
frame.pts = pts
frame.time_base = Fraction(1, 30)
pts += duration
for packet in video_stream.encode(frame):
packet.duration = duration # exact stts in the source
container.mux(packet)
container.mux(video_stream.encode(None))
result = transcode_and_probe(VideoFromFile(buffer))
assert result["frames"] == len(durations)
assert result["video_seconds"] == pytest.approx(sum(durations) / 30, abs=0.05)
trimmed = transcode_and_probe(VideoFromFile(buffer, duration=10))
assert trimmed["frames"] == 8 # frames at 12.167 s+ fall outside the window
assert trimmed["video_seconds"] == pytest.approx(10.0, abs=0.05)
def test_save_to_transcode_trim_survives_missing_leading_pts():
"""A trim should survive pts-less kept frames followed by a real-pts frame past the window."""
nulled_frames = 0
class _PacketProxy:
def __init__(self, packet):
self._packet = packet
def __getattr__(self, name):
return getattr(self._packet, name)
@property
def stream(self):
return self._packet.stream
def decode(self):
nonlocal nulled_frames
frames = self._packet.decode()
for frame in frames:
if nulled_frames < 2:
frame.pts = None
nulled_frames += 1
return frames
class _ContainerProxy:
def __init__(self, real):
self._real = real
def __getattr__(self, name):
return getattr(self._real, name)
def demux(self, *streams):
for packet in self._real.demux(*streams):
yield _PacketProxy(packet)
file_path = create_transcode_source(frames=10, audio_streams=0)
try:
buffer = io.BytesIO()
with av.open(file_path) as container:
# 0.05 s window: both pts-less frames are kept (synthesized pts 0 and 512),
# and the first real-pts frame (1024 ticks) already lies past end_pts (768)
VideoFromFile(file_path, duration=0.05)._save_transcoded(
_ContainerProxy(container), buffer, VideoContainer.MP4, VideoCodec.H264, None, 8
)
assert nulled_frames == 2
buffer.seek(0)
with av.open(buffer) as container:
video_stream = container.streams.video[0]
frames = [f for p in container.demux(video_stream) for f in p.decode()]
assert len(frames) == 2
assert float(video_stream.duration * video_stream.time_base) == pytest.approx(2 / 30, abs=0.01)
finally:
os.unlink(file_path)
def test_save_to_transcode_bakes_rotation():
"""A 90-degree display-matrix rotation swaps the output dimensions (portrait video)"""
file_path = create_transcode_source(width=64, height=32, rotation=True)
try:
result = transcode_and_probe(VideoFromFile(file_path))
assert (result["width"], result["height"]) == (32, 64)
assert result["frames"] == 30
finally:
os.unlink(file_path)
def test_save_to_transcode_skips_undecodable_audio():
"""Streaming transcode keeps the decodable audio track and drops undecodable ones;
with no decodable audio at all the output is video-only instead of crashing."""
mixed = all_bad = None
try:
mixed = create_transcode_source(audio_streams=1, undecodable_audio=1)
all_bad = create_transcode_source(audio_streams=0, undecodable_audio=2)
result = transcode_and_probe(VideoFromFile(mixed))
assert result["audio_codecs"] == ["aac"]
assert result["audio_seconds"] == pytest.approx(1.0, abs=0.1)
assert transcode_and_probe(VideoFromFile(all_bad))["audio_codecs"] == []
finally:
for path in (mixed, all_bad):
if path:
os.unlink(path)

View File

@@ -0,0 +1,186 @@
"""SeedVR2 conditioning node regression tests."""
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
from comfy.ldm.seedvr.constants import SEEDVR2_LATENT_CHANNELS
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:
importlib.import_module("comfy")
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):
def __init__(self):
super().__init__()
self.freqs = nn.Parameter(torch.zeros(4))
class _Block(nn.Module):
def __init__(self):
super().__init__()
self.rope = _Rope()
class _DiffusionModel(nn.Module):
def __init__(self, n_blocks=3, conditioning_dtype=torch.float32):
super().__init__()
self.blocks = nn.ModuleList([_Block() for _ in range(n_blocks)])
self.register_buffer("positive_conditioning", torch.ones((2, 4), dtype=conditioning_dtype))
self.register_buffer("negative_conditioning", torch.zeros((3, 4), dtype=conditioning_dtype))
class _ModelInner:
def __init__(self, diffusion_model):
self.diffusion_model = diffusion_model
class _ModelPatcher:
def __init__(self, diffusion_model):
self.model = _ModelInner(diffusion_model)
def test_seedvr2_conditioning_schema_exposes_conditioning_outputs():
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] == [
"positive",
"negative",
]
finally:
restore()
def test_seedvr2_conditioning_rejects_wrong_latent_channels():
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
try:
patcher = _ModelPatcher(_DiffusionModel())
vae_conditioning = {"samples": torch.zeros(1, 8, 2, 2, 2)}
with pytest.raises(ValueError, match=f"{SEEDVR2_LATENT_CHANNELS} channels"):
nodes_seedvr.SeedVR2Conditioning.execute(patcher, vae_conditioning)
finally:
restore()
def test_seedvr2_conditioning_returns_conditioning_deterministically():
nodes_seedvr, restore = _import_nodes_seedvr_isolated()
try:
diffusion_model = _DiffusionModel()
patcher = _ModelPatcher(diffusion_model)
samples = torch.arange(
1,
1 + SEEDVR2_LATENT_CHANNELS * 3 * 2 * 2,
dtype=torch.float32,
).reshape(1, SEEDVR2_LATENT_CHANNELS, 3, 2, 2)
vae_conditioning = {"samples": samples}
first_positive, first_negative = (
nodes_seedvr.SeedVR2Conditioning.execute(
patcher,
vae_conditioning,
)
)
second_positive, second_negative = (
nodes_seedvr.SeedVR2Conditioning.execute(
patcher,
vae_conditioning,
)
)
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)
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()

View 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):
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)

View File

@@ -0,0 +1,51 @@
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
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)
with patch.object(nodes_seedvr, "lab_color_transfer", _lab):
with pytest.raises(RuntimeError) as excinfo:
nodes_seedvr.SeedVR2PostProcessing._color_transfer_chunked(
decoded, reference, torch.device("cpu"), "lab",
)
assert "color_correction_method=lab" in str(excinfo.value)
assert " method=lab" not in str(excinfo.value)
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)
with pytest.raises(ValueError) as excinfo:
nodes_seedvr.SeedVR2PostProcessing.execute(decoded, original, "bogus")
assert "color_correction_method" in str(excinfo.value)

View File

@@ -0,0 +1,77 @@
"""SeedVR2 temporal chunk/merge node regression tests."""
import pytest
import torch
from comfy.cli_args import args as cli_args
from comfy.ldm.seedvr.constants import (
BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE,
SEEDVR2_CHUNK_GIB_PER_MPX_FRAME,
SEEDVR2_CHUNK_RESERVED_GIB,
SEEDVR2_CHUNK_SIGMA_GIB,
SEEDVR2_CHUNK_SIGMA_K,
SEEDVR2_LATENT_CHANNELS,
)
if not torch.cuda.is_available():
cli_args.cpu = True
import comfy.model_management # noqa: E402
from comfy_extras.nodes_seedvr import SeedVR2TemporalChunk, SeedVR2TemporalMerge, _seedvr2_chunk_crossfade_weights # noqa: E402
def _latent(t_latent, h=8, w=8, b=1):
g = torch.Generator().manual_seed(7)
return {"samples": torch.randn(b, SEEDVR2_LATENT_CHANNELS, t_latent, h, w, generator=g)}
def _split(latent, frames_per_chunk, temporal_overlap, chunking_mode="manual"):
combo = {"chunking_mode": chunking_mode}
if chunking_mode != "auto":
combo["frames_per_chunk"] = frames_per_chunk
return SeedVR2TemporalChunk.execute(latent, temporal_overlap, combo).args
def _merge(chunks, temporal_overlap):
return SeedVR2TemporalMerge.execute(chunks, [temporal_overlap]).args[0]
def test_chunk_temporal_windows_and_validation():
with pytest.raises(ValueError, match="4n\\+1"):
_split(_latent(9), 20, 0)
with pytest.raises(ValueError, match="5-D"):
_split({"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS * 9, 8, 8)}, 21, 0)
with pytest.raises(ValueError, match="chunking_mode"):
_split(_latent(13), 21, 0, "adaptive")
latent = _latent(13)
chunks, overlap = _split(latent, 21, 2) # chunk_latent=6, step=4 -> [0:6], [4:10], [8:13]
assert overlap == 2 and [c["samples"].shape[2] for c in chunks] == [6, 6, 5]
assert all(torch.equal(c["samples"], latent["samples"][:, :, s:e]) for c, (s, e) in zip(chunks, [(0, 6), (4, 10), (8, 13)]))
assert len(_split(_latent(13), 21, 999)[0]) == 8 # overlap clamps to chunk_latent-1 -> step=1
assert (r := _split(_latent(5), 21, 3)) and len(r[0]) == 1 and r[1] == 0 # t_pixel <= 21: passthrough
def test_chunk_auto_mode_applies_vram_law(monkeypatch):
mpx_per_frame = (32 * 32) * (BYTEDANCE_VAE_SPATIAL_DOWNSAMPLE ** 2) / 1e6
free_gb = (
SEEDVR2_CHUNK_RESERVED_GIB
+ SEEDVR2_CHUNK_SIGMA_K * SEEDVR2_CHUNK_SIGMA_GIB
+ 5.1 * SEEDVR2_CHUNK_GIB_PER_MPX_FRAME * mpx_per_frame
)
monkeypatch.setattr(comfy.model_management, "get_free_memory", lambda dev=None: free_gb * (1024 ** 3))
assert [c["samples"].shape[2] for c in _split(_latent(13, h=32, w=32), 1, 0, "auto")[0]] == [5, 5, 3]
assert _split(_latent(13, h=32, w=32, b=2), 1, 0, "auto")[0][0]["samples"].shape[2] == 2 # batch halves the chunk
def test_merge_crossfade_and_reassembly():
latent = _latent(13)
latent["noise_mask"] = torch.rand(1, 1, 13, 8, 8)
latent["batch_index"] = [0]
merged = _merge(_split(latent, 21, 0)[0], 0)
assert torch.equal(merged["samples"], latent["samples"])
assert "noise_mask" not in merged and merged["batch_index"] == [0]
assert torch.allclose(_merge(_split(latent, 21, 3)[0], 3)["samples"], latent["samples"], atol=1e-6)
w = _seedvr2_chunk_crossfade_weights(3, merged["samples"].device, merged["samples"].dtype)
assert w[0] == 1.0 and w[-1] == 0.0 and torch.all(w[:-1] >= w[1:])
ones, zeros = {"samples": torch.ones(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}, {"samples": torch.zeros(1, SEEDVR2_LATENT_CHANNELS, 6, 8, 8)}
fused = _merge([ones, zeros], 3)["samples"] # overlap equals w: prev fades out, next fades in
assert torch.equal(fused[:, :, 3:6], w.view(1, 1, 3, 1, 1).expand(1, SEEDVR2_LATENT_CHANNELS, 3, 8, 8))
assert torch.equal(fused[:, :, :3], ones["samples"][:, :, :3]) and torch.equal(fused[:, :, 6:], zeros["samples"][:, :, :3])
short = _split(latent, 21, 2)[0]
short[0]["samples"] = short[0]["samples"][:, :, :4]
with pytest.raises(ValueError, match="only the final chunk may be shorter"):
_merge(short, 2)

View File

@@ -15,7 +15,7 @@ if not has_gpu():
args.cpu = True
from comfy import ops
from comfy.quant_ops import QuantizedTensor
from comfy.quant_ops import QUANT_ALGOS, QuantizedTensor
import comfy.utils
@@ -283,7 +283,59 @@ class TestMixedPrecisionOps(unittest.TestCase):
saved = model.state_dict()
saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes())
self.assertTrue(saved_conf["convrot"])
def test_convrot_w4a4_loads_into_params(self):
"""ConvRot W4A4 checkpoints must load as the dedicated kitchen layout."""
if "convrot_w4a4" not in QUANT_ALGOS:
self.skipTest("comfy_kitchen does not provide ConvRot W4A4")
torch.manual_seed(456)
layer_quant_config = {
"layer": {
"format": "convrot_w4a4",
"convrot_groupsize": 256,
"linear_dtype": "int8",
}
}
weight = torch.randn(16, 256, dtype=torch.bfloat16)
bias = torch.randn(16, dtype=torch.bfloat16)
q_weight = QuantizedTensor.from_float(
weight,
"TensorCoreConvRotW4A4Layout",
convrot_groupsize=256,
quant_group_size=64,
)
state_dict = {
"layer.weight": q_weight._qdata,
"layer.bias": bias,
"layer.weight_scale": q_weight._params.scale,
}
state_dict, _ = comfy.utils.convert_old_quants(
state_dict,
metadata={"_quantization_metadata": json.dumps({"layers": layer_quant_config})},
)
model = torch.nn.Module()
model.layer = ops.mixed_precision_ops({}).Linear(256, 16, device="cpu", dtype=torch.bfloat16)
model.load_state_dict(state_dict, strict=False)
self.assertIsInstance(model.layer.weight, QuantizedTensor)
self.assertEqual(model.layer.weight._layout_cls, "TensorCoreConvRotW4A4Layout")
self.assertEqual(model.layer.weight._params.convrot_groupsize, 256)
self.assertEqual(model.layer.weight._params.quant_group_size, 64)
self.assertEqual(model.layer.weight._params.linear_dtype, "int8")
input_tensor = torch.randn(4, 256, dtype=torch.bfloat16)
loaded_out = model.layer(input_tensor)
ref_out = torch.nn.functional.linear(input_tensor, q_weight, bias)
self.assertTrue(torch.equal(loaded_out, ref_out))
saved = model.state_dict()
saved_conf = json.loads(saved["layer.comfy_quant"].numpy().tobytes())
self.assertEqual(saved_conf["format"], "convrot_w4a4")
self.assertEqual(saved_conf["convrot_groupsize"], 256)
self.assertEqual(saved_conf["linear_dtype"], "int8")
self.assertNotIn("quant_group_size", saved_conf)
if __name__ == "__main__":
unittest.main()

View File

@@ -2,7 +2,7 @@ from collections import defaultdict
import torch
from comfy.model_detection import detect_unet_config, model_config_from_unet_config
from comfy.model_detection import detect_unet_config, model_config_from_unet, model_config_from_unet_config
import comfy.supported_models
@@ -73,6 +73,60 @@ def _make_flux_schnell_comfyui_sd():
return sd
def _make_seedvr2_7b_separate_mm_sd():
return {
"blocks.35.mlp.vid.proj_out.weight": torch.empty(3072, 1),
"positive_conditioning": torch.empty(58, 5120),
"negative_conditioning": torch.empty(64, 5120),
}
def _make_seedvr2_7b_shared_mm_sd():
return {
"blocks.35.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
"positive_conditioning": torch.empty(58, 5120),
"negative_conditioning": torch.empty(64, 5120),
}
def _make_seedvr2_3b_shared_mm_sd():
return {
"blocks.31.mlp.all.proj_in_gate.weight": torch.empty(1, 1),
"positive_conditioning": torch.empty(58, 5120),
"negative_conditioning": torch.empty(64, 5120),
}
def _make_pid_v1_5_sd(latent_proj_channels=16):
sd = {
"pixel_embedder.proj.weight": torch.empty(16, 3, device="meta"),
"lq_proj.latent_proj.0.weight": torch.empty(1024, latent_proj_channels, 3, 3, device="meta"),
"lq_proj.pit_head.weight": torch.empty(1536, 1024, device="meta"),
"lq_proj.gate_modules.0.content_proj.weight": torch.empty(1, 3072, device="meta"),
"pixel_blocks.0.attn.q_norm.weight": torch.empty(72, device="meta"),
"pixel_blocks.0.adaLN_modulation.0.weight": torch.empty(24576, 1536, device="meta"),
"pixel_blocks.0.adaLN_modulation.0.bias": torch.empty(24576, device="meta"),
}
for i in range(7):
sd[f"lq_proj.gate_modules.{i}.log_alpha"] = torch.empty((), device="meta")
return sd
def _make_joyimage_edit_plus_sd():
sd = {
"img_in.weight": torch.empty(4096, 16, 1, 2, 2, device="meta"),
"condition_embedder.time_embedder.linear_1.weight": torch.empty(1, device="meta"),
"double_blocks.0.attn.img_attn_q_norm.weight": torch.empty(128, device="meta"),
}
for i in range(40):
sd[f"double_blocks.{i}.attn.img_attn_qkv.weight"] = torch.empty(1, device="meta")
return sd
def _add_model_diffusion_prefix(sd):
return {f"model.diffusion_model.{k}": v for k, v in sd.items()}
class TestModelDetection:
"""Verify that first-match model detection selects the correct model
based on list ordering and unet_config specificity."""
@@ -125,6 +179,116 @@ 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["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["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"
def test_seedvr2_model_match_requires_conditioning_tensors(self):
sd = _make_seedvr2_7b_shared_mm_sd()
unet_config = detect_unet_config(sd, "")
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "SeedVR2"
del sd["positive_conditioning"]
assert model_config_from_unet_config(unet_config, sd) is None
def test_seedvr2_model_match_accepts_full_checkpoint_prefix(self):
sd = _add_model_diffusion_prefix(_make_seedvr2_7b_shared_mm_sd())
assert type(model_config_from_unet(sd, "model.diffusion_model.")).__name__ == "SeedVR2"
def test_pid_v1_5_detection(self):
sd = _make_pid_v1_5_sd()
unet_config = detect_unet_config(sd, "")
assert unet_config == {
"image_model": "pid",
"lq_latent_channels": 16,
"lq_hidden_dim": 1024,
"latent_spatial_down_factor": 8,
"lq_interval": 2,
"lq_latent_unpatchify_factor": 1,
"lq_conv_padding_mode": "replicate",
"lq_gate_per_token": True,
"pit_lq_inject": True,
"rope_ref_h": 2048,
"rope_ref_w": 2048,
}
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "PiD"
def test_pid_v1_5_flux2_detection(self):
unet_config = detect_unet_config(_make_pid_v1_5_sd(latent_proj_channels=32), "")
assert unet_config["lq_latent_channels"] == 128
assert unet_config["latent_spatial_down_factor"] == 16
assert unet_config["lq_latent_unpatchify_factor"] == 2
def test_pid_v1_5_pixel_adaln_conversion(self):
sd = _make_pid_v1_5_sd()
model_config = model_config_from_unet_config(detect_unet_config(sd, ""), sd)
processed = model_config.process_unet_state_dict(sd)
assert processed["pixel_blocks.0.attn.q_norm.weight"].shape == (72,)
assert processed["pixel_blocks.0.adaLN_modulation_msa.weight"].shape == (12288, 1536)
assert processed["pixel_blocks.0.adaLN_modulation_mlp.weight"].shape == (12288, 1536)
assert processed["pixel_blocks.0.adaLN_modulation_msa.bias"].shape == (12288,)
assert processed["pixel_blocks.0.adaLN_modulation_mlp.bias"].shape == (12288,)
def test_joyimage_edit_plus_detection(self):
sd = _make_joyimage_edit_plus_sd()
unet_config = detect_unet_config(sd, "")
assert unet_config == {
"image_model": "joyimage",
"in_channels": 16,
"hidden_size": 4096,
"patch_size": [1, 2, 2],
"num_layers": 40,
"num_attention_heads": 32,
"text_dim": 4096,
}
assert type(model_config_from_unet_config(unet_config, sd)).__name__ == "JoyImage"
def test_incomplete_joyimage_signature_is_not_detected(self):
sd = _make_joyimage_edit_plus_sd()
del sd["double_blocks.0.attn.img_attn_q_norm.weight"]
assert detect_unet_config(sd, "") 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

View File

@@ -0,0 +1,74 @@
"""Regression tests for the SeedVR2 VAE forward return contract."""
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
from comfy.ldm.seedvr.vae import SEEDVR2_LATENT_CHANNELS, VideoAutoencoderKL # noqa: E402
_LATENT_SHAPE = (1, SEEDVR2_LATENT_CHANNELS, 2, 2, 2)
_DECODED_SHAPE = (1, 3, 5, 16, 16)
_INPUT_ENCODE_SHAPE = (1, 3, 5, 16, 16)
_INPUT_DECODE_SHAPE = _LATENT_SHAPE
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):
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)
def test_forward_rejects_unknown_mode():
vae = _StubVAE()
with pytest.raises(ValueError, match="Unknown SeedVR2 VAE forward mode"):
vae.forward(torch.zeros(*_INPUT_ENCODE_SHAPE), mode="bogus")

View File

@@ -0,0 +1,79 @@
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.sd
import comfy.supported_models
import comfy.ldm.seedvr.model as seedvr_model
import comfy.ldm.seedvr.vae as seedvr_vae
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():
wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper)
latent_channels = seedvr_vae.SEEDVR2_LATENT_CHANNELS
estimate = wrapper.comfy_memory_used_decode((1, latent_channels, 26, 120, 160))
old_estimate = latent_channels * 120 * 160 * (4 * 8 * 8) * 2
assert estimate == 101 * 960 * 1280 * 160
assert estimate > 15 * 1024 ** 3
assert estimate > old_estimate * 100
def test_seedvr2_vae_encode_preserves_compute_dtype(monkeypatch):
wrapper = seedvr_vae.VideoAutoencoderKLWrapper.__new__(seedvr_vae.VideoAutoencoderKLWrapper)
nn.Module.__init__(wrapper)
wrapper._dummy = nn.Parameter(torch.empty(1, dtype=torch.float16))
input_dtype = None
def encode(self, x):
nonlocal input_dtype
input_dtype = x.dtype
return x
monkeypatch.setattr(seedvr_vae.VideoAutoencoderKL, "encode", encode)
x = torch.zeros((1, 3, 1, 8, 8), dtype=torch.float32)
wrapper._encode_with_raw_latent(x)
assert input_dtype == torch.float32
def test_seedvr2_vae_ops_cast_weights_to_compute_dtype():
attention = seedvr_vae.Attention(query_dim=4, heads=1, dim_head=4).to(torch.float16)
hidden_states = torch.zeros((1, 2, 4), dtype=torch.float32)
output = attention(hidden_states)
assert output.dtype == torch.float32

View File

@@ -0,0 +1,169 @@
"""SeedVR2 internals regression tests."""
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.vae import ( # noqa: E402
causal_norm_wrapper,
set_norm_limit,
)
from comfy.ldm.seedvr.attention import var_attention_optimized_split # noqa: E402
_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)
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=comfy_ops.disable_weight_init.RMSNorm,
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
assert call["cu_seqlens_q"] == [0, 7, 14]
assert call["cu_seqlens_k"] == [0, 7, 14]
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 = [0, 2, 5]
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)

View File

@@ -0,0 +1,320 @@
"""SeedVR2 model, latent-format, and VAE graph regression tests."""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
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 as seedvr_model # noqa: E402
import comfy.ldm.seedvr.vae as seedvr_vae_mod # noqa: E402
import comfy.model_management # noqa: E402
import comfy.ops as comfy_ops # 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
_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS
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()
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,
num_layers=4,
mlp_type="normal",
vid_dim=vid_dim,
txt_in_dim=txt_in_dim,
heads=24,
mm_layers=3,
operations=comfy_ops.disable_weight_init,
)
return flags
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
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 // _LATENT_CHANNELS
else:
b, _, t, h, w = z.shape
return torch.zeros(b, 3, t, h * 8, w * 8, dtype=z.dtype, device=z.device)
def test_seedvr2_wrapper_public_encode_returns_tensor(monkeypatch):
raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 2.0)
seen_shapes = []
def base_encode(self, x):
seen_shapes.append(tuple(x.shape))
return raw_latent.to(device=x.device, dtype=x.dtype)
monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode)
vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper)
nn.Module.__init__(vae)
vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32))
latent = vae.encode(torch.zeros(1, 3, 32, 40))
assert type(latent) is torch.Tensor
assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5)
assert seen_shapes == [(1, 3, 1, 32, 40)]
def test_seedvr2_wrapper_private_encode_helper_keeps_raw_latent(monkeypatch):
raw_latent = torch.full((1, _LATENT_CHANNELS, 1, 4, 5), 3.0)
def base_encode(self, x):
return raw_latent.to(device=x.device, dtype=x.dtype)
monkeypatch.setattr(seedvr_vae_mod.VideoAutoencoderKL, "encode", base_encode)
vae = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(seedvr_vae_mod.VideoAutoencoderKLWrapper)
nn.Module.__init__(vae)
vae._dummy = nn.Parameter(torch.zeros((), dtype=torch.float32))
latent, raw = vae._encode_with_raw_latent(torch.zeros(1, 3, 32, 40))
assert tuple(latent.shape) == (1, _LATENT_CHANNELS, 4, 5)
assert tuple(raw.shape) == (1, _LATENT_CHANNELS, 1, 4, 5)
assert torch.equal(raw, raw_latent)
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 = _LATENT_CHANNELS
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.handles_tiling = isinstance(wrapper, seedvr_vae_mod.VideoAutoencoderKLWrapper)
vae.format_encoded = wrapper.comfy_format_encoded
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
def test_missing_context_falls_back_to_positive_buffer():
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
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)
def test_seedvr2_forward_requires_conditioning_latents():
model = NaDiT.__new__(NaDiT)
x = torch.zeros(1, _LATENT_CHANNELS, 1, 4, 5)
with pytest.raises(ValueError, match="requires conditioning latents"):
NaDiT.forward(model, x, timestep=torch.tensor([1.0]), context=None)
def test_seedvr2_latent_format_uses_native_video_latent_shape():
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 == _LATENT_CHANNELS
assert latent_format.latent_dimensions == 3
assert fixed.shape == (1, _LATENT_CHANNELS, 1, 4, 5)
def test_seedvr2_model_requires_native_5d_latent():
latent = torch.zeros(1, _LATENT_CHANNELS, 2, 4, 5)
assert NaDiT._check_seedvr2_video_latent(latent, _LATENT_CHANNELS, "latent") is latent
with pytest.raises(ValueError, match="5-D native latent"):
NaDiT._check_seedvr2_video_latent(torch.zeros(1, _LATENT_CHANNELS * 2, 4, 5), _LATENT_CHANNELS, "latent")
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, _LATENT_CHANNELS, 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, _LATENT_CHANNELS, 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 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR))
tiled = torch.full((1, _LATENT_CHANNELS, 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, _LATENT_CHANNELS, 2, 4, 5)
assert tiled_latent.dtype == torch.float32
assert torch.equal(tiled_latent, torch.full_like(tiled_latent, 3.0 * seedvr_vae_mod.BYTEDANCE_VAE_SCALING_FACTOR))
def test_vaedecode_tiled_spatial_applies_temporal_discarded(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, _LATENT_CHANNELS, 2, 4, 5)},
tile_size=512,
overlap=64,
temporal_size=16,
temporal_overlap=4,
)
# Spatial inputs flow through; temporal inputs are discarded as public tiling
# knobs, but SeedVR2's internal MemoryState causal slicing is left intact.
assert vae.first_stage_model.calls == [
{
"shape": (1, _LATENT_CHANNELS, 2, 4, 5),
"seedvr2_tiling": {
"enable_tiling": True,
"tile_size": (512, 512),
"tile_overlap": (64, 64),
"temporal_size": None,
"temporal_overlap": None,
},
}
]

View File

@@ -0,0 +1,94 @@
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
_LATENT_CHANNELS = vae_mod.SEEDVR2_LATENT_CHANNELS
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, _LATENT_CHANNELS * 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, _LATENT_CHANNELS, 2, 4, 5))
assert tuple(out.shape) == (1, 3, 2, 32, 40)
assert wrapper.calls == [(1, _LATENT_CHANNELS, 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, _LATENT_CHANNELS, 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)

View File

@@ -0,0 +1,407 @@
from contextlib import ExitStack
from unittest.mock import MagicMock, 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
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
_LATENT_CHANNELS = seedvr_vae_mod.SEEDVR2_LATENT_CHANNELS
def test_runtime_decode_zero_temporal_size_preserves_model_slicing():
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 vae_mod.VideoAutoencoderKL.slicing_decode(self, t_chunk)
def _decode(self, z, memory_state=MemoryState.DISABLED, memory_cache=None):
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, _LATENT_CHANNELS, 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 == [2]
assert vae.memory_states == [MemoryState.INITIALIZING, MemoryState.ACTIVE]
assert vae.slicing_latent_min_size == 2
def test_zero_temporal_size_preserves_min_size_when_encode_raises():
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)
with pytest.raises(RuntimeError, match="simulated encode failure"):
tiled_vae(
x,
vae,
tile_size=(64, 64),
tile_overlap=(0, 0),
temporal_size=0,
temporal_overlap=0,
encode=True,
)
assert vae.slicing_sample_min_size == 4
def test_tiled_vae_encode_uses_tensor_return_without_indexing():
class TensorEncodeVAEModel(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))
self.calls = []
def encode(self, t_chunk):
self.calls.append(tuple(t_chunk.shape))
b, _, _, h, w = t_chunk.shape
return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=t_chunk.dtype)
vae = TensorEncodeVAEModel()
x = torch.zeros((2, 3, 1, 64, 64), dtype=torch.float32)
out = tiled_vae(
x,
vae,
tile_size=(64, 64),
tile_overlap=(0, 0),
temporal_size=0,
temporal_overlap=0,
encode=True,
)
assert vae.calls == [(2, 3, 1, 64, 64)]
assert tuple(out.shape) == (2, _LATENT_CHANNELS, 1, 8, 8)
def test_tiled_vae_preserves_compute_dtype_with_different_parameter_dtype():
class DummyVAE(nn.Module):
spatial_downsample_factor = 8
temporal_downsample_factor = 4
slicing_sample_min_size = 8
def __init__(self):
super().__init__()
self.device = torch.device("cpu")
self._dummy = nn.Parameter(torch.zeros(1, dtype=torch.float16))
self.input_dtype = None
def encode(self, t_chunk):
self.input_dtype = t_chunk.dtype
b, _, _, h, w = t_chunk.shape
return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=t_chunk.dtype)
vae = DummyVAE()
x = torch.zeros((1, 3, 1, 64, 64), dtype=torch.float32)
tiled_vae(x, vae, tile_size=(64, 64), tile_overlap=(16, 16), encode=True)
assert vae.input_dtype == torch.float32
def test_tiled_vae_preserves_input_dtype_on_single_tile():
class FloatOutputVAEModel(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):
b, _, _, h, w = t_chunk.shape
return torch.ones((b, _LATENT_CHANNELS, 1, h // 8, w // 8), dtype=torch.float32)
out = tiled_vae(
torch.zeros((1, 3, 1, 64, 64), dtype=torch.float16),
FloatOutputVAEModel(),
tile_size=(64, 64),
tile_overlap=(0, 0),
temporal_size=0,
temporal_overlap=0,
encode=True,
)
assert out.dtype == torch.float16
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, memory_cache=None):
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(
_LATENT_CHANNELS * 5 * 8 * 8,
dtype=torch.float32,
).reshape(1, _LATENT_CHANNELS, 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, _LATENT_CHANNELS, 2, 2), seedvr2_tiling=seedvr2_tiling)
assert captured["temporal_overlap"] == 7
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.handles_tiling = isinstance(first_stage_model, seedvr_vae_mod.VideoAutoencoderKLWrapper)
vae.format_encoded = None
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_owned", 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_owned_decode_tiled():
wrapper = seedvr_vae_mod.VideoAutoencoderKLWrapper.__new__(
seedvr_vae_mod.VideoAutoencoderKLWrapper)
vae = _make_vae(wrapper, latent_channels=_LATENT_CHANNELS, 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, _LATENT_CHANNELS * 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
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 = _LATENT_CHANNELS
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.handles_tiling = isinstance(vae.first_stage_model, seedvr_vae_mod.VideoAutoencoderKLWrapper)
vae.format_encoded = 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_owned_encode_tiled_on_oom():
vae = _make_seedvr2_vae_fallback()
pixel_samples = torch.zeros((1, 8, 64, 64, 3))
seedvr2_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 2, 8, 8))
generic_call = MagicMock(return_value=torch.zeros(1, _LATENT_CHANNELS, 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_owned", seedvr2_call), \
patch.object(sd_mod.VAE, "encode_tiled_3d", generic_call):
vae.encode(pixel_samples)
assert seedvr2_call.call_count == 1, (
f"Expected _encode_tiled_owned 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, _LATENT_CHANNELS, 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)

View File

@@ -11,6 +11,11 @@ from comfy_api.feature_flags import (
_coerce_flag_value,
_parse_cli_feature_flags,
)
from comfy.comfy_api_env import (
environment_overrides_for_base,
get_environment_overrides,
normalize_comfy_api_base,
)
class TestFeatureFlags:
@@ -183,3 +188,65 @@ class TestCliFeatureFlagRegistry:
assert "type" in info, f"{key} missing 'type'"
assert "default" in info, f"{key} missing 'default'"
assert "description" in info, f"{key} missing 'description'"
class TestComfyApiEnv:
"""--comfy-api-base staging-tier detection + testenv main-host -> -registry rewrite."""
@pytest.mark.parametrize(
"url, expected",
[
# testenv friendly main host -> comfy-api -registry sibling (slash trimmed)
("https://pr-4398.testenvs.comfy.org", "https://pr-4398-registry.testenvs.comfy.org"),
("https://pr-4398.testenvs.comfy.org/", "https://pr-4398-registry.testenvs.comfy.org"),
("https://pr-4398-registry.testenvs.comfy.org", "https://pr-4398-registry.testenvs.comfy.org"),
# staging + everything else -> unchanged (no -registry split)
("https://stagingapi.comfy.org", "https://stagingapi.comfy.org"),
("https://api.comfy.org", "https://api.comfy.org"),
("https://pr-1.testenvs.comfy.org.evil.com", "https://pr-1.testenvs.comfy.org.evil.com"),
("", ""),
],
)
def test_normalize_comfy_api_base(self, url, expected):
assert normalize_comfy_api_base(url) == expected
def test_config_for_staging_tier_else_none(self):
# ephemeral testenv: friendly main host -> -registry, staging platform, dev Firebase env
eph = environment_overrides_for_base("https://pr-1234.testenvs.comfy.org/")
assert eph["comfy_api_base_url"] == "https://pr-1234-registry.testenvs.comfy.org"
assert eph["comfy_platform_base_url"] == "https://stagingplatform.comfy.org"
assert eph["firebase_env"] == "dev"
# staging api host: emitted as-is
stg = environment_overrides_for_base("https://stagingapi.comfy.org")
assert stg["comfy_api_base_url"] == "https://stagingapi.comfy.org"
assert stg["comfy_platform_base_url"] == "https://stagingplatform.comfy.org"
assert stg["firebase_env"] == "dev"
# prod / unknown: nothing
assert environment_overrides_for_base("https://api.comfy.org") is None
def test_environment_overrides_only_for_staging_tier(self, monkeypatch):
def set_base(url):
monkeypatch.setattr(
"comfy.comfy_api_env.args",
type("Args", (), {"comfy_api_base": url})(),
)
# The overrides merged into the HTTP /features response are present for staging-tier bases...
set_base("https://stagingapi.comfy.org")
assert "comfy_api_base_url" in get_environment_overrides()
set_base("https://pr-7.testenvs.comfy.org")
assert "comfy_api_base_url" in get_environment_overrides()
# ...but never for prod.
set_base("https://api.comfy.org")
assert get_environment_overrides() is None
def test_server_features_never_carry_env_overrides(self, monkeypatch):
"""The WebSocket capability handshake must stay free of routing keys."""
monkeypatch.setattr(
"comfy.comfy_api_env.args",
type("Args", (), {"comfy_api_base": "https://pr-7.testenvs.comfy.org"})(),
)
features = get_server_features()
assert "comfy_api_base_url" not in features
assert "comfy_platform_base_url" not in features
assert "firebase_env" not in features