Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3251822a7f |
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
Reference in New Issue
Block a user