Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ee607bd0d0 | ||
|
|
633d393568 | ||
|
|
5854aec2ce |
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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`.
|
||||
"""
|
||||
|
||||
@@ -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
@@ -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"
|
||||
Executable
+170
@@ -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)
|
||||
@@ -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"])
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user