Compare commits

...
Author SHA1 Message Date
Will Lin 3251822a7f update 2026-01-06 22:10:07 -08:00
8 changed files with 480 additions and 35 deletions
+19
View File
@@ -70,6 +70,12 @@ class SamplingParam:
# Misc
save_video: bool = True
return_frames: bool = False
# Return controls for non-`return_frames` code paths.
# When `return_frames=True`, `generate_video()` returns the raw frames list.
# Otherwise, it returns a dict; these flags control which bulky fields are
# included in that dict.
return_frames_in_dict: bool = False
return_samples: bool = False
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
@@ -216,6 +222,19 @@ class SamplingParam:
default=SamplingParam.return_frames,
help="Whether to return the raw frames",
)
parser.add_argument(
"--return-frames-in-dict",
action="store_true",
default=SamplingParam.return_frames_in_dict,
help=
"Include frames in the returned dict (ignored if --return-frames is set)",
)
parser.add_argument(
"--return-samples",
action="store_true",
default=SamplingParam.return_samples,
help="Include model samples tensor in the returned dict",
)
parser.add_argument(
"--image-path",
type=str,
+63
View File
@@ -0,0 +1,63 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import numpy as np
import torch
@dataclass(frozen=True)
class OutputOptions:
"""Controls what `VideoGenerator.generate_video()` returns and saves.
Notes:
- This is intentionally separate from `SamplingParam` (sampling settings).
- When `return_format="legacy"`, `return_frames=True` preserves the old
behavior of returning a list of frames instead of a result object.
- When `return_format="dataclass"`, `return_frames` only controls whether
frames are produced; the return type is still `GenerationResult`.
"""
save_video: bool | None = None
return_frames: bool = False
include_frames: bool = False
include_samples: bool = False
include_trajectory_latents: bool = False
include_trajectory_decoded: bool = False
@dataclass
class GenerationResult:
output_path: str
prompt: str
size: tuple[int, int, int]
generation_time: float
logging_info: Any
frames: list[np.ndarray] | None = None
samples: torch.Tensor | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_decoded: list[torch.Tensor] | None = None
def to_legacy_dict(self) -> dict[str, Any]:
out: dict[str, Any] = {
"output_path": self.output_path,
"prompt": self.prompt,
"prompts": self.prompt,
"size": self.size,
"generation_time": self.generation_time,
"logging_info": self.logging_info,
}
if self.frames is not None:
out["frames"] = self.frames
if self.samples is not None:
out["samples"] = self.samples
if self.trajectory_latents is not None:
out["trajectory"] = self.trajectory_latents
out["trajectory_timesteps"] = self.trajectory_timesteps
if self.trajectory_decoded is not None:
out["trajectory_decoded"] = self.trajectory_decoded
return out
+107 -27
View File
@@ -23,6 +23,7 @@ from fastvideo.configs.sample import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ForwardBatch
from fastvideo.entrypoints.generation_types import GenerationResult, OutputOptions
from fastvideo.utils import align_to, shallow_asdict
from fastvideo.worker.executor import Executor
@@ -107,6 +108,8 @@ class VideoGenerator:
keyboard_cond: torch.Tensor | None = None,
grid_sizes: tuple[int, int, int] | list[int] | torch.Tensor
| None = None,
output_options: OutputOptions | None = None,
return_format: str = "legacy",
**kwargs,
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]]:
"""
@@ -137,6 +140,11 @@ class VideoGenerator:
sampling_param = SamplingParam.from_pretrained(
self.fastvideo_args.model_path)
if return_format not in ("legacy", "dataclass"):
raise ValueError(
"return_format must be 'legacy' or 'dataclass', got "
f"{return_format!r}")
# Add action control inputs to kwargs if provided
if mouse_cond is not None:
kwargs['mouse_cond'] = mouse_cond
@@ -145,7 +153,18 @@ class VideoGenerator:
if grid_sizes is not None:
kwargs['grid_sizes'] = grid_sizes
sampling_param.update(kwargs)
# Only apply generation-time overrides that actually exist on
# SamplingParam. Other kwargs (e.g. monitoring, callbacks) should not
# be treated as sampling fields.
sampling_updates = {
key: value
for key, value in kwargs.items() if hasattr(sampling_param, key)
}
sampling_param.update(sampling_updates)
sampling_param.check_sampling_param()
output_options = self._resolve_output_options(sampling_param,
output_options)
if self.fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
prompt_txt_path = sampling_param.prompt_path or self.fastvideo_args.prompt_txt
@@ -174,10 +193,14 @@ class VideoGenerator:
result = self._generate_single_video(
prompt=batch_prompt,
sampling_param=sampling_param,
output_options=output_options,
return_format=return_format,
**kwargs)
# Add prompt info to result
if isinstance(result, dict):
if isinstance(result, GenerationResult):
pass
elif isinstance(result, dict):
result["prompt_index"] = i
result["prompt"] = batch_prompt
@@ -189,6 +212,10 @@ class VideoGenerator:
logger.error("Failed to generate video for prompt %d: %s",
i + 1, e)
continue
monitor = kwargs.get("monitor")
if monitor is not None and hasattr(monitor, "print_stats"):
monitor.print_stats(
f"INTERNAL: After prompt {i + 1}/{len(prompts)}")
logger.info(
"Completed batch processing. Generated %d videos successfully.",
@@ -203,8 +230,26 @@ class VideoGenerator:
kwargs["output_path"] = output_path
return self._generate_single_video(prompt=prompt,
sampling_param=sampling_param,
output_options=output_options,
return_format=return_format,
**kwargs)
def _resolve_output_options(
self, sampling_param: SamplingParam,
output_options: OutputOptions | None) -> OutputOptions:
if output_options is not None:
return output_options
return OutputOptions(
save_video=None,
return_frames=sampling_param.return_frames,
include_frames=getattr(sampling_param, "return_frames_in_dict",
False),
include_samples=getattr(sampling_param, "return_samples", False),
include_trajectory_latents=sampling_param.return_trajectory_latents,
include_trajectory_decoded=sampling_param.return_trajectory_decoded,
)
def _prepare_output_path(
self,
output_path: str,
@@ -270,8 +315,10 @@ class VideoGenerator:
self,
prompt: str,
sampling_param: SamplingParam | None = None,
output_options: OutputOptions | None = None,
return_format: str = "legacy",
**kwargs,
) -> dict[str, Any] | list[np.ndarray]:
) -> dict[str, Any] | list[np.ndarray] | GenerationResult:
"""Internal method for single video generation"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
@@ -368,6 +415,17 @@ class VideoGenerator:
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
if output_options is None:
output_options = self._resolve_output_options(sampling_param, None)
if output_options.save_video is not None:
batch.save_video = output_options.save_video
batch.return_frames = output_options.return_frames
batch.return_frames_in_dict = output_options.include_frames
batch.return_samples = output_options.include_samples
batch.return_trajectory_latents = output_options.include_trajectory_latents
batch.return_trajectory_decoded = output_options.include_trajectory_decoded
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
@@ -377,33 +435,55 @@ class VideoGenerator:
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
frames: list[np.ndarray] | None = None
need_frames = batch.save_video or batch.return_frames or batch.return_frames_in_dict
if need_frames:
# Process outputs into CPU frames (potentially large).
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
if batch.save_video:
imageio.mimsave(output_path,
frames,
fps=batch.fps,
format="mp4")
logger.info("Saved video to %s", output_path)
result = GenerationResult(
output_path=output_path,
prompt=prompt,
size=(target_height, target_width, batch.num_frames),
generation_time=gen_time,
logging_info=logging_info,
)
if batch.return_frames or batch.return_frames_in_dict:
if frames is None:
raise RuntimeError(
"Frames requested but frames were not produced")
result.frames = frames
if batch.return_samples:
result.samples = samples
if batch.return_trajectory_latents:
result.trajectory_latents = output_batch.trajectory_latents
result.trajectory_timesteps = output_batch.trajectory_timesteps
if batch.return_trajectory_decoded:
result.trajectory_decoded = output_batch.trajectory_decoded
if return_format == "dataclass":
return result
if batch.return_frames:
return frames
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
assert result.frames is not None
return result.frames
out_dict = result.to_legacy_dict()
if not batch.return_frames_in_dict:
out_dict.pop("frames", None)
return out_dict
def set_lora_adapter(self,
lora_nickname: str,
@@ -179,6 +179,8 @@ class ForwardBatch:
# Misc
save_video: bool = True
return_frames: bool = False
return_frames_in_dict: bool = False
return_samples: bool = False
# TeaCache parameters
enable_teacache: bool = False
@@ -14,7 +14,7 @@ try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
except Exception:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
@@ -22,7 +22,7 @@ try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
except Exception:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
@@ -494,4 +494,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result
+3 -3
View File
@@ -32,21 +32,21 @@ try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
except Exception:
st_attn_available = False
try:
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
from fastvideo.utils import is_vmoba_available
vmoba_attn_available = is_vmoba_available()
except ImportError:
except Exception:
vmoba_attn_available = False
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
except Exception:
vsa_available = False
logger = init_logger(__name__)
@@ -19,7 +19,7 @@ try:
from fastvideo.attention.backends.sliding_tile_attn import (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
except Exception:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
@@ -27,7 +27,7 @@ try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
except Exception:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
@@ -0,0 +1,281 @@
from types import SimpleNamespace
import numpy as np
import pytest
import torch
from fastvideo.configs.sample import SamplingParam
from fastvideo.entrypoints.generation_types import GenerationResult, OutputOptions
from fastvideo.entrypoints.video_generator import VideoGenerator
class _DummyExecutor:
def __init__(self, samples: torch.Tensor):
self._samples = samples
self.calls = 0
def execute_forward(self, batch, fastvideo_args):
self.calls += 1
return SimpleNamespace(
output=self._samples,
logging_info=SimpleNamespace(),
trajectory_latents="trajectory_latents",
trajectory_timesteps="trajectory_timesteps",
trajectory_decoded="trajectory_decoded",
)
def _new_stub_generator(samples: torch.Tensor) -> VideoGenerator:
generator = VideoGenerator.__new__(VideoGenerator)
generator.fastvideo_args = SimpleNamespace(
model_path="dummy",
num_gpus=1,
prompt_txt=None,
VSA_sparsity=0.0,
pipeline_config=SimpleNamespace(
vae_config=SimpleNamespace(
arch_config=SimpleNamespace(temporal_compression_ratio=4),
use_temporal_scaling_frames=True,
),
flow_shift=0.0,
embedded_cfg_scale=0.0,
),
)
generator.executor = _DummyExecutor(samples)
return generator
def _base_params(tmp_path) -> SamplingParam:
return SamplingParam(
num_frames=2,
height=64,
width=64,
fps=8,
num_inference_steps=1,
seed=0,
save_video=False,
output_path=str(tmp_path / "outputs"),
)
def test_generate_single_minimal_dict_does_not_build_frames(tmp_path,
monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
def _fail_make_grid(*args, **kwargs):
raise AssertionError("make_grid should not be called")
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fail_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
)
assert isinstance(result, dict)
assert "output_path" in result
assert "frames" not in result
assert "samples" not in result
def test_generate_single_dataclass_minimal_does_not_build_frames(
tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
def _fail_make_grid(*args, **kwargs):
raise AssertionError("make_grid should not be called")
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fail_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
return_format="dataclass",
)
assert isinstance(result, GenerationResult)
assert result.frames is None
assert result.samples is None
def test_generate_single_return_frames_in_dict(tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
called = {"count": 0}
def _fake_make_grid(x, nrow=6):
called["count"] += 1
return x[0]
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fake_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
return_frames_in_dict=True,
)
assert called["count"] > 0
assert isinstance(result, dict)
assert isinstance(result["frames"], list)
assert isinstance(result["frames"][0], np.ndarray)
assert "samples" not in result
def test_generate_single_dataclass_with_output_options(tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
called = {"count": 0}
def _fake_make_grid(x, nrow=6):
called["count"] += 1
return x[0]
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fake_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
output_options=OutputOptions(
include_frames=True,
include_samples=True,
),
return_format="dataclass",
)
assert called["count"] > 0
assert isinstance(result, GenerationResult)
assert isinstance(result.frames, list)
assert isinstance(result.frames[0], np.ndarray)
assert result.samples is samples
def test_generate_single_return_samples(tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
def _fail_make_grid(*args, **kwargs):
raise AssertionError("make_grid should not be called")
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fail_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
return_samples=True,
)
assert isinstance(result, dict)
assert result["samples"] is samples
assert "frames" not in result
def test_generate_single_return_frames_list(tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
def _fake_make_grid(x, nrow=6):
return x[0]
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fake_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
return_frames=True,
)
assert isinstance(result, list)
assert isinstance(result[0], np.ndarray)
def test_generate_single_trajectory_flags(tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
def _fail_make_grid(*args, **kwargs):
raise AssertionError("make_grid should not be called")
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fail_make_grid,
)
result = generator.generate_video(
prompt="hello",
sampling_param=_base_params(tmp_path),
return_trajectory_latents=True,
return_trajectory_decoded=True,
)
assert result["trajectory"] == "trajectory_latents"
assert result["trajectory_timesteps"] == "trajectory_timesteps"
assert result["trajectory_decoded"] == "trajectory_decoded"
def test_generate_batch_prompt_path_minimal_outputs(tmp_path, monkeypatch):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
def _fail_make_grid(*args, **kwargs):
raise AssertionError("make_grid should not be called")
monkeypatch.setattr(
"fastvideo.entrypoints.video_generator.torchvision.utils.make_grid",
_fail_make_grid,
)
prompt_path = tmp_path / "prompts.txt"
prompt_path.write_text("first prompt\n\nsecond prompt\n", encoding="utf-8")
results = generator.generate_video(
sampling_param=_base_params(tmp_path),
prompt_path=str(prompt_path),
monitor=object(),
)
assert isinstance(results, list)
assert len(results) == 2
for idx, result in enumerate(results):
assert isinstance(result, dict)
assert result["prompt_index"] == idx
assert result["prompt"] in {"first prompt", "second prompt"}
assert "frames" not in result
assert "samples" not in result
def test_generate_prompt_path_requires_txt_extension(tmp_path):
samples = torch.rand(1, 3, 2, 4, 4)
generator = _new_stub_generator(samples)
prompt_path = tmp_path / "prompts.json"
prompt_path.write_text("hello\n", encoding="utf-8")
with pytest.raises(ValueError, match="prompt_path must be a txt file"):
generator.generate_video(
sampling_param=_base_params(tmp_path),
prompt_path=str(prompt_path),
)