[misc]: add CPU regression tests for the on-device uint8 frame path
Two tests certifying #1362's frame semantics without a GPU: - frames_match_legacy_cpu_loop: the quantize-then-grid path reproduces the legacy make_grid -> *255 -> uint8 per-frame loop bit-exactly for in-range fp32 pixels (batch>1 nrow=6 grid layout, odd frame count, uint8 HWC contract). Verified to pass against main's legacy loop too, so it pins both sides of the equivalence. - frames_clamp_out_of_range_pixels: out-of-[0,1] VAE output saturates at 0/255 instead of wrapping mod 256; fails on the pre-#1362 loop. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
f56f567042
commit
aa95a4c18e
@@ -2,8 +2,10 @@ import os
|
||||
from types import SimpleNamespace
|
||||
import warnings
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
import fastvideo.entrypoints.video_generator as video_generator_module
|
||||
from fastvideo.api import (
|
||||
@@ -271,6 +273,70 @@ def test_generate_single_video_return_frames_still_materializes_output(tmp_path)
|
||||
assert result["video_path"] is None
|
||||
|
||||
|
||||
def test_generate_single_video_frames_match_legacy_cpu_loop(tmp_path):
|
||||
"""The on-device quantize path (#1362) must reproduce the legacy
|
||||
per-frame CPU loop (make_grid -> permute -> *255 -> uint8) bit-exactly
|
||||
for in-range fp32 pixels: same uint8 dtype, same HWC grid layout with
|
||||
nrow=6 (batch>1), odd frame count. CPU-only: on CUDA the float->uint8
|
||||
cast may differ by <=1 LSB, but on CPU both orderings run identical
|
||||
fp32 ops, so exact equality is required."""
|
||||
torch.manual_seed(0)
|
||||
output = torch.rand((2, 3, 3, 16, 16), dtype=torch.float32)
|
||||
output_batch = _single_video_output_batch(output)
|
||||
fastvideo_args = _single_video_args()
|
||||
generator = _single_video_generator(output_batch, fastvideo_args)
|
||||
sampling_param = _small_sampling_param(save_video=False, return_frames=True)
|
||||
sampling_param.num_frames = 3
|
||||
sampling_param.num_videos_per_prompt = 2
|
||||
|
||||
result = generator._generate_single_video(
|
||||
prompt="grid parity",
|
||||
sampling_param=sampling_param,
|
||||
fastvideo_args=fastvideo_args,
|
||||
output_path=str(tmp_path / "unused.mp4"),
|
||||
)
|
||||
|
||||
legacy_frames = []
|
||||
for x in rearrange(output, "b c t h w -> t b c h w"):
|
||||
grid = video_generator_module.torchvision.utils.make_grid(x, nrow=6)
|
||||
grid = grid.permute(1, 2, 0).squeeze(-1)
|
||||
legacy_frames.append((grid * 255).to(torch.uint8).contiguous().cpu().numpy())
|
||||
|
||||
torch.testing.assert_close(result["samples"], output)
|
||||
assert len(result["frames"]) == 3
|
||||
for got, want in zip(result["frames"], legacy_frames, strict=True):
|
||||
assert got.dtype == np.uint8
|
||||
assert got.shape == want.shape
|
||||
np.testing.assert_array_equal(got, want)
|
||||
|
||||
|
||||
def test_generate_single_video_frames_clamp_out_of_range_pixels(tmp_path):
|
||||
"""VAE output slightly outside [0, 1] must saturate at 0/255 in the
|
||||
uint8 frames. The pre-#1362 unclamped cast wrapped mod 256 (e.g.
|
||||
1.5 -> 126). CPU-only."""
|
||||
output = torch.full((1, 3, 2, 16, 16), 1.5, dtype=torch.float32)
|
||||
output[:, :, 1] = -0.5
|
||||
output_batch = _single_video_output_batch(output)
|
||||
fastvideo_args = _single_video_args()
|
||||
generator = _single_video_generator(output_batch, fastvideo_args)
|
||||
|
||||
result = generator._generate_single_video(
|
||||
prompt="clamp",
|
||||
sampling_param=_small_sampling_param(save_video=False, return_frames=True),
|
||||
fastvideo_args=fastvideo_args,
|
||||
output_path=str(tmp_path / "unused.mp4"),
|
||||
)
|
||||
|
||||
frames = result["frames"]
|
||||
assert len(frames) == 2
|
||||
# make_grid passes a single image through without grid padding, so
|
||||
# every pixel comes from the (clamped) output tensor.
|
||||
assert frames[0].dtype == np.uint8
|
||||
assert frames[0].shape == (16, 16, 3)
|
||||
assert (frames[0] == 255).all()
|
||||
assert (frames[1] == 0).all()
|
||||
|
||||
|
||||
def test_generate_single_video_save_video_still_builds_frames(monkeypatch, tmp_path):
|
||||
output = torch.ones((1, 3, 2, 16, 16), dtype=torch.float32) * 0.5
|
||||
output_batch = _single_video_output_batch(output)
|
||||
|
||||
Reference in New Issue
Block a user