Compare commits

..
3 Commits
Author SHA1 Message Date
SolitaryThinker ee607bd0d0 feat(api): typed omni request plane (OmniRequest/Output/Event)
First stacked PR of plan.md M1 (the request-plane cut). Adds the typed
omni vocabulary from design.md §6.1 as a new, additive, torch-free module
(fastvideo/api/omni.py):

- OmniRequest: declared TaskType, typed ModalPart inputs, AR SamplingParams
  vs DiffusionParams, OutputSpec, per-node NodeOverrides.
- OmniOutput: named Artifact slots with provenance (kills extra["audio"]).
- OmniEvent: one progress/chunk/final streaming union.
- Adapters to/from GenerationRequest / GenerationResult / VideoEvent so the
  new types run through today's VideoGenerator unchanged.

Later stacked PRs evolve GenerationRequest into OmniRequest in place and
retire the adapters. No existing surface changed.
2026-06-13 01:36:24 -07:00
alexzmsandmergify[bot] 633d393568 [ci] layer-0 grad-norm regression for per-method training tests (5a-ii) (#1396)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-12 04:45:07 +00:00
Junda SuandPeiyuan Zhang 5854aec2ce [feat] Add Wan RL DiffusionNFT training (#1450)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2026-06-11 21:18:59 -07:00
43 changed files with 4502 additions and 42 deletions
+32
View File
@@ -0,0 +1,32 @@
---
name: add-reward-model
description: Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.
---
# Add Reward Model
Use for reward models consumed by RL methods.
## Placement
- Put reusable reward code under `fastvideo/train/methods/rl/rewards/`.
- Expose public builders from `fastvideo/train/methods/rl/rewards/__init__.py`.
- Keep method-specific aggregation or advantage logic out of reward classes.
## Media Inputs
- Reward callables receive decoded media tensors.
- Accept single-frame tensors as `[B, C, H, W]` and multi-frame tensors as `[B, C, T, H, W]` when practical.
- Frame selection is reward-specific. Frame scorers such as PickScore and CLIPScore should explicitly select frame `0`; temporal rewards should inspect whichever frames they need.
- Return one scalar reward per prompt/sample.
## Attribution
- If code is ported or closely adapted from another repo, add a short comment or docstring naming the source file/function.
- Preserve SPDX headers used by FastVideo files.
## Tests
- Unit-test tensor layout handling without loading large reward checkpoints.
- Allow fake scorer injection for multi-reward tests.
- Test weighted reward aggregation and metric keys.
+38
View File
@@ -0,0 +1,38 @@
---
name: add-rl-method
description: Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.
---
# Add RL Method
Use for new RL methods in the modular `fastvideo/train` stack.
## Required Shape
- Add the method under `fastvideo/train/methods/rl/`.
- Subclass `TrainingMethod`.
- Keep model-family logic in `ModelBase` wrappers.
- Decode generated latents through `ModelBase.decode_latents`; add that hook to the new model wrapper instead of decoding inside the RL method.
- Use `fastvideo/train/methods/rl/common/sampling.py` for generation unless the method has a documented reason to avoid sampling.
- Use `fastvideo/train/methods/rl/common/prompt_sampling.py` for reusable grouped prompt sampling patterns such as DiffusionNFT K-repeat.
- Use `fastvideo/train/methods/rl/rewards/` for reward models.
## Optimization
- Return `manages_optimization() == True` only when the method must own a nonstandard outer/inner loop.
- If using managed optimization, implement `managed_train_step(data_stream, iteration)`.
- Existing trainer callbacks, checkpointing, tracking, and validation should still work.
## Config
- Put method knobs under `method`.
- Put sampler knobs under `method.sampling`.
- Do not put scheduler or trajectory policy into model configs.
- Do not split a diffusers-style scheduler from its built-in `step()` solver in YAML; use `trajectory` only for higher-level ODE vs re-noise behavior.
- Avoid fixed timestep lists in examples unless reproducing a known baseline; prefer scheduler-generated defaults.
## Tests
- Add fake-model tests for sampler/method behavior.
- Add config parse tests for the public YAML.
- Confirm existing train methods stay on the default Trainer path.
+3
View File
@@ -8,3 +8,6 @@
{"name": "decompose-pipeline-pr", "description": "Decompose an oversized FastVideo pipeline PR into a stack of independently-reviewable PRs. Tiers the diff by blast radius (invisible / dead code / cross-cutting infra / activation), produces a branch graph and worktree bootstrap, drafts the AGENTS.md manifest, flags missing tests on cross-cutting infra changes, and extracts lessons from the PR body. Worked example: PR #1280 daVinci-MagiHuman (9.8k LOC) decomposed into 10 stacked PRs.", "path": "decompose-pipeline-pr/SKILL.md", "status": "tested", "trust": "medium"}
{"name": "reseed-performance-baseline", "description": "Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline must be advanced by replicating one reviewed shifted source result into three success=true records, or five records when explicitly requested", "path": "reseed-performance-baseline/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-model", "description": "Add a new model (or variant) to FastVideo: DiT + configs + pipeline + presets + registry + tests. Walks through FastVideo's single stage-based pipeline architecture with exact file paths and registration hooks.", "path": "add-model/SKILL.md", "status": "draft", "trust": "low"}
{"name": "rlhf-training-abstractions", "description": "Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.", "path": "rlhf-training-abstractions/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-rl-method", "description": "Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.", "path": "add-rl-method/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-reward-model", "description": "Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.", "path": "add-reward-model/SKILL.md", "status": "draft", "trust": "low"}
@@ -0,0 +1,41 @@
---
name: rlhf-training-abstractions
description: Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.
---
# RLHF Training Abstractions
Use this skill before editing RLHF-style training code in `fastvideo/train`.
## Boundaries
- RL methods live under `fastvideo/train/methods/rl/` and own algorithm logic: reward collection, advantage computation, policy loss, KL/reference terms, and optimizer cadence.
- Rewards live under `fastvideo/train/methods/rl/rewards/` and must be reusable across RL methods.
- RL methods pass decoded media to rewards; each reward decides whether to use the first frame, sampled frames, or the full video.
- Sampling lives under `fastvideo/train/methods/rl/common/` and must use `ModelBase` primitives plus scheduler math, not model-family inference pipelines.
- Model wrappers under `fastvideo/train/models/` own model-specific forward details.
- Model wrappers also own model-specific latent decoding via `ModelBase.decode_latents`; RL methods should not reach into VAE normalization internals.
- Shared RL helpers such as K-repeat prompt sampling belong under `fastvideo/train/methods/rl/common/` when they are reusable across RL methods.
## Anti-Patterns
- Do not bind RL methods to inference pipeline classes such as `WanDMDPipeline`.
- Do not hardcode timestep lists in a method when the scheduler can generate them.
- Do not put reward-model code inside one RL method.
- Do not make existing non-RL methods use method-managed optimization unless explicitly requested.
## Sampling Policy
- Prefer YAML-configured `method.sampling` with `scheduler`, `trajectory`, `num_steps`, `timesteps`, and `sigmas`.
- Treat diffusers-style scheduler classes as owning both the timestep schedule and their `step()` update rule; avoid a separate `solver` field unless a new sampler truly implements solver math outside the scheduler object.
- Missing `timesteps` means “ask the scheduler”; explicit `timesteps` or `sigmas` are overrides.
- ODE-style trajectories should not re-noise between denoising steps.
- SDE/re-noise behavior must be explicit in config.
## Validation
- Run focused local tests for sampler config and Trainer opt-in behavior.
- Verify existing train methods still report `manages_optimization() == False`.
- Keep fixed-prompt validation helpers in `fastvideo/train/methods/rl/common/validation.py` so new RL methods can reuse sharding and captions.
- Test distributed prompt grouping helpers separately from heavyweight model loading.
- Run `pre-commit run --files <changed paths>`; respect configured excludes.
@@ -0,0 +1,52 @@
{
"data": [
{
"caption": "gold tip pyramid in the night, extremely detailed , rain, stars"
},
{
"caption": "a fruit stacking in the shape of a dog, stock image, shutterstock"
},
{
"caption": "Landscape, By Lee madgwick, by Luis Royo, by Louise nevelson"
},
{
"caption": "A colorful poster that says \"philo is a weird\""
},
{
"caption": "Danish male with blue eyes, realistic, viking"
},
{
"caption": "futuristic, cityscape, flying cars, neon lights, towering skyscrapers, glowing purple sky."
},
{
"caption": "a crow with cameras for eyes, sitting on a mans shoulder, anime, studio ghibli, fantasy, fairytale, sketch, digital art, watercolor, dnd, rustic, professional photograph, medieval, hd, 4k"
},
{
"caption": "a background image mixing the matrix and AI"
},
{
"caption": "Golden sunset, a bright orange and yellow sky is visible, lit up by the setting sun, the horizon is a mix of bright colors and deep shadows"
},
{
"caption": "Grim reaper playing an electric guitar"
},
{
"caption": "an epic view of a demonic Rose-ringed parakeet cyborg inside an ironmaiden robot,wearing a noble robe,large view,a surrealist painting, aralan bean and Philippe Druillet,hiromu arakawa,volumetric lighting,detailed shadows"
},
{
"caption": "Ben Shapiro as the cover of ministry's filth pig album, but covered in milk"
},
{
"caption": "king charles spaniel with , ethereal, midjourney style lighting and shadows, insanely detailed, 8k, photorealistic"
},
{
"caption": "A website for a party resort service"
},
{
"caption": "full shot of a steampunk horse"
},
{
"caption": "60s psycedelic spiritual jazz album art"
}
]
}
+1
View File
@@ -87,6 +87,7 @@ training:
# --- training.data [TYPED] -> DataConfig ---
data:
data_path: data/my_dataset # default: ""
preprocessed_data_type: t2v # default: "t2v" ("text_only" for simulate-only DMD text prompts)
train_batch_size: 1 # default: 1
dataloader_num_workers: 4 # default: 0
training_cfg_rate: 0.1 # default: 0.0
@@ -0,0 +1,116 @@
# DiffusionNFT multi-reward single-frame RL: Wan 2.1 T2V 1.3B on text-only PickScore prompts.
#
# Single-frame RL is represented as a one-latent-frame Wan run:
# num_latent_t: 1
# num_frames: 1
#
# The method trains the full transformer (no LoRA) and keeps an old-policy
# transformer plus a frozen reference transformer, matching the non-LoRA
# DiffusionNFT loss path.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
old:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
reference:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod
reward_fn:
pickscore: 1.0
clipscore: 1.0
sampling:
num_steps: 25
scheduler: flow_match_euler
trajectory: ode
flow_shift: inherit
validation:
every_steps: 10
num_steps: 40
num_prompts: 16
batch_size: 16
log_samples: true
seed: 42
# Null reuses training.data.data_path. Override this with a held-out
# preprocessed parquet path when one is available.
data_path:
# DiffusionNFT sd3_multi_reward on 4 GPUs resolves to per-GPU sample batch
# size 6, 48 sample batches per outer epoch, and grad accumulation 48.
sample_train_batch_size: 6
train_batch_size: 6
num_batches_per_epoch: 48
num_video_per_prompt: 24
num_inner_epochs: 1
timestep_fraction: 0.99
beta: 0.1
kl_beta: 0.0001
decay_type: 1
adv_mode: all
adv_clip_max: 5
max_grad_norm: 1.0
ema:
enabled: true
decay: 0.9
update_after_step: 0
validation: true
terminal_progress: true
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/pickscore_text_only_preprocessed
preprocessed_data_type: text_only
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 1
num_height: 448
num_width: 832
num_frames: 1
optimizer:
learning_rate: 3.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0001
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 100000
gradient_accumulation_steps: 48
checkpoint:
output_dir: outputs/wan2.1_diffusion_nft_pick_clip
training_state_checkpointing_steps: 30
checkpoints_total_limit: 3
tracker:
project_name: diffusion_nft_wan
run_name: wan2.1_diffusion_nft_pick_clip
model:
enable_gradient_checkpointing_type: full
pipeline:
flow_shift: 8
+624
View File
@@ -0,0 +1,624 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed omni request plane (design.md §6.1).
This module introduces the request-plane vocabulary the next-generation
runtime is built on:
- :class:`OmniRequest` — a typed multimodal request whose ``task`` is
*declared, never inferred*, and whose inputs are typed
:class:`ModalPart` s instead of per-model fields on a god-object.
- :class:`OmniOutput` — a typed multimodal output whose modalities are
named :class:`Artifact` slots carrying provenance, replacing the
``extra["audio"]`` escape hatch (design.md P3).
- :data:`OmniEvent` — one streaming-event union (progress / chunk /
final), the single channel that ``LoopStage.step`` emits through.
It is deliberately additive and engine-agnostic: nothing here imports
``torch`` or the pipeline machinery, and adapters bridge to/from today's
:class:`~fastvideo.api.schema.GenerationRequest` /
:class:`~fastvideo.api.results.GenerationResult` so the new types are
usable through the existing ``VideoGenerator`` before the engine itself
is rebuilt (plan.md M1). Later milestones evolve ``GenerationRequest``
into ``OmniRequest`` in place and drop the adapters.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, ClassVar
from uuid import uuid4
from fastvideo.api.results import (
GenerationResult,
VideoEvent,
VideoFinalEvent,
VideoPartialEvent,
VideoProgressEvent,
)
from fastvideo.api.schema import (
GenerationRequest,
InputConfig,
OutputConfig,
RequestRuntimeConfig,
SamplingConfig,
)
class Modality(str, Enum):
"""A media modality. ``str`` mixin keeps it JSON/serialization friendly."""
TEXT = "text"
IMAGE = "image"
VIDEO = "video"
AUDIO = "audio"
ACTION = "action"
LATENT = "latent"
class TaskType(str, Enum):
"""The declared task (design.md §6.1 / kills P7).
The pipeline graph branches on ``request.task``. Heuristics may only
*suggest* a default at the API boundary (see :func:`infer_task`); they
never decide control flow inside the runtime.
"""
T2V = "t2v" # text -> video
I2V = "i2v" # image (+text) -> video
TI2V = "ti2v" # text + init image -> video
V2V = "v2v" # video -> video (edit / restyle)
V2W = "v2w" # video -> world (continue a world-model rollout)
T2I = "t2i" # text -> image
I2I = "i2i" # image -> image (edit)
T2A = "t2a" # text -> audio
T2VS = "t2vs" # text -> video + sound (joint A/V)
A2W = "a2w" # action -> world (interactive world model)
REASON = "reason" # text -> text (AR reasoner / prompt upsampling)
# ---------------------------------------------------------------------------
# Inputs: typed modality parts (replaces InputConfig's per-model fields, P3).
# ---------------------------------------------------------------------------
@dataclass
class ModalPart:
"""Base for a typed input part.
``modality`` is intrinsic to the concrete subclass (a ``ClassVar``, not
an instance field). ``role`` disambiguates several parts of one
modality — e.g. ``"prompt"`` vs ``"negative"`` text, ``"init"`` vs
``"conditioning"`` image — and is keyword-only so subclass payloads stay
positional.
"""
modality: ClassVar[Modality]
role: str | None = field(default=None, kw_only=True)
@dataclass
class TextPart(ModalPart):
modality: ClassVar[Modality] = Modality.TEXT
text: str = ""
@dataclass
class ImagePart(ModalPart):
modality: ClassVar[Modality] = Modality.IMAGE
image: Any | None = None # PIL.Image / ndarray / tensor
path: str | None = None
@dataclass
class VideoPart(ModalPart):
modality: ClassVar[Modality] = Modality.VIDEO
video: Any | None = None
path: str | None = None
fps: float | None = None
@dataclass
class AudioPart(ModalPart):
modality: ClassVar[Modality] = Modality.AUDIO
audio: Any | None = None
path: str | None = None
sample_rate: int | None = None
@dataclass
class ActionPart(ModalPart):
modality: ClassVar[Modality] = Modality.ACTION
action: Any | None = None # mouse / keyboard / camera tensors
kind: str | None = None # "mouse" | "keyboard" | "camera" | ...
@dataclass
class LatentPart(ModalPart):
modality: ClassVar[Modality] = Modality.LATENT
latents: Any | None = None
of_modality: Modality = Modality.VIDEO # modality these latents decode to
# ---------------------------------------------------------------------------
# Per-call parameters: AR sampling vs diffusion, separated (design.md §6.1).
# ---------------------------------------------------------------------------
@dataclass
class SamplingParams:
"""AR decode knobs.
Distinct from the legacy diffusion ``fastvideo.api.SamplingParam``: these
drive ``ARDecodeLoop`` (Cosmos3 reasoner, omni thinkers/talkers, codec
decode), not denoising.
"""
max_tokens: int = 512
temperature: float = 1.0
top_p: float = 1.0
top_k: int | None = None
stop: list[str] = field(default_factory=list)
seed: int | None = None
@dataclass
class DiffusionParams:
"""Denoise knobs. ``guidance_per_modality`` carries per-modality CFG
scales for joint A/V denoise (LTX-2, Cosmos3 t2vs); ``guidance_scale`` is
the scalar default."""
steps: int = 50
guidance_scale: float = 1.0
guidance_per_modality: dict[Modality, float] = field(default_factory=dict)
sigmas: list[float] | None = None
flow_shift: float | None = None
height: int | None = None
width: int | None = None
num_frames: int | None = None
fps: int | None = None
seed: int | None = None
# ---------------------------------------------------------------------------
# Outputs spec: requested modalities + streaming + capture flags.
# ---------------------------------------------------------------------------
@dataclass
class StreamSpec:
"""Per-modality streaming policy. ``chunk_ms`` applies to audio chunks,
``per_chunk`` to chunked-causal video; ``enabled`` alone covers token
text."""
enabled: bool = True
chunk_ms: int | None = None
per_chunk: bool = False
@dataclass
class OutputSpec:
"""Requested output modalities, streaming, and capture flags."""
modalities: list[Modality] = field(default_factory=lambda: [Modality.VIDEO])
stream: dict[Modality, StreamSpec] = field(default_factory=dict)
return_latents: bool = False
return_trajectory: bool = False
@property
def streaming(self) -> bool:
"""True if any modality is requested as a stream."""
return any(spec.enabled for spec in self.stream.values())
@dataclass
class NodeOverrides:
"""Per-graph-node parameter overrides (design.md §6.1).
For a multi-loop graph, ``node_params["refine"].steps`` overrides the
refine loop's step count without leaking ``refine_*`` onto the universal
schema. Validation against each node's declared schema arrives with
``PipelineSpec`` (a later milestone); for now this is a typed bag with
``get`` / item / attribute access.
"""
params: dict[str, Any] = field(default_factory=dict)
def get(self, key: str, default: Any = None) -> Any:
return self.params.get(key, default)
def __getitem__(self, key: str) -> Any:
return self.params[key]
def __getattr__(self, key: str) -> Any:
# Only invoked when normal attribute lookup fails, so ``params``
# itself resolves through ``__dict__`` and never recurses.
try:
return self.__dict__["params"][key]
except KeyError as exc:
raise AttributeError(key) from exc
@dataclass
class OmniRequest:
"""A typed multimodal request (design.md §6.1)."""
task: TaskType
inputs: list[ModalPart] = field(default_factory=list)
sampling: SamplingParams = field(default_factory=SamplingParams)
diffusion: DiffusionParams = field(default_factory=DiffusionParams)
outputs: OutputSpec = field(default_factory=OutputSpec)
node_params: dict[str, NodeOverrides] = field(default_factory=dict)
priority: int = 0
request_id: str = field(default_factory=lambda: uuid4().hex)
# -- accessors ----------------------------------------------------------
def parts(self, modality: Modality) -> list[ModalPart]:
return [p for p in self.inputs if p.modality is modality]
@property
def prompt(self) -> str | None:
for part in self.inputs:
if isinstance(part, TextPart) and part.role in (None, "prompt"):
return part.text
return None
@property
def negative_prompt(self) -> str | None:
for part in self.inputs:
if isinstance(part, TextPart) and part.role in ("negative", "negative_prompt"):
return part.text
return None
# -- constructors / adapters -------------------------------------------
@classmethod
def from_prompt(
cls,
prompt: str | None,
task: TaskType,
*,
negative_prompt: str | None = None,
**kwargs: Any,
) -> OmniRequest:
"""Build a request from a bare prompt — what ``generate_video`` calls
internally so the offline shim constructs an ``OmniRequest`` (G5)."""
inputs: list[ModalPart] = []
if prompt is not None:
inputs.append(TextPart(prompt))
if negative_prompt is not None:
inputs.append(TextPart(negative_prompt, role="negative"))
return cls(task=task, inputs=inputs, **kwargs)
@classmethod
def from_generation_request(
cls,
request: GenerationRequest,
task: TaskType | None = None,
) -> OmniRequest:
"""Lift a legacy typed :class:`GenerationRequest` into an
``OmniRequest``; ``task`` defaults to the boundary heuristic."""
resolved = task if task is not None else infer_task(request)
inputs: list[ModalPart] = []
prompt = request.prompt
if isinstance(prompt, list):
prompt = prompt[0] if prompt else None
if prompt is not None:
inputs.append(TextPart(prompt))
if request.negative_prompt is not None:
inputs.append(TextPart(request.negative_prompt, role="negative"))
inp = request.inputs
image_path = inp.image_path if isinstance(inp.image_path, str) else None
if image_path is not None or inp.pil_image is not None:
inputs.append(ImagePart(image=inp.pil_image, path=image_path))
video_path = inp.video_path if isinstance(inp.video_path, str) else None
if video_path is not None:
inputs.append(VideoPart(path=video_path))
if inp.mouse_cond is not None or inp.keyboard_cond is not None:
inputs.append(ActionPart(action=inp.mouse_cond, kind="mouse"))
sampling = request.sampling
diffusion = DiffusionParams(
steps=sampling.num_inference_steps,
guidance_scale=sampling.guidance_scale,
sigmas=sampling.sigmas,
height=sampling.height,
width=sampling.width,
num_frames=sampling.num_frames,
fps=sampling.fps,
seed=sampling.seed,
)
outputs = OutputSpec(
return_trajectory=request.runtime.return_trajectory_latents,
return_latents=request.runtime.return_trajectory_decoded,
)
node_params = {
node: NodeOverrides(params=dict(overrides))
for node, overrides in request.stage_overrides.items() if isinstance(overrides, dict)
}
return cls(
task=resolved,
inputs=inputs,
diffusion=diffusion,
outputs=outputs,
node_params=node_params,
)
def to_generation_request(self) -> GenerationRequest:
"""Lower to a legacy :class:`GenerationRequest` so today's
``VideoGenerator`` can execute an ``OmniRequest`` unchanged (M1)."""
diff = self.diffusion
sampling = SamplingConfig()
sampling.num_inference_steps = diff.steps
sampling.guidance_scale = diff.guidance_scale
sampling.sigmas = diff.sigmas
if diff.seed is not None:
sampling.seed = diff.seed
if diff.num_frames is not None:
sampling.num_frames = diff.num_frames
if diff.height is not None:
sampling.height = diff.height
if diff.width is not None:
sampling.width = diff.width
if diff.fps is not None:
sampling.fps = diff.fps
image_path: str | None = None
pil_image: Any | None = None
video_path: str | None = None
for part in self.inputs:
if isinstance(part, ImagePart):
image_path = image_path or part.path
pil_image = pil_image if pil_image is not None else part.image
elif isinstance(part, VideoPart):
video_path = video_path or part.path
inputs = InputConfig(image_path=image_path, pil_image=pil_image, video_path=video_path)
runtime = RequestRuntimeConfig(
return_trajectory_latents=self.outputs.return_trajectory,
return_trajectory_decoded=self.outputs.return_latents,
)
output = OutputConfig(return_frames=Modality.VIDEO in self.outputs.modalities)
stage_overrides = {node: dict(ov.params) for node, ov in self.node_params.items()}
return GenerationRequest(
prompt=self.prompt,
negative_prompt=self.negative_prompt,
inputs=inputs,
sampling=sampling,
runtime=runtime,
output=output,
stage_overrides=stage_overrides,
)
# ---------------------------------------------------------------------------
# Outputs: named artifacts with provenance (kills extra["audio"], P3).
# ---------------------------------------------------------------------------
@dataclass
class Artifact:
"""Base output artifact. ``source_node`` records which graph node
produced it (provenance, design.md §6.1)."""
modality: ClassVar[Modality]
source_node: str | None = field(default=None, kw_only=True)
@dataclass
class VideoArtifact(Artifact):
modality: ClassVar[Modality] = Modality.VIDEO
frames: Any | None = None # numpy (N, H, W, 3) uint8
tensor: Any | None = None # raw sample tensor
path: str | None = None
fps: float | None = None
@dataclass
class AudioArtifact(Artifact):
modality: ClassVar[Modality] = Modality.AUDIO
audio: Any | None = None
sample_rate: int | None = None
@dataclass
class TextArtifact(Artifact):
modality: ClassVar[Modality] = Modality.TEXT
text: str = ""
token_ids: list[int] | None = None
@dataclass
class TensorArtifact(Artifact):
"""Action tensors and other raw tensor outputs."""
modality: ClassVar[Modality] = Modality.ACTION
tensor: Any | None = None
@dataclass
class LatentArtifact(Artifact):
modality: ClassVar[Modality] = Modality.LATENT
latents: Any | None = None
timesteps: Any | None = None
of_modality: Modality = Modality.VIDEO
@dataclass
class RequestMetrics:
generation_time: float | None = None
peak_memory_mb: float | None = None
extra: dict[str, Any] = field(default_factory=dict)
@dataclass
class OmniOutput:
"""Typed multimodal output (design.md §6.1)."""
request_id: str
artifacts: dict[str, Artifact] = field(default_factory=dict)
metrics: RequestMetrics = field(default_factory=RequestMetrics)
def get(self, name: str) -> Artifact | None:
return self.artifacts.get(name)
@property
def video(self) -> VideoArtifact | None:
art = self.artifacts.get("video")
return art if isinstance(art, VideoArtifact) else None
@property
def audio(self) -> AudioArtifact | None:
art = self.artifacts.get("audio")
return art if isinstance(art, AudioArtifact) else None
@classmethod
def from_generation_result(
cls,
result: GenerationResult,
request_id: str = "",
) -> OmniOutput:
"""Map a legacy :class:`GenerationResult` into named artifacts.
Audio becomes a first-class :class:`AudioArtifact` carrying its
sample rate, instead of riding in ``extra["audio"]`` (P3).
"""
artifacts: dict[str, Artifact] = {}
if result.frames is not None or result.samples is not None or result.video_path is not None:
artifacts["video"] = VideoArtifact(
frames=result.frames,
tensor=result.samples,
path=result.video_path,
source_node="decode",
)
if result.audio is not None:
artifacts["audio"] = AudioArtifact(
audio=result.audio,
sample_rate=result.audio_sample_rate,
source_node="audio_decode",
)
if result.trajectory is not None or result.trajectory_decoded is not None:
artifacts["latents"] = LatentArtifact(
latents=result.trajectory,
timesteps=result.trajectory_timesteps,
source_node="denoise",
)
return cls(
request_id=request_id,
artifacts=artifacts,
metrics=RequestMetrics(
generation_time=result.generation_time,
peak_memory_mb=result.peak_memory_mb,
),
)
# ---------------------------------------------------------------------------
# Streaming: one event union (evolves api/results.py's Video*Event).
# ---------------------------------------------------------------------------
@dataclass
class OmniProgressEvent:
"""Per-step progress telemetry."""
step: int
total_steps: int
node: str = "denoise"
@dataclass
class OmniChunkEvent:
"""A streamed artifact chunk — the universal ``StepResult.emit`` channel
(design.md §6.2.2): text tokens, audio chunks, or decoded frame chunks.
``pts`` optionally carries a presentation timestamp for raw-frame
streaming over WebRTC (design.md §9.1)."""
modality: Modality
index: int
payload: Any = None
pts: float | None = None
node: str | None = None
@dataclass
class OmniFinalEvent:
"""Terminal event carrying the full :class:`OmniOutput`."""
output: OmniOutput
OmniEvent = OmniProgressEvent | OmniChunkEvent | OmniFinalEvent
"""Union of every event the engine streams; consumers match by ``isinstance``."""
def infer_task(request: GenerationRequest) -> TaskType:
"""Best-effort boundary heuristic for a legacy request (design.md §6.1).
A *suggestion* only — the runtime always branches on the explicit
``OmniRequest.task``, never on this. Single-frame requests are images;
a video input implies edit; an image input implies image-to-video.
"""
inp = request.inputs
has_image = bool(inp.image_path or inp.pil_image)
has_video = bool(inp.video_path or inp.stage1_video)
if request.sampling.num_frames == 1:
return TaskType.I2I if has_image else TaskType.T2I
if has_video:
return TaskType.V2V
if has_image:
return TaskType.I2V
return TaskType.T2V
def omni_event_from_video_event(event: VideoEvent, request_id: str = "") -> OmniEvent:
"""Adapt a legacy :data:`~fastvideo.api.results.VideoEvent` to an
:data:`OmniEvent` so the streaming surface can migrate incrementally."""
if isinstance(event, VideoProgressEvent):
return OmniProgressEvent(step=event.step, total_steps=event.total_steps, node=event.stage)
if isinstance(event, VideoPartialEvent):
return OmniChunkEvent(modality=Modality.VIDEO, index=event.index, payload=event.frames, node="decode")
if isinstance(event, VideoFinalEvent):
if event.result is not None:
output = OmniOutput.from_generation_result(event.result, request_id)
else:
output = OmniOutput(request_id=request_id)
if event.frames is not None:
output.artifacts["video"] = VideoArtifact(frames=event.frames, source_node="decode")
return OmniFinalEvent(output=output)
raise TypeError(f"unknown VideoEvent type: {type(event).__name__}")
__all__ = [
"ActionPart",
"Artifact",
"AudioArtifact",
"AudioPart",
"DiffusionParams",
"ImagePart",
"LatentArtifact",
"LatentPart",
"Modality",
"ModalPart",
"NodeOverrides",
"OmniChunkEvent",
"OmniEvent",
"OmniFinalEvent",
"OmniOutput",
"OmniProgressEvent",
"OmniRequest",
"OutputSpec",
"RequestMetrics",
"SamplingParams",
"StreamSpec",
"TaskType",
"TensorArtifact",
"TextArtifact",
"TextPart",
"VideoArtifact",
"VideoPart",
"infer_task",
"omni_event_from_video_event",
]
+168
View File
@@ -0,0 +1,168 @@
# SPDX-License-Identifier: Apache-2.0
"""Tests for the typed omni request plane (fastvideo/api/omni.py, design.md §6.1)."""
from __future__ import annotations
import pytest
from fastvideo.api.omni import (
AudioArtifact,
DiffusionParams,
ImagePart,
Modality,
NodeOverrides,
OmniChunkEvent,
OmniFinalEvent,
OmniOutput,
OmniProgressEvent,
OmniRequest,
OutputSpec,
StreamSpec,
TaskType,
TextPart,
VideoArtifact,
infer_task,
omni_event_from_video_event,
)
from fastvideo.api.results import (
GenerationResult,
VideoFinalEvent,
VideoPartialEvent,
VideoProgressEvent,
)
from fastvideo.api.schema import GenerationRequest, InputConfig, SamplingConfig
def test_enums_are_strings():
assert TaskType.T2V == "t2v"
assert Modality.AUDIO == "audio"
# str mixin keeps them usable anywhere a string is expected
assert f"{TaskType.REASON}" == "TaskType.REASON" or TaskType.REASON.value == "reason"
assert TaskType("i2v") is TaskType.I2V
def test_request_id_autogenerated_and_unique():
a = OmniRequest(task=TaskType.T2V)
b = OmniRequest(task=TaskType.T2V)
assert a.request_id and b.request_id
assert a.request_id != b.request_id
def test_prompt_and_negative_accessors():
req = OmniRequest.from_prompt("a cat", TaskType.T2V, negative_prompt="blurry")
assert req.prompt == "a cat"
assert req.negative_prompt == "blurry"
assert req.parts(Modality.TEXT)
assert req.parts(Modality.IMAGE) == []
def test_positional_text_part_is_text_not_role():
# role is keyword-only, so the positional arg is the payload
part = TextPart("hello")
assert part.text == "hello"
assert part.role is None
def test_round_trip_through_generation_request():
req = OmniRequest(
task=TaskType.T2V,
inputs=[TextPart("a fox"), TextPart("ugly", role="negative")],
diffusion=DiffusionParams(steps=30, guidance_scale=6.0, height=480, width=832, num_frames=49),
)
legacy = req.to_generation_request()
assert isinstance(legacy, GenerationRequest)
assert legacy.prompt == "a fox"
assert legacy.negative_prompt == "ugly"
assert legacy.sampling.num_inference_steps == 30
assert legacy.sampling.guidance_scale == 6.0
assert legacy.sampling.height == 480
assert legacy.sampling.num_frames == 49
back = OmniRequest.from_generation_request(legacy, task=TaskType.T2V)
assert back.prompt == "a fox"
assert back.negative_prompt == "ugly"
assert back.diffusion.steps == 30
assert back.diffusion.guidance_scale == 6.0
assert back.diffusion.height == 480
assert back.diffusion.num_frames == 49
def test_image_input_lowers_and_lifts():
req = OmniRequest(task=TaskType.I2V, inputs=[TextPart("dance"), ImagePart(path="/tmp/x.png")])
legacy = req.to_generation_request()
assert legacy.inputs.image_path == "/tmp/x.png"
back = OmniRequest.from_generation_request(legacy)
assert back.task == TaskType.I2V # inferred from the image input
assert back.parts(Modality.IMAGE)[0].path == "/tmp/x.png" # type: ignore[attr-defined]
@pytest.mark.parametrize(
("inputs", "sampling", "expected"),
[
(InputConfig(), SamplingConfig(num_frames=49), TaskType.T2V),
(InputConfig(image_path="/a.png"), SamplingConfig(num_frames=49), TaskType.I2V),
(InputConfig(video_path="/a.mp4"), SamplingConfig(num_frames=49), TaskType.V2V),
(InputConfig(), SamplingConfig(num_frames=1), TaskType.T2I),
(InputConfig(image_path="/a.png"), SamplingConfig(num_frames=1), TaskType.I2I),
],
)
def test_infer_task_heuristic(inputs, sampling, expected):
req = GenerationRequest(prompt="x", inputs=inputs, sampling=sampling)
assert infer_task(req) == expected
def test_named_artifacts_replace_extra_audio():
result = GenerationResult(
prompt="song",
frames="FRAMES",
audio="WAV",
audio_sample_rate=44100,
generation_time=1.5,
peak_memory_mb=1234.0,
)
out = OmniOutput.from_generation_result(result, request_id="r1")
assert out.request_id == "r1"
assert isinstance(out.video, VideoArtifact)
assert out.video.frames == "FRAMES"
# audio is a first-class artifact carrying its sample rate, not extra["audio"]
assert isinstance(out.audio, AudioArtifact)
assert out.audio.sample_rate == 44100
assert out.audio.source_node == "audio_decode"
assert out.metrics.generation_time == 1.5
assert out.metrics.peak_memory_mb == 1234.0
def test_omni_event_from_video_event():
prog = omni_event_from_video_event(VideoProgressEvent(step=3, total_steps=50, stage="refine"))
assert isinstance(prog, OmniProgressEvent)
assert prog.step == 3 and prog.node == "refine"
chunk = omni_event_from_video_event(VideoPartialEvent(frames="CHUNK", index=2))
assert isinstance(chunk, OmniChunkEvent)
assert chunk.modality == Modality.VIDEO and chunk.index == 2 and chunk.payload == "CHUNK"
final = omni_event_from_video_event(
VideoFinalEvent(result=GenerationResult(frames="F", audio="A", audio_sample_rate=24000)),
request_id="rid",
)
assert isinstance(final, OmniFinalEvent)
assert final.output.request_id == "rid"
assert final.output.audio.sample_rate == 24000
def test_node_overrides_access():
req = OmniRequest(task=TaskType.T2V, node_params={"refine": NodeOverrides(params={"steps": 8})})
assert req.node_params["refine"].steps == 8 # attribute access
assert req.node_params["refine"]["steps"] == 8 # item access
assert req.node_params["refine"].get("missing", 0) == 0
with pytest.raises(AttributeError):
_ = req.node_params["refine"].nonexistent
# node_params survive the round-trip into stage_overrides
legacy = req.to_generation_request()
assert legacy.stage_overrides == {"refine": {"steps": 8}}
def test_output_spec_streaming_flag():
spec = OutputSpec(modalities=[Modality.VIDEO, Modality.AUDIO])
assert spec.streaming is False
spec.stream[Modality.AUDIO] = StreamSpec(enabled=True, chunk_ms=200)
assert spec.streaming is True
+29
View File
@@ -286,6 +286,35 @@ def run_train_framework_tests():
)
@app.function(gpu="L40S:1",
image=image,
timeout=1800,
secrets=[
modal.Secret.from_dict(
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
],
volumes={"/root/data": model_vol})
def seed_grad_norm_references():
"""Record the per-method grad-norm reference for the **CI GPU (L40S only)**.
Phase 2 / 5a-ii one-off seeding entrypoint. Pinned to ``gpu="L40S:1"`` (the
Modal CI runner), so this function only seeds the ``L40S`` key in
``fastvideo/tests/train/methods/grad_norm_refs.json``.
``FASTVIDEO_GRADNORM_UPDATE=1`` makes ``check_grad_norm_regression`` record
the measured norm instead of asserting; ``-rs`` surfaces the recorded value
in the log so it can be copied into the JSON.
To seed any other device (e.g. our local Blackwell dev box → ``GB200``
key), run the same env-var + pytest invocation directly on that
workstation — see the module docstring of ``grad_norm_regression.py`` for
the local command and the ``_DEVICE_MAPPINGS`` table.
"""
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
)
@app.function(gpu="L40S:1",
image=image,
timeout=3600,
@@ -67,6 +67,8 @@ class TestConstructor:
assert cb.num_frames is None
assert cb.sampling_timesteps is None
assert cb.output_dir is None
assert cb.offload_training_state is False
assert cb.unload_pipeline_after_validation is False
# Lazy fields not yet populated.
assert cb._pipeline is None
assert cb._sampling_param is None
@@ -83,12 +85,16 @@ class TestConstructor:
guidance_scale="4.5", # type: ignore[arg-type]
num_frames="77", # type: ignore[arg-type]
sampling_timesteps=["1000", "500"],
offload_training_state="1", # type: ignore[arg-type]
unload_pipeline_after_validation="false", # type: ignore[arg-type]
)
assert cb.every_steps == 50
assert cb.sampling_steps == [20, 40]
assert cb.guidance_scale == 4.5
assert cb.num_frames == 77
assert cb.sampling_timesteps == [1000, 500]
assert cb.offload_training_state is True
assert cb.unload_pipeline_after_validation is False
def test_pipeline_kwargs_collected(self) -> None:
cb = ValidationCallback(
@@ -0,0 +1,10 @@
{
"test_wan_causal_dfsft": {
"GB200": 2.9781,
"L40S": 3.2562
},
"test_wan_finetune": {
"GB200": 1.6486,
"L40S": 1.6467
}
}
@@ -0,0 +1,158 @@
# SPDX-License-Identifier: Apache-2.0
"""Layer-0 grad-norm regression for the per-method training smoke tests.
Phase 2 / 5a-ii: layers a device-keyed grad-norm check on top of the
finite/non-zero grad assertions established in 5a-i. After one
``single_train_step`` + ``backward``, the L2 norm of transformer block 0's
trainable gradients is compared against a reference value pinned per GPU in
``grad_norm_refs.json`` (next to this module).
Determinism: the harness seeds both the global RNG and the method's
``cuda_generator`` via ``method.on_train_start()`` (``training.data.seed`` in the
fixture), and the synthetic ``raw_batch`` is built *after* that call, so the
forward/backward is reproducible within bf16 reduction noise on a given GPU.
Why device-keyed: grad norms differ across GPU architectures (kernels,
accumulation order), so a single golden value can't cover every runner. The
JSON currently carries refs for the two GPUs we actually run on — ``L40S`` (CI)
and ``GB200`` (our Blackwell dev box; ``B200`` maps to the same key).
Seeding a reference for the current device:
- **CI / L40S** — invoke ``modal run`` against ``seed_grad_norm_references`` in
``fastvideo/tests/modal/pr_test.py`` (pinned to ``gpu="L40S:1"``), then copy
the recorded value from the log into ``grad_norm_refs.json``.
- **Local / non-L40S GPUs** — on that workstation::
FASTVIDEO_GRADNORM_UPDATE=1 \\
pytest fastvideo/tests/train/methods -vs -rs
The harness writes the measured norm into ``grad_norm_refs.json`` under the
device's key and skips the assertion for that run. Append a new substring
entry to ``_DEVICE_MAPPINGS`` first for any device not already listed.
"""
from __future__ import annotations
import json
import os
from pathlib import Path
import pytest
import torch
_REFS_PATH = Path(__file__).resolve().parent / "grad_norm_refs.json"
_UPDATE_ENV = "FASTVIDEO_GRADNORM_UPDATE"
# bf16 single-step smoke: catch gross breakage (wrong wiring, dead grads,
# scale regressions), not micro-drift from reduction nondeterminism.
_DEFAULT_RTOL = 0.10
# GPU-name substring -> reference key. First match wins. Only devices with
# seeded references in ``grad_norm_refs.json`` are listed here — to add a new
# GPU, append an entry, then seed the reference (see module docstring).
_DEVICE_MAPPINGS: tuple[tuple[str, str], ...] = (
("L40S", "L40S"),
("GB200", "GB200"),
("B200", "GB200"), # same Blackwell arch as GB200
)
def _device_name() -> str:
if not torch.cuda.is_available():
return "CPU"
return torch.cuda.get_device_name(0)
def resolve_device_key(device_name: str | None = None) -> str | None:
"""Map a CUDA device name to its reference key, or None if unsupported.
The substring match is case-insensitive so it survives driver/environment
differences in how ``torch.cuda.get_device_name`` capitalizes the model.
"""
name = device_name if device_name is not None else _device_name()
name_lower = name.lower()
for pattern, key in _DEVICE_MAPPINGS:
if pattern.lower() in name_lower:
return key
return None
def layer0_grad_norm(transformer) -> float:
"""Global L2 norm of transformer block 0's trainable gradients.
Block 0 is the reference surface 5a-i already isolates: its grad is the
*last* one produced during backprop, so a healthy value implies the whole
forward + chain-rule path is intact.
Accumulates the squared sums on the GPU and does a single CPU-GPU sync
(``.item()``) at the end, rather than one per parameter.
"""
blocks = getattr(transformer, "blocks", None)
assert blocks is not None and len(blocks) > 0, (
"transformer is expected to expose a non-empty ``.blocks``")
grads = [
p.grad for p in blocks[0].parameters()
if p.requires_grad and p.grad is not None
]
if not grads:
return 0.0
sq_sum = torch.zeros((), device=grads[0].device, dtype=torch.float32)
for g in grads:
sq_sum += g.detach().float().pow(2).sum()
return sq_sum.sqrt().item()
def _load_refs() -> dict[str, dict[str, float]]:
if _REFS_PATH.exists():
return json.loads(_REFS_PATH.read_text(encoding="utf-8"))
return {}
def _save_refs(refs: dict[str, dict[str, float]]) -> None:
_REFS_PATH.write_text(
json.dumps(refs, indent=2, sort_keys=True) + "\n",
encoding="utf-8")
def check_grad_norm_regression(
test_name: str,
transformer,
*,
rtol: float = _DEFAULT_RTOL,
) -> None:
"""Assert block-0 grad norm matches the device-keyed reference within rtol.
- Skips when the current GPU has no reference (unsupported device, or not
yet seeded) so a new runner never hard-fails before its golden exists.
- With ``FASTVIDEO_GRADNORM_UPDATE=1`` records/updates the reference for the
current device instead of asserting.
"""
norm = layer0_grad_norm(transformer)
device_key = resolve_device_key()
if os.environ.get(_UPDATE_ENV) == "1":
if device_key is None:
pytest.skip(
f"{_UPDATE_ENV}=1 but GPU '{_device_name()}' has no reference "
"key; add it to _DEVICE_MAPPINGS first")
refs = _load_refs()
refs.setdefault(test_name, {})[device_key] = round(norm, 4)
_save_refs(refs)
pytest.skip(
f"recorded grad-norm reference {test_name}[{device_key}] = "
f"{norm:.4f} (assertion skipped under {_UPDATE_ENV}=1)")
ref = _load_refs().get(test_name, {}).get(device_key) \
if device_key is not None else None
if ref is None:
pytest.skip(
f"no grad-norm reference for {test_name} on '{_device_name()}' "
f"(device_key={device_key}); run with {_UPDATE_ENV}=1 to seed it")
rel = abs(norm - ref) / (abs(ref) + 1e-12)
assert rel <= rtol, (
f"{test_name}[{device_key}] grad-norm regression: got {norm:.4f}, "
f"reference {ref:.4f}, relative error {rel:.3%} exceeds rtol "
f"{rtol:.0%}. If this is an intentional change, refresh the reference "
f"with {_UPDATE_ENV}=1 and explain why in the PR.")
@@ -28,6 +28,8 @@ from fastvideo.train.methods.fine_tuning.dfsft import (
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -122,3 +124,7 @@ def test_wan_causal_dfsft_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
@@ -36,6 +36,8 @@ from fastvideo.train.methods.fine_tuning.finetune import (
from fastvideo.train.models.wan import WanModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -139,3 +141,7 @@ def test_wan_finetune_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_finetune", model.transformer)
+186 -8
View File
@@ -68,6 +68,8 @@ class ValidationCallback(Callback):
num_frames: int | None = None,
output_dir: str | None = None,
sampling_timesteps: list[int] | None = None,
offload_training_state: bool = False,
unload_pipeline_after_validation: bool = False,
**pipeline_kwargs: Any,
) -> None:
self.pipeline_target = str(pipeline_target)
@@ -78,6 +80,8 @@ class ValidationCallback(Callback):
self.num_frames = (int(num_frames) if num_frames is not None else None)
self.output_dir = (str(output_dir) if output_dir is not None else None)
self.sampling_timesteps = ([int(s) for s in sampling_timesteps] if sampling_timesteps is not None else None)
self.offload_training_state = self._coerce_bool(offload_training_state)
self.unload_pipeline_after_validation = self._coerce_bool(unload_pipeline_after_validation)
self.pipeline_kwargs = dict(pipeline_kwargs)
# Set after on_train_start.
@@ -88,6 +92,12 @@ class ValidationCallback(Callback):
self.validation_random_generator: (torch.Generator | None) = None
self.seed: int = 0
@staticmethod
def _coerce_bool(value: Any) -> bool:
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
return bool(value)
# ----------------------------------------------------------
# Callback hooks
# ----------------------------------------------------------
@@ -140,16 +150,183 @@ class ValidationCallback(Callback):
) -> None:
transformer = method.student.transformer
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
try:
with self._validation_memory_context(
method,
validation_transformer=transformer,
):
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
method,
step,
t,
)
finally:
if self.unload_pipeline_after_validation:
self._clear_pipeline_cache()
@contextlib.contextmanager
def _validation_memory_context(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
):
if not self.offload_training_state:
yield
return
optimizer_tensor_records: list[tuple[Any, Any, torch.device]] = []
module_records: list[tuple[str, torch.nn.Module, torch.device]] = []
try:
self._offload_optimizer_states_to_cpu(
method,
step,
t,
optimizer_tensor_records,
)
self._offload_inactive_role_modules_to_cpu(
method,
validation_transformer=validation_transformer,
module_records=module_records,
)
self._empty_cuda_cache()
yield
finally:
self._restore_inactive_role_modules(module_records)
self._restore_optimizer_states(optimizer_tensor_records)
self._empty_cuda_cache()
def _offload_optimizer_states_to_cpu(
self,
method: TrainingMethod,
records: list[tuple[Any, Any, torch.device]],
) -> None:
optimizers = getattr(method, "_optimizer_dict", {})
if not optimizers:
return
moved = 0
for optimizer in optimizers.values():
state = getattr(optimizer, "state", None)
if not isinstance(state, dict):
continue
for param_state in state.values():
moved += self._offload_tensor_container_to_cpu(
param_state,
records,
)
if moved:
logger.info(
"Offloaded %d optimizer state tensors to CPU for validation.",
moved,
)
def _offload_tensor_container_to_cpu(
self,
obj: Any,
records: list[tuple[Any, Any, torch.device]],
) -> int:
moved = 0
if isinstance(obj, dict):
for key, value in list(obj.items()):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, key, value.device))
obj[key] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
if isinstance(obj, list):
for idx, value in enumerate(list(obj)):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, idx, value.device))
obj[idx] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
def _restore_optimizer_states(
self,
records: list[tuple[Any, Any, torch.device]],
) -> None:
for container, key, device in reversed(records):
value = container[key]
if torch.is_tensor(value):
container[key] = value.to(device=device)
if records:
logger.info(
"Restored %d optimizer state tensors after validation.",
len(records),
)
def _offload_inactive_role_modules_to_cpu(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
role_models = getattr(method, "_role_models", {})
if not isinstance(role_models, dict):
return
for role, model in role_models.items():
module = getattr(model, "transformer", None)
if not isinstance(module, torch.nn.Module):
continue
if module is validation_transformer:
continue
device = self._first_cuda_tensor_device(module)
if device is None:
continue
try:
module.to("cpu")
except Exception as exc:
logger.warning(
"Could not offload role %r transformer to CPU before validation: %s",
role,
exc,
)
continue
module_records.append((str(role), module, device))
logger.info(
"Offloaded role %r transformer from %s to CPU for validation.",
role,
device,
)
def _restore_inactive_role_modules(
self,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
for role, module, device in reversed(module_records):
module.to(device)
logger.info(
"Restored role %r transformer to %s after validation.",
role,
device,
)
@staticmethod
def _first_cuda_tensor_device(module: torch.nn.Module) -> torch.device | None:
for tensor in list(module.parameters(recurse=True)) + list(module.buffers(recurse=True)):
device = getattr(tensor, "device", None)
if isinstance(device, torch.device) and device.type == "cuda":
return device
return None
def _clear_pipeline_cache(self) -> None:
self._pipeline = None
self._pipeline_key = None
self._empty_cuda_cache()
@staticmethod
def _empty_cuda_cache() -> None:
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _find_ema_callback(self) -> Any | None:
"""Find the EMA callback in the callback dict."""
@@ -293,6 +470,7 @@ class ValidationCallback(Callback):
}
if flow_shift is not None:
kwargs["flow_shift"] = float(flow_shift)
kwargs.update(self.pipeline_kwargs)
self._pipeline = PipelineCls.from_pretrained(
tc.model_path,
+40
View File
@@ -197,6 +197,46 @@ class TrainingMethod(torch.nn.Module, ABC):
# -- Shared hooks (override in subclasses as needed) --
def manages_optimization(self) -> bool:
"""Whether the method owns backward/optimizer stepping internally.
Most methods return loss tensors and let :class:`Trainer` handle
gradient accumulation, callbacks, optimizer stepping, and scheduler
stepping. RL-style methods such as DiffusionNFT need to preserve their
own sample-then-inner-train loop, so they can provide a specific
``managed_train_step``.
"""
return False
def managed_train_step(
self,
data_stream: Any,
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
"""Run one method-managed step.
Subclasses that return ``True`` from :meth:`manages_optimization`
should override this. The fallback consumes one dataloader batch and
delegates to ``single_train_step`` so tests can exercise the hook with
tiny fake methods.
"""
return self.single_train_step(next(data_stream), iteration)
def on_validation_begin(self, iteration: int = 0) -> dict[str, LogScalar]:
"""Run method-owned validation, if any.
Pipeline-style validation should remain in callbacks. Methods that
intentionally avoid inference pipelines, such as RL methods with their
own sampler/reward loop, can override this hook and return metrics for
the trainer to log at ``iteration``.
"""
del iteration
return {}
def get_grad_clip_targets(
self,
iteration: int,
@@ -55,6 +55,8 @@ class DMD2Method(TrainingMethod):
raise ValueError("DMD2Method requires critic to be trainable")
self._cfg_uncond = self._parse_cfg_uncond()
self._rollout_mode = self._parse_rollout_mode()
self._validate_preprocessed_data_type()
self._configure_student_negative_conditioning()
self._denoising_step_list: torch.Tensor | None = (None)
# Initialize preprocessors on student.
@@ -206,6 +208,13 @@ class DMD2Method(TrainingMethod):
return targets
def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]:
"""Parse how DMD2 obtains the latent point used for rollout.
``simulate`` starts from fresh noise and lets the student create an
artificial latent trajectory, so it can run with text-only data.
``data_latent`` starts from preprocessed VAE latents and perturbs them
at a sampled denoising timestep.
"""
raw = self.method_config.get("rollout_mode", None)
if raw is None:
raise ValueError("method_config.rollout_mode must be set "
@@ -223,6 +232,34 @@ class DMD2Method(TrainingMethod):
"{simulate, data_latent}, got "
f"{raw!r}")
def _validate_preprocessed_data_type(self) -> None:
data_type = str(getattr(
self.training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
if data_type == "text_only" and self._rollout_mode != "simulate":
raise ValueError("training.data.preprocessed_data_type='text_only' "
"requires method.rollout_mode='simulate'; "
"data_latent rollout requires vae_latent data.")
def _uses_negative_prompt_conditioning(self) -> bool:
if self._cfg_uncond is None:
return True
text_policy = self._cfg_uncond.get("text", None)
if text_policy is None:
return True
return str(text_policy).strip().lower() == "negative_prompt"
def _configure_student_negative_conditioning(self) -> None:
setter = getattr(
self.student,
"set_requires_negative_conditioning",
None,
)
if setter is not None:
setter(self._uses_negative_prompt_conditioning())
def _parse_cfg_uncond(self, ) -> dict[str, Any] | None:
raw = self.method_config.get("cfg_uncond", None)
if raw is None:
+6
View File
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""RL training methods."""
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
__all__ = ["DiffusionNFTMethod"]
@@ -0,0 +1,30 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable RL training primitives."""
from fastvideo.train.methods.rl.common.sampling import (
DiffusionSampler,
SamplingConfig,
SamplingResult,
)
from fastvideo.train.methods.rl.common.prompt_sampling import (
KRepeatSample,
distributed_k_repeat_indices,
)
from fastvideo.train.methods.rl.common.validation import (
RLValidationConfig,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
__all__ = [
"DiffusionSampler",
"KRepeatSample",
"RLValidationConfig",
"SamplingConfig",
"SamplingResult",
"distributed_k_repeat_indices",
"media_to_video_array",
"validation_caption",
"validation_shard_indices",
]
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""Prompt-row sampling helpers for online RL methods.
This module chooses and repeats dataset prompt rows across ranks for RL training
batches. Here, "sampling" means selection, not generator sampling.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True, slots=True)
class KRepeatSample:
"""Local prompt indices for one distributed K-repeat sampling batch."""
local_indices: list[int]
unique_prompt_count: int
def distributed_k_repeat_indices(
*,
dataset_length: int,
batch_size: int,
repeats_per_prompt: int,
world_size: int,
rank: int,
seed: int,
) -> KRepeatSample:
"""Mirror DiffusionNFT's distributed K-repeat prompt sampler.
Adapted from DiffusionNFT's
``scripts/train_nft_sd3.py::DistributedKRepeatSampler``.
"""
dataset_length = int(dataset_length)
batch_size = int(batch_size)
repeats_per_prompt = int(repeats_per_prompt)
world_size = int(world_size)
rank = int(rank)
if dataset_length <= 0:
raise ValueError("dataset_length must be positive")
if batch_size <= 0:
raise ValueError("batch_size must be positive")
if repeats_per_prompt <= 0:
raise ValueError("repeats_per_prompt must be positive")
if world_size <= 0:
raise ValueError("world_size must be positive")
if rank < 0 or rank >= world_size:
raise ValueError(f"rank must be in [0, {world_size}), got {rank}")
total_samples = world_size * batch_size
if total_samples % repeats_per_prompt != 0:
raise ValueError("world_size * batch_size must be divisible by repeats_per_prompt "
f"({world_size} * {batch_size} vs {repeats_per_prompt})")
unique_prompt_count = total_samples // repeats_per_prompt
if unique_prompt_count > dataset_length:
raise ValueError("K-repeat sampling needs at least as many rows as unique prompts "
f"per sampling batch ({dataset_length} < {unique_prompt_count})")
generator = torch.Generator()
generator.manual_seed(int(seed))
indices = torch.randperm(dataset_length, generator=generator)[:unique_prompt_count].tolist()
repeated_indices = [idx for idx in indices for _ in range(repeats_per_prompt)]
shuffled_order = torch.randperm(len(repeated_indices), generator=generator).tolist()
shuffled_samples = [int(repeated_indices[idx]) for idx in shuffled_order]
start = rank * batch_size
end = start + batch_size
return KRepeatSample(
local_indices=shuffled_samples[start:end],
unique_prompt_count=unique_prompt_count,
)
@@ -0,0 +1,223 @@
# SPDX-License-Identifier: Apache-2.0
"""Configurable diffusion samplers for RL training methods."""
from __future__ import annotations
import copy
from dataclasses import dataclass
from typing import Any, Literal
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.base import ModelBase
SchedulerName = Literal["flow_match_euler", "model_default"]
TrajectoryName = Literal["ode", "sde_reflow"]
@dataclass(slots=True)
class SamplingConfig:
"""YAML-backed sampling knobs shared by RL methods."""
num_steps: int = 25
scheduler: SchedulerName = "model_default"
trajectory: TrajectoryName = "ode"
flow_shift: float | None = None
timesteps: list[float] | None = None
sigmas: list[float] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> SamplingConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.sampling must be a mapping, got {type(raw).__name__}")
supported_keys = {
"flow_shift",
"num_steps",
"scheduler",
"sigmas",
"timesteps",
"trajectory",
}
unsupported_keys = sorted(set(raw) - supported_keys)
if unsupported_keys:
raise ValueError(f"Unsupported method.sampling key(s): {unsupported_keys}. "
f"Supported keys: {sorted(supported_keys)}")
scheduler = str(raw.get("scheduler", "model_default") or "model_default").strip().lower()
if scheduler not in {"flow_match_euler", "model_default"}:
raise ValueError("method.sampling.scheduler must be one of "
"{flow_match_euler, model_default}, got "
f"{raw.get('scheduler')!r}")
trajectory = str(raw.get("trajectory", "ode") or "ode").strip().lower()
if trajectory not in {"ode", "sde_reflow"}:
raise ValueError("method.sampling.trajectory must be one of "
"{ode, sde_reflow}, got "
f"{raw.get('trajectory')!r}")
timesteps = raw.get("timesteps")
sigmas = raw.get("sigmas")
if timesteps is not None:
if not isinstance(timesteps, list) or not timesteps:
raise ValueError("method.sampling.timesteps must be a non-empty list when set")
timesteps = [float(t) for t in timesteps]
if sigmas is not None:
if not isinstance(sigmas, list) or not sigmas:
raise ValueError("method.sampling.sigmas must be a non-empty list when set")
sigmas = [float(s) for s in sigmas]
if timesteps is not None and sigmas is not None and len(timesteps) != len(sigmas):
raise ValueError("method.sampling.timesteps and method.sampling.sigmas must have the same length")
num_steps = int(raw.get("num_steps", 25) or 25)
if num_steps <= 0:
raise ValueError("method.sampling.num_steps must be positive")
return cls(
num_steps=num_steps,
scheduler=scheduler, # type: ignore[arg-type]
trajectory=trajectory, # type: ignore[arg-type]
flow_shift=(None if raw.get("flow_shift", None) in (None, "inherit") else float(raw["flow_shift"])),
timesteps=timesteps,
sigmas=sigmas,
)
@dataclass(slots=True)
class SamplingResult:
latents: torch.Tensor
timesteps: torch.Tensor
sigmas: torch.Tensor
class DiffusionSampler:
"""Thin model/scheduler sampler used by RL methods.
This intentionally does not call FastVideo's full inference pipelines.
RL training needs a reusable sampling primitive that works with
``ModelBase`` wrappers and scheduler math without binding a method to
model-family pipeline classes such as ``WanDMDPipeline``.
"""
def __init__(self, config: SamplingConfig) -> None:
self.config = config
@torch.no_grad()
def sample(
self,
model: ModelBase,
batch: TrainingBatch,
*,
generator: torch.Generator | None,
) -> SamplingResult:
latents = batch.latents
if latents is None:
raise RuntimeError("TrainingBatch.latents is required for RL sampling")
current = torch.randn(
latents.shape,
device=latents.device,
dtype=latents.dtype,
generator=generator,
)
scheduler = self._prepare_scheduler(model, current.device)
timesteps = scheduler.timesteps.to(device=current.device)
sigmas = scheduler.sigmas.to(device=current.device)
original_timesteps = batch.timesteps
try:
if self.config.trajectory == "ode":
pred_clean = current
for timestep in timesteps:
model_timestep = self._model_timestep(timestep, current)
batch.timesteps = model_timestep
pred_noise = model.predict_noise(
current,
model_timestep,
batch,
conditional=True,
attn_kind="dense",
)
current = scheduler.step(
pred_noise.flatten(0, 1),
timestep,
current.flatten(0, 1),
return_dict=False,
)[0].unflatten(0, pred_noise.shape[:2])
pred_clean = current
return SamplingResult(latents=pred_clean, timesteps=timesteps, sigmas=sigmas)
return SamplingResult(
latents=self._sample_sde_reflow(
model,
batch,
current,
timesteps,
generator=generator,
),
timesteps=timesteps,
sigmas=sigmas,
)
finally:
batch.timesteps = original_timesteps
def _prepare_scheduler(
self,
model: ModelBase,
device: torch.device,
) -> Any:
if self.config.scheduler == "flow_match_euler":
shift = self.config.flow_shift
if shift is None:
shift = float(getattr(model.noise_scheduler, "shift", 1.0))
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(shift))
else:
scheduler = copy.deepcopy(model.noise_scheduler)
kwargs: dict[str, Any] = {"device": device}
if self.config.timesteps is not None:
kwargs["timesteps"] = self.config.timesteps
kwargs["num_inference_steps"] = len(self.config.timesteps)
if self.config.sigmas is not None:
kwargs["sigmas"] = self.config.sigmas
kwargs["num_inference_steps"] = len(self.config.sigmas)
if "num_inference_steps" not in kwargs:
kwargs["num_inference_steps"] = self.config.num_steps
scheduler.set_timesteps(**kwargs)
return scheduler
def _sample_sde_reflow(
self,
model: ModelBase,
batch: TrainingBatch,
current: torch.Tensor,
timesteps: torch.Tensor,
*,
generator: torch.Generator | None,
) -> torch.Tensor:
pred_clean = current
for step_idx, timestep in enumerate(timesteps):
timestep_tensor = self._model_timestep(timestep, current)
batch.timesteps = timestep_tensor
pred_clean = model.predict_x0(
current,
timestep_tensor,
batch,
conditional=True,
attn_kind="dense",
)
if step_idx < len(timesteps) - 1:
next_timestep = timesteps[step_idx + 1].reshape(1).to(device=current.device)
noise = torch.randn(
pred_clean.shape,
device=pred_clean.device,
dtype=pred_clean.dtype,
generator=generator,
)
current = model.add_noise(pred_clean, noise, next_timestep)
return pred_clean
@staticmethod
def _model_timestep(
timestep: torch.Tensor,
current: torch.Tensor,
) -> torch.Tensor:
return timestep.reshape(1).to(device=current.device).expand(current.shape[0]).contiguous()
@@ -0,0 +1,81 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared validation helpers for RL training methods."""
from __future__ import annotations
from dataclasses import dataclass
import math
from typing import Any
import torch
@dataclass(slots=True)
class RLValidationConfig:
every_steps: int = 0
num_steps: int = 40 # Reference DiffusionNFT sampling num steps for best visual quality
num_prompts: int = 16
batch_size: int = 16
log_samples: bool = True
seed: int = 42
data_path: str | None = None
sampling: dict[str, Any] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> RLValidationConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.validation must be a mapping, got {type(raw).__name__}")
data_path = raw.get("data_path", None)
sampling = raw.get("sampling", None)
if sampling is not None and not isinstance(sampling, dict):
raise ValueError(f"method.validation.sampling must be a mapping, got {type(sampling).__name__}")
return cls(
every_steps=max(0, int(raw.get("every_steps", 0) or 0)),
num_steps=max(1, int(raw.get("num_steps", 40) or 40)),
num_prompts=max(1, int(raw.get("num_prompts", 16) or 16)),
batch_size=max(1, int(raw.get("batch_size", 16) or 16)),
log_samples=bool(raw.get("log_samples", True)),
seed=int(raw.get("seed", 42) or 42),
data_path=(None if data_path in (None, "") else str(data_path)),
sampling=(dict(sampling) if sampling is not None else None),
)
def validation_shard_indices(
num_prompts: int,
*,
rank: int,
world_size: int,
) -> list[tuple[int, bool]]:
"""Return fixed validation prompt indices for one distributed rank."""
num_prompts = max(1, int(num_prompts))
world_size = max(1, int(world_size))
per_rank = int(math.ceil(num_prompts / world_size))
padded_total = per_rank * world_size
return [((idx % num_prompts), idx < num_prompts) for idx in range(rank, padded_total, world_size)]
def validation_caption(
prompt: str,
rewards: dict[str, float],
) -> str:
reward_parts = [f"{key}: {float(rewards[key]):.4f}" for key in sorted(rewards)]
return f"{' | '.join(reward_parts)} | {prompt[:1000]}"
def media_to_video_array(media: torch.Tensor) -> Any:
"""Convert decoded media to a tracker video array.
Accepts ``[C, T, H, W]`` tensors. ``[C, H, W]`` tensors are treated as
``T=1`` media. Output follows the existing tracker convention used
elsewhere in FastVideo: ``[T, C, H, W]`` uint8.
"""
if media.ndim == 3:
media = media.unsqueeze(1)
if media.ndim != 4:
raise ValueError("media must have shape [C, T, H, W] or [C, H, W], "
f"got {tuple(media.shape)}")
video = (media.detach().float().clamp(0, 1) * 255).round().to(torch.uint8)
return video.permute(1, 0, 2, 3).contiguous().cpu().numpy()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,37 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable reward models for training methods."""
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
select_first_frame,
)
def build_multi_reward_scorer(
reward_weights,
*,
device="cuda",
scorers: dict[str, RewardScorer] | None = None,
) -> MultiRewardScorer:
available: dict[str, RewardScorer] = dict(scorers or {})
if not available:
available = {
"pickscore": PickScoreScorer(device=device),
"clipscore": ClipScoreScorer(device=device),
}
return MultiRewardScorer(reward_weights, scorers=available)
__all__ = [
"ClipScoreScorer",
"MultiRewardScorer",
"PickScoreScorer",
"RewardScorer",
"build_multi_reward_scorer",
"select_first_frame",
]
@@ -0,0 +1,130 @@
# SPDX-License-Identifier: Apache-2.0
"""Frame-based reward scorers used by RL training methods."""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from PIL import Image
import torch
from fastvideo.train.methods.rl.rewards.media import select_first_frame
class PickScoreScorer(torch.nn.Module):
"""PickScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/pickscore_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
from transformers import AutoModel, AutoProcessor
processor_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
model_path = "yuvalkirstain/PickScore_v1"
self.device = torch.device(device)
self.dtype = dtype
self.processor = AutoProcessor.from_pretrained(processor_path)
self.model = AutoModel.from_pretrained(model_path).eval().to(self.device)
self.model = self.model.to(dtype=dtype)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media)
frame_np = (frame_tensor.detach().float().clamp(0, 1) * 255).round()
frame_np = frame_np.to(torch.uint8).cpu().numpy().transpose(0, 2, 3, 1)
pil_frames = [Image.fromarray(frame) for frame in frame_np]
frame_inputs = self.processor(
images=pil_frames,
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
frame_inputs = {k: v.to(device=self.device) for k, v in frame_inputs.items()}
text_inputs = self.processor(
text=list(prompts),
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
text_inputs = {k: v.to(device=self.device) for k, v in text_inputs.items()}
text_embs = self.model.get_text_features(**text_inputs)
text_embs = text_embs / text_embs.norm(p=2, dim=-1, keepdim=True)
frame_embs = self.model.get_image_features(**frame_inputs)
frame_embs = frame_embs / frame_embs.norm(p=2, dim=-1, keepdim=True)
scores = self.model.logit_scale.exp() * (text_embs @ frame_embs.T)
return scores.diag().float() / 26.0
class ClipScoreScorer(torch.nn.Module):
"""CLIPScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/clip_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
) -> None:
super().__init__()
import torch.nn as nn
import torchvision.transforms as T
from transformers import CLIPModel, CLIPProcessor
def get_size(size: Any) -> Any:
if isinstance(size, int):
return (size, size)
if isinstance(size, Mapping) and "height" in size and "width" in size:
return (size["height"], size["width"])
if isinstance(size, Mapping) and "shortest_edge" in size:
return size["shortest_edge"]
raise ValueError(f"Invalid processor size: {size!r}")
def get_frame_transform(processor: Any) -> torch.nn.Module:
config = processor.to_dict()
resize = T.Resize(get_size(config.get("size"))) if config.get("do_resize") else nn.Identity()
crop = T.CenterCrop(get_size(config.get("crop_size"))) if config.get("do_center_crop") else nn.Identity()
normalize = (T.Normalize(mean=processor.image_mean, std=processor.image_std)
if config.get("do_normalize") else nn.Identity())
return T.Compose([resize, crop, normalize])
self.device = torch.device(device)
self.model = CLIPModel.from_pretrained("openai/clip-vit-large-patch14").to(self.device).eval()
self.processor = CLIPProcessor.from_pretrained("openai/clip-vit-large-patch14")
self.transform = get_frame_transform(self.processor.image_processor)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media).detach().float().clamp(0, 1)
texts = self.processor(
text=list(prompts),
padding="max_length",
truncation=True,
return_tensors="pt",
).to(self.device)
pixels = self.transform(frame_tensor).to(device=self.device, dtype=frame_tensor.dtype)
outputs = self.model(pixel_values=pixels, **texts)
return outputs.logits_per_image.diagonal().float() / 100.0
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Generic media reward composition utilities."""
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
import torch
RewardScorer = Callable[[torch.Tensor, Sequence[str]], torch.Tensor]
def select_first_frame(media: torch.Tensor) -> torch.Tensor:
"""Return first-frame media as ``[B, C, H, W]``.
This is a helper for reward models that are intrinsically frame-based
(for example PickScore and CLIPScore). Video-aware rewards should inspect
the full ``[B, C, T, H, W]`` tensor themselves.
"""
if not torch.is_tensor(media):
raise TypeError(f"media must be a torch.Tensor, got {type(media).__name__}")
if media.ndim == 5:
return media[:, :, 0]
if media.ndim == 4:
return media
raise ValueError("media must have shape [B, C, H, W] or [B, C, T, H, W], "
f"got {tuple(media.shape)}")
class MultiRewardScorer:
"""Weighted sum of reusable media reward scorers.
Mirrors DiffusionNFT's ``flow_grpo/rewards.py::multi_score`` behavior,
while leaving frame selection to each concrete reward.
"""
def __init__(
self,
reward_weights: Mapping[str, float],
*,
scorers: Mapping[str, RewardScorer],
) -> None:
self.reward_weights = {str(k): float(v) for k, v in reward_weights.items()}
if not self.reward_weights:
raise ValueError("reward_weights must contain at least one reward")
self.scorers = dict(scorers)
unsupported = sorted(set(self.reward_weights) - set(self.scorers))
if unsupported:
raise ValueError(f"Unsupported reward(s): {unsupported}. "
f"Available rewards: {sorted(self.scorers)}")
@torch.no_grad()
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> dict[str, torch.Tensor]:
prompt_count = len(prompts)
if media.shape[0] != prompt_count:
raise ValueError(f"media batch size ({media.shape[0]}) must match prompt count ({prompt_count})")
total: torch.Tensor | None = None
details: dict[str, torch.Tensor] = {}
for name, weight in self.reward_weights.items():
scores = self.scorers[name](media, prompts).detach().float()
if scores.ndim != 1 or int(scores.shape[0]) != prompt_count:
raise ValueError(f"Reward {name!r} must return shape [{prompt_count}], got {tuple(scores.shape)}")
details[name] = scores
weighted = scores * float(weight)
total = weighted if total is None else total.to(weighted.device) + weighted
assert total is not None
details["avg"] = total
return details
+11
View File
@@ -91,6 +91,17 @@ class ModelBase(ABC):
def on_train_start(self) -> None: # noqa: B027
"""Called once before the training loop begins."""
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
"""Decode ``[B, T, C, H, W]`` latents to ``[B, C, T, H, W]`` media.
RL reward methods call this hook instead of reaching into
model-specific VAE normalization details.
"""
raise NotImplementedError(f"{type(self).__name__} does not implement decode_latents()")
# ------------------------------------------------------------------
# Timestep helpers
# ------------------------------------------------------------------
+47 -5
View File
@@ -101,6 +101,7 @@ class WanModel(ModelBase):
self.negative_prompt_embeds: (torch.Tensor | None) = None
self.negative_prompt_attention_mask: (torch.Tensor | None) = None
self._requires_negative_conditioning = True
# Timestep mechanics.
self.timestep_shift: float = float(flow_shift)
@@ -160,17 +161,31 @@ class WanModel(ModelBase):
self._init_timestep_mechanics()
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_t2v, )
pyarrow_schema_t2v,
pyarrow_schema_text_only,
)
from fastvideo.train.utils.dataloader import (
build_parquet_t2v_train_dataloader, )
preprocessed_data_type = str(getattr(
training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
parquet_schema = pyarrow_schema_t2v
if preprocessed_data_type == "text_only":
parquet_schema = pyarrow_schema_text_only
elif preprocessed_data_type != "t2v":
raise ValueError("Unsupported Wan preprocessed_data_type: "
f"{preprocessed_data_type!r}")
text_len = (
training_config.pipeline_config.text_encoder_configs[ # type: ignore[union-attr]
0].arch_config.text_len)
self.dataloader = build_parquet_t2v_train_dataloader(
training_config.data,
text_len=int(text_len),
parquet_schema=pyarrow_schema_t2v,
parquet_schema=parquet_schema,
)
self.start_step = 0
@@ -178,6 +193,9 @@ class WanModel(ModelBase):
def num_train_timesteps(self) -> int:
return int(self.num_train_timestep)
def set_requires_negative_conditioning(self, requires: bool) -> None:
self._requires_negative_conditioning = bool(requires)
def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
timestep = shift_timestep(
timestep,
@@ -187,7 +205,25 @@ class WanModel(ModelBase):
return timestep.clamp(self.min_timestep, self.max_timestep)
def on_train_start(self) -> None:
self.ensure_negative_conditioning()
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
@torch.no_grad()
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
if self.vae is None:
raise RuntimeError("Wan VAE is not initialized")
latents = latents_b_t_c_h_w.permute(0, 2, 1, 3, 4).float()
if bool(getattr(self.vae, "handles_latent_denorm", False)):
denorm = latents
else:
mean = torch.tensor(self.vae.latents_mean, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
std = torch.tensor(self.vae.latents_std, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
denorm = latents * std + mean
media = self.vae.to(latents.device).decode(denorm)
return (media / 2 + 0.5).clamp(0, 1)
# ------------------------------------------------------------------
# Runtime primitives
@@ -200,7 +236,8 @@ class WanModel(ModelBase):
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
self.ensure_negative_conditioning()
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
assert self.training_config is not None
tc = self.training_config
@@ -285,7 +322,7 @@ class WanModel(ModelBase):
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = noisy_latents.dtype
dtype = self._get_training_dtype()
if conditional:
text_dict = batch.conditional_dict
if text_dict is None:
@@ -301,6 +338,11 @@ class WanModel(ModelBase):
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
# Keep Wan training autocast tied to the model's training dtype, not
# to caller-created intermediates that may accidentally be fp32.
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
+3 -1
View File
@@ -128,7 +128,9 @@ class WanCausalModel(WanModel, CausalModelBase):
}
device_type = self.device.type
dtype = noisy_latents.dtype
dtype = self._get_training_dtype()
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
if conditional:
text_dict = batch.conditional_dict
+67 -25
View File
@@ -12,7 +12,7 @@ from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
if TYPE_CHECKING:
@@ -82,6 +82,22 @@ class Trainer:
batch = next(data_iter)
yield batch
def _run_method_validation(
self,
method: TrainingMethod,
iteration: int,
) -> None:
hook = getattr(method, "on_validation_begin", None)
if hook is None:
return
validation_metrics: dict[str, LogScalar] = hook(iteration)
validation_metrics = {
k: float(_coerce_log_scalar(v, where=(f"method.on_validation_begin().metrics[{k!r}]")))
for k, v in validation_metrics.items()
}
if self.global_rank == 0 and validation_metrics:
self.tracker.log(validation_metrics, iteration)
def run(
self,
method: TrainingMethod,
@@ -115,6 +131,7 @@ class Trainer:
method,
iteration=start_step,
)
self._run_method_validation(method, start_step)
method.optimizers_zero_grad(start_step)
data_stream = self._iter_dataloader(dataloader)
@@ -130,6 +147,8 @@ class Trainer:
desc="Steps",
disable=self.local_rank > 0,
)
# Allow method-specific optimization flow (e.g. DiffusionNFT).
method_manages_optimization = bool(method.manages_optimization())
for step in progress:
t0 = time.perf_counter()
@@ -137,47 +156,69 @@ class Trainer:
# to CPU once per step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
if method_manages_optimization:
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
))
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.single_train_step()"
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors
# to float right before logging.
metrics = {k: float(v) / grad_accum for k, v in loss_sums.items()}
metrics.update({k: float(v) / grad_accum for k, v in metric_sums.items()})
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
@@ -196,6 +237,7 @@ class Trainer:
method,
iteration=step,
)
self._run_method_validation(method, step)
self.callbacks.on_validation_end(
method,
iteration=step,
+7
View File
@@ -343,6 +343,12 @@ def _build_training_config(
if init_from is not None:
model_path = str(init_from)
preprocessed_data_type = str(da.get("preprocessed_data_type", "t2v") or "t2v").strip().lower()
if preprocessed_data_type not in {"t2v", "text_only"}:
raise ValueError("training.data.preprocessed_data_type must be one of "
"{'t2v', 'text_only'}, got "
f"{preprocessed_data_type!r}")
return TrainingConfig(
distributed=DistributedConfig(
num_gpus=num_gpus,
@@ -354,6 +360,7 @@ def _build_training_config(
),
data=DataConfig(
data_path=str(da.get("data_path", "") or ""),
preprocessed_data_type=preprocessed_data_type,
train_batch_size=int(da.get("train_batch_size", 1) or 1),
dataloader_num_workers=int(da.get("dataloader_num_workers", 0) or 0),
training_cfg_rate=float(da.get("training_cfg_rate", 0.0) or 0.0),
+1
View File
@@ -23,6 +23,7 @@ class DistributedConfig:
@dataclass(slots=True)
class DataConfig:
data_path: str = ""
preprocessed_data_type: str = "t2v"
train_batch_size: int = 1
dataloader_num_workers: int = 0
training_cfg_rate: float = 0.0
+1 -1
View File
@@ -112,7 +112,7 @@ class BaseTracker:
self._timed_metrics = {}
def log_artifacts(self, artifacts: dict[str, Any], step: int) -> None:
"""Log artifacts such as videos or images.
"""Log tracker artifacts such as sampled media.
By default this is treated the same as :meth:`log`.
"""
+2 -2
View File
@@ -1694,7 +1694,7 @@ class EMA_FSDP:
if p_local.numel() == 0:
# Nothing to swap on this rank for this param
continue
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
self.saved[name] = p_local.clone().to("cpu")
if name in self.ema.shadow:
ema_cpu = self.ema.shadow[name]
if ema_cpu.numel() != p_local.numel():
@@ -1714,7 +1714,7 @@ class EMA_FSDP:
saved_local = self.saved[name]
if saved_local.numel() != p_local.numel():
continue
p_local.copy_(saved_local)
p_local.copy_(saved_local.to(dtype=p_local.dtype, device=p_local.device))
self.saved.clear()
return False
@@ -0,0 +1,180 @@
#!/usr/bin/env bash
set -euo pipefail
# Build a text-only Parquet dataset from a one-prompt-per-line file. The script
# shards prompts across GPU_NUM single-GPU torchrun workers, runs
# v1_preprocess.py with --preprocess_task text_only, and writes prompt
# embeddings/captions under OUTPUT_DIR for DMD2/DiffusionNFT text-only runs.
INPUT_FILE="${1:-train.txt}"
OUTPUT_DIR="${OUTPUT_DIR:-data/train_text_only_dmd_preprocessed}"
MODEL_PATH="${MODEL_PATH:-Wan-AI/Wan2.1-T2V-1.3B-Diffusers}"
GPU_NUM="${GPU_NUM:-2}"
BATCH_SIZE="${BATCH_SIZE:-1}"
SAMPLES_PER_FILE="${SAMPLES_PER_FILE:-8}"
FLUSH_FREQUENCY="${FLUSH_FREQUENCY:-8}"
TEXT_MAX_LENGTH="${TEXT_MAX_LENGTH:-512}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
MIN_FREE_GPU_MB="${MIN_FREE_GPU_MB:-22000}"
if [[ ! -f "$INPUT_FILE" ]]; then
echo "Input text file not found: $INPUT_FILE" >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
visible_gpus=$(nvidia-smi --query-gpu=index --format=csv,noheader 2>/dev/null | wc -l | tr -d ' ')
if [[ "$visible_gpus" -lt "$GPU_NUM" ]]; then
echo "Expected at least $GPU_NUM GPUs, found $visible_gpus" >&2
exit 1
fi
for gpu_id in $(seq 0 $((GPU_NUM - 1))); do
free_mb=$(nvidia-smi --id="$gpu_id" --query-gpu=memory.free --format=csv,noheader,nounits | tr -d ' ')
if [[ "$free_mb" -lt "$MIN_FREE_GPU_MB" ]]; then
echo "GPU $gpu_id has only ${free_mb} MiB free; text-only Wan preprocessing needs" \
"about ${MIN_FREE_GPU_MB} MiB." >&2
echo "Free the GPU or lower MIN_FREE_GPU_MB if you know this run will fit." >&2
exit 1
fi
done
MARKER="$OUTPUT_DIR/.fastvideo_text_only_dmd_output"
if [[ -d "$OUTPUT_DIR" && ! -f "$MARKER" ]]; then
echo "Refusing to overwrite existing non-script output directory: $OUTPUT_DIR" >&2
echo "Set OUTPUT_DIR to a new path or remove the directory manually." >&2
exit 1
fi
rm -rf "$OUTPUT_DIR"
mkdir -p "$OUTPUT_DIR"
touch "$MARKER"
echo "Text-only DMD preprocessing config:"
echo " input: $INPUT_FILE"
echo " output: $OUTPUT_DIR"
echo " model: $MODEL_PATH"
echo " gpus: $GPU_NUM"
echo " batch size per GPU: $BATCH_SIZE"
echo " text max length: $TEXT_MAX_LENGTH"
echo " samples per parquet file: $SAMPLES_PER_FILE"
echo " flush frequency: $FLUSH_FREQUENCY"
SHARD_DIR="$OUTPUT_DIR/_text_shards"
mkdir -p "$SHARD_DIR"
python - "$INPUT_FILE" "$SHARD_DIR" "$GPU_NUM" <<'PY'
from pathlib import Path
import sys
input_path = Path(sys.argv[1])
shard_dir = Path(sys.argv[2])
num_shards = int(sys.argv[3])
prompts = [line.rstrip("\n") for line in input_path.read_text(encoding="utf-8").splitlines() if line.strip()]
if not prompts:
raise SystemExit(f"No non-empty prompts found in {input_path}")
for shard_idx in range(num_shards):
shard_prompts = prompts[shard_idx::num_shards]
shard_path = shard_dir / f"train_text_shard_{shard_idx}.txt"
shard_path.write_text("\n".join(shard_prompts) + "\n", encoding="utf-8")
print(f"Wrote {len(shard_prompts)} prompts to {shard_path}")
PY
run_preprocess_worker() {
local gpu_id="$1"
local shard_file="$2"
local shard_output="$3"
local log_file="$4"
local master_port="$5"
local -a cmd=(
torchrun
--nnodes=1
--nproc_per_node=1
--master_port "$master_port"
fastvideo/pipelines/preprocess/v1_preprocess.py
--model_path "$MODEL_PATH"
--data_merge_path "$shard_file"
--preprocess_video_batch_size "$BATCH_SIZE"
--seed 42
--max_height 448
--max_width 832
--num_frames 77
--dataloader_num_workers 0
--output_dir "$shard_output"
--train_fps 16
--samples_per_file "$SAMPLES_PER_FILE"
--flush_frequency "$FLUSH_FREQUENCY"
--text_max_length "$TEXT_MAX_LENGTH"
--video_length_tolerance_range 5
--preprocess_task text_only
)
{
echo "[gpu${gpu_id}] log file: $log_file"
echo "[gpu${gpu_id}] command: CUDA_VISIBLE_DEVICES=${gpu_id} ${cmd[*]}"
} | tee "$log_file"
CUDA_VISIBLE_DEVICES="$gpu_id" "${cmd[@]}" 2>&1 \
| sed -u "s/^/[gpu${gpu_id}] /" \
| tee -a "$log_file"
local status=${PIPESTATUS[0]}
if [[ "$status" -ne 0 ]]; then
echo "[gpu${gpu_id}] preprocessing failed with exit code $status" | tee -a "$log_file"
fi
return "$status"
}
pids=()
for gpu_id in $(seq 0 $((GPU_NUM - 1))); do
shard_file="$SHARD_DIR/train_text_shard_${gpu_id}.txt"
shard_output="$OUTPUT_DIR/shard_${gpu_id}"
mkdir -p "$shard_output"
log_file="$OUTPUT_DIR/preprocess_gpu_${gpu_id}.log"
echo "Launching text-only preprocessing on GPU ${gpu_id}: ${shard_file}"
run_preprocess_worker "$gpu_id" "$shard_file" "$shard_output" "$log_file" "$((29610 + gpu_id))" &
pids+=("$!")
done
failed=0
for pid in "${pids[@]}"; do
if ! wait "$pid"; then
failed=1
fi
done
if [[ "$failed" -ne 0 ]]; then
echo "One or more preprocessing workers failed. Check logs under $OUTPUT_DIR." >&2
exit 1
fi
num_parquet=$(find "$OUTPUT_DIR" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files were produced under $OUTPUT_DIR" >&2
exit 1
fi
echo "Text-only preprocessing complete."
echo "Parquet files: $num_parquet"
echo "Use this training data_path:"
echo "$OUTPUT_DIR"
+118
View File
@@ -0,0 +1,118 @@
#!/usr/bin/env bash
set -euo pipefail
CONFIG="${CONFIG:-examples/train/configs/rl/wan/diffusion_nft_pick_clip.yaml}"
DATA_PATH="${DATA_PATH:-data/pickscore_text_only_preprocessed}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/wan2.1_diffusion_nft_pick_clip}"
NUM_GPUS="${NUM_GPUS:-4}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29531}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$NUM_GPUS}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
NUM_FRAMES="${NUM_FRAMES:-1}"
NUM_LATENT_T="${NUM_LATENT_T:-1}"
PROJECT_NAME="${PROJECT_NAME:-diffusion_nft_wan}"
RUN_NAME="${RUN_NAME:-wan2.1_diffusion_nft_pick_clip}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
LOG_DIR="${LOG_DIR:-logs/train}"
if [[ ! -f "$CONFIG" ]]; then
echo "Training config not found: $CONFIG" >&2
exit 1
fi
if [[ ! -d "$DATA_PATH" ]]; then
echo "Preprocessed dataset directory not found: $DATA_PATH" >&2
echo "Run preprocessing first, for example:" >&2
echo " GPU_NUM=4 BATCH_SIZE=1 OUTPUT_DIR=$DATA_PATH \\" >&2
echo " bash scripts/preprocess/preprocess_train_text_only_dmd.sh DiffusionNFT/dataset/pickscore/train.txt" >&2
exit 1
fi
num_parquet=$(find "$DATA_PATH" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files found under $DATA_PATH" >&2
echo "Run preprocessing first, for example:" >&2
echo " GPU_NUM=4 BATCH_SIZE=1 OUTPUT_DIR=$DATA_PATH \\" >&2
echo " bash scripts/preprocess/preprocess_train_text_only_dmd.sh DiffusionNFT/dataset/pickscore/train.txt" >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-online}"
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_diffusion_nft_wan}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
mkdir -p "$LOG_DIR" "$OUTPUT_DIR"
timestamp="$(date +%Y%m%d_%H%M%S)"
log_file="$LOG_DIR/diffusion_nft_wan_pick_clip_${timestamp}.log"
cmd=(
torchrun
--nnodes "$NNODES"
--node_rank "$NODE_RANK"
--nproc_per_node "$NUM_GPUS"
--master_addr "$MASTER_ADDR"
--master_port "$MASTER_PORT"
-m fastvideo.train.entrypoint.train
--config "$CONFIG"
--training.data.data_path "$DATA_PATH"
--training.data.preprocessed_data_type text_only
--training.data.dataloader_num_workers "$DATALOADER_NUM_WORKERS"
--training.data.num_frames "$NUM_FRAMES"
--training.data.num_latent_t "$NUM_LATENT_T"
--training.distributed.num_gpus "$NUM_GPUS"
--training.distributed.sp_size "$SP_SIZE"
--training.distributed.tp_size "$TP_SIZE"
--training.distributed.hsdp_replicate_dim "$HSDP_REPLICATE_DIM"
--training.distributed.hsdp_shard_dim "$HSDP_SHARD_DIM"
--training.checkpoint.output_dir "$OUTPUT_DIR"
--training.tracker.project_name "$PROJECT_NAME"
--training.tracker.run_name "$RUN_NAME"
)
echo "DiffusionNFT Wan single-frame RL training config:"
echo " config: $CONFIG"
echo " data path: $DATA_PATH"
echo " parquet files: $num_parquet"
echo " output dir: $OUTPUT_DIR"
echo " frames / latent T: $NUM_FRAMES / $NUM_LATENT_T"
echo " rewards: pickscore + clipscore"
echo " learning rate: 3e-5"
echo " GPUs: $NUM_GPUS"
echo " SP/TP: $SP_SIZE/$TP_SIZE"
echo " HSDP replicate/shard: $HSDP_REPLICATE_DIM/$HSDP_SHARD_DIM"
echo " W&B mode: $WANDB_MODE"
echo " log file: $log_file"
echo "Command:"
printf ' %q' "${cmd[@]}" "$@"
echo
"${cmd[@]}" "$@" 2>&1 | tee "$log_file"
+170
View File
@@ -0,0 +1,170 @@
#!/usr/bin/env bash
set -euo pipefail
CONFIG="${CONFIG:-examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml}"
DATA_PATH="${DATA_PATH:-data/train_text_only_dmd_preprocessed}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/wan2.1_dmd2_text_only}"
NUM_GPUS="${NUM_GPUS:-2}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29521}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$NUM_GPUS}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
NUM_FRAMES="${NUM_FRAMES:-1}"
NUM_LATENT_T="${NUM_LATENT_T:-1}"
VALIDATION_NUM_FRAMES="${VALIDATION_NUM_FRAMES:-$NUM_FRAMES}"
VALIDATION_PROMPT_FILE="${VALIDATION_PROMPT_FILE:-}"
VALIDATION_FILE="${VALIDATION_FILE:-examples/train/configs/distribution_matching/wan/dmd2_text_only_validation.json}"
VALIDATION_OFFLOAD_TRAINING_STATE="${VALIDATION_OFFLOAD_TRAINING_STATE:-true}"
VALIDATION_UNLOAD_PIPELINE_AFTER="${VALIDATION_UNLOAD_PIPELINE_AFTER:-true}"
CFG_UNCOND_TEXT="${CFG_UNCOND_TEXT:-zero}"
CFG_UNCOND_ON_MISSING="${CFG_UNCOND_ON_MISSING:-ignore}"
PROJECT_NAME="${PROJECT_NAME:-distillation_wan_text_only}"
RUN_NAME="${RUN_NAME:-wan2.1_dmd2_text_only}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
LOG_DIR="${LOG_DIR:-logs/train}"
if [[ ! -f "$CONFIG" ]]; then
echo "Training config not found: $CONFIG" >&2
exit 1
fi
if [[ ! -d "$DATA_PATH" ]]; then
echo "Preprocessed dataset directory not found: $DATA_PATH" >&2
echo "Run scripts/preprocess/preprocess_train_text_only_dmd.sh first." >&2
exit 1
fi
num_parquet=$(find "$DATA_PATH" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files found under $DATA_PATH" >&2
echo "Run scripts/preprocess/preprocess_train_text_only_dmd.sh first." >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-offline}"
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_dmd2_text_only}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
mkdir -p "$LOG_DIR" "$OUTPUT_DIR"
if [[ -n "$VALIDATION_PROMPT_FILE" ]]; then
if [[ ! -f "$VALIDATION_PROMPT_FILE" ]]; then
echo "Validation prompt file not found: $VALIDATION_PROMPT_FILE" >&2
echo "Set VALIDATION_PROMPT_FILE to a text file, or leave it empty and set VALIDATION_FILE=<validation.json>." >&2
exit 1
fi
python - "$VALIDATION_PROMPT_FILE" "$VALIDATION_FILE" <<'PY'
import json
import os
import sys
prompt_file, validation_file = sys.argv[1:3]
with open(prompt_file, encoding="utf-8") as f:
prompts = [line.strip() for line in f if line.strip()]
if not prompts:
raise SystemExit(f"No validation prompts found in {prompt_file}")
validation_dir = os.path.dirname(os.path.abspath(validation_file))
if validation_dir:
os.makedirs(validation_dir, exist_ok=True)
with open(validation_file, "w", encoding="utf-8") as f:
json.dump(
{"data": [{"caption": prompt} for prompt in prompts]},
f,
indent=2,
ensure_ascii=False,
)
f.write("\n")
print(f"Wrote {len(prompts)} validation prompts to {validation_file}")
PY
elif [[ ! -f "$VALIDATION_FILE" ]]; then
echo "Validation dataset file not found: $VALIDATION_FILE" >&2
exit 1
fi
timestamp="$(date +%Y%m%d_%H%M%S)"
log_file="$LOG_DIR/dmd2_t2v_text_only_${timestamp}.log"
cmd=(
torchrun
--nnodes "$NNODES"
--node_rank "$NODE_RANK"
--nproc_per_node "$NUM_GPUS"
--master_addr "$MASTER_ADDR"
--master_port "$MASTER_PORT"
-m fastvideo.train.entrypoint.train
--config "$CONFIG"
--training.data.data_path "$DATA_PATH"
--training.data.preprocessed_data_type text_only
--training.data.dataloader_num_workers "$DATALOADER_NUM_WORKERS"
--training.data.num_frames "$NUM_FRAMES"
--training.data.num_latent_t "$NUM_LATENT_T"
--callbacks.validation.dataset_file "$VALIDATION_FILE"
--callbacks.validation.num_frames "$VALIDATION_NUM_FRAMES"
--callbacks.validation.offload_training_state "$VALIDATION_OFFLOAD_TRAINING_STATE"
--callbacks.validation.unload_pipeline_after_validation "$VALIDATION_UNLOAD_PIPELINE_AFTER"
--method.cfg_uncond.text "$CFG_UNCOND_TEXT"
--method.cfg_uncond.on_missing "$CFG_UNCOND_ON_MISSING"
--training.distributed.num_gpus "$NUM_GPUS"
--training.distributed.sp_size "$SP_SIZE"
--training.distributed.tp_size "$TP_SIZE"
--training.distributed.hsdp_replicate_dim "$HSDP_REPLICATE_DIM"
--training.distributed.hsdp_shard_dim "$HSDP_SHARD_DIM"
--training.checkpoint.output_dir "$OUTPUT_DIR"
--training.tracker.project_name "$PROJECT_NAME"
--training.tracker.run_name "$RUN_NAME"
)
echo "DMD2 T2V text-only training config:"
echo " config: $CONFIG"
echo " data path: $DATA_PATH"
echo " parquet files: $num_parquet"
echo " output dir: $OUTPUT_DIR"
echo " frames / latent T: $NUM_FRAMES / $NUM_LATENT_T"
echo " validation prompt file: ${VALIDATION_PROMPT_FILE:-<none>}"
echo " validation dataset: $VALIDATION_FILE"
echo " validation frames: $VALIDATION_NUM_FRAMES"
echo " validation offload training state: $VALIDATION_OFFLOAD_TRAINING_STATE"
echo " validation unload pipeline after: $VALIDATION_UNLOAD_PIPELINE_AFTER"
echo " GPUs: $NUM_GPUS"
echo " SP/TP: $SP_SIZE/$TP_SIZE"
echo " HSDP replicate/shard: $HSDP_REPLICATE_DIM/$HSDP_SHARD_DIM"
echo " CFG uncond text/on_missing: $CFG_UNCOND_TEXT/$CFG_UNCOND_ON_MISSING"
echo " W&B mode: $WANDB_MODE"
echo " log file: $log_file"
echo "Command:"
printf ' %q' "${cmd[@]}" "$@"
echo
"${cmd[@]}" "$@" 2>&1 | tee "$log_file"
@@ -0,0 +1,99 @@
from types import SimpleNamespace
import torch
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
class _FakeEMA:
def __init__(self):
self.updates = 0
def update(self, module):
del module
self.updates += 1
def test_reward_diagnostic_metrics_match_per_prompt_groups():
method = object.__new__(DiffusionNFTMethod)
method._trained_prompt_hashes = set()
sample_items = [{
"prompts": ["a", "a"],
}, {
"prompts": ["b", "b"],
}]
rewards = {"avg": torch.tensor([1.0, 3.0, 2.0, 6.0])}
metrics = method._reward_diagnostic_metrics(sample_items, rewards)
assert metrics["group_size"] == 2.0
assert metrics["trained_prompt_num"] == 2.0
assert torch.isclose(metrics["zero_std_ratio"], torch.tensor(0.0))
assert torch.isclose(metrics["reward_std_mean"], torch.tensor(1.5))
assert torch.isclose(metrics["mean_reward_100"], torch.tensor(3.0))
assert torch.isclose(metrics["mean_reward_50"], torch.tensor(4.5))
method._reward_diagnostic_metrics(sample_items, rewards)
assert len(method._trained_prompt_hashes) == 2
def test_update_ema_honors_update_after_step():
method = object.__new__(DiffusionNFTMethod)
method._ema_enabled = True
method._student_ema = _FakeEMA()
method._ema_update_count = 0
method._ema_update_after_step = 1
method.student = SimpleNamespace(transformer=object())
method._update_ema()
assert method._student_ema.updates == 0
assert method._ema_update_count == 1
method._update_ema()
assert method._student_ema.updates == 1
assert method._ema_update_count == 2
def test_num_train_timesteps_uses_explicit_schedule_length():
method = object.__new__(DiffusionNFTMethod)
method._sample_steps = 25
method._timestep_fraction = 0.5
method._sampling_config = SimpleNamespace(
timesteps=[900, 800, 700, 600, 500, 400, 300, 200, 100, 10],
sigmas=None,
)
assert method._num_train_timesteps() == 5
def test_checkpoint_state_saves_frozen_old_policy_weights():
student = torch.nn.Linear(2, 2)
old = torch.nn.Linear(2, 2)
for param in old.parameters():
param.requires_grad_(False)
method = object.__new__(DiffusionNFTMethod)
method._role_models = {
"student": SimpleNamespace(
transformer=student,
_trainable=True,
),
"old": SimpleNamespace(
transformer=old,
_trainable=False,
),
}
method.student = method._role_models["student"]
method.old = method._role_models["old"]
method._student_optimizer = None
method._student_lr_scheduler = None
method._ema_enabled = False
states = method.checkpoint_state()
old_state = states["roles.old.transformer"].state_dict()
assert "weight" in old_state
assert "bias" in old_state
assert torch.equal(old_state["weight"], old.weight)
+57
View File
@@ -0,0 +1,57 @@
import torch
import pytest
from fastvideo.train.methods.rl.rewards import MultiRewardScorer, select_first_frame
def test_select_first_frame_for_video_tensor():
video = torch.arange(2 * 3 * 4 * 5 * 6).reshape(2, 3, 4, 5, 6)
frame = select_first_frame(video)
assert frame.shape == (2, 3, 5, 6)
torch.testing.assert_close(frame, video[:, :, 0])
def test_select_first_frame_keeps_frame_tensor():
frame = torch.randn(2, 3, 5, 6)
selected = select_first_frame(frame)
assert selected is frame
def test_multi_reward_weighted_sum_with_injected_scorers():
def pickscore(media, prompts):
assert media.shape == (2, 3, 4, 5, 6)
assert prompts == ["a", "b"]
return torch.tensor([1.0, 2.0])
def clipscore(media, prompts):
assert media.shape == (2, 3, 4, 5, 6)
assert prompts == ["a", "b"]
return torch.tensor([0.5, 1.5])
scorer = MultiRewardScorer(
{"pickscore": 2.0, "clipscore": 3.0},
scorers={
"pickscore": pickscore,
"clipscore": clipscore,
},
)
scores = scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
torch.testing.assert_close(scores["pickscore"], torch.tensor([1.0, 2.0]))
torch.testing.assert_close(scores["clipscore"], torch.tensor([0.5, 1.5]))
torch.testing.assert_close(scores["avg"], torch.tensor([3.5, 8.5]))
def test_multi_reward_validates_score_shape():
scorer = MultiRewardScorer(
{"pickscore": 1.0},
scorers={"pickscore": lambda media, prompts: torch.tensor([[1.0], [2.0]])},
)
with pytest.raises(ValueError, match="must return shape"):
scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
+238
View File
@@ -0,0 +1,238 @@
import torch
import pytest
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.methods.rl.common import (
DiffusionSampler,
SamplingConfig,
distributed_k_repeat_indices,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
from fastvideo.train.utils.config import load_run_config
class _FakeScheduler:
def __init__(self):
self.num_train_timesteps = 1000
self.set_timesteps_calls = []
self.timesteps = torch.empty(0)
self.sigmas = torch.empty(0)
self.step_calls = 0
def set_timesteps(self, num_inference_steps=None, device=None, timesteps=None, sigmas=None):
self.set_timesteps_calls.append({
"num_inference_steps": num_inference_steps,
"timesteps": timesteps,
"sigmas": sigmas,
})
if timesteps is not None:
self.timesteps = torch.tensor(timesteps, device=device, dtype=torch.float32)
else:
self.timesteps = torch.linspace(1000, 0, int(num_inference_steps), device=device)
if sigmas is not None:
self.sigmas = torch.tensor(sigmas, device=device, dtype=torch.float32)
else:
self.sigmas = torch.cat([self.timesteps / 1000.0, torch.zeros(1, device=device)])
def step(self, model_output, timestep, sample, return_dict=False):
del timestep
self.step_calls += 1
prev = sample + model_output
return (prev, ) if not return_dict else {"prev_sample": prev}
class _FakeModel:
def __init__(self):
self.noise_scheduler = _FakeScheduler()
self.add_noise_calls = 0
self.timestep_shapes = []
def predict_noise(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return torch.zeros_like(noisy_latents)
def predict_x0(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return noisy_latents
def add_noise(self, clean_latents, noise, timestep):
del timestep
self.add_noise_calls += 1
return clean_latents + noise
def _batch():
batch = TrainingBatch()
batch.latents = torch.zeros(2, 1, 3, 4, 4)
return batch
def test_sampler_preserves_latent_dtype():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
batch = _batch()
batch.latents = batch.latents.to(torch.bfloat16)
result = sampler.sample(model, batch, generator=torch.Generator().manual_seed(0))
assert result.latents.dtype is torch.bfloat16
def test_sampler_uses_scheduler_generated_timesteps_by_default():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [1000.0, 666.6666259765625, 333.3333435058594, 0.0]
def test_sampler_honors_explicit_timestep_override():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=3, timesteps=[900, 300, 10]))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [900.0, 300.0, 10.0]
def test_sampler_honors_explicit_timesteps_without_matching_num_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(timesteps=[900, 300, 10]))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [900.0, 300.0, 10.0]
assert model.noise_scheduler.set_timesteps_calls == []
def test_sampling_config_rejects_unknown_keys():
with pytest.raises(ValueError, match="Unsupported method.sampling key"):
SamplingConfig.from_mapping({"solver": "dpm2"})
def test_sampler_restores_original_batch_timestep_after_sampling():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=2))
batch = _batch()
original_timesteps = torch.tensor([123.0])
batch.timesteps = original_timesteps
sampler.sample(batch=batch, model=model, generator=torch.Generator().manual_seed(0))
assert batch.timesteps is original_timesteps
def test_euler_sampler_does_not_renoise_between_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert model.add_noise_calls == 0
assert model.timestep_shapes == [(2,), (2,), (2,), (2,)]
def test_sde_reflow_sampler_renoises_between_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4, trajectory="sde_reflow"))
sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert model.add_noise_calls == 3
def test_diffusion_nft_config_uses_rl_sampler_not_dmd_pipeline():
config_path = "examples/train/configs/rl/wan/diffusion_nft_pick_clip.yaml"
cfg = load_run_config(config_path)
raw_text = open(config_path, encoding="utf-8").read()
assert cfg.method["_target_"] == "fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod"
assert cfg.training.optimizer.learning_rate == 3.0e-5
assert cfg.training.data.num_latent_t == 1
assert cfg.training.data.num_frames == 1
assert "sampling_timesteps" not in raw_text
assert "WanDMDPipeline" not in raw_text
assert "solver" not in cfg.method["sampling"]
assert cfg.method["sampling"]["scheduler"] == "flow_match_euler"
assert cfg.method["sampling"]["trajectory"] == "ode"
assert cfg.method["sampling"]["flow_shift"] == "inherit"
assert "deterministic" not in cfg.method["sampling"]
assert "noise_level" not in cfg.method["sampling"]
assert cfg.method["validation"]["every_steps"] == 10
assert cfg.method["validation"]["num_steps"] == 40
assert cfg.method["validation"]["num_prompts"] == 16
assert cfg.method["validation"]["log_samples"] is True
def test_validation_shard_indices_are_stable_and_padded():
rank0 = validation_shard_indices(5, rank=0, world_size=2)
rank1 = validation_shard_indices(5, rank=1, world_size=2)
assert rank0 == [(0, True), (2, True), (4, True)]
assert rank1 == [(1, True), (3, True), (0, False)]
def test_distributed_k_repeat_indices_repeats_prompts_globally():
rank0 = distributed_k_repeat_indices(
dataset_length=100,
batch_size=6,
repeats_per_prompt=24,
world_size=4,
rank=0,
seed=123,
)
all_indices = []
for rank in range(4):
sample = distributed_k_repeat_indices(
dataset_length=100,
batch_size=6,
repeats_per_prompt=24,
world_size=4,
rank=rank,
seed=123,
)
all_indices.extend(sample.local_indices)
assert rank0.unique_prompt_count == 1
assert len(all_indices) == 24
assert len(set(all_indices)) == 1
def test_validation_caption_puts_rewards_first():
caption = validation_caption(
"a small blue cube",
{
"avg": 0.75,
"pickscore": 0.5,
},
)
assert caption.startswith("avg: 0.7500 | pickscore: 0.5000 | ")
assert caption.endswith("a small blue cube")
def test_media_to_video_array_treats_frame_as_single_frame_video():
frame = torch.ones(3, 4, 5)
video = media_to_video_array(frame)
assert video.shape == (1, 3, 4, 5)
assert video.dtype.name == "uint8"
def test_media_to_video_array_preserves_video_frames():
media = torch.ones(3, 2, 4, 5)
video = media_to_video_array(media)
assert video.shape == (2, 3, 4, 5)
@@ -0,0 +1,157 @@
from dataclasses import dataclass, field
import torch
from fastvideo.train.trainer import Trainer
from fastvideo.train.utils.training_config import (
CheckpointConfig,
DistributedConfig,
ModelTrainingConfig,
OptimizerConfig,
TrackerConfig,
TrainingConfig,
TrainingLoopConfig,
)
class _Tracker:
def __init__(self):
self.logged = []
self.finished = False
def log(self, metrics, step):
self.logged.append((step, metrics))
def finish(self):
self.finished = True
class _Callbacks:
def __init__(self):
self.before_optimizer_steps = 0
self.training_step_ends = 0
def on_train_start(self, method, iteration=0):
pass
def on_before_optimizer_step(self, method, iteration=0):
self.before_optimizer_steps += 1
def on_training_step_end(self, method, metrics, iteration=0):
self.training_step_ends += 1
def on_validation_begin(self, method, iteration=0):
pass
def on_validation_end(self, method, iteration=0):
pass
def on_train_end(self, method, iteration=0):
pass
class _World:
rank = 0
local_rank = 0
class _Method(torch.nn.Module):
def __init__(self):
super().__init__()
self.calls = 0
self.backward_calls = 0
self.optimizer_steps = 0
self.tracker = None
def set_tracker(self, tracker):
self.tracker = tracker
def on_train_start(self):
pass
def manages_optimization(self):
return True
def managed_train_step(self, data_stream, iteration):
batch = next(data_stream)
self.calls += 1
return (
{"total_loss": torch.tensor(float(batch["x"]))},
{},
{"managed_metric": float(iteration)},
)
def backward(self, *args, **kwargs):
self.backward_calls += 1
def optimizers_schedulers_step(self, iteration):
self.optimizer_steps += 1
def optimizers_zero_grad(self, iteration):
pass
class _MethodWithValidation(_Method):
def __init__(self):
super().__init__()
self.validation_iterations = []
def on_validation_begin(self, iteration=0):
self.validation_iterations.append(iteration)
return {"validation/fake": float(iteration)}
def test_trainer_skips_default_optimizer_path_for_managed_methods(monkeypatch):
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.build_tracker", lambda *args, **kwargs: _Tracker())
cfg = TrainingConfig(
distributed=DistributedConfig(),
optimizer=OptimizerConfig(),
loop=TrainingLoopConfig(max_train_steps=1, gradient_accumulation_steps=3),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
)
trainer = Trainer(cfg)
trainer.callbacks = _Callbacks()
method = _Method()
dataloader = [{"x": 2}]
trainer.run(method, dataloader=dataloader, max_steps=1)
assert method.calls == 1
assert method.backward_calls == 0
assert method.optimizer_steps == 0
assert trainer.callbacks.before_optimizer_steps == 0
assert trainer.callbacks.training_step_ends == 1
assert trainer.tracker.logged[0][1]["total_loss"] == 2.0
def test_trainer_logs_method_validation_at_step_zero(monkeypatch):
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.build_tracker", lambda *args, **kwargs: _Tracker())
cfg = TrainingConfig(
distributed=DistributedConfig(),
optimizer=OptimizerConfig(),
loop=TrainingLoopConfig(max_train_steps=1, gradient_accumulation_steps=1),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
)
trainer = Trainer(cfg)
trainer.callbacks = _Callbacks()
method = _MethodWithValidation()
dataloader = [{"x": 2}]
trainer.run(method, dataloader=dataloader, max_steps=1)
assert method.validation_iterations == [0, 1]
assert trainer.tracker.logged[0] == (0, {"validation/fake": 0.0})
@@ -0,0 +1,62 @@
import torch
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.train.models.wan.wan import WanModel
class _CPUWanModel(WanModel):
@property
def device(self) -> torch.device:
return torch.device("cpu")
class _AutocastProbe(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.autocast_enabled: bool | None = None
self.autocast_dtype: torch.dtype | None = None
def forward(
self,
*,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor,
timestep: torch.Tensor,
return_dict: bool,
) -> torch.Tensor:
del encoder_hidden_states, encoder_attention_mask, timestep, return_dict
self.autocast_enabled = torch.is_autocast_enabled("cpu")
self.autocast_dtype = torch.get_autocast_dtype("cpu")
self.hidden_states_dtype = hidden_states.dtype
return hidden_states
def test_wan_predict_noise_uses_training_dtype_autocast_for_fp32_inputs():
model = object.__new__(_CPUWanModel)
model.transformer = _AutocastProbe()
batch = TrainingBatch()
batch.timesteps = torch.tensor([1], dtype=torch.long)
batch.conditional_dict = {
"encoder_hidden_states": torch.randn(1, 4, 8, dtype=torch.float32),
"encoder_attention_mask": torch.ones(1, 4, dtype=torch.float32),
}
noisy_latents = torch.randn(1, 1, 2, 4, 4, dtype=torch.float32)
timestep = torch.tensor([1], dtype=torch.long)
pred_noise = model.predict_noise(
noisy_latents,
timestep,
batch,
conditional=True,
)
assert pred_noise.shape == noisy_latents.shape
assert pred_noise.dtype is torch.bfloat16
assert model.transformer.hidden_states_dtype is torch.bfloat16
assert model.transformer.autocast_enabled is True
assert model.transformer.autocast_dtype is torch.bfloat16