Compare commits

...
4 Commits
Author SHA1 Message Date
SolitaryThinker 5a01741ea0 update 2026-06-15 16:09:35 -07:00
SolitaryThinker de4d430d87 design 2026-06-13 15:39:21 -07:00
alexzmsandmergify[bot] 633d393568 [ci] layer-0 grad-norm regression for per-method training tests (5a-ii) (#1396)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-12 04:45:07 +00:00
Junda SuandPeiyuan Zhang 5854aec2ce [feat] Add Wan RL DiffusionNFT training (#1450)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2026-06-11 21:18:59 -07:00
47 changed files with 9292 additions and 42 deletions
+32
View File
@@ -0,0 +1,32 @@
---
name: add-reward-model
description: Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.
---
# Add Reward Model
Use for reward models consumed by RL methods.
## Placement
- Put reusable reward code under `fastvideo/train/methods/rl/rewards/`.
- Expose public builders from `fastvideo/train/methods/rl/rewards/__init__.py`.
- Keep method-specific aggregation or advantage logic out of reward classes.
## Media Inputs
- Reward callables receive decoded media tensors.
- Accept single-frame tensors as `[B, C, H, W]` and multi-frame tensors as `[B, C, T, H, W]` when practical.
- Frame selection is reward-specific. Frame scorers such as PickScore and CLIPScore should explicitly select frame `0`; temporal rewards should inspect whichever frames they need.
- Return one scalar reward per prompt/sample.
## Attribution
- If code is ported or closely adapted from another repo, add a short comment or docstring naming the source file/function.
- Preserve SPDX headers used by FastVideo files.
## Tests
- Unit-test tensor layout handling without loading large reward checkpoints.
- Allow fake scorer injection for multi-reward tests.
- Test weighted reward aggregation and metric keys.
+38
View File
@@ -0,0 +1,38 @@
---
name: add-rl-method
description: Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.
---
# Add RL Method
Use for new RL methods in the modular `fastvideo/train` stack.
## Required Shape
- Add the method under `fastvideo/train/methods/rl/`.
- Subclass `TrainingMethod`.
- Keep model-family logic in `ModelBase` wrappers.
- Decode generated latents through `ModelBase.decode_latents`; add that hook to the new model wrapper instead of decoding inside the RL method.
- Use `fastvideo/train/methods/rl/common/sampling.py` for generation unless the method has a documented reason to avoid sampling.
- Use `fastvideo/train/methods/rl/common/prompt_sampling.py` for reusable grouped prompt sampling patterns such as DiffusionNFT K-repeat.
- Use `fastvideo/train/methods/rl/rewards/` for reward models.
## Optimization
- Return `manages_optimization() == True` only when the method must own a nonstandard outer/inner loop.
- If using managed optimization, implement `managed_train_step(data_stream, iteration)`.
- Existing trainer callbacks, checkpointing, tracking, and validation should still work.
## Config
- Put method knobs under `method`.
- Put sampler knobs under `method.sampling`.
- Do not put scheduler or trajectory policy into model configs.
- Do not split a diffusers-style scheduler from its built-in `step()` solver in YAML; use `trajectory` only for higher-level ODE vs re-noise behavior.
- Avoid fixed timestep lists in examples unless reproducing a known baseline; prefer scheduler-generated defaults.
## Tests
- Add fake-model tests for sampler/method behavior.
- Add config parse tests for the public YAML.
- Confirm existing train methods stay on the default Trainer path.
+3
View File
@@ -8,3 +8,6 @@
{"name": "decompose-pipeline-pr", "description": "Decompose an oversized FastVideo pipeline PR into a stack of independently-reviewable PRs. Tiers the diff by blast radius (invisible / dead code / cross-cutting infra / activation), produces a branch graph and worktree bootstrap, drafts the AGENTS.md manifest, flags missing tests on cross-cutting infra changes, and extracts lessons from the PR body. Worked example: PR #1280 daVinci-MagiHuman (9.8k LOC) decomposed into 10 stacked PRs.", "path": "decompose-pipeline-pr/SKILL.md", "status": "tested", "trust": "medium"}
{"name": "reseed-performance-baseline", "description": "Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline must be advanced by replicating one reviewed shifted source result into three success=true records, or five records when explicitly requested", "path": "reseed-performance-baseline/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-model", "description": "Add a new model (or variant) to FastVideo: DiT + configs + pipeline + presets + registry + tests. Walks through FastVideo's single stage-based pipeline architecture with exact file paths and registration hooks.", "path": "add-model/SKILL.md", "status": "draft", "trust": "low"}
{"name": "rlhf-training-abstractions", "description": "Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.", "path": "rlhf-training-abstractions/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-rl-method", "description": "Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.", "path": "add-rl-method/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-reward-model", "description": "Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.", "path": "add-reward-model/SKILL.md", "status": "draft", "trust": "low"}
@@ -0,0 +1,41 @@
---
name: rlhf-training-abstractions
description: Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.
---
# RLHF Training Abstractions
Use this skill before editing RLHF-style training code in `fastvideo/train`.
## Boundaries
- RL methods live under `fastvideo/train/methods/rl/` and own algorithm logic: reward collection, advantage computation, policy loss, KL/reference terms, and optimizer cadence.
- Rewards live under `fastvideo/train/methods/rl/rewards/` and must be reusable across RL methods.
- RL methods pass decoded media to rewards; each reward decides whether to use the first frame, sampled frames, or the full video.
- Sampling lives under `fastvideo/train/methods/rl/common/` and must use `ModelBase` primitives plus scheduler math, not model-family inference pipelines.
- Model wrappers under `fastvideo/train/models/` own model-specific forward details.
- Model wrappers also own model-specific latent decoding via `ModelBase.decode_latents`; RL methods should not reach into VAE normalization internals.
- Shared RL helpers such as K-repeat prompt sampling belong under `fastvideo/train/methods/rl/common/` when they are reusable across RL methods.
## Anti-Patterns
- Do not bind RL methods to inference pipeline classes such as `WanDMDPipeline`.
- Do not hardcode timestep lists in a method when the scheduler can generate them.
- Do not put reward-model code inside one RL method.
- Do not make existing non-RL methods use method-managed optimization unless explicitly requested.
## Sampling Policy
- Prefer YAML-configured `method.sampling` with `scheduler`, `trajectory`, `num_steps`, `timesteps`, and `sigmas`.
- Treat diffusers-style scheduler classes as owning both the timestep schedule and their `step()` update rule; avoid a separate `solver` field unless a new sampler truly implements solver math outside the scheduler object.
- Missing `timesteps` means “ask the scheduler”; explicit `timesteps` or `sigmas` are overrides.
- ODE-style trajectories should not re-noise between denoising steps.
- SDE/re-noise behavior must be explicit in config.
## Validation
- Run focused local tests for sampler config and Trainer opt-in behavior.
- Verify existing train methods still report `manages_optimization() == False`.
- Keep fixed-prompt validation helpers in `fastvideo/train/methods/rl/common/validation.py` so new RL methods can reuse sharding and captions.
- Test distributed prompt grouping helpers separately from heavyweight model loading.
- Run `pre-commit run --files <changed paths>`; respect configured excludes.
+2294
View File
File diff suppressed because it is too large Load Diff
+71
View File
@@ -0,0 +1,71 @@
# FastVideo Next-Gen Runtime — Two-Page Summary
**Companion to** `design.md` (v19, 2026-06-12) · **Status:** draft for discussion · **Ask:** read this, then dive into the sections you own.
---
## The problem
FastVideo's pipeline abstraction has been outgrown by its own model zoo. Four facts, all on `main` today:
- **The denoise/sampling loop exists in four copies** — inference stages (`pipelines/stages/denoising.py`, a 1,381-line file), `train/` distillation methods, legacy `training/` monoliths, and the just-landed RL work. The fourth copy documents the cause in its own docstring: `DiffusionSampler` (PR #1450) *"intentionally does not call FastVideo's full inference pipelines"* because the only consumable units are family-bound pipeline classes. Every new post-training method must pick between a wrong dependency and a private loop.
- **There is no serving runtime.** One request at a time, no queue, no cross-request batching. Dreamverse — our shipping product — hand-rolls its own GPU pool, queue, warmup, and streaming relay at a cost of **one B200 per user session**.
- **Cosmos3 outgrows both the stage abstraction and the alternatives.** Its AR text reasoner and multimodal diffusion denoiser are *the same resident weights* driven by two loop types within one request — our full port (incl. action modality, on `feat/cosmos3-reasoning`) runs it only by bypassing stages with one monolithic block. Multi-engine DAG stacks (vllm-omni, sglang-omni) compose separable stages with disjoint weights; none can express a mixture-of-transformers model. This is the forcing function.
- **RL has arrived and pays the tax already** — likelihood-free DiffusionNFT for Wan (#1450): in-process rollouts with zero serving-grade optimizations (no CFG, dense attention, full 25-step ODE), a vendored sampler, and a parallel validation path built *because* inference pipelines aren't consumable as a library.
## The design
```
Request plane OpenAI-compatible server (videos/images/audio/chat) · AsyncEngine
OmniRequest ─► admission ─► queue ─► OmniOutput (typed modality parts)
Pipeline plane PipelineSpec: declarative graph per family
nodes: Stage | LoopStage (DenoiseLoop, ARDecodeLoop) · typed Artifacts
policies: CFG, ExpertRouting, AttnMetadata, Precision, FlowShift
Execution plane StepScheduler (multiplexes denoise + AR steps) · worker pools (TP/SP + CFG-parallel)
CacheManager (paged text-KV ▪ chunked causal-video KV ▪ feature caches) · connectors
```
Default deployment is exactly today's: one SPMD pool, co-located nodes, synchronous call. Serving is additive configuration, not a different code path. Five load-bearing decisions:
1. **Loop inversion.** Loops become `LoopStage` nodes exposing `init / step / finalize`; the runtime owns iteration, families own step bodies (with a custom-step escape hatch — the runtime never dictates step factoring). This is the enabler for step-level scheduling, streaming, MoT interleaving, and one shared loop across inference, distillation, and RL. Stated honestly: **no surveyed system does this at scheduler granularity** — the risk is retired by Phase-1 bit-identical parity gates and a measured falsifier, not borrowed validation.
2. **Cost-model scheduling.** Denoise steps and AR tokens are incommensurable (bidirectional attention is O(L²) per step with zero KV amortization; steps differ ~1000×) — the budget currency is **predicted GPU-time** from a per-(model, phase, shape) cost model, calibrated by the profiler and published to Dynamo's router/Planner as the same artifact.
3. **One substrate for inference, training, and RL** — models, loaders, configs, schedulers, parallel state, loop step bodies — under a strict `engine never imports train` rule. Trainers keep their internals; their embedded sampling paths migrate onto the shared loops.
4. **Consistency is a declared, measured contract**, not a hope: **C0** corrected (profiles differ, TIS/MIS fixes it) / **C1** kernel-pinned (RL default; drift gated in CI) / **C2** bitwise (batch-invariant kernels + Behavior Record, for goldens and MoE parity). One repo ≠ automatic parity — the ladder is what makes the single-runtime bet honest.
5. **Extensions, never monkeypatching**: read-only observers (ParityAligner, ActivationTrace, Profiler, NaNWatch) and compute-altering interceptors at declared points — **cache-dit** is the first interceptor. **Dynamo is the fleet layer** (first-class partner, not a dependency we rebuild): registration/health/cost contract in-engine, seven concrete upstream asks (affinity key spaces, cost interface, media streaming, role-graph disagg, RL weight plane, KVBM generalization, sessions) — each with a fallback.
## Why this is the moat
Unlike LLMs — where inference optimization is post-hoc on frozen weights — **a usable video model is itself a post-training artifact**. Every inference capability we ship is a *(recipe, runtime)* pair: step distillation ↔ few-step samplers; self-forcing ↔ causal KV streaming; QAT-NVFP4 ↔ FP4 kernels; VSA ↔ sparse attention backend; RL ↔ samplers + capture. The training loop *embeds* the inference loop, so whoever owns both sides of the pair owns the optimization frontier. The industry's RL pain proves the converse: verl-omni re-implements Wan2.2 inside vLLM-Omni and corrects the numerics afterward; miles' headline features are all mismatch patches for two runtimes with different kernels. We answer with one model definition, one kernel set, one measured ladder — at FastVideo's 1–30B FSDP2 scale, where the bet is viable.
## What's pulling on it
| Customer | Pull | Proof point |
|---|---|---|
| **Cosmos3 / omni** | MoT loops, packed sequences, reasoner KV, world-model rollout | 150-test parity suite on `feat/cosmos3-reasoning` |
| **Dreamverse** | Engine-client replaces hand-rolled pool; capacity = duty cycle + cost-model admission + distillation | today 1 B200/session; Phase-3 gate: ≥2 sessions/GPU on a recorded duty-cycle trace, p95 within SLO |
| **RL (landed)** | #1450 migrates onto shared loops (Phase 1), engine-client rollouts (Phase 2+); GRPO-class next | the vendored-sampler docstring; C1-by-construction discipline |
| **ComfyUI funnel** | embed (nodes) → **compile** (workflow→PipelineSpec, accelerated cloud) → productize (Studio) | tier-1 ~20-node static sublanguage maps onto PipelineSpec |
## Migration — seven phases, each independently shippable
| Phase | Ships | Gate |
|---|---|---|
| **−1** | Merge cosmos3 chain; seed SSIM for uncovered families | baselines exist |
| **0** | Typed omni I/O; config freeze (`compat.py` shrinks monotonically to zero) | all SSIM suites unchanged |
| **1** | **Loop inversion** + policies + extension core (cache-dit, ParityAligner); RL migrates off its vendored sampler | old vs new loop **bit-identical**; per-method grad-norm refs (#1396) extended |
| **2** | AsyncEngine + StepScheduler; LTX-2 linear graph; Dynamo worker (stock); colocated weight sync | ≤2% batch-1 latency regression; Dreamverse single-session parity; RL engine-client parity |
| **3** | PipelineSpec graphs, role pools, declarative parallelism, ComfyUI compiler MVP, general WeightSyncPlan | ≥2 Dreamverse sessions/GPU on recorded duty-cycle trace |
| **4** | Cosmos3 native; AR continuous batching + paged KV (arriving *with* their workload, per N5); RL hardening (C1/C2, Behavior Record) | Cosmos3 parity suite on new runtime; drift ≈ 0 on a Wan RL run |
| **5** | Deletion: legacy `training/` retires, then `ComposedPipelineBase`, legacy loop, `forward_context.py`, `compat.py`, `RayDistributedExecutor` | the deletion diff — **4 loop copies → 1** |
Enforcement the last freeze lacked (it was broken 19×): `compat.py` frozen from Phase 0; CI path gates + CODEOWNERS once Phase 1 lands; new families land on new abstractions from Phase-1 completion.
## What we are deliberately not doing
Datacenter orchestration (Dynamo's job) · trainer internals (frozen `training/`; `train/` is a consumer) · replacing the bit-exact porting methodology · migrating 20+ families at once · **standalone LLM-serving excellence** — AR machinery arrives only at the sophistication omni workloads pull (N5).
## Decisions we need from this review
1. **sglang `multimodal_gen` relationship** — upstream, friendly fork, or shared core (decide by Phase 2; drift is a strategic cost either way).
2. **Dynamo asks** — green-light proposing the Phase-2 asks (A2 cost interface, A3 media streaming) to the team first, with A5 (RL weight plane) queued behind them?
3. **Phase −1 start** — merge the cosmos3 chain and seed SSIM baselines now; it blocks everything else.
+830
View File
@@ -0,0 +1,830 @@
# FastVideo v3 — A Model-Native Runtime for the (Recipe, Runtime) Era
**Status:** unconstrained north-star. This document assumes we are free to build a brand-new architecture with no
backward-compatibility, no migration tax, and no obligation to the current code. It exists to define the *ceiling*:
the system FastVideo should be if nothing held it back. Migration is a separate, later question — deliberately out of
scope here.
**Lineage:** this is the synthesis of `design.md` (the strategic thesis, product pull, and hard-won serving realism)
and `designv2.md` (the model-native center and typed contracts), with the open tensions of both resolved rather than
hedged.
---
## Table of contents
1. [The thesis](#1-the-thesis)
2. [The two signature ideas](#2-the-two-signature-ideas)
3. [Planes and their dependency order](#3-planes-and-their-dependency-order)
4. [Model Plane — the center](#4-model-plane--the-center)
5. [The loop contract — driven loops](#5-the-loop-contract--driven-loops)
6. [Runtime and scheduler — one currency, one WorkUnit](#6-runtime-and-scheduler)
7. [Memory, cache, transport, compile](#7-memory-cache-transport-compile)
8. [Parallelism as a model contract](#8-parallelism)
9. [Correctness — parity as a typed gate](#9-correctness)
10. [Training and RL on the same loops](#10-training-and-rl)
11. [Extensions — observers and interceptors](#11-extensions)
12. [Request, session, artifact, stream](#12-request-session-artifact-stream)
13. [Programs and workflows](#13-programs-and-workflows)
14. [Deployment and fleet](#14-deployment-and-fleet)
15. [Worked examples](#15-worked-examples)
16. [What this unlocks](#16-what-this-unlocks)
17. [Honest unknowns and falsifiers](#17-honest-unknowns-and-falsifiers)
18. [Package layout](#18-package-layout)
19. [Reference synthesis](#19-reference-synthesis)
---
## 1. The thesis
Three facts about video generation, taken together, dictate the architecture.
**A deployable video model is a post-training artifact.** Unlike an LLM — where inference optimizes frozen weights
post-hoc — a *usable* video model is *created* by training: step distillation is mandatory for usable latency, low
precision needs QAT, and causal/world models are made by distillation plus self-forcing. Every inference capability is
therefore a **(recipe, runtime) pair**: the weights and the loop that produced-and-assumes them are inseparable. A
"4-step NVFP4 FastWan" is not weights plus a flag; it is a distillation recipe, a sampler, a precision path, and a
parity contract that are one object.
**Video systems are loop systems.** Denoise timesteps, AR decode, chunked world-model rollout, VAE tiles, encoder
chunks, audio tokens, reward batches, optimizer steps, media chunks — the work is iteration, not a single `forward`. A
runtime that reduces everything to `forward()` cannot schedule, batch, cancel, stream, reserve memory for, or capture
the behavior of the thing that actually runs.
**Omni models share weights across loop types within one request.** Cosmos3's text reasoner and multimodal denoiser
are the *same resident weights*, driven by an AR loop and then a diffusion loop in one request. This cannot be a
DAG of separate engines (that doubles 30B+ of weights and severs the shared KV/denoise state); it must be one resident
model instance running many loop types. This is achievable — vllm-omni's `bagel_single_stage`/`lance` already run one
resident MoT instance doing AR `generate_text` and diffusion `generate_image` on co-resident experts in a single
request — *but* they bury that interleaving inside one opaque `DIFFUSION` stage their scheduler never sees inside,
request-scheduled with `max_num_running_reqs` forced to 1. The hard part, and the differentiation, is not *expressing*
the shared-weight loops — it is making them **runtime-visible, step-scheduled, batchable, and cost-priced.**
The architecture that falls out:
> **The atomic unit is the (recipe, runtime) pair, owned by a typed `ModelCard`.** Everything — serving, training, RL,
> products, deployment — is a *view* over that card. The **runtime owns loop *lifecycle*** (admission, scheduling,
> batching, caching, cancellation, streaming, behavior capture); the **model owns loop *semantics*** (typed state
> transitions and kernel execution). One resident model instance runs many loops; one scheduler schedules the
> *steps* of all of them in a single currency; one parity contract binds the train-forward to the serve-forward so
> the recipe and the runtime never silently drift apart.
The single invariant, stated once:
```text
Model cards own components, loops, recipes, and parity.
Programs compose loops into tasks.
The scheduler executes the steps of loops as WorkUnits under one budget.
Caches are correct by key, not by hope.
Training records behavior on the same loops it serves.
Deployment places and routes; it does not define semantics.
Products stream artifacts; they do not reach into the model.
```
Everything below is the elaboration of that invariant.
---
## 2. The two signature ideas
Two ideas do most of the work and are what an unconstrained design can reach that an incremental one cannot.
### 2.1 The (recipe, runtime) pair is a first-class, versioned, typed object
A model in v3 is not a checkpoint. It is a `ModelCard` that owns, as one versioned unit:
- the **components** (weights, loaders, layouts),
- the **loops** it can run (the runtime semantics),
- the **recipe** that produced the weights (distillation/QAT/RL config, teacher, data contract, the sampler the
recipe assumes), and
- the **parity contract** asserting that the train-forward and the serve-forward agree to a declared level.
You cannot ship the weights without the loop they assume, and you cannot change the loop without re-proving parity.
This makes design.md's "(recipe, runtime) pair" *literal*: the deployable artifact carries its own provenance and its
own correctness obligation. It is the thing that turns "we do training and inference in one repo" from an org chart
into a guarantee — and it is essentially un-retrofittable, which is exactly why it belongs in a clean design.
### 2.2 Driven loops: the model owns control flow, the runtime owns execution
Loop inversion, done right, is not a heavy `plan_step`/`run_step` contract and not a hidden `for t in timesteps`. It
is a **driven loop**: the model describes the *next step it needs*, the runtime *decides when and with whom that step
runs*, and the model folds the result back into its own state and decides what to do next. The model keeps its control
flow (so content-adaptive decisions — cache-dit skips, EOS, VSA tile selection — are natural); the runtime keeps the
`await` (so admission, batching, cancellation, streaming, and behavior capture are universal). Per-request state lives
in the loop's own typed `LoopState`, never in module globals, so interleaving requests through one model instance
cannot smear state — the failure mode that makes naive loop-inversion dangerous is *structurally* excluded.
These two ideas are developed in §4–§5. The rest of the system is their consequence.
---
## 3. Planes and their dependency order
```text
Products: Python · CLI · OpenAI API · ComfyUI · Dreamverse · RTC · Trainer
│ (thin: validate intent, make requests/sessions, subscribe)
Request / Session / Artifact / Stream
│ (typed runtime objects, cancellation, streaming)
Program Plane ← typed loop programs + compiled workflows
│
┌──────────────── Model Plane (CENTER) ────────────────┐
│ ModelCard: components · loops · recipe · parity │
│ capabilities · caches · parallelism · precision │
└───────────────────────┬──────────────────────────────┘
│
┌─────────────────────────────┼─────────────────────────────┐
│ Runtime / Scheduler │ Training / RL │ (same loops, different capture)
│ WorkUnits · GPU-time budget │ rollout · reward · weight-sync│
└─────────────────────────────┼─────────────────────────────┘
│
Memory · Cache · Transport · Compile (typed CacheKey, per-class pools, CuMem sleep/wake)
│
Parallelism (named axes → DeviceMesh, validated, part of the cache key)
│
Deployment / Fleet (DeploymentCard → Dynamo; never the core)
```
Dependency rules (enforced at the package boundary, §18):
- Products do not define model semantics. Workflows do not define model semantics. Deployment does not define model
semantics. Training does not redefine model semantics. **All of them reference the Model Plane.**
- The runtime *executes* model loops but does not *own* their math. Training *captures* behavior on serving loops but
does not *fork* them. Cross-cutting concerns — extensions (§11), parallelism (§8), and parity (§9) — are contracts
declared on the card, not features bolted onto the runtime.
---
## 4. Model Plane — the center
### 4.1 ModelCard
```python
class ModelCard:
model_id: str # "fastwan-1.3b-nvfp4-4step"
family: str # "wan"
components: dict[str, ComponentSpec]
loops: dict[str, LoopSpec]
capabilities: CapabilityMatrix # text_to_video, image_to_video, reasoning_text, vae_decode, ...
recipe: RecipeSpec # ← what produced these weights (signature idea §2.1)
parity: ParitySpec # ← train-forward ≡ serve-forward, to a declared level (§9)
caches: dict[str, CacheContract]
parallelism: ParallelismContract
precision: PrecisionContract
checkpoint: CheckpointManifest # explicit components, layouts, key maps — no name-detector guessing
```
The card is both a **declarative contract** (strict enough to validate before any GPU touches it) and a **runtime
factory** (it knows how to instantiate components, bind loops, and resolve caches). It is hub-interchange compatible
with diffusers' `modular_model_index.json` / `ComponentSpec` so models published either way load both ways.
`CheckpointManifest` replaces today's implicit `model_index.json` + name-detector resolution with explicit declared
components and `required_for` / `optional_for` task sets (the Cosmos3 lazy-sound-VAE problem becomes a declaration, not
an `if env_var` inside `forward`).
### 4.2 RecipeSpec — the provenance half of the pair
```python
class RecipeSpec:
method: str # "dmd2" | "self_forcing" | "attn_qat_nvfp4" | "diffusion_nft" | "base"
parents: list[str] # teacher / base model_ids this was distilled or RL'd from
data_contract: DataRef # what the recipe trained on (for governance and reproduction)
assumes_loop: str # the loop_id this recipe's weights require at serve time
assumes_precision: str # the precision the QAT recipe baked in
consistency_required: str # the minimum parity level this recipe's outputs are valid under (§9)
```
`assumes_loop` and `assumes_precision` are the teeth: a 4-step distilled model whose `assumes_loop = "ddim_4step"`
cannot be served under a 50-step sampler without a typed mismatch error. The recipe and the runtime are bound.
### 4.3 ComponentSpec and LoopSpec
```python
class ComponentSpec:
component_id: str
kind: str # dit | vae | text_encoder | reasoner_tower | reward_head | ...
load_id: str
config_schema: type
io_schema: tuple[type, type]
precision_policy: PrecisionPolicy
placement_policy: PlacementPolicy
parallel_constraints: ParallelConstraint
parity_tests: list[ParityTestSpec]
class LoopSpec:
loop_id: str # diffusion_denoise | ar_decode | chunk_rollout | vae_tile | ...
state_schema: type # the typed LoopState (no dicts)
step_schema: type # the typed WorkPlan a step emits
result_schema: type # the typed StepResult a step returns
behavior_schema: type | None # what to capture for RL (None if not training-relevant)
step_cost_model: CostModel # predicted GPU-time per step at (shape, precision, policy) — §6
valid_parallel_plans: list[ParallelPlanPattern]
graph_capture: GraphCapturePolicy
cache_policy: CachePolicy
```
A `ModelInstance` is a resident, loaded card: component instances, model state, caches, compiled graphs, and a
parallel plan. **A request may run several of the card's loops against one `ModelInstance`.** That single sentence is
the difference between this design and a stage-only design, and it is what makes omni native:
```text
one Cosmos3 ModelInstance, one request:
ar_decode(reasoner) → pack → diffusion_denoise(vision[+action][+sound]) → vae_tile_decode → audio_decode
└────────────── same resident weights, shared packed state, scheduled as steps ──────────────┘
```
---
## 5. The loop contract — driven loops
### 5.1 The contract
A loop is a **serializable state machine** the runtime drives:
```python
class Loop(Protocol):
def init(self, req: Request, model: ModelState, ctx: LoopContext) -> LoopState: ...
def next(self, state: LoopState) -> WorkPlan | Done: ... # describe the next step; NO GPU kernels here
def advance(self, state: LoopState, result: StepResult) -> LoopState: ... # fold result in; decide what's next
def finalize(self, state: LoopState) -> LoopResult: ...
```
The runtime's driver — the only place iteration lives:
```python
state = loop.init(req, model_state, ctx)
while True:
plan = loop.next(state) # typed WorkPlan: resources, cache reads/writes, shape, sinks, cancel-scope
if isinstance(plan, Done):
break
result = await ctx.execute(plan) # ← THE INVERSION POINT: runtime admits, batches, places, runs, returns
state = loop.advance(state, result) # content-adaptive: next() can branch on everything in state, incl. result
for chunk in plan.emits:
ctx.emit(chunk) # streaming falls out
return loop.finalize(state)
```
Why this is the right contract, point by point against the failure modes:
- **Content-adaptive steps are natural.** `next()` reads `state`, and `advance()` has already folded in the last
`StepResult` — so cache-dit's skip decision (a residual comparison from the prior step), AR's EOS, and VSA's
content-dependent tile selection are ordinary control flow in the model. This is the tension `designv2.md`'s
"pure plan_step that pre-declares shape" could not resolve; here it dissolves, because planning the *next* step is
allowed to depend on the *previous* result. `next()` is still kernel-free (it *describes* work; it does not run it),
which is all the scheduler needs.
- **Cross-request state safety is structural.** All per-request state is in the loop's typed `LoopState`. There are no
module-level residual/KV globals (the bug that silently corrupts cache-dit and TeaCache forks under concurrency).
Interleaving requests through one `ModelInstance` cannot smear state because there is no shared mutable state to
smear. This is the safety property naive loop inversion lacks, made impossible-to-get-wrong by construction.
- **The runtime owns everything it needs and nothing it doesn't.** `execute(plan)` is the single seam for admission,
memory reservation, cross-request batching, placement, graph dispatch, cancellation, and behavior capture. The model
never sees the scheduler; the scheduler never sees the model's math.
- **Serializable, therefore migratable and resumable.** `LoopState` is typed and serializable, so a half-finished
1000-step job is a resume point: preempt by stopping the driver, migrate by shipping `LoopState` to another worker,
recover from a crash by replaying from the last serialized state. (A coroutine that keeps state in a suspended Python
frame — the tempting sugar — cannot do this; the explicit state machine is the price of resumability, and it is
worth paying.)
Custom step bodies are first-class, not an escape hatch with an asterisk: a family whose math is genuinely braided
(Cosmos's EDM coefficients consumed inside the CFG branch with an x0-space combine; LTX-2's 1–4 runtime-decided
guidance passes) writes `next()`/`advance()` by hand using samplers and CFG utilities as a *library*. The runtime
requires only the four methods; *how* a step body is factored is the model's business. Policies (CFG, expert routing,
precision, flow-shift, conditioning) are the *default* decomposition that deletes duplication for the families that
fit — never an admission requirement.
### 5.2 Loop granularity
Chosen by runtime value, not purity:
- too coarse → cannot cancel/batch/reserve/stream/record at useful points;
- too fine → scheduler overhead dominates, graph capture fragments;
- good default → one denoise step (or window), one AR decode batch, one encoder chunk, one VAE-tile batch, one
reward/logprob batch.
The runtime may **fuse** adjacent compatible WorkPlans after planning (an optimization); the unfused boundary remains
the semantic model, so parity and behavior capture are defined on the unfused loop.
### 5.3 CFG is a policy over *one* shared denoise body (verified)
A natural worry: CFG changes the *shape* of the step (one forward vs two vs a batched pair vs a data-parallel split),
so can one shared denoise loop really host all of it by swapping a policy? **Yes — and it is proven by existing code,
not aspiration.** vllm-omni's `CFGParallelMixin.predict_noise_maybe_with_cfg` + `combine_cfg_noise`
(`diffusion/.../cfg_parallel.py:76-212`) already runs sequential-2-forward, batched-1-forward, *and* cfg-parallel
through **one** pair, with the loop body unaware of which. The clean cut is three layers:
- **In-loop `CFGPolicy`** — branch vocabulary (`[cond]`, `[cond, uncond]`, per-modality, STG-perturbed),
the combine formula (standard `uncond + s·(cond−uncond)`, CFG-zero `st_star`, `cfg_normalize`/`guidance_rescale`),
and **per-request mutable state** (the adaptive-gate cached delta with model-id self-invalidation is the canonical
state case — and exactly why state lives in `LoopState`, §5.1). **Batched-vs-two-forward is a *dispatch detail
inside one policy*, not a separate mechanism.** This covers classic / batched / adaptive-gate / per-modality.
- **`cfg`-parallel is a *parallelism axis*, not a policy** — it shards the policy's branches across ranks and runs the
*same rank-invariant `combine`* on every rank. It composes *under* any `CFGPolicy`; you own a `BatchedCFG` policy
**or** a `cfg` group, never both (the §9 build-guard).
- **Companions are an *orchestrator pattern*, not in the loop** — splitting a request into companion sub-requests
upstream of diffusion (the conditioning is precomputed and bundled in; the loop is unchanged).
Two caveats keep the first pass honest: the `combine` runs in *the step body's* numeric space (Cosmos combines in
x0-space after EDM preconditioning, not noise-space — the body fixes the space, the policy fixes the algebra), and
embedded-guidance (Flux) is a **degenerate single-branch identity-combine policy** (guidance rides inside the forward
kwarg), kept *inside* the same abstraction rather than special-cased as "no CFG." This is the same shared denoise body
that RL rollout reuses (§10) — one CFG taxonomy serves both serving and rollout.
---
## 6. Runtime and scheduler
### 6.1 One WorkUnit, one currency
Every `await ctx.execute(plan)` produces a **WorkUnit**: the smallest schedulable action with a resource reservation
and a loop boundary. Kinds: `ar_prefill`, `ar_token`, `diffusion_step`, `diffusion_window`, `chunk_step`,
`encoder_chunk`, `vae_tile`, `audio_chunk`, `reward_batch`, `logprob_batch`, `transfer`, `cache_io`, `graph_capture`.
Tokens are *one kind*, not the scheduler — this is the generalization of vLLM's token scheduler that diffusion forces.
**The budget currency is predicted GPU-time, not counts.** A bidirectional denoise step re-attends the full latent at
O(L²) with zero KV amortization (every step pays full price); an AR decode step is ~O(context) against a cache; a
chunked-causal step sits between. Counting "steps" or "tokens" puts items three orders of magnitude apart in one
bucket. So each WorkUnit converts to GPU-seconds via its `LoopSpec.step_cost_model`, calibrated online by the Profiler
(§11). **The same cost model is the interface published to the fleet** (§14): the scheduler's internal budget and
Dynamo's routing/autoscaling input are one object, built once.
Two honesty caveats, kept from design.md's contact with reality:
- **Admission uses the conservative baseline.** The design's own flagship features make realized cost unknowable in
advance — cache-dit skips are residual comparisons, VSA tiles are content-dependent, AR length is unbounded
(budgeted at the `max_tokens` cap, refunded on early EOS). Telemetry refines calibration; it never licenses
admission optimism.
- **A denoise step is indivisible.** A 30s-1080p step bounds iteration latency no matter the budget. Mitigations are
first-class, not afterthoughts: cost-class pools (jumbo steps don't co-schedule with latency-class work), SP within
a node to shrink jumbo wall-time, and admission-time SLO classes so the fleet planner scales pools per class. This
is why the scheduler is *cost-aware*, not just *count-aware*.
### 6.2 WorkPlan and admission
```python
class WorkPlan:
loop_id: str
instance_id: str
kind: str
shape_sig: ShapeSignature # for batch compatibility + graph capture key
resources: ResourceRequest # compute (GPU-s), resident bytes, peak-activation bytes, cache blocks, xfer bw, sinks
cache: CachePlan # typed reads/writes (§7)
placement: PlacementHint
cancel_scope: CancelScope
emits: list[StreamChunk]
class Done: result: LoopResult
```
**Admission rule (the soundness condition of multiplexing):** *do not admit a waiting WorkUnit unless every resource
it requests can be reserved* — compute budget **and** memory (resident + worst-case peak) **and** cache blocks **and**
transfer bandwidth **and** graph-capture shape **and** output sinks. Two requests that fit individually but jointly OOM
are rejected at admission, not discovered at step 37. This is vLLM's "token budget is half the story, `allocate_slots`
is the other half," generalized.
### 6.3 The scheduler, in layers (each testable on a fake pool, no GPU)
1. **RequestScheduler** — accepts requests/sessions, selects programs, starts loop drivers.
2. **LoopScheduler** — drives `next()`, collects pending WorkPlans.
3. **BatchScheduler** — groups compatible WorkPlans by `(instance, loop_kind, shape_sig, precision, parallel_plan,
graph_key)`; image diffusion and AR decode batch across requests, jumbo video stays batch-of-1, Cosmos-style
token-budget packing is an opt-in.
4. **PlacementScheduler** — worker, role pool, instance, device mesh.
5. **TransferScheduler** — tensor / cache / artifact movement as scheduled WorkUnits.
6. **AdmissionController** — the reservation gate of §6.2.
Policies: running loops first (vLLM); preempt only at loop step boundaries; cancel only at declared scopes; prefer
cache hits when latency/fairness allow; never starve a long denoise loop behind short AR requests.
### 6.4 SPMD consistency and failure isolation
All ranks of a pool must make identical scheduling decisions or NCCL deadlocks. **Rank-0 decides, broadcasts** — the
existing discipline, now also the channel for the **abort broadcast** (failure isolation and scheduling share one
consistency mechanism). Failure classes: *request-fatal* (NaN flagged by NaNWatch, one request's step error) →
SPMD-consistent abort of that request, deliver partial artifacts with a structured error; *pool-fatal* (illegal
access, NCCL desync) → pool re-init, invalidate pool caches, resume requests from serialized `LoopState` where one
exists. **Cancellation is common-path, not exceptional** — vibe-directing makes abandoning in-flight work the *normal*
user action; cancel takes effect at the next step boundary, drops queued WorkUnits, releases `LoopState` and cache
handles, reports `cancelled`.
---
## 7. Memory, cache, transport, compile
Video and omni inference are memory systems as much as compute systems. This plane is explicit and typed.
### 7.1 Cache correctness is a contract
```python
class CacheKey:
model_id: str; component_id: str; loop_id: str | None
weights_version: str; adapter_versions: dict[str, str]
precision: str; parallel_plan_hash: str
shape_sig: str; layout_sig: str
scheduler_sig: str | None; guidance_sig: str | None; seed: int | None
input_hashes: dict[str, str]; step_index: int | None
contract_version: str
```
**If a field can change output semantics, it is in the key.** Incorrect reuse is worse than no reuse. The serving
hazard this kills: a workflow-cloud request that shares a prompt but differs in te-LoRA stack must not serve stale
embeddings — so the key is *partitioned* by `adapter_versions`, not flushed. An RL `update_weights` bumps
`weights_version` and invalidates wholesale.
### 7.2 Per-class pools (the granularity reality)
There is **no single unified block pool**, because a unified pool requires uniform bytes-per-block and our cache
classes differ by 150–500× in natural granularity (a text-KV page ≈ 64 KB/layer; a causal-video latent-chunk slab is
9.6–32 MB/layer) and their demand is workload-decoupled. Each class gets a statically budgeted pool behind one
`CacheHandle`: paged text-KV (`ar_decode`), slab chunk-KV (`chunk_rollout`, with a declared training mode that
disables mid-rollout recycling and keeps grad-aware index snapshots), feature caches (text/vision-encoder, content-hash
keyed, reference-counted FIFO), residual caches (cache-dit, scoped per `LoopState`), weight/adapter cache
(disk→CPU→GPU LRU for the workflow cloud). MoT falls out: the und pathway draws paged KV, the gen pathway draws slab or
nothing — independent budgets, no interference. **KV is the minority case** — a pure bidirectional deployment allocates
none of it; the machinery materializes only when a card declares KV-bearing loops.
### 7.3 Memory, transport, compile
- **Memory** — tagged pools, sleep/wake by tag (CuMem-style; tags are component names), reservation before admission,
per-role budgets, host-pinned staging. Sleep/wake is component-granular for RL (drop DiT + caches, keep
VAE/text-encoder resident).
- **Transport** — manifest-based and pluggable: in-proc reference → SHM → CUDA IPC → NCCL/UCXX/NIXL/RDMA →
object-store. KV/cache-bearing edges speak a `KVConnector`-shaped protocol (scheduler-side query/alloc/finish +
worker-side async load/save) so NIXL/LMCache/Mooncake/KVBM implement it directly. Transfers are scheduled WorkUnits,
not side effects.
- **Compile** — CUDA graphs and `torch.compile` managed by a `CompileCache` keyed on `(model, component, loop,
work_kind, shape_sig, precision, parallel_plan, backend)`. **Never full-graph across the engine** (per vLLM's own
reversal): per-block compile where it pays, manual fused ops permitted in model code, breakable CUDA graphs as an
*optimization tier* over an always-correct eager baseline. Graph capture is planned by the scheduler (padding,
bucketing, capture sizes affect admission and batching).
---
## 8. Parallelism as a model contract
Parallelism is not a launch flag; it affects cache keys, scheduling, transport, capture, and parity, so it lives on
the card.
```python
class ParallelPlan:
axes: dict[str, int] # dp, tp, sp(=ulysses×ring), cp, cfgp(≤2), pp_patch, vae, ep, fsdp, role, replica
mesh_order: list[str]
placement: PlacementSpec
communication: CommunicationSpec
```
Declarative, validated, compiled to a PyTorch `DeviceMesh` via a `ParallelDims`-style builder
(product-of-degrees validation, cached submeshes). **Pre-flight or it fails at load, never halfway.** Ownership
conflicts are build errors (CFG owned by a `BatchedCFG` *policy* or a `cfgp` *group*, never both). Applicability
conditions travel with axes: `pp_patch` (PipeFusion displaced-patch pipelining) is **invalid for causal/AR** (stale KV
breaks causality) and the validator enforces it per card. Degree-one axes exist as trivial groups so component code
needs no special cases. **Pools are single-node**; multi-node scale is *multiple pools* fronted by the fleet (§14) —
the engine never owns a cross-node NCCL mesh inside one pool.
---
## 9. Correctness — parity as a typed gate
This is the section both prior documents needed and neither fully had.
### 9.1 The parity contract
Every card carries a `ParitySpec`. Parity is **measured, never assumed**, by a `ParityAligner` observer (§11): record
named taps per step/block from a reference (the official framework, or a pre-change build); compare-mode replays with
fixed seeds and reports the first divergence beyond per-tap tolerance. This is the engine behind the "old loop vs new
loop, bit-identical" gate and the standing instrument for every port, precision change, and kernel swap.
### 9.2 The consistency ladder (with the rung both prior docs missed)
```text
C0 component parity — VAE, encoder, transformer block, scheduler step in isolation
C1 loop parity — full denoise trajectory / AR logits, fixed seed
C2 behavioral identity — the train-forward and serve-forward agree on the quantity the RL objective uses:
· likelihood-based methods (GRPO-class): per-step log-prob identity
· likelihood-free methods (DiffusionNFT-class): seeded final-sample +
prediction-space identity (old_deviate / ref-MSE) — there are NO log-probs to match
C3 distribution parity — rollout distribution under allowed nondeterminism
C4 artifact quality — SSIM-class, reward agreement, human-preference (gates product claims; needs the eval system)
```
The C2 split is load-bearing and is the lesson of the landed RL stack: the shipped Wan DiffusionNFT is
**likelihood-free** — it captures only final clean latents and contrasts the student against an implicit negative
policy in prediction space, so "log-prob identity" is *undefined* for it. A ladder that assumes log-probs (as both
`design.md`'s and `designv2.md`'s early framings did) cannot describe the only RL method actually in the tree. RL
methods declare their required level on the `RecipeSpec`.
### 9.3 The gate that catches what batch-of-1 cannot
Loop inversion's real hazard is **cross-request state smearing under interleaving** — and a batch-of-1 parity gate is
*structurally blind* to it, because the corruption only manifests when two requests share a loop. §5.1 excludes the
hazard by construction (state in `LoopState`, never globals), but construction-arguments need a test. So v3 makes a
**batch-of-N interleave parity test** a *required* gate: two (or more) concurrent requests, interleaved at step
granularity, must be bit-identical to the same requests run serially. This is the test the whole loop-inversion bet
lives or dies on, and it is named here as a first-class obligation, not left implicit.
### 9.4 Three execution profiles, one definition
Even in one runtime there are three forwards: the **serve** forward (no-grad, graphed, cached, possibly quantized), the
**rollout** forward (serve profile + behavior capture), and the **train** forward (grad, checkpointed, FSDP-gathered).
They share *one* loop definition; they differ only in grad mode and capture. The ladder measures the gap; the recipe
declares the level it needs. "Train BF16, serve FP8" is legal only at C2-corrected with importance-sampling, and the
card says so. This is how the (recipe, runtime) pair stays honest: the contract is typed and tested, not trusted.
---
## 10. Training and RL on the same loops
```text
serve : request → program → loop → WorkUnits → artifacts
rollout : prompt batch → program → loop → WorkUnits → BehaviorRecords → rewards → update
```
The loop kernel is shared; the only difference is output capture and training policy — not a second interpretation of
the model. This is design.md's §8 thesis and v2's training plane, with the dependency rule kept absolute: **`training`
may require behavior records but must not fork serving loop logic; the engine never imports `training`.** The engine
*is* the rollout engine (it already runs the loops); the trainer is a client.
**This is the moat — and it is the one place a serving-only runtime structurally cannot follow.** vllm-omni proves
omni serving can be production-grade, but it has *no* training/RL plane at all; verl-omni and miles prove the
alternative — a standalone trainer-side sampler on a *different* runtime than serving — costs the two-runtime tax
forever. The whole point of collocation is that **the rollout forward *is* the serve forward plus capture**: same loop,
same caches, same batcher, same numerics. Three consequences nothing else gets:
- **Every serving optimization is automatically a rollout optimization.** Distilled few-step samplers, cache-dit
skips, CFG-parallel, paged/feature caches, step batching — the recipe team builds them once for serving and the RL
rollout inherits them for free. FastVideo's *own* landed DiffusionNFT is the negative example that proves the
point: it vendors a bare-model `for`-loop (`rl/common/sampling.py`, whose docstring says it "intentionally does not
call FastVideo's full inference pipelines"), and DMD2 vendors a *second* one (`dmd2.py::_student_rollout`) — so
today's rollout runs with **zero serving-grade optimizations** (no CFG, dense attention, full 25-step ODE,
one-sample-at-a-time). Collocation deletes both private loops.
- **RL rollout is a *better* batching case than open-world serving — not a worse one.** A GRPO/NFT group is K
*identical-config* samples of one prompt: same shape, same schedule, same CFG branch. The landed config is K=24
(`num_video_per_prompt: 24`), 6 prompts/batch × 48 batches = **288 prompt-slots/GPU/epoch, each a 24-wide homogeneous
denoise batch** — zero bucketing required (serving must bucket heterogeneous resolutions/steps/CFG across users; a
GRPO group is homogeneous *by construction*). And all K samples share one prompt embedding, so the content-hash
feature cache computes the text encoder **once per group instead of 24×**. The vendored sampler captures none of
this; it carries the embedding per sample and runs one shape at a time.
- **One numerics surface.** Serve-forward and rollout-forward differ only in grad mode and capture (§9.4), so there is
no rollout-vs-train kernel gap to patch — the consistency ladder *measures* the gap rather than a correction layer
*papering over* it. For the landed likelihood-free NFT, "reuse holds" means it holds at the **C2 behavioral rung**
(seeded sample + prediction-space identity), under a `CFGPolicy` that is conditional-only and a `WeightSyncPlan`
whose role is the decay-blended old policy — all of which the card already declares.
- **BehaviorRecord** — captured at generation time (reconstructing later is fragile): seeds, scheduler trajectory,
timesteps, latents-or-refs, logprobs *where applicable*, sampled/action tokens, guidance, reward in/out, cache
assumptions, precision, parallel plan, attention backend, deterministic flags, `weights_version`. Sized honestly:
full MoE-routing capture is GB/sample for Cosmos3-class requests, so it is an **opt-in instrument** for goldens and
debugging, not always-on.
- **Weight-sync lifecycle** — freeze admission for the affected role/version → drain or boundary-stop in-flight loops →
transfer weights/deltas → bump `weights_version` → invalidate incompatible caches and graphs → publish version →
resume. A `WeightSyncPlan` is three inputs (mesh specs + per-model layout adapters + transport), validated
pre-flight, CPU-testable on fake pools. RL ships a *role*, not "the weights": student / EMA / decay-blended old
policy is declared (the landed NFT behavior policy is the *old* copy, not the student — the plan must carry that).
- **Roles** (policy, rollout, reference, reward, critic, evaluator, data, coordinator) reference the same cards and
loops; they are deployment concerns, scaled by the fleet.
- **The industry tax we delete:** verl-omni re-implements Wan inside vLLM-Omni and corrects numerics afterward; miles'
headline features (TIS/MIS, bitwise logprobs, R3 routing replay, unified FP8) are all mismatch patches for *two
runtimes with different kernels*. One model definition, one kernel set, one measured ladder is the answer — viable
at FastVideo's 1–30B FSDP2 scale (the boundary condition: a Megatron-class trainer at 100B+ re-enters the
two-runtime world, and the ladder is the fallback there).
---
## 11. Extensions — observers and interceptors
The optimization, debugging, and parity surface, as versioned hook points assembled at loop build (an unused hook is
*literally absent* from the hot path). It composes with §5 cleanly: the hooks wrap `ctx.execute(plan)`.
- **Observers (read-only):** `ParityAligner` (§9), `Profiler` (per-step wall+CUDA, calibrates the cost model),
`NaNWatch` (first-NaN localization), `ActivationTrace`. They cannot mutate state.
- **Interceptors (compute-altering):** `StepInterceptor` (step-skip / cached-prediction) and `BlockInterceptor`
(cache-dit's DBCache/FBCache/TaylorSeer). State lives in `LoopState.plugin_state[id]`, keyed **per request and per
CFG branch** — the structural fix for the module-global residual state that silently corrupts cache-dit/TeaCache
forks under concurrency. cache-dit is the reference integration (the library sglang's serving already uses);
conflicting interceptors are rejected pre-flight; a 4-step distilled card *rejects* step-skip caches rather than
producing garbage.
**Trust boundary:** plugins are enabled at deploy scope only (never a per-request `plugins=[...]` field that would wire
third-party code selection into the public API); requests only *parameterize* pre-enabled plugins through validated
schemas, and exact-mode requests reject `distribution_altering` parameterization outright.
---
## 12. Request, session, artifact, stream
Typed runtime objects, not IDs in a batch (the Dreamverse/LiveKit lesson):
- `Request` — one generation, scoring, encoding, training-sample, or conversion job.
- `Session` — a long-lived interactive context: prompt memory, media streams, cancellation, partial updates,
cross-request chunk-KV that persists for a game/scene session.
- `Artifact` — a *named, typed* output with provenance (which node produced it): `VideoArtifact`, `AudioArtifact(
sample_rate)`, `TextArtifact(token_ids, text)`, `TensorArtifact`, `LatentArtifact`. This kills the `extra["audio"]`
pattern — audio carries its sample rate as a first-class artifact, not a dict passenger.
- `Stream` — one ordered event channel for previews, media chunks, progress, logs, finals.
- `CancelScope` — structured cancellation target (request / loop / stream / session).
Typed event taxonomy (`request.*`, `session.*`, `artifact.*`, `media.{init,chunk,complete}`, `trace.*`). A
`media.chunk` must know its stream, byte-range or shared-buffer ref, codec/container, timestamp range, and
preview-vs-final — invalid combinations are unrepresentable.
The **request is the only currency crossing the product boundary.** A typed `Request` carries `task: TaskType`
(declared, never inferred), `inputs: list[ModalPart]` (Text/Image/Video/Audio/Action/Latent), AR `sampling` vs
`diffusion` params, an `OutputSpec` (requested modalities + streaming + capture flags), and per-node overrides. Task is
declared; heuristics may only *suggest* a default at the boundary.
---
## 13. Programs and workflows
A **Program** composes a card's loops into a task; the card says what loops *exist*, the program says how to *run* them
for this request. Kinds: `InlineProgram` (many loops, one resident instance — the omni default), `DisaggregatedProgram`
(encoder→denoiser→decoder role pools), `WorkflowProgram` (compiled from ComfyUI), `TrainingProgram`,
`RealtimeProgram`. Nodes: `ModelLoopNode`, `ComponentNode`, `ExternalNode`, `ArtifactNode`, `ControlNode`,
`StreamNode`, `TransferNode`. Edges are typed (`TensorEdge`, `ArtifactEdge`, `StreamEdge`, `ControlEdge`, `CacheEdge`,
`BehaviorEdge`). Linear pipelines are the degenerate case; branches/fan-out/fan-in are real (video and audio decode in
parallel after a joint denoise). A separate deploy config maps nodes → pools/devices/parallelism, defaulting to "one
pool, everything colocated."
**Workflows compile, they are not the runtime.** A ComfyUI workflow's tier-1/tier-2 static sublanguage maps onto a
`Program` (`CheckpointLoaderSimple→card`, `KSampler→diffusion_denoise` with sampler/CFG policies,
`LoraLoader→adapter hot-swap`, `ControlNetApply→ConditioningInjector`); unknown nodes become `ExternalNode`s or a
coverage rejection — never silent wrongness. The moat: an orchestrator can run stock workflows on rented silicon;
substituting a *credibly faster* model requires owning the recipe (§2.1) — which an orchestrator structurally cannot
do. Equivalence is a quality-metric vs a reference render (C4), never a bit-parity claim.
---
## 14. Deployment and fleet
The engine exports a `DeploymentCard` and lets a fleet orchestrator (Dynamo) route — Dynamo orchestrates engines, it
is never the engine core.
```python
class DeploymentCard:
engine_id: str; model_cards: list[str]
capabilities: CapabilityMatrix; role_pools: list[RolePoolSpec]
supported_programs: list[str]; supported_parallel_plans: list[ParallelPlan]
cache_events: list[CacheEventSpec]; transfer_endpoints: list[TransferEndpoint]
cost_model: CostModel # the SAME §6 cost model — one object, two consumers
health: HealthSchema; slo: SLOSchema
```
Clean line: the **fleet** owns global routing, tenant policy, cold start, role-pool scaling, cross-node transfer,
placement-by-SLO, health/failover, multi-engine upgrades, global cache routing. The **engine** owns model load, loop
execution, local scheduling, local memory/cache, model-specific behavior, parity, WorkUnit batching. The asks of the
fleet are concrete and each has a fallback: generic affinity key-spaces (checkpoint/session/lora/weight_version beyond
token prefixes), a heterogeneous request-cost interface (the §6 cost model), chunked media streaming through the
frontend, role-graph disagg (N roles, not two), an RL weight plane (versioned broadcast + staleness-aware routing),
cache-object tiering (KVBM generalized to latent/session caches), and session lifecycle as a routing primitive.
---
## 15. Worked examples
**(a) Text → video, one instance.** `Request(T2V)` → `InlineProgram` → `diffusion_denoise` loop. Driver: `init`
builds sigmas/latents; `next` emits a `diffusion_step` WorkPlan (batch-of-1, SP+CFG-parallel); `advance` folds the
model output and the CFG combine; cache-dit's `BlockInterceptor` may skip blocks based on the prior residual; at
`Done`, a `vae_tile_decode` loop runs; output is a named `VideoArtifact`. Compiles/captures exactly like today's inner
loop — nothing tensor-level changes for batch-of-1.
**(b) Cosmos3 omni, one request, shared weights.** `InlineProgram` over one `ModelInstance`: `ar_decode(reasoner)`
yields `ar_token` WorkUnits that join the AR continuous-batching group → `pack` → `diffusion_denoise(vision+action+
sound)` yields `diffusion_step` WorkUnits → fan-out `vae_tile_decode` + `audio_decode`. The reasoner's tokens and the
denoiser's steps hit the *same resident weights*; the scheduler is the mode multiplexer. AR decode runs data-parallel
across the cfg×sp weight-replica axes (decode is sequence-length-1; SP has nothing to shard). This is the workload no
DAG-of-engines can express.
**(c) Image serving at scale.** Many `Request(T2I)` → the `BatchScheduler` groups `diffusion_step` WorkUnits by
resolution bucket and batches across requests every step — the case where cross-request batching pays most. The
*same* scheduler that runs (a) and (b).
**(d) RL rollout.** A `TrainingProgram` drives the *same* `diffusion_denoise` loop with `OutputSpec(capture=behavior)`;
each step emits a `BehaviorRecord` slice; rollouts run C2 by construction (in-process, trainer kernels, pinned
attention). For likelihood-free NFT the behavior is seeded final latents + prediction-space deviations; for a future
GRPO-class method it is per-step log-probs — the loop is identical, the capture differs, the ladder rung is declared.
**(e) Dreamverse session.** A `Session(realtime_video_continue)` holds chunk-KV across 5s segments; `push_text`
updates prompt memory; `stream` yields `media.chunk` previews from loop `emit`s; a direction change throws
`Cancelled` at the next step boundary and starts a new segment. Capacity comes from duty cycle + cost-model admission +
distillation — interleaving is fairness, not throughput.
**(f) ComfyUI compile.** `workflow.compile(json)` → a `WorkflowProgram` of `ModelLoopNode`/`ComponentNode` over a
weight-fleet-cached card, with stacked-LoRA patch/unpatch priced by the §6 weight-transition cost — same runtime, new
frontend.
---
## 16. What this unlocks (the unconstrained payoff)
Things no incremental design — and neither prior document — could actually claim:
- **True omni/MoT serving.** One resident model, many loop types, scheduled at step granularity, in one request. Not a
monolith bypassing the abstraction (the Cosmos3 port's necessary hack), not a DAG doubling weights — native.
- **Train ≡ serve by construction.** Because rollout and serve are the *same loop*, the (recipe, runtime) flywheel is
real and measured, not aspirational: Dreamverse's directing sessions emit preference data → the RL plane → faster
distilled cards → a better product, with the ladder guaranteeing the preferences collected under the serving profile
transfer into training.
- **Real-time interactive omni.** The driven-loop contract + sessions + WebRTC frame/PTS streaming + step-boundary
cancellation make the <100ms motion-to-photon interactive world-model loop expressible in the same runtime that
serves batch T2V.
- **One substrate, three personas.** A research vehicle (new ports land as cards), a product engine (Dreamverse, the
workflow cloud), and an RL rollout engine — without three codebases. The dependency rules keep them from fusing into
mud.
- **Correctness you can sign.** A deployable card is a *(recipe, runtime)* pair with a typed parity obligation; "this
fast model is equivalent" is a claim with a test behind it, which is the one thing an orchestrator-without-recipes
can never say.
---
## 17. Honest unknowns and falsifiers
An unconstrained design is not an unfalsifiable one. The bets, stated with the experiment that kills each:
- **The novelty is concentrated and real.** Runtime-owned diffusion iteration has a **narrow, opt-in precedent** —
vllm-omni's `SupportsStepExecution` (`prepare_encode/denoise_step/step_scheduler/post_decode`,
`diffusion/models/interface.py:44-67`) is exactly runtime-owned diffusion iteration at step granularity, and maps
almost 1:1 onto our `init/next/advance/finalize` — but it is Qwen-Image-only and off in every shipped deploy. What
is unprecedented is making it the **always-on universal contract** *and* a fully general WorkUnit scheduler over
heterogeneous units. The risk is not the loop contract (a state machine is well-understood, and now demonstrably
shippable); it is whether step-level cross-request scheduling *pays* for video. **Falsifier:** publish a load profile and targets from a real duty-cycle trace; if step-level scheduling does
not beat a request-level baseline (≥2 concurrent sessions/GPU, p95 within SLO), the scheduler degrades to
request-level dispatch and the loop contract keeps only its streaming/cancellation/behavior seams — which still
justify it. The contract is safe even if the scheduling bet loses; that is the design's insurance.
- **The general WorkUnit scheduler may be over-general.** Scheduling VAE tiles, transfers, and graph-captures through
the *same* admission machinery as denoise steps is elegant and unproven. **Falsifier:** if, after Phase 2, the
non-diffusion/non-AR WorkUnit kinds (tile, transfer, cache_io) gain nothing from unified scheduling over a simple
in-loop call, collapse them back to in-loop operations and keep WorkUnits for the step-bearing kinds only.
- **Cost-model admission is a modeling bet.** It converges toward cost-class pool routing once the indivisible-step
reality is respected — which is close to what request-level pooling + a fleet planner already do. The fine-grained
interleave win has a *narrow* window (many small concurrent jobs); it should be argued on that window, measured.
- **The clean-slate premise is the elephant.** This document deliberately ignores migration. The org that would build
it broke its own freeze 19 times and ships 20+ families, a live product, and a landed RL stack. A clean-slate
rebuild is the highest-risk path that exists for *this* org; the responsible realization is to build v3 as a
*parallel* engine around one forcing-function card (Cosmos3), prove it on the parity ladder, then migrate families
onto it behind an adapter while everything keeps shipping — i.e., reach this architecture incrementally. That plan is
out of scope here by request; it is non-optional in reality.
- **Quality is unmeasured.** C4 (artifact quality / human preference) and the eval system it needs do not exist yet,
and they gate every product claim ("fast mode is equivalent", RL reward validity, distillation comparisons). Named
as a required, currently-absent subsystem, not assumed.
---
## 18. Package layout
```text
fastvideo/
card/ specs, components, loops, recipes, parity, checkpoints, capabilities # the Model Plane
loop/ driver, loopstate, workplan, policies (cfg, expert, precision, flowshift, conditioning)
runtime/ engine, scheduler/{request,loop,batch,placement,transfer,admission}, workers, events
cache/ keys, classes/{paged_kv, slab_kv, feature, residual, weight_fleet}, policies
memory/ allocator, sleep_wake, reservations
transport/ manifests, backends/{shm, cuda_ipc, nccl, nixl, kvbm}, relay
parallel/ plans, mesh, process_groups, validation
parity/ aligner, ladder, interleave_gate # §9 is its own home
extend/ observers, interceptors, cache_dit, registry, trust
program/ specs, compiler, workflows
request/ requests, sessions, artifacts, streams, cancel
training/ rollout, behavior, rewards, weight_sync, methods # imports card/loop/runtime; never imported by them
deploy/ cards, role_pools, dynamo_adapter
integrations/ comfyui, dreamverse, livekit, diffusers
```
Enforced boundaries: `card/` imports no product/runtime; `runtime/` executes `card/` loops but defines no semantics;
`training/` may require behavior records but forks no loop; `integrations/` adapt external systems into core specs and
events, never bypass them. **`parity/` is a first-class package**, not a test folder — it is how the (recipe, runtime)
pair is kept honest.
---
## 19. Reference synthesis
| Source | Take | Constrain / reject |
|---|---|---|
| Cosmos3 (official + port) | Shared model instance across reasoning/diffusion/action/sound; packed multimodal sequences; component+scheduler parity matrices | A strong `ModelCard`, not the framework; no Cosmos-specific branching in global runtime |
| vLLM core | Running-first scheduling, reservation-before-admission, model-owned state, encoder/KV cache managers, CuMem sleep/wake, CUDA-graph dispatch, KV-connector split | Token scheduling is one WorkUnit kind; never full-graph compile |
| sglang `multimodal_gen` | Role pools, request lifecycle, capacity dispatch, transfer manifests, disagg state machine, cache-dit integration | No large mutable `Req`/`ForwardBatch` as the stable API; not single-item diffusion scheduling |
| vLLM-Omni | Frozen pipeline spec separate from deploy YAML (verified, adopt); `OmniConnectorBase` + `chunk_ready` readiness; **`SupportsStepExecution` as loop-inversion prior art** (opt-in, Qwen-Image-only — we generalize to always-on); TP-rank- and CFG-branch-aware KV-copy transfer; **`CFGParallelMixin` proves CFG-as-policy over one shared denoise body** (§5.3); 3 separate cache subsystems confirm per-class pools | Expresses shared-weight MoT (`bagel`/`lance`) only as **one opaque request-scheduled stage** the scheduler never sees inside — no step visibility, no cross-request batching by default; cross-stage KV is a *copy*, not a shared live cache; **no cost model** (per-stage count budgets); readiness-parking, not credit flow; RDMA = Mooncake/Mori/Yuanrong, not NIXL/NCCL |
| sglang-omni | The `next/wait_for/merge_fn/stream_to` edge vocabulary; Relay transport + **credit-based flow control** (this is sglang-omni's, not vllm-omni's) | Stages own disjoint weights; hybrid AR+diffusion only as AR-stage → DiT-stage; per-model bootstrap duplication |
| Dynamo | Fleet routing, disagg role pools, KV-aware routing, KVBM, SLA planner, ModelExpress cold-start/weight streaming | Orchestrates engines; never the engine core. Export a `DeploymentCard` + cost model to it |
| diffusers Modular | `ComponentSpec`/`modular_model_index.json` interchange; Guiders ≈ CFG policies | A Python pipeline interpreter is not the performance boundary; import is lossy |
| xDiT | DiT parallelism catalog (USP, ring/ulysses, PipeFusion, CFG-parallel, DistVAE) + world-size validation | Parallelism lives in the runtime + card, not a wrapper-per-model library; `pp_patch` invalid for causal |
| TorchTitan | Named mesh axes, `ParallelDims` validation, ModelSpec discipline, TorchStore weight-sync, batch-invariance utils | Adopt the discipline, not the stack; DCP/TorchStore don't reshard — `WeightSyncPlan` owns layout |
| verl-omni / miles / cosmos-rl | Rollout adapters, per-step capture, async rewards, group-relative advantage, TIS/MIS, deterministic/batch-invariant modes, per-payload `weight_version`, AIPO/off-policy masking | The two-runtime tax is the thing to delete; capture behavior *in* the serving loop, not after the fact |
| ComfyUI | Workflow graph, node-signature cache, model memory management, App-Mode (workflows-as-products) | Compile to `Program`; dynamic node execution is not the serving/training core; GPL hygiene |
| Dreamverse | Sessions, prompt memory, typed media IPC, cancellation, the duty-cycle capacity reality, the preference-data flywheel | Product/session behavior is first-class in the request plane, never merged into the model core |
| LiveKit | Realtime sessions, push audio/video, interruptions, turn/activity state, frame+PTS streaming | Realtime triggers only when they fire (<100ms interactive); don't force RTC onto offline jobs |
| Thinking Machines (batch-invariance) | The C2/C3 mechanism: batch-invariant kernels for bitwise rollout↔train identity | Scoped to goldens; the conservative baseline governs admission |
---
## Final position
```text
A model card is a (recipe, runtime) pair with a parity obligation.
The model owns loop semantics; the runtime owns loop lifecycle.
One resident instance runs many loops; one scheduler runs their steps in one currency.
Caches are correct by key; parity is correct by test; the interleave gate is non-negotiable.
Training records behavior on the same loops it serves.
Deployment places and routes; products stream artifacts; neither defines the model.
```
This is the ceiling: a model-native runtime where omni is native, train and serve are the same loops by construction,
correctness is a typed contract you can sign, and the (recipe, runtime) flywheel is real. The constraint we removed to
see it was migration. Putting that constraint back is the next document, not this one.
+1895
View File
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,52 @@
{
"data": [
{
"caption": "gold tip pyramid in the night, extremely detailed , rain, stars"
},
{
"caption": "a fruit stacking in the shape of a dog, stock image, shutterstock"
},
{
"caption": "Landscape, By Lee madgwick, by Luis Royo, by Louise nevelson"
},
{
"caption": "A colorful poster that says \"philo is a weird\""
},
{
"caption": "Danish male with blue eyes, realistic, viking"
},
{
"caption": "futuristic, cityscape, flying cars, neon lights, towering skyscrapers, glowing purple sky."
},
{
"caption": "a crow with cameras for eyes, sitting on a mans shoulder, anime, studio ghibli, fantasy, fairytale, sketch, digital art, watercolor, dnd, rustic, professional photograph, medieval, hd, 4k"
},
{
"caption": "a background image mixing the matrix and AI"
},
{
"caption": "Golden sunset, a bright orange and yellow sky is visible, lit up by the setting sun, the horizon is a mix of bright colors and deep shadows"
},
{
"caption": "Grim reaper playing an electric guitar"
},
{
"caption": "an epic view of a demonic Rose-ringed parakeet cyborg inside an ironmaiden robot,wearing a noble robe,large view,a surrealist painting, aralan bean and Philippe Druillet,hiromu arakawa,volumetric lighting,detailed shadows"
},
{
"caption": "Ben Shapiro as the cover of ministry's filth pig album, but covered in milk"
},
{
"caption": "king charles spaniel with , ethereal, midjourney style lighting and shadows, insanely detailed, 8k, photorealistic"
},
{
"caption": "A website for a party resort service"
},
{
"caption": "full shot of a steampunk horse"
},
{
"caption": "60s psycedelic spiritual jazz album art"
}
]
}
+1
View File
@@ -87,6 +87,7 @@ training:
# --- training.data [TYPED] -> DataConfig ---
data:
data_path: data/my_dataset # default: ""
preprocessed_data_type: t2v # default: "t2v" ("text_only" for simulate-only DMD text prompts)
train_batch_size: 1 # default: 1
dataloader_num_workers: 4 # default: 0
training_cfg_rate: 0.1 # default: 0.0
@@ -0,0 +1,116 @@
# DiffusionNFT multi-reward single-frame RL: Wan 2.1 T2V 1.3B on text-only PickScore prompts.
#
# Single-frame RL is represented as a one-latent-frame Wan run:
# num_latent_t: 1
# num_frames: 1
#
# The method trains the full transformer (no LoRA) and keeps an old-policy
# transformer plus a frozen reference transformer, matching the non-LoRA
# DiffusionNFT loss path.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
old:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
reference:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod
reward_fn:
pickscore: 1.0
clipscore: 1.0
sampling:
num_steps: 25
scheduler: flow_match_euler
trajectory: ode
flow_shift: inherit
validation:
every_steps: 10
num_steps: 40
num_prompts: 16
batch_size: 16
log_samples: true
seed: 42
# Null reuses training.data.data_path. Override this with a held-out
# preprocessed parquet path when one is available.
data_path:
# DiffusionNFT sd3_multi_reward on 4 GPUs resolves to per-GPU sample batch
# size 6, 48 sample batches per outer epoch, and grad accumulation 48.
sample_train_batch_size: 6
train_batch_size: 6
num_batches_per_epoch: 48
num_video_per_prompt: 24
num_inner_epochs: 1
timestep_fraction: 0.99
beta: 0.1
kl_beta: 0.0001
decay_type: 1
adv_mode: all
adv_clip_max: 5
max_grad_norm: 1.0
ema:
enabled: true
decay: 0.9
update_after_step: 0
validation: true
terminal_progress: true
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/pickscore_text_only_preprocessed
preprocessed_data_type: text_only
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 1
num_height: 448
num_width: 832
num_frames: 1
optimizer:
learning_rate: 3.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0001
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 100000
gradient_accumulation_steps: 48
checkpoint:
output_dir: outputs/wan2.1_diffusion_nft_pick_clip
training_state_checkpointing_steps: 30
checkpoints_total_limit: 3
tracker:
project_name: diffusion_nft_wan
run_name: wan2.1_diffusion_nft_pick_clip
model:
enable_gradient_checkpointing_type: full
pipeline:
flow_shift: 8
+29
View File
@@ -286,6 +286,35 @@ def run_train_framework_tests():
)
@app.function(gpu="L40S:1",
image=image,
timeout=1800,
secrets=[
modal.Secret.from_dict(
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
],
volumes={"/root/data": model_vol})
def seed_grad_norm_references():
"""Record the per-method grad-norm reference for the **CI GPU (L40S only)**.
Phase 2 / 5a-ii one-off seeding entrypoint. Pinned to ``gpu="L40S:1"`` (the
Modal CI runner), so this function only seeds the ``L40S`` key in
``fastvideo/tests/train/methods/grad_norm_refs.json``.
``FASTVIDEO_GRADNORM_UPDATE=1`` makes ``check_grad_norm_regression`` record
the measured norm instead of asserting; ``-rs`` surfaces the recorded value
in the log so it can be copied into the JSON.
To seed any other device (e.g. our local Blackwell dev box → ``GB200``
key), run the same env-var + pytest invocation directly on that
workstation — see the module docstring of ``grad_norm_regression.py`` for
the local command and the ``_DEVICE_MAPPINGS`` table.
"""
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
)
@app.function(gpu="L40S:1",
image=image,
timeout=3600,
@@ -67,6 +67,8 @@ class TestConstructor:
assert cb.num_frames is None
assert cb.sampling_timesteps is None
assert cb.output_dir is None
assert cb.offload_training_state is False
assert cb.unload_pipeline_after_validation is False
# Lazy fields not yet populated.
assert cb._pipeline is None
assert cb._sampling_param is None
@@ -83,12 +85,16 @@ class TestConstructor:
guidance_scale="4.5", # type: ignore[arg-type]
num_frames="77", # type: ignore[arg-type]
sampling_timesteps=["1000", "500"],
offload_training_state="1", # type: ignore[arg-type]
unload_pipeline_after_validation="false", # type: ignore[arg-type]
)
assert cb.every_steps == 50
assert cb.sampling_steps == [20, 40]
assert cb.guidance_scale == 4.5
assert cb.num_frames == 77
assert cb.sampling_timesteps == [1000, 500]
assert cb.offload_training_state is True
assert cb.unload_pipeline_after_validation is False
def test_pipeline_kwargs_collected(self) -> None:
cb = ValidationCallback(
@@ -0,0 +1,10 @@
{
"test_wan_causal_dfsft": {
"GB200": 2.9781,
"L40S": 3.2562
},
"test_wan_finetune": {
"GB200": 1.6486,
"L40S": 1.6467
}
}
@@ -0,0 +1,158 @@
# SPDX-License-Identifier: Apache-2.0
"""Layer-0 grad-norm regression for the per-method training smoke tests.
Phase 2 / 5a-ii: layers a device-keyed grad-norm check on top of the
finite/non-zero grad assertions established in 5a-i. After one
``single_train_step`` + ``backward``, the L2 norm of transformer block 0's
trainable gradients is compared against a reference value pinned per GPU in
``grad_norm_refs.json`` (next to this module).
Determinism: the harness seeds both the global RNG and the method's
``cuda_generator`` via ``method.on_train_start()`` (``training.data.seed`` in the
fixture), and the synthetic ``raw_batch`` is built *after* that call, so the
forward/backward is reproducible within bf16 reduction noise on a given GPU.
Why device-keyed: grad norms differ across GPU architectures (kernels,
accumulation order), so a single golden value can't cover every runner. The
JSON currently carries refs for the two GPUs we actually run on — ``L40S`` (CI)
and ``GB200`` (our Blackwell dev box; ``B200`` maps to the same key).
Seeding a reference for the current device:
- **CI / L40S** — invoke ``modal run`` against ``seed_grad_norm_references`` in
``fastvideo/tests/modal/pr_test.py`` (pinned to ``gpu="L40S:1"``), then copy
the recorded value from the log into ``grad_norm_refs.json``.
- **Local / non-L40S GPUs** — on that workstation::
FASTVIDEO_GRADNORM_UPDATE=1 \\
pytest fastvideo/tests/train/methods -vs -rs
The harness writes the measured norm into ``grad_norm_refs.json`` under the
device's key and skips the assertion for that run. Append a new substring
entry to ``_DEVICE_MAPPINGS`` first for any device not already listed.
"""
from __future__ import annotations
import json
import os
from pathlib import Path
import pytest
import torch
_REFS_PATH = Path(__file__).resolve().parent / "grad_norm_refs.json"
_UPDATE_ENV = "FASTVIDEO_GRADNORM_UPDATE"
# bf16 single-step smoke: catch gross breakage (wrong wiring, dead grads,
# scale regressions), not micro-drift from reduction nondeterminism.
_DEFAULT_RTOL = 0.10
# GPU-name substring -> reference key. First match wins. Only devices with
# seeded references in ``grad_norm_refs.json`` are listed here — to add a new
# GPU, append an entry, then seed the reference (see module docstring).
_DEVICE_MAPPINGS: tuple[tuple[str, str], ...] = (
("L40S", "L40S"),
("GB200", "GB200"),
("B200", "GB200"), # same Blackwell arch as GB200
)
def _device_name() -> str:
if not torch.cuda.is_available():
return "CPU"
return torch.cuda.get_device_name(0)
def resolve_device_key(device_name: str | None = None) -> str | None:
"""Map a CUDA device name to its reference key, or None if unsupported.
The substring match is case-insensitive so it survives driver/environment
differences in how ``torch.cuda.get_device_name`` capitalizes the model.
"""
name = device_name if device_name is not None else _device_name()
name_lower = name.lower()
for pattern, key in _DEVICE_MAPPINGS:
if pattern.lower() in name_lower:
return key
return None
def layer0_grad_norm(transformer) -> float:
"""Global L2 norm of transformer block 0's trainable gradients.
Block 0 is the reference surface 5a-i already isolates: its grad is the
*last* one produced during backprop, so a healthy value implies the whole
forward + chain-rule path is intact.
Accumulates the squared sums on the GPU and does a single CPU-GPU sync
(``.item()``) at the end, rather than one per parameter.
"""
blocks = getattr(transformer, "blocks", None)
assert blocks is not None and len(blocks) > 0, (
"transformer is expected to expose a non-empty ``.blocks``")
grads = [
p.grad for p in blocks[0].parameters()
if p.requires_grad and p.grad is not None
]
if not grads:
return 0.0
sq_sum = torch.zeros((), device=grads[0].device, dtype=torch.float32)
for g in grads:
sq_sum += g.detach().float().pow(2).sum()
return sq_sum.sqrt().item()
def _load_refs() -> dict[str, dict[str, float]]:
if _REFS_PATH.exists():
return json.loads(_REFS_PATH.read_text(encoding="utf-8"))
return {}
def _save_refs(refs: dict[str, dict[str, float]]) -> None:
_REFS_PATH.write_text(
json.dumps(refs, indent=2, sort_keys=True) + "\n",
encoding="utf-8")
def check_grad_norm_regression(
test_name: str,
transformer,
*,
rtol: float = _DEFAULT_RTOL,
) -> None:
"""Assert block-0 grad norm matches the device-keyed reference within rtol.
- Skips when the current GPU has no reference (unsupported device, or not
yet seeded) so a new runner never hard-fails before its golden exists.
- With ``FASTVIDEO_GRADNORM_UPDATE=1`` records/updates the reference for the
current device instead of asserting.
"""
norm = layer0_grad_norm(transformer)
device_key = resolve_device_key()
if os.environ.get(_UPDATE_ENV) == "1":
if device_key is None:
pytest.skip(
f"{_UPDATE_ENV}=1 but GPU '{_device_name()}' has no reference "
"key; add it to _DEVICE_MAPPINGS first")
refs = _load_refs()
refs.setdefault(test_name, {})[device_key] = round(norm, 4)
_save_refs(refs)
pytest.skip(
f"recorded grad-norm reference {test_name}[{device_key}] = "
f"{norm:.4f} (assertion skipped under {_UPDATE_ENV}=1)")
ref = _load_refs().get(test_name, {}).get(device_key) \
if device_key is not None else None
if ref is None:
pytest.skip(
f"no grad-norm reference for {test_name} on '{_device_name()}' "
f"(device_key={device_key}); run with {_UPDATE_ENV}=1 to seed it")
rel = abs(norm - ref) / (abs(ref) + 1e-12)
assert rel <= rtol, (
f"{test_name}[{device_key}] grad-norm regression: got {norm:.4f}, "
f"reference {ref:.4f}, relative error {rel:.3%} exceeds rtol "
f"{rtol:.0%}. If this is an intentional change, refresh the reference "
f"with {_UPDATE_ENV}=1 and explain why in the PR.")
@@ -28,6 +28,8 @@ from fastvideo.train.methods.fine_tuning.dfsft import (
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -122,3 +124,7 @@ def test_wan_causal_dfsft_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
@@ -36,6 +36,8 @@ from fastvideo.train.methods.fine_tuning.finetune import (
from fastvideo.train.models.wan import WanModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -139,3 +141,7 @@ def test_wan_finetune_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_finetune", model.transformer)
+186 -8
View File
@@ -68,6 +68,8 @@ class ValidationCallback(Callback):
num_frames: int | None = None,
output_dir: str | None = None,
sampling_timesteps: list[int] | None = None,
offload_training_state: bool = False,
unload_pipeline_after_validation: bool = False,
**pipeline_kwargs: Any,
) -> None:
self.pipeline_target = str(pipeline_target)
@@ -78,6 +80,8 @@ class ValidationCallback(Callback):
self.num_frames = (int(num_frames) if num_frames is not None else None)
self.output_dir = (str(output_dir) if output_dir is not None else None)
self.sampling_timesteps = ([int(s) for s in sampling_timesteps] if sampling_timesteps is not None else None)
self.offload_training_state = self._coerce_bool(offload_training_state)
self.unload_pipeline_after_validation = self._coerce_bool(unload_pipeline_after_validation)
self.pipeline_kwargs = dict(pipeline_kwargs)
# Set after on_train_start.
@@ -88,6 +92,12 @@ class ValidationCallback(Callback):
self.validation_random_generator: (torch.Generator | None) = None
self.seed: int = 0
@staticmethod
def _coerce_bool(value: Any) -> bool:
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
return bool(value)
# ----------------------------------------------------------
# Callback hooks
# ----------------------------------------------------------
@@ -140,16 +150,183 @@ class ValidationCallback(Callback):
) -> None:
transformer = method.student.transformer
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
try:
with self._validation_memory_context(
method,
validation_transformer=transformer,
):
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
method,
step,
t,
)
finally:
if self.unload_pipeline_after_validation:
self._clear_pipeline_cache()
@contextlib.contextmanager
def _validation_memory_context(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
):
if not self.offload_training_state:
yield
return
optimizer_tensor_records: list[tuple[Any, Any, torch.device]] = []
module_records: list[tuple[str, torch.nn.Module, torch.device]] = []
try:
self._offload_optimizer_states_to_cpu(
method,
step,
t,
optimizer_tensor_records,
)
self._offload_inactive_role_modules_to_cpu(
method,
validation_transformer=validation_transformer,
module_records=module_records,
)
self._empty_cuda_cache()
yield
finally:
self._restore_inactive_role_modules(module_records)
self._restore_optimizer_states(optimizer_tensor_records)
self._empty_cuda_cache()
def _offload_optimizer_states_to_cpu(
self,
method: TrainingMethod,
records: list[tuple[Any, Any, torch.device]],
) -> None:
optimizers = getattr(method, "_optimizer_dict", {})
if not optimizers:
return
moved = 0
for optimizer in optimizers.values():
state = getattr(optimizer, "state", None)
if not isinstance(state, dict):
continue
for param_state in state.values():
moved += self._offload_tensor_container_to_cpu(
param_state,
records,
)
if moved:
logger.info(
"Offloaded %d optimizer state tensors to CPU for validation.",
moved,
)
def _offload_tensor_container_to_cpu(
self,
obj: Any,
records: list[tuple[Any, Any, torch.device]],
) -> int:
moved = 0
if isinstance(obj, dict):
for key, value in list(obj.items()):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, key, value.device))
obj[key] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
if isinstance(obj, list):
for idx, value in enumerate(list(obj)):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, idx, value.device))
obj[idx] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
def _restore_optimizer_states(
self,
records: list[tuple[Any, Any, torch.device]],
) -> None:
for container, key, device in reversed(records):
value = container[key]
if torch.is_tensor(value):
container[key] = value.to(device=device)
if records:
logger.info(
"Restored %d optimizer state tensors after validation.",
len(records),
)
def _offload_inactive_role_modules_to_cpu(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
role_models = getattr(method, "_role_models", {})
if not isinstance(role_models, dict):
return
for role, model in role_models.items():
module = getattr(model, "transformer", None)
if not isinstance(module, torch.nn.Module):
continue
if module is validation_transformer:
continue
device = self._first_cuda_tensor_device(module)
if device is None:
continue
try:
module.to("cpu")
except Exception as exc:
logger.warning(
"Could not offload role %r transformer to CPU before validation: %s",
role,
exc,
)
continue
module_records.append((str(role), module, device))
logger.info(
"Offloaded role %r transformer from %s to CPU for validation.",
role,
device,
)
def _restore_inactive_role_modules(
self,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
for role, module, device in reversed(module_records):
module.to(device)
logger.info(
"Restored role %r transformer to %s after validation.",
role,
device,
)
@staticmethod
def _first_cuda_tensor_device(module: torch.nn.Module) -> torch.device | None:
for tensor in list(module.parameters(recurse=True)) + list(module.buffers(recurse=True)):
device = getattr(tensor, "device", None)
if isinstance(device, torch.device) and device.type == "cuda":
return device
return None
def _clear_pipeline_cache(self) -> None:
self._pipeline = None
self._pipeline_key = None
self._empty_cuda_cache()
@staticmethod
def _empty_cuda_cache() -> None:
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _find_ema_callback(self) -> Any | None:
"""Find the EMA callback in the callback dict."""
@@ -293,6 +470,7 @@ class ValidationCallback(Callback):
}
if flow_shift is not None:
kwargs["flow_shift"] = float(flow_shift)
kwargs.update(self.pipeline_kwargs)
self._pipeline = PipelineCls.from_pretrained(
tc.model_path,
+40
View File
@@ -197,6 +197,46 @@ class TrainingMethod(torch.nn.Module, ABC):
# -- Shared hooks (override in subclasses as needed) --
def manages_optimization(self) -> bool:
"""Whether the method owns backward/optimizer stepping internally.
Most methods return loss tensors and let :class:`Trainer` handle
gradient accumulation, callbacks, optimizer stepping, and scheduler
stepping. RL-style methods such as DiffusionNFT need to preserve their
own sample-then-inner-train loop, so they can provide a specific
``managed_train_step``.
"""
return False
def managed_train_step(
self,
data_stream: Any,
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
"""Run one method-managed step.
Subclasses that return ``True`` from :meth:`manages_optimization`
should override this. The fallback consumes one dataloader batch and
delegates to ``single_train_step`` so tests can exercise the hook with
tiny fake methods.
"""
return self.single_train_step(next(data_stream), iteration)
def on_validation_begin(self, iteration: int = 0) -> dict[str, LogScalar]:
"""Run method-owned validation, if any.
Pipeline-style validation should remain in callbacks. Methods that
intentionally avoid inference pipelines, such as RL methods with their
own sampler/reward loop, can override this hook and return metrics for
the trainer to log at ``iteration``.
"""
del iteration
return {}
def get_grad_clip_targets(
self,
iteration: int,
@@ -55,6 +55,8 @@ class DMD2Method(TrainingMethod):
raise ValueError("DMD2Method requires critic to be trainable")
self._cfg_uncond = self._parse_cfg_uncond()
self._rollout_mode = self._parse_rollout_mode()
self._validate_preprocessed_data_type()
self._configure_student_negative_conditioning()
self._denoising_step_list: torch.Tensor | None = (None)
# Initialize preprocessors on student.
@@ -206,6 +208,13 @@ class DMD2Method(TrainingMethod):
return targets
def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]:
"""Parse how DMD2 obtains the latent point used for rollout.
``simulate`` starts from fresh noise and lets the student create an
artificial latent trajectory, so it can run with text-only data.
``data_latent`` starts from preprocessed VAE latents and perturbs them
at a sampled denoising timestep.
"""
raw = self.method_config.get("rollout_mode", None)
if raw is None:
raise ValueError("method_config.rollout_mode must be set "
@@ -223,6 +232,34 @@ class DMD2Method(TrainingMethod):
"{simulate, data_latent}, got "
f"{raw!r}")
def _validate_preprocessed_data_type(self) -> None:
data_type = str(getattr(
self.training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
if data_type == "text_only" and self._rollout_mode != "simulate":
raise ValueError("training.data.preprocessed_data_type='text_only' "
"requires method.rollout_mode='simulate'; "
"data_latent rollout requires vae_latent data.")
def _uses_negative_prompt_conditioning(self) -> bool:
if self._cfg_uncond is None:
return True
text_policy = self._cfg_uncond.get("text", None)
if text_policy is None:
return True
return str(text_policy).strip().lower() == "negative_prompt"
def _configure_student_negative_conditioning(self) -> None:
setter = getattr(
self.student,
"set_requires_negative_conditioning",
None,
)
if setter is not None:
setter(self._uses_negative_prompt_conditioning())
def _parse_cfg_uncond(self, ) -> dict[str, Any] | None:
raw = self.method_config.get("cfg_uncond", None)
if raw is None:
+6
View File
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""RL training methods."""
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
__all__ = ["DiffusionNFTMethod"]
@@ -0,0 +1,30 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable RL training primitives."""
from fastvideo.train.methods.rl.common.sampling import (
DiffusionSampler,
SamplingConfig,
SamplingResult,
)
from fastvideo.train.methods.rl.common.prompt_sampling import (
KRepeatSample,
distributed_k_repeat_indices,
)
from fastvideo.train.methods.rl.common.validation import (
RLValidationConfig,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
__all__ = [
"DiffusionSampler",
"KRepeatSample",
"RLValidationConfig",
"SamplingConfig",
"SamplingResult",
"distributed_k_repeat_indices",
"media_to_video_array",
"validation_caption",
"validation_shard_indices",
]
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""Prompt-row sampling helpers for online RL methods.
This module chooses and repeats dataset prompt rows across ranks for RL training
batches. Here, "sampling" means selection, not generator sampling.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True, slots=True)
class KRepeatSample:
"""Local prompt indices for one distributed K-repeat sampling batch."""
local_indices: list[int]
unique_prompt_count: int
def distributed_k_repeat_indices(
*,
dataset_length: int,
batch_size: int,
repeats_per_prompt: int,
world_size: int,
rank: int,
seed: int,
) -> KRepeatSample:
"""Mirror DiffusionNFT's distributed K-repeat prompt sampler.
Adapted from DiffusionNFT's
``scripts/train_nft_sd3.py::DistributedKRepeatSampler``.
"""
dataset_length = int(dataset_length)
batch_size = int(batch_size)
repeats_per_prompt = int(repeats_per_prompt)
world_size = int(world_size)
rank = int(rank)
if dataset_length <= 0:
raise ValueError("dataset_length must be positive")
if batch_size <= 0:
raise ValueError("batch_size must be positive")
if repeats_per_prompt <= 0:
raise ValueError("repeats_per_prompt must be positive")
if world_size <= 0:
raise ValueError("world_size must be positive")
if rank < 0 or rank >= world_size:
raise ValueError(f"rank must be in [0, {world_size}), got {rank}")
total_samples = world_size * batch_size
if total_samples % repeats_per_prompt != 0:
raise ValueError("world_size * batch_size must be divisible by repeats_per_prompt "
f"({world_size} * {batch_size} vs {repeats_per_prompt})")
unique_prompt_count = total_samples // repeats_per_prompt
if unique_prompt_count > dataset_length:
raise ValueError("K-repeat sampling needs at least as many rows as unique prompts "
f"per sampling batch ({dataset_length} < {unique_prompt_count})")
generator = torch.Generator()
generator.manual_seed(int(seed))
indices = torch.randperm(dataset_length, generator=generator)[:unique_prompt_count].tolist()
repeated_indices = [idx for idx in indices for _ in range(repeats_per_prompt)]
shuffled_order = torch.randperm(len(repeated_indices), generator=generator).tolist()
shuffled_samples = [int(repeated_indices[idx]) for idx in shuffled_order]
start = rank * batch_size
end = start + batch_size
return KRepeatSample(
local_indices=shuffled_samples[start:end],
unique_prompt_count=unique_prompt_count,
)
@@ -0,0 +1,223 @@
# SPDX-License-Identifier: Apache-2.0
"""Configurable diffusion samplers for RL training methods."""
from __future__ import annotations
import copy
from dataclasses import dataclass
from typing import Any, Literal
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.base import ModelBase
SchedulerName = Literal["flow_match_euler", "model_default"]
TrajectoryName = Literal["ode", "sde_reflow"]
@dataclass(slots=True)
class SamplingConfig:
"""YAML-backed sampling knobs shared by RL methods."""
num_steps: int = 25
scheduler: SchedulerName = "model_default"
trajectory: TrajectoryName = "ode"
flow_shift: float | None = None
timesteps: list[float] | None = None
sigmas: list[float] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> SamplingConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.sampling must be a mapping, got {type(raw).__name__}")
supported_keys = {
"flow_shift",
"num_steps",
"scheduler",
"sigmas",
"timesteps",
"trajectory",
}
unsupported_keys = sorted(set(raw) - supported_keys)
if unsupported_keys:
raise ValueError(f"Unsupported method.sampling key(s): {unsupported_keys}. "
f"Supported keys: {sorted(supported_keys)}")
scheduler = str(raw.get("scheduler", "model_default") or "model_default").strip().lower()
if scheduler not in {"flow_match_euler", "model_default"}:
raise ValueError("method.sampling.scheduler must be one of "
"{flow_match_euler, model_default}, got "
f"{raw.get('scheduler')!r}")
trajectory = str(raw.get("trajectory", "ode") or "ode").strip().lower()
if trajectory not in {"ode", "sde_reflow"}:
raise ValueError("method.sampling.trajectory must be one of "
"{ode, sde_reflow}, got "
f"{raw.get('trajectory')!r}")
timesteps = raw.get("timesteps")
sigmas = raw.get("sigmas")
if timesteps is not None:
if not isinstance(timesteps, list) or not timesteps:
raise ValueError("method.sampling.timesteps must be a non-empty list when set")
timesteps = [float(t) for t in timesteps]
if sigmas is not None:
if not isinstance(sigmas, list) or not sigmas:
raise ValueError("method.sampling.sigmas must be a non-empty list when set")
sigmas = [float(s) for s in sigmas]
if timesteps is not None and sigmas is not None and len(timesteps) != len(sigmas):
raise ValueError("method.sampling.timesteps and method.sampling.sigmas must have the same length")
num_steps = int(raw.get("num_steps", 25) or 25)
if num_steps <= 0:
raise ValueError("method.sampling.num_steps must be positive")
return cls(
num_steps=num_steps,
scheduler=scheduler, # type: ignore[arg-type]
trajectory=trajectory, # type: ignore[arg-type]
flow_shift=(None if raw.get("flow_shift", None) in (None, "inherit") else float(raw["flow_shift"])),
timesteps=timesteps,
sigmas=sigmas,
)
@dataclass(slots=True)
class SamplingResult:
latents: torch.Tensor
timesteps: torch.Tensor
sigmas: torch.Tensor
class DiffusionSampler:
"""Thin model/scheduler sampler used by RL methods.
This intentionally does not call FastVideo's full inference pipelines.
RL training needs a reusable sampling primitive that works with
``ModelBase`` wrappers and scheduler math without binding a method to
model-family pipeline classes such as ``WanDMDPipeline``.
"""
def __init__(self, config: SamplingConfig) -> None:
self.config = config
@torch.no_grad()
def sample(
self,
model: ModelBase,
batch: TrainingBatch,
*,
generator: torch.Generator | None,
) -> SamplingResult:
latents = batch.latents
if latents is None:
raise RuntimeError("TrainingBatch.latents is required for RL sampling")
current = torch.randn(
latents.shape,
device=latents.device,
dtype=latents.dtype,
generator=generator,
)
scheduler = self._prepare_scheduler(model, current.device)
timesteps = scheduler.timesteps.to(device=current.device)
sigmas = scheduler.sigmas.to(device=current.device)
original_timesteps = batch.timesteps
try:
if self.config.trajectory == "ode":
pred_clean = current
for timestep in timesteps:
model_timestep = self._model_timestep(timestep, current)
batch.timesteps = model_timestep
pred_noise = model.predict_noise(
current,
model_timestep,
batch,
conditional=True,
attn_kind="dense",
)
current = scheduler.step(
pred_noise.flatten(0, 1),
timestep,
current.flatten(0, 1),
return_dict=False,
)[0].unflatten(0, pred_noise.shape[:2])
pred_clean = current
return SamplingResult(latents=pred_clean, timesteps=timesteps, sigmas=sigmas)
return SamplingResult(
latents=self._sample_sde_reflow(
model,
batch,
current,
timesteps,
generator=generator,
),
timesteps=timesteps,
sigmas=sigmas,
)
finally:
batch.timesteps = original_timesteps
def _prepare_scheduler(
self,
model: ModelBase,
device: torch.device,
) -> Any:
if self.config.scheduler == "flow_match_euler":
shift = self.config.flow_shift
if shift is None:
shift = float(getattr(model.noise_scheduler, "shift", 1.0))
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(shift))
else:
scheduler = copy.deepcopy(model.noise_scheduler)
kwargs: dict[str, Any] = {"device": device}
if self.config.timesteps is not None:
kwargs["timesteps"] = self.config.timesteps
kwargs["num_inference_steps"] = len(self.config.timesteps)
if self.config.sigmas is not None:
kwargs["sigmas"] = self.config.sigmas
kwargs["num_inference_steps"] = len(self.config.sigmas)
if "num_inference_steps" not in kwargs:
kwargs["num_inference_steps"] = self.config.num_steps
scheduler.set_timesteps(**kwargs)
return scheduler
def _sample_sde_reflow(
self,
model: ModelBase,
batch: TrainingBatch,
current: torch.Tensor,
timesteps: torch.Tensor,
*,
generator: torch.Generator | None,
) -> torch.Tensor:
pred_clean = current
for step_idx, timestep in enumerate(timesteps):
timestep_tensor = self._model_timestep(timestep, current)
batch.timesteps = timestep_tensor
pred_clean = model.predict_x0(
current,
timestep_tensor,
batch,
conditional=True,
attn_kind="dense",
)
if step_idx < len(timesteps) - 1:
next_timestep = timesteps[step_idx + 1].reshape(1).to(device=current.device)
noise = torch.randn(
pred_clean.shape,
device=pred_clean.device,
dtype=pred_clean.dtype,
generator=generator,
)
current = model.add_noise(pred_clean, noise, next_timestep)
return pred_clean
@staticmethod
def _model_timestep(
timestep: torch.Tensor,
current: torch.Tensor,
) -> torch.Tensor:
return timestep.reshape(1).to(device=current.device).expand(current.shape[0]).contiguous()
@@ -0,0 +1,81 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared validation helpers for RL training methods."""
from __future__ import annotations
from dataclasses import dataclass
import math
from typing import Any
import torch
@dataclass(slots=True)
class RLValidationConfig:
every_steps: int = 0
num_steps: int = 40 # Reference DiffusionNFT sampling num steps for best visual quality
num_prompts: int = 16
batch_size: int = 16
log_samples: bool = True
seed: int = 42
data_path: str | None = None
sampling: dict[str, Any] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> RLValidationConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.validation must be a mapping, got {type(raw).__name__}")
data_path = raw.get("data_path", None)
sampling = raw.get("sampling", None)
if sampling is not None and not isinstance(sampling, dict):
raise ValueError(f"method.validation.sampling must be a mapping, got {type(sampling).__name__}")
return cls(
every_steps=max(0, int(raw.get("every_steps", 0) or 0)),
num_steps=max(1, int(raw.get("num_steps", 40) or 40)),
num_prompts=max(1, int(raw.get("num_prompts", 16) or 16)),
batch_size=max(1, int(raw.get("batch_size", 16) or 16)),
log_samples=bool(raw.get("log_samples", True)),
seed=int(raw.get("seed", 42) or 42),
data_path=(None if data_path in (None, "") else str(data_path)),
sampling=(dict(sampling) if sampling is not None else None),
)
def validation_shard_indices(
num_prompts: int,
*,
rank: int,
world_size: int,
) -> list[tuple[int, bool]]:
"""Return fixed validation prompt indices for one distributed rank."""
num_prompts = max(1, int(num_prompts))
world_size = max(1, int(world_size))
per_rank = int(math.ceil(num_prompts / world_size))
padded_total = per_rank * world_size
return [((idx % num_prompts), idx < num_prompts) for idx in range(rank, padded_total, world_size)]
def validation_caption(
prompt: str,
rewards: dict[str, float],
) -> str:
reward_parts = [f"{key}: {float(rewards[key]):.4f}" for key in sorted(rewards)]
return f"{' | '.join(reward_parts)} | {prompt[:1000]}"
def media_to_video_array(media: torch.Tensor) -> Any:
"""Convert decoded media to a tracker video array.
Accepts ``[C, T, H, W]`` tensors. ``[C, H, W]`` tensors are treated as
``T=1`` media. Output follows the existing tracker convention used
elsewhere in FastVideo: ``[T, C, H, W]`` uint8.
"""
if media.ndim == 3:
media = media.unsqueeze(1)
if media.ndim != 4:
raise ValueError("media must have shape [C, T, H, W] or [C, H, W], "
f"got {tuple(media.shape)}")
video = (media.detach().float().clamp(0, 1) * 255).round().to(torch.uint8)
return video.permute(1, 0, 2, 3).contiguous().cpu().numpy()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,37 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable reward models for training methods."""
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
select_first_frame,
)
def build_multi_reward_scorer(
reward_weights,
*,
device="cuda",
scorers: dict[str, RewardScorer] | None = None,
) -> MultiRewardScorer:
available: dict[str, RewardScorer] = dict(scorers or {})
if not available:
available = {
"pickscore": PickScoreScorer(device=device),
"clipscore": ClipScoreScorer(device=device),
}
return MultiRewardScorer(reward_weights, scorers=available)
__all__ = [
"ClipScoreScorer",
"MultiRewardScorer",
"PickScoreScorer",
"RewardScorer",
"build_multi_reward_scorer",
"select_first_frame",
]
@@ -0,0 +1,130 @@
# SPDX-License-Identifier: Apache-2.0
"""Frame-based reward scorers used by RL training methods."""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from PIL import Image
import torch
from fastvideo.train.methods.rl.rewards.media import select_first_frame
class PickScoreScorer(torch.nn.Module):
"""PickScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/pickscore_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
from transformers import AutoModel, AutoProcessor
processor_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
model_path = "yuvalkirstain/PickScore_v1"
self.device = torch.device(device)
self.dtype = dtype
self.processor = AutoProcessor.from_pretrained(processor_path)
self.model = AutoModel.from_pretrained(model_path).eval().to(self.device)
self.model = self.model.to(dtype=dtype)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media)
frame_np = (frame_tensor.detach().float().clamp(0, 1) * 255).round()
frame_np = frame_np.to(torch.uint8).cpu().numpy().transpose(0, 2, 3, 1)
pil_frames = [Image.fromarray(frame) for frame in frame_np]
frame_inputs = self.processor(
images=pil_frames,
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
frame_inputs = {k: v.to(device=self.device) for k, v in frame_inputs.items()}
text_inputs = self.processor(
text=list(prompts),
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
text_inputs = {k: v.to(device=self.device) for k, v in text_inputs.items()}
text_embs = self.model.get_text_features(**text_inputs)
text_embs = text_embs / text_embs.norm(p=2, dim=-1, keepdim=True)
frame_embs = self.model.get_image_features(**frame_inputs)
frame_embs = frame_embs / frame_embs.norm(p=2, dim=-1, keepdim=True)
scores = self.model.logit_scale.exp() * (text_embs @ frame_embs.T)
return scores.diag().float() / 26.0
class ClipScoreScorer(torch.nn.Module):
"""CLIPScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/clip_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
) -> None:
super().__init__()
import torch.nn as nn
import torchvision.transforms as T
from transformers import CLIPModel, CLIPProcessor
def get_size(size: Any) -> Any:
if isinstance(size, int):
return (size, size)
if isinstance(size, Mapping) and "height" in size and "width" in size:
return (size["height"], size["width"])
if isinstance(size, Mapping) and "shortest_edge" in size:
return size["shortest_edge"]
raise ValueError(f"Invalid processor size: {size!r}")
def get_frame_transform(processor: Any) -> torch.nn.Module:
config = processor.to_dict()
resize = T.Resize(get_size(config.get("size"))) if config.get("do_resize") else nn.Identity()
crop = T.CenterCrop(get_size(config.get("crop_size"))) if config.get("do_center_crop") else nn.Identity()
normalize = (T.Normalize(mean=processor.image_mean, std=processor.image_std)
if config.get("do_normalize") else nn.Identity())
return T.Compose([resize, crop, normalize])
self.device = torch.device(device)
self.model = CLIPModel.from_pretrained("openai/clip-vit-large-patch14").to(self.device).eval()
self.processor = CLIPProcessor.from_pretrained("openai/clip-vit-large-patch14")
self.transform = get_frame_transform(self.processor.image_processor)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media).detach().float().clamp(0, 1)
texts = self.processor(
text=list(prompts),
padding="max_length",
truncation=True,
return_tensors="pt",
).to(self.device)
pixels = self.transform(frame_tensor).to(device=self.device, dtype=frame_tensor.dtype)
outputs = self.model(pixel_values=pixels, **texts)
return outputs.logits_per_image.diagonal().float() / 100.0
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Generic media reward composition utilities."""
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
import torch
RewardScorer = Callable[[torch.Tensor, Sequence[str]], torch.Tensor]
def select_first_frame(media: torch.Tensor) -> torch.Tensor:
"""Return first-frame media as ``[B, C, H, W]``.
This is a helper for reward models that are intrinsically frame-based
(for example PickScore and CLIPScore). Video-aware rewards should inspect
the full ``[B, C, T, H, W]`` tensor themselves.
"""
if not torch.is_tensor(media):
raise TypeError(f"media must be a torch.Tensor, got {type(media).__name__}")
if media.ndim == 5:
return media[:, :, 0]
if media.ndim == 4:
return media
raise ValueError("media must have shape [B, C, H, W] or [B, C, T, H, W], "
f"got {tuple(media.shape)}")
class MultiRewardScorer:
"""Weighted sum of reusable media reward scorers.
Mirrors DiffusionNFT's ``flow_grpo/rewards.py::multi_score`` behavior,
while leaving frame selection to each concrete reward.
"""
def __init__(
self,
reward_weights: Mapping[str, float],
*,
scorers: Mapping[str, RewardScorer],
) -> None:
self.reward_weights = {str(k): float(v) for k, v in reward_weights.items()}
if not self.reward_weights:
raise ValueError("reward_weights must contain at least one reward")
self.scorers = dict(scorers)
unsupported = sorted(set(self.reward_weights) - set(self.scorers))
if unsupported:
raise ValueError(f"Unsupported reward(s): {unsupported}. "
f"Available rewards: {sorted(self.scorers)}")
@torch.no_grad()
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> dict[str, torch.Tensor]:
prompt_count = len(prompts)
if media.shape[0] != prompt_count:
raise ValueError(f"media batch size ({media.shape[0]}) must match prompt count ({prompt_count})")
total: torch.Tensor | None = None
details: dict[str, torch.Tensor] = {}
for name, weight in self.reward_weights.items():
scores = self.scorers[name](media, prompts).detach().float()
if scores.ndim != 1 or int(scores.shape[0]) != prompt_count:
raise ValueError(f"Reward {name!r} must return shape [{prompt_count}], got {tuple(scores.shape)}")
details[name] = scores
weighted = scores * float(weight)
total = weighted if total is None else total.to(weighted.device) + weighted
assert total is not None
details["avg"] = total
return details
+11
View File
@@ -91,6 +91,17 @@ class ModelBase(ABC):
def on_train_start(self) -> None: # noqa: B027
"""Called once before the training loop begins."""
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
"""Decode ``[B, T, C, H, W]`` latents to ``[B, C, T, H, W]`` media.
RL reward methods call this hook instead of reaching into
model-specific VAE normalization details.
"""
raise NotImplementedError(f"{type(self).__name__} does not implement decode_latents()")
# ------------------------------------------------------------------
# Timestep helpers
# ------------------------------------------------------------------
+47 -5
View File
@@ -101,6 +101,7 @@ class WanModel(ModelBase):
self.negative_prompt_embeds: (torch.Tensor | None) = None
self.negative_prompt_attention_mask: (torch.Tensor | None) = None
self._requires_negative_conditioning = True
# Timestep mechanics.
self.timestep_shift: float = float(flow_shift)
@@ -160,17 +161,31 @@ class WanModel(ModelBase):
self._init_timestep_mechanics()
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_t2v, )
pyarrow_schema_t2v,
pyarrow_schema_text_only,
)
from fastvideo.train.utils.dataloader import (
build_parquet_t2v_train_dataloader, )
preprocessed_data_type = str(getattr(
training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
parquet_schema = pyarrow_schema_t2v
if preprocessed_data_type == "text_only":
parquet_schema = pyarrow_schema_text_only
elif preprocessed_data_type != "t2v":
raise ValueError("Unsupported Wan preprocessed_data_type: "
f"{preprocessed_data_type!r}")
text_len = (
training_config.pipeline_config.text_encoder_configs[ # type: ignore[union-attr]
0].arch_config.text_len)
self.dataloader = build_parquet_t2v_train_dataloader(
training_config.data,
text_len=int(text_len),
parquet_schema=pyarrow_schema_t2v,
parquet_schema=parquet_schema,
)
self.start_step = 0
@@ -178,6 +193,9 @@ class WanModel(ModelBase):
def num_train_timesteps(self) -> int:
return int(self.num_train_timestep)
def set_requires_negative_conditioning(self, requires: bool) -> None:
self._requires_negative_conditioning = bool(requires)
def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
timestep = shift_timestep(
timestep,
@@ -187,7 +205,25 @@ class WanModel(ModelBase):
return timestep.clamp(self.min_timestep, self.max_timestep)
def on_train_start(self) -> None:
self.ensure_negative_conditioning()
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
@torch.no_grad()
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
if self.vae is None:
raise RuntimeError("Wan VAE is not initialized")
latents = latents_b_t_c_h_w.permute(0, 2, 1, 3, 4).float()
if bool(getattr(self.vae, "handles_latent_denorm", False)):
denorm = latents
else:
mean = torch.tensor(self.vae.latents_mean, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
std = torch.tensor(self.vae.latents_std, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
denorm = latents * std + mean
media = self.vae.to(latents.device).decode(denorm)
return (media / 2 + 0.5).clamp(0, 1)
# ------------------------------------------------------------------
# Runtime primitives
@@ -200,7 +236,8 @@ class WanModel(ModelBase):
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
self.ensure_negative_conditioning()
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
assert self.training_config is not None
tc = self.training_config
@@ -285,7 +322,7 @@ class WanModel(ModelBase):
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = noisy_latents.dtype
dtype = self._get_training_dtype()
if conditional:
text_dict = batch.conditional_dict
if text_dict is None:
@@ -301,6 +338,11 @@ class WanModel(ModelBase):
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
# Keep Wan training autocast tied to the model's training dtype, not
# to caller-created intermediates that may accidentally be fp32.
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
+3 -1
View File
@@ -128,7 +128,9 @@ class WanCausalModel(WanModel, CausalModelBase):
}
device_type = self.device.type
dtype = noisy_latents.dtype
dtype = self._get_training_dtype()
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
if conditional:
text_dict = batch.conditional_dict
+67 -25
View File
@@ -12,7 +12,7 @@ from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
if TYPE_CHECKING:
@@ -82,6 +82,22 @@ class Trainer:
batch = next(data_iter)
yield batch
def _run_method_validation(
self,
method: TrainingMethod,
iteration: int,
) -> None:
hook = getattr(method, "on_validation_begin", None)
if hook is None:
return
validation_metrics: dict[str, LogScalar] = hook(iteration)
validation_metrics = {
k: float(_coerce_log_scalar(v, where=(f"method.on_validation_begin().metrics[{k!r}]")))
for k, v in validation_metrics.items()
}
if self.global_rank == 0 and validation_metrics:
self.tracker.log(validation_metrics, iteration)
def run(
self,
method: TrainingMethod,
@@ -115,6 +131,7 @@ class Trainer:
method,
iteration=start_step,
)
self._run_method_validation(method, start_step)
method.optimizers_zero_grad(start_step)
data_stream = self._iter_dataloader(dataloader)
@@ -130,6 +147,8 @@ class Trainer:
desc="Steps",
disable=self.local_rank > 0,
)
# Allow method-specific optimization flow (e.g. DiffusionNFT).
method_manages_optimization = bool(method.manages_optimization())
for step in progress:
t0 = time.perf_counter()
@@ -137,47 +156,69 @@ class Trainer:
# to CPU once per step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
if method_manages_optimization:
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
))
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.single_train_step()"
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors
# to float right before logging.
metrics = {k: float(v) / grad_accum for k, v in loss_sums.items()}
metrics.update({k: float(v) / grad_accum for k, v in metric_sums.items()})
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
@@ -196,6 +237,7 @@ class Trainer:
method,
iteration=step,
)
self._run_method_validation(method, step)
self.callbacks.on_validation_end(
method,
iteration=step,
+7
View File
@@ -343,6 +343,12 @@ def _build_training_config(
if init_from is not None:
model_path = str(init_from)
preprocessed_data_type = str(da.get("preprocessed_data_type", "t2v") or "t2v").strip().lower()
if preprocessed_data_type not in {"t2v", "text_only"}:
raise ValueError("training.data.preprocessed_data_type must be one of "
"{'t2v', 'text_only'}, got "
f"{preprocessed_data_type!r}")
return TrainingConfig(
distributed=DistributedConfig(
num_gpus=num_gpus,
@@ -354,6 +360,7 @@ def _build_training_config(
),
data=DataConfig(
data_path=str(da.get("data_path", "") or ""),
preprocessed_data_type=preprocessed_data_type,
train_batch_size=int(da.get("train_batch_size", 1) or 1),
dataloader_num_workers=int(da.get("dataloader_num_workers", 0) or 0),
training_cfg_rate=float(da.get("training_cfg_rate", 0.0) or 0.0),
+1
View File
@@ -23,6 +23,7 @@ class DistributedConfig:
@dataclass(slots=True)
class DataConfig:
data_path: str = ""
preprocessed_data_type: str = "t2v"
train_batch_size: int = 1
dataloader_num_workers: int = 0
training_cfg_rate: float = 0.0
+1 -1
View File
@@ -112,7 +112,7 @@ class BaseTracker:
self._timed_metrics = {}
def log_artifacts(self, artifacts: dict[str, Any], step: int) -> None:
"""Log artifacts such as videos or images.
"""Log tracker artifacts such as sampled media.
By default this is treated the same as :meth:`log`.
"""
+2 -2
View File
@@ -1694,7 +1694,7 @@ class EMA_FSDP:
if p_local.numel() == 0:
# Nothing to swap on this rank for this param
continue
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
self.saved[name] = p_local.clone().to("cpu")
if name in self.ema.shadow:
ema_cpu = self.ema.shadow[name]
if ema_cpu.numel() != p_local.numel():
@@ -1714,7 +1714,7 @@ class EMA_FSDP:
saved_local = self.saved[name]
if saved_local.numel() != p_local.numel():
continue
p_local.copy_(saved_local)
p_local.copy_(saved_local.to(dtype=p_local.dtype, device=p_local.device))
self.saved.clear()
return False
+225
View File
@@ -0,0 +1,225 @@
# FastVideo Runtime — Aggressive Implementation Plan
**Companion to** `design.md` (v19) and `design_summary.md` · **Stance:** this plan trades interface stability for
speed. Where it deviates from design.md's conservative migration (§10), the deviation is flagged with **⚡**.
design.md remains the architectural authority; this is the execution order.
---
## 1. Rules of engagement
**We break, freely and early:**
- The public Python API: `generate_video(**kwargs)` and `SamplingParam` are **deleted**, not deprecated.
- Config schemas: `FastVideoArgs` (1,272 lines, 81 fields) stops being a public or threaded surface.
- `fastvideo/api/compat.py` (651 lines): **deleted in M1** ⚡ (design.md §6.6 shrinks it monotonically to Phase 5 —
that policy existed only to honor signatures we are now licensed to break).
- CLI flags, YAML schemas, package layout, `fastvideo.api` exports, ComfyUI node params, every example.
- In-repo dependents (`apps/dreamverse`, `comfyui/`, `examples/`, `scripts/`) get **fixed in the same PR train** —
we own them; no deprecation period, no shims.
**We never break, at any speed:**
- **Numerics.** Bit-identical loop parity and SSIM gates are not "conservative" — they are the definition of
correct. Aggression applies to interfaces, never to outputs.
- **Model coverage** for the families that matter (tier list in §6 — the tail is a decision, not a casualty).
- The frozen legacy `fastvideo/training/` stack (N2) and the bit-exact porting methodology (N3).
- External users get **batched breakage**: all user-visible breaks land in at most two releases (R1 = request/config
cut, R2 = engine default), each with a migration guide and a `fastvideo migrate` codemod — never a drip.
## 2. The sequencing argument (answering "fix omni request first, then separate the planes?")
**Yes to the first half. The second half should not be a project.** The three planes are not separated by moving
code into plane-named directories — today's monolithic stages would just get reshuffled and then rewritten when
loops invert. The planes are *born* from two cuts, and a third that is really a config change:
1. **The request-plane cut (M1)** — `OmniRequest` becomes the only currency crossing the boundary. Everything
behind it is implementation. This is your "fix the omni input and request first," and it goes first because it
is low-risk, it defines the vocabulary every later stage consumes, and it gets the user-facing pain over with
while the codebase is still familiar.
2. **Loop inversion (M2)** — this *is* the pipeline/execution plane separation. Once families expose
`init/step/finalize` step bodies, something other than the family must own iteration; that owner is the
executor, and the execution plane exists by construction. Before inversion there is nothing for an execution
plane to schedule — "separating" it would be an empty directory.
3. **The config cut (inside M1)** — the real coupling between planes today is `FastVideoArgs`: one 81-field object
threaded through entrypoints, pipelines, stages, and executors, mixing deploy-time, model-time, and
request-time concerns. Splitting it into `DeployConfig` / `ModelSpec` / `OmniRequest` (design.md §6.6's four
layers) is the single highest-leverage "separation" action, and it's schema work, not architecture work.
So the order is: **M1 request+config cut → M2 loop inversion (planes now exist) → M3 engine on top.** Plane
separation is the *outcome* of M1+M2, not a milestone.
## 3. Milestones
Timeline assumes 3–4 engineers on the runtime critical path. Overlap is deliberate; gates are not ⚡-able.
### M0 — Baselines, harness, enforcement (weeks 0–3, overlaps M1)
The license for everything aggressive afterward. Not skippable, not shrinkable. M0 does not block M1 (which
changes no numerics) — the only hard rule is **no family's M2 migration starts before its baseline exists**.
- Merge the `feat/cosmos3-reasoning` chain (design.md sizes this alone at 2–3 engineer-months — it runs as its own
track); seed SSIM references for the ~7 uncovered families.
- ParityAligner v0: record/compare named taps on *current* pipelines (it must exist before anything changes).
- **The enforcement package, on day one** (design.md §10 — the prior freeze was broken 19× for lack of exactly
this): CI path gates (reject new `fastvideo/training/` files now; reject new `DenoisingStage` subclasses once the
first M2 family lands), CODEOWNERS on the frozen and migrating paths, a named owner per milestone, and the
inflow rule — new model families land on the new abstractions from the first Wan/Flux2 landing onward.
- Announce the M1 freeze window for in-flight PRs touching `fastvideo/api/`, `fastvideo_args.py`, entrypoints.
*Gate: every tier-A family has a recorded SSIM + activation baseline; CI gates live.*
### M1 — The request-plane cut (weeks 1–4) → **breaking release R1**
The typed API is partway there: `VideoGenerator.generate(GenerationRequest)` is already the documented primary
entrypoint (`generate_video` carries a deprecation warning), and `fastvideo/entrypoints/openai/` already serves
`POST /v1/videos` and `POST /v1/images`. But the legacy path is still what's *used*: Dreamverse calls
`generate_video(**kwargs)` (`apps/dreamverse/dreamverse/video_generation.py:508`), as do ComfyUI and most
examples. M1 finishes the cut instead of bridging it:
- **`OmniRequest` / `OmniOutput` / `OmniEvent`**: evolve `api/schema.py`'s `GenerationRequest` in place — typed
modality parts, `TaskType`, per-model `ModelOptions` registered blocks (formalizing the `api/matrixgame2.py`
pattern), seeds/priority/streaming flags. `api/results.py`'s `Video*Event` types become `OmniEvent`.
- **Config: four layers, one owner each** (§6.6): extract `DeployConfig` (placement, parallelism axes, memory/
offload, compile, plugins) from the runtime third of `FastVideoArgs` + `EngineConfig`/`ParallelismConfig`;
`ModelSpec` manifest v0 (manifest-first component resolution; today's name-detectors as fallback);
`OmniRequest` absorbs every per-call field. CLI flags, OpenAI protocol models, and presets are **generated**
from the schema.
- **Delete** ⚡: `compat.py` (651), `sampling_param.py` (411), `generate_video()`, the `fastvideo.api` legacy
exports, `FastVideoArgs` as a *public* type. Internally it survives as a boundary-constructed shim for as long
as anything still receives it: migrated families drop it per-family in M2, but unmigrated tier-B stages
(`LegacyPipelineNode`) and the frozen `training/` stack (whose `TrainingArgs` subclasses it) carry it until M6 —
it dies as a type with the tail, not before.
- **Fix in-train**: ComfyUI nodes (legacy-API callers), all `examples/` (~75 files, mostly mechanical),
`scripts/`, docs. Ship `fastvideo migrate` (codemod: old kwargs/YAML → `OmniRequest`/`DeployConfig`).
- Internals unchanged: `ForwardBatch` is built *from* `OmniRequest` at the boundary; the executor and stages are
untouched in M1.
*Gate: all SSIM suites unchanged; Dreamverse + ComfyUI + examples green on the new surface; R1 notes + codemod
published.*
### M2 — Loop inversion (weeks 4–10): the pipeline plane is born
- `DenoiseLoop` / `ARDecodeLoop` with `init/step/finalize`; runtime owns iteration; custom-step escape hatch from
day one (the Cosmos3-port and self-forcing pattern is legitimate, §6.2.3).
- **Family order** (each lands step body + policies, and **deletes its legacy stage code in the same PR** ⚡ —
continuous deletion, no end-of-plan cliff): **Wan 2.1/2.2 + Flux2 first, jointly** — design.md's rationale
stands: together they exercise CFG variants, expert routing, chunk-KV, and the image path, so the step-body
contract freezes only after all four are exercised → Wan-causal (self-forcing student) → LTX-2 → HunyuanVideo →
Stable Audio → remaining image families. Unmigrated families keep running via `LegacyPipelineNode`.
- Policies: `CFGPolicy` (absorbs the 3 CFG copies), `AttnMetadataProvider`, `FlowShiftPolicy`, `PrecisionPolicy`.
- Extension core lands with the loop (it's why the loop is being rebuilt): observer bus, ParityAligner promoted to
per-request observer, Profiler/NaNWatch, and **cache-dit as the first interceptor** (retiring `enable_teacache`).
- `forward_context.py` off the *migrated* inference path (194 references across ~68 files today: ~8 importer files
in frozen `training/`, the rest spread across train/ models, tier-B inference stages, quantization, and tests);
the module survives as a shim for frozen `training/` **and unmigrated tier-B stages** until M6 — what M2
guarantees is that no migrated family and no new code touches it.
- **`train/` migrates per-family, immediately behind inference**: DMD2 and the landed DiffusionNFT (#1450) adopt
the shared step functions as each family's body lands — `rl/common/sampling.py`'s loop is deleted, #1396
grad-norm refs extended to the migrated methods (RL included).
*Gate, per family: old-vs-new loop bit-identical (ParityAligner) + SSIM + a recorded loop-overhead / batch-of-1
latency measurement (the baseline M3 gates against); for train/: seeded rollout latents identical, reward metrics
- grad-norms neutral. No family is ever dual-maintained.*
### M3 — Execution plane: engine + scheduler (weeks 8–14, overlaps M2) → **breaking release R2**
- `AsyncEngine` (queue, admission, cancellation-as-common-path, failure isolation); offline `VideoGenerator` keeps
its name, becomes a thin sync wrapper that can bypass the queue.
- `StepScheduler` v0: multiplexes denoise steps across requests in a pool; budget currency = **predicted GPU-time**
from a calibrated per-(model, phase, shape) cost table (the cost *model* matures later; the currency is right
from day one). Carries the `ARDecodeLoop` contract; AR batching itself waits for its workload (N5).
- CacheManager v0: per-request chunk-KV slabs behind `KVHandle`; CFG-parallel axis (2-branch in practice).
- **Dynamo stock worker** (registration, health/drain, cost metrics), retiring the locked
`dynamo/examples/diffusers/worker.py` pattern.
- **Dreamverse hard-cut** (per design.md Phase 2; the aggressive delta is doing it in one PR): `gpu_pool.py`,
queue, warmup, and stream relay deleted and replaced by engine-client calls; the duty-cycle concurrency study
runs on the result.
- Colocated weight-sync RPC + component-granular sleep/wake + `RolloutClient` (engine-client RL mode for #1450).
*Gate: serving load tests; batch-of-1 latency regression ≤ 2% vs the M2-recorded measurement; Dreamverse
single-session parity; RL engine-client seeded final-latent parity vs in-process; deploys under stock Dynamo.*
### M4 — Graphs, parallelism, multi-session (weeks 14–20)
- `PipelineSpec` graph IR: per-family pipeline classes shrink to **spec + step body + policies**
(`create_pipeline_stages()` retires); LTX-2 and Hunyuan15+SR land as real fan-out graphs.
- Role pools + connectors (port `multimodal_gen`'s disagg state machine); declarative stacked-parallelism axes
compiled to DeviceMesh; general cross-mesh `WeightSyncPlan`.
- ComfyUI workflow→spec compiler MVP (tier-1 ~20-node vocabulary) + weight/adapter fleet cache.
*Gate (design.md Phase 3's, in full): ≥2 Dreamverse sessions/GPU on the recorded duty-cycle trace, p95 within SLO
— this is also where the loop-inversion **falsifier** is evaluated (see §7); LTX-2 A/V full-fan-out end-to-end;
disaggregated-vs-colocated throughput benchmark; CPU-only topology validation suite; ComfyUI tier-1 workflows
compile and run with equivalence reports; spec-built pipelines SSIM-identical to M2 loop versions.*
### M5 — Omni/MoT native + RL hardening (weeks 20–30)
- Cosmos3 re-port onto specs: packed factored sequences, dual-pathway attention, reasoner paged KV, joint denoise,
world-model `ChunkRollout`; `/v1/chat/completions`; AR continuous batching arrives **with** this workload (N5).
- Consistency ladder enforced end-to-end: C1 default in CI, C2 bitwise mode for goldens, Behavior Record opt-in;
first GRPO-class method lands on the engine-client rollout path (log-prob drift becomes the gated metric).
*Gate: Cosmos3 150-test parity suite on the new runtime; reasoner pool efficiency — tokens/s/GPU at target
concurrent denoise throughput, with the ≥10×-vs-re-prefill sanity floor; C1 drift ≈ 0 on a Wan RL run with the
drift dashboard live.*
### M6 — The tail and the precondition (week 30+)
Continuous deletion (M1/M2) shrinks the final phase but does not eliminate it: what remains by M5 is the tier-B
tail on `LegacyPipelineNode` and the frozen `training/` stack — which is a *live consumer* of
`ComposedPipelineBase` and `forward_context`, so its retirement is the precondition, exactly as design.md Phase 5
states. M6 = execute the §6 tail decision (migrate or deprecate each tier-B family), retire `training/` per the
checklist, then delete `ComposedPipelineBase`, the legacy `DenoisingStage`, `forward_context.py`,
`FastVideoArgs`/`TrainingArgs`, and `RayDistributedExecutor` together. **4 loop copies → 1.**
## 4. Breakage manifest (user-visible)
| Release | What breaks | Replacement | Aid |
|---|---|---|---|
| **R1** (M1) | `generate_video(prompt, **kwargs)`, `SamplingParam`, `FastVideoArgs` as public type, `fastvideo.api` legacy exports, CLI flag names, YAML config schema, streaming event types (`Video*Event` → `OmniEvent`, `schema_version`'d from day one) | `VideoGenerator.generate(OmniRequest)`, `DeployConfig`, generated CLI/protocol, `OmniEvent` | `fastvideo migrate` codemod, migration guide, R0 pinned |
| **R2** (M3) | Default execution path becomes the engine (offline bypass preserved); server lifecycle (queue/admission semantics, job states) | `AsyncEngine` | guide; `OmniEvent` schema unchanged from R1 |
| after R2 | nothing user-visible — M4/M5 are additive | — | — |
## 5. Deviations from design.md §10, stated honestly
| design.md | this plan | why it's safe now |
|---|---|---|
| Phase 0 keeps `VideoGenerator`/CLI signatures; `compat.py` shrinks to Phase 5 | M1 breaks signatures, deletes `compat.py` ⚡ | the only argument for the shim was signature stability — explicitly revoked |
| Legacy code deleted at Phase 5 | per-family deletion at parity, M2 onward ⚡ | parity gate is per-family anyway; carrying dead code to a final phase only invites the 19×-broken-freeze failure mode |
| Phases strictly sequential | M2/M3 overlap ⚡ | the engine consumes step bodies, not finished families; the step-body contract freezes at the Wan+Flux2 landing |
| Phases −1 through 4 sized at 36–54 engineer-months | ~21–28 engineer-months (3–4 eng × 30 wks) ⚡ | the delta is real deleted work — no compat maintenance, no adapter upkeep, no dual-stack carry — plus M2/M3 overlap; treat 30 weeks as the aggressive case and 36–40 as the planning case |
| Unchanged | parity/SSIM gates (restored in full at every milestone), enforcement package (CI path gates, CODEOWNERS, inflow rule — now at M0), train/RL migration timing (design.md Phase 1 already migrates NFT), Dreamverse hard-cut (Phase 2 already prescribes it), N2/N3/N5, cost-model currency, Dynamo asks + fallbacks, schema versioning | aggression budget is spent on interfaces only |
## 6. Decisions needed before M0
1. **Tier the model zoo.** Tier A (migrated, coverage guaranteed): Wan 2.1/2.2, Wan-causal/self-forcing, LTX-2,
Flux2, HunyuanVideo, Stable Audio, Cosmos3 (contingent on the M0 merge — it is not on `main` today), image
families. Tier B (runs on `LegacyPipelineNode` until someone claims it, candidate for deprecation at M6):
gen3c, matrixgame2/3, longcat, the rest. **Approve or edit the split** — it bounds M2.
2. **Release framing.** R1 as `v0.3.0` (pre-1.0 semantics, loud notes) vs holding breaks for a `v1.0` story.
Recommendation: `v0.3.0` now — waiting taxes every milestone.
3. **Freeze windows.** M1 freezes `api/`/args/entrypoints PRs ~2 weeks; M2 freezes per-family stage PRs while that
family migrates (days each). Needs maintainer sign-off.
4. **Staffing.** Critical path is M2's per-family step bodies — parallelizable per family after the Wan+Flux2
reference lands. 3–4 engineers ≈ 30 weeks to M5 in the aggressive case (design.md's own sizing implies 36–40
weeks at the same staffing — see §5); 2 engineers ≈ stretch ~1.5×. The Cosmos3-chain merge (M0) is its own
2–3 engineer-month track and should be staffed separately from the runtime critical path.
## 7. Risks specific to the aggressive posture
- **In-flight PR collisions** with layout/schema moves → freeze windows (above) + landing schema cuts at
milestone *starts*, not ends.
- **Community churn at R1** (ComfyUI users, script users) → codemod covers the mechanical 90%; the 10% that isn't
mechanical (kwargs with changed semantics) is enumerated in the guide; previous version stays pinned and
installable.
- **Parity harness becomes the bottleneck** — every aggressive deletion is licensed by it. Mitigation: it is the
*first* deliverable (M0), and per-family migration PRs are template-driven (record → port → compare → delete).
- **Overlap risk (M2/M3)**: the engine team building against a moving step-body contract → the contract
(`init/step/finalize` + `StepResult`) freezes at the *first* family (Wan), enforced by the same schema-version
discipline as external surfaces.
- **The known unknown**: loop inversion at scheduler granularity has no production precedent (design.md §1). The
falsifier stands, on design.md §11.6's schedule: the M3 duty-cycle study *publishes the targets*; the falsifier
is **evaluated at the M4 gate** — if step-level multiplexing can't beat request-level serialization on real
Dreamverse traces, StepScheduler retreats to request-level dispatch and the loop contract keeps only its
streaming/preemption seams, with no family code changing — step bodies and the M1/M2 cuts retain full value.
+267
View File
@@ -0,0 +1,267 @@
# Adversarial Review of `design.md` (v12) — FastVideo Next-Generation Inference Runtime
**Date:** 2026-06-11
**Method:** Multi-agent adversarial review. 9 fact-check agents verified 70 concrete claims against the repo, the local reference checkouts (`cosmos-framework/`, `dynamo/`, `vllm-omni/`, `~/sglang`, `~/vllm`, `~/miles`, `~/verl-omni`, `~/diffusers`, `~/torchtitan`, `~/xDiT`, `~/ComfyUI`, `~/sglang-omni`, `~/cosmos-rl`), and GitHub. 9 attack lenses (abstractions, scheduler/perf, memory/cache, training/RL, strategy, migration, internal consistency, omissions, external borrowings) plus a completeness critic raised 70 findings; every finding went to a refute-by-default verifier. 36 findings were refuted; this document contains only the 34 that survived (1 critical, 26 major — consolidated below where lenses converged — 7 minor), plus fact-check corrections.
---
## Verdict
The architecture survives its strongest attacks — loop inversion's expressibility, the typed-state hybrid, the N1/N5 scope discipline, the clean-room GPL posture, and the C2-for-batch-1-video argument all held under refutation attempts. What does not survive is:
1. **The migration plan**, which consumes its own substrate two phases before building it and rests on a "frozen legacy stack" premise this repo has already empirically falsified.
2. **Two load-bearing factual errors** about reference systems (vLLM's BlockPool page sizes, diffusers' loop ownership) that each drove a recorded design decision.
3. **A family of undesigned failure/memory/trust paths** that the multiplexing bet itself creates. One is critical.
---
## Critical
### C1. No failure-isolation or cancellation semantics for the multiplexed pool — the blast-radius problem the architecture itself creates
**Where:** §6.3.1; absent from §12.
Today one request per pool means one request's CUDA error is its own problem. Step-multiplexing changes the failure class categorically: a mid-step OOM/illegal-access/NaN from one request poisons the CUDA context and desyncs in-flight NCCL collectives for *every* co-scheduled tenant on the pool, including resident Dreamverse session caches. The doc designs none of the machinery: no SPMD-consistent abort broadcast (the dual of its scheduling broadcast), no request-fatal vs pool-fatal classification, no pool re-init + cache-invalidation policy, no partial-artifact semantics for fan-out graphs. "OOM" and request cancellation appear nowhere in 1799 lines; "abort" appears once (RL stragglers).
Ordinary cancellation is also missing — and vibe directing makes abandoning in-flight generations the *common* path. Worse, Phase 2 retires Dreamverse's `gpu_pool.py`, which today has a working sentinel-fd worker-death watch (`gpu_pool.py:542-586`), into engine-client calls — a reliability regression for the flagship customer if the gate ships as written. vLLM v1, the doc's own scheduler template, needed first-class machinery for exactly this (`ENGINE_CORE_DEAD`, `EngineDeadError`, `abort_requests`). Risk 4 covers only scheduling-decision divergence; the long-job-resilience known-gap is single-job-framed.
The abort path shapes the StepScheduler loop, the worker RPC surface, and CacheManager handle lifetimes — it must be designed *with* Phase 2, and by the doc's own standard ("absence reads as a decision"), this absence is an oversight.
---
## Major — reference-system misreads that drove recorded decisions
### M1. The single-BlockPool CacheManager rests on a property vLLM explicitly does not have: per-group page sizes
**Where:** §6.3.2 lines 555-559.
The sentence asserts two mutually exclusive properties. vLLM's one-pool/no-fragmentation guarantee exists *only because* physical bytes-per-block are uniform across all groups: `kv_cache_utils.py` asserts a single page size (`get_uniform_page_size`), and its docstring says verbatim that breaking this "is non-trivial due to memory fragmentation concerns." Groups differ only in tokens-per-block at equal byte size; the unification mechanism inflates the smaller group's `block_size`.
Apply that to FastVideo's groups: a text-KV page (~64 KB/layer) vs a latent-frame slab (9.6–32 MB/layer for 1.3B/14B causal Wan) is a 150–500× ratio — unification means a 500-token reasoner prompt strands a multi-MB slab per layer-group. The one vLLM path with multiple page sizes (DeepseekV4) statically partitions capacity at startup over a single global block-id free list, which is harmless when group demand is token-coupled (every token passes through all layer groups) but wasteful exactly when demand is workload-decoupled — FastVideo's regime, where text-KV and chunk-KV demand vary independently with request mix.
Since this misread is what reversed the two-pool sketch (recorded at line 280), the decision rests on a false premise: either chunk-KV stays uniformly fine-paged (losing the slab semantics the MoT "falls out naturally" story depends on), or the two-pool design returns and needs its own fragmentation/deadlock argument.
### M2. diffusers Modular is not loop inversion — the "strongest external validation" of the keystone doesn't validate it
**Where:** §5 line 277, §6.2.3 lines 426-428.
`LoopSequentialPipelineBlocks.__call__` raises `NotImplementedError`; every concrete family hand-writes `for i, t in enumerate(timesteps)` inside its own blocking wrapper (`wan/denoise.py:434`, `stable_diffusion_xl/denoise.py:701` — SDXL ships four such wrappers, the subclass forest again). The iteration is block-owned, invisible to any runtime — no init/step/finalize, no external driver, none of the properties §6.2.2 says inversion exists for (scheduling, interleaving, preemption, streaming, fair sharing). In scheduling terms it is the current `DenoisingStage` with a refactored body — i.e., it validates the Guiders/policy pillar but as evidence for inversion it is *equally consistent with the alternative the design rejects* ("keep loops in stages, make bodies pluggable"). The class also carries an explicit experimental warning.
Consequence: no surveyed system — vLLM, sglang, multimodal_gen, diffusers — implements runtime-owned diffusion iteration at scheduler granularity. Loop inversion is the design's most novel element with zero production precedent, and risk 3 (which admits novelty only for the hybrid AR+denoise slice) should say so instead of borrowing validation the reference doesn't provide.
### M3. Cost-currency scheduling drops the memory half of vLLM's admission — and memory is never a scheduling resource anywhere in the design
**Where:** §6.3.1 (lines 476-547), §6.3.2; two lenses converged here.
vLLM's token budget is not a prediction — it is an exact cap checked *in the same loop as memory admission* (`allocate_slots` per request, preempt on allocation failure; activation memory separately bounded by a profiled worst case). The design takes the accounting structure, swaps the currency for a *forecast* (predicted GPU-time), and drops the memory dimension entirely: latents, conditioning sets, CFG duplicates, and activation peaks live in `RequestState`, explicitly outside the CacheManager, and nothing bounds how many concurrent LoopStates a pool admits — for a workload the doc itself calls memory-bound (line 499). Two items that each fit alone can jointly OOM, and a GPU-seconds currency cannot see it; combined with C1, that OOM is a pool-wide event. "Preemption only at step boundaries" never defines what happens to a preempted request's multi-GB resident state (offload? drop-and-resume-from-LoopState? — different economics from KV recompute).
Related internal contradiction, verified: cost is "static and known at admission... a table lookup" (line 538), but the same cost model is cache-dit-aware (line 520) — DBCache skip decisions are runtime data-dependent residual comparisons, unknowable at admission.
**Fix:** the budget needs a memory axis (resident-state + peak-activation per schedulable item), admission needs a memory planner over RequestState, and preemption semantics must be specified. The Phase-2 "≥2 sessions per GPU" gate rests on unaccounted memory until then.
### M4. Punica cannot express ComfyUI LoRA semantics
**Where:** §9.4 lines 1376-1380 (also §6.3.2 lines 575-579). *Verifier rated minor-to-major; grouped here with the borrowings cluster.*
vLLM's `LoRARequest` carries one `lora_int_id` and no strength field; scaling is baked into `lora_b` at registration; the Punica wrapper maps one adapter index per token. ComfyUI traffic — the workload §9.4 names — is N stacked LoRAs per request with continuous user-set `strength_model` *and* `strength_clip`, routinely tweaked per generation. Pushing that through Punica means registering each (ordered-set, strengths) tuple as a synthetic concatenated adapter: near-zero cache-hit rate across strength tweaks, registration churn in the stacked GPU weight slots, and concatenated ranks colliding with `max_lora_rank`. "Strictly better than hot-swap-only" is unsupported without a composition layer that doesn't exist anywhere, including in vLLM.
---
## Major — execution-plane gaps
### M5. MoT mode multiplexing has no parallelism answer
**Where:** §6.3.1 lines 502-503 vs §6.3.4; Phase 4 gate.
The "mode multiplexer" claim assumes both loop types share one static pool layout (`parallel: [dp, cfg, sp, tp]`), but their optimal layouts are disjoint: denoise wants SP+CFG; AR decode is sequence-length-1 — SP has nothing to shard and CFG doesn't exist. On a `[cfg(2), sp(4)]` 8-GPU pool the reasoner either runs replicated (1/8 useful work, paged KV duplicated 8×) or needs TP — and TP-everywhere regresses the bread-and-butter denoise workload on the flagship pool. Per-phase re-layout of the same resident weights is not expressible in the §6.3.4 spec (one static stack per pool), and resharding machinery exists only for train↔rollout weight sync (§8.6). §6.3.1's own jumbo-step mitigation (split cost classes across pools) is structurally unavailable for MoT — AR steps and denoise steps are the same weights — so concurrent reasoner token latency is gated by indivisible 50–500 ms denoise steps.
A workable resolution exists (AR continuous batching data-parallel across the cfg×sp weight-replica axes onto TP subgroups, plus §6.3.3 per-pathway TP, plus routing pure-REASON traffic to differently-shaped pools), but the doc never states one, and the Phase-4 gate ("reasoner ≥10× faster than re-prefill") is measured against an O(n²) strawman baseline that certifies nothing about pool efficiency. Risk 3's "prototype early in Phase 4" defers a *design contradiction*, not an implementation unknown.
### M6. The engine's own multi-node story is unstated, and the Ray executor silently disappears
**Where:** N1 line 143, §6.3.5 line 673, §6.0 line 304.
Whether one worker pool may span nodes is a load-bearing decision the doc never makes — Dynamo routes *between* workers; it does not own the NCCL mesh *inside* one. If pools are single-node by fiat, SP degree caps at ~8 GPUs, directly contradicting line 543's jumbo-step mitigation ("shrink jumbo step wall-time with SP"), capping MoT model scale — and `RayDistributedExecutor`, today's shipping multi-node path, is silently dropped: it appears in the §3.1 diagram and then never again in §6, §10, §11, or §12 (violating the plan's own "every phase deletes or freezes what it replaces" discipline). If pools may span nodes, the engine owns cross-node collective bring-up, NCCL-timeout fault domains, and a multi-node health/drain contract — none designed, and C1's recovery problem becomes a multi-node recovery problem. Either answer changes Phase 2/3 scope. "Node-group" appears once, undefined.
### M7. Policies carry per-request mutable state with no state-scoping contract — and the doc contradicts itself on when policies are resolved
**Where:** §6.2.3 lines 412-417 vs §6.2.2 line 387 vs risk 2 line 1621; §6.4 lines 837-840.
The doc says policies are resolved at pipeline build (lines 412-413; risk 2: "resolved to bound methods at build time") *and* in `DenoiseLoop.init` (line 387) — a genuine contradiction on a load-bearing contract. It matters: AdaptiveGateCFG — a named CFGPolicy example and the Wan2.2 worked-example default — is per-request mutable state in shipped code (`denoising.py:338-343, 507-551`: `delta_cached`, `delta_cached_model_id`, gate counters). Build-time-resolved singletons mean request A's cached CFG delta gets applied to request B the moment Phase 2 interleaving lands — silent quality corruption no Phase-2 gate (load tests, latency budget) can catch. This is the *exact* failure mode §6.4 cites to justify interceptor state scoping ("silently corrupts under concurrent requests") — the contract was designed for the plugin tier and forgotten for the policy tier, which sits on a hotter path. Cheap fix (policy state into LoopState, same as plugins), but it must be in the spec.
### M8. The six-policy taxonomy does not factor the shipped step bodies — no step skeleton or cross-policy interaction contract is defined
**Where:** §6.2.3 (policy table, line 424 claim); §6.2.2 lines 386-389; §6.4 lines 837-844.
The proposed step is three phases (forward → CFG combine → scheduler step); the shipped loops need ~six, with dependencies that cross policy boundaries. Verified examples:
- **Cosmos** conditioning-frame injection consumes the *sampler's* EDM coefficients, applies per-CFG-branch both pre-forward (input mix) and post-forward (x0 clamp), and the CFG combine runs in x0 space — ConditioningInjector × Sampler × CFGPolicy interleaved inside each branch, unownable by any one of them (`denoising.py:845-933`).
- **TI2V** clamps latents *after* `scheduler.step` — a post-step constraint with no policy slot (`denoising.py:570-573`).
- **Cosmos2.5** builds per-frame timestep vectors with a conditioned-frame override and re-clamps GT every step pre-forward.
- **CausalDMD** renoises between steps choosing `add_noise` vs `add_noise_high` by expert boundary — Sampler × ExpertRouting (`causal_denoising.py:268-301`).
- **AdaptiveGateCFG** must observe ExpertRouting's switch to invalidate its delta (today an inline `id(current_model)` check) — yet no channel for one policy to observe another is defined anywhere.
- **LTX2** guidance is 1–4 runtime-decided passes whose branches alter the network via forward kwargs (`skip_cross_modal_attn`, `skip_video/audio_self_attn_blocks`) — colliding with BlockInterceptor's domain in a way the "two block-skippers conflict" pre-flight check cannot see, and breaking §6.4's per-CFG-branch state scoping, which assumes a fixed cond/uncond branch vocabulary (`ltx2_denoising.py:503-605, 620-631`).
None of the six policies covers prediction-space conversion, per-token timestep construction, post-step latent constraints, inter-step renoising, or chunk-boundary refresh. The fix is not abandoning policies — the Sampler registry is the natural home for some of this, and composition still strips the duplicated offload/attn-metadata/autocast/trajectory plumbing — but the design needs the fixed step skeleton with ordered, typed extension points and an explicit policy-interaction contract, worked through Cosmos2.5 and LTX2 *in the doc*. Until then, "a new model contributes policies + a graph spec; it does not edit shared loop code" (line 424) is asserted, not demonstrated.
### M9. OmniRequest cannot parameterize multi-loop graphs
**Where:** §6.1 lines 318-334; §6.6 line 905; worked examples (c)(d) lines 943-950.
One flat `SamplingParams` + one flat `DiffusionParams` per request, while the design's own flagship examples are multi-loop graphs needing per-node knobs: LTX-2's refine loop has its own step count and guidance scale *today* as first-class fields (`fastvideo_args.py:204-205`, threaded through `compat.py` and `dynamo/examples/diffusers/worker.py:201-203`); a thinker and talker need different `max_tokens`/`temperature`/`stop`. No request→graph-node parameter binding is defined anywhere; the only escape hatch is line 905's per-model `ModelOptions` blocks — i.e., the `ltx2_*` field-leakage pattern the doc indicts at P3, with a type wrapper, regenerated into the OpenAI/CLI views that derive from the request schema (line 907). Needs a real decision — parameters keyed by graph-node id, or per-node override blocks validated against the PipelineSpec — made in Phase 0, because that schema ships first and external consumers build against it.
---
## Major — caches and weights
### M10. No feature-cache invalidation story under LoRA hot-swap — te-LoRAs make the embedding cache serve stale embeddings in the workflow cloud
**Where:** §6.3.2 lines 570-574 vs §9.4 lines 1349-1380.
The only invalidation rule in the document is RL `update_weights` → `reset()`. But ComfyUI-grade LoRAs routinely patch the *text encoder* alongside the DiT (`comfy/lora.py` maintains `lora_te/lora_te1/lora_te2` key maps; `load_lora_for_models` takes a separate `strength_clip`), so a content-hash-keyed embedding cache returns embeddings computed under the wrong adapter state the moment two workflows share a prompt but differ in te-LoRA stacks — silent wrong output in the exact product (§9.4 "exact mode") whose trust claim is reproducibility. §11.8 even makes cross-request embedding reuse load-bearing as the radix-cache substitute. And once Punica-style batched multi-LoRA lands, requests with different adapter stacks coexist concurrently on one pool, so the cache must be key-*partitioned* by (encoder identity × adapter set × strengths), not flushed — a different design from the `EncoderCacheManager` reset() semantics being adopted, which come from a world where encoders are never patched per request. The key schema needs a weight-state epoch / adapter-set hash as a mandatory component, decided before Phase 3.
### M11. Checkpoint/LoRA patching mutates pool-shared weights — a pool-quiescing barrier the StepScheduler has no vocabulary for
**Where:** §9.4 lines 1371-1380 vs §6.3.1 and §6.0 line 299.
Components are "one resident copy per worker pool"; patch/unpatch mutates that copy, which is global to every loop interleaved on the pool — yet step-interleaving is the engine's core Phase-2 value. Two interleaved loops requiring different patch states cannot coexist, so every cross-group transition is a drain barrier: finish in-flight steps, apply/undo `W += scale·BA` across 14–28 GB shard-consistently across TP/SP ranks (ComfyUI keeps weight backups for the undo — 2× weight memory or a CPU→GPU restore at PCIe seconds), re-admit. Under workflow-cloud traffic (long-tail checkpoints, per-request adapter stacks), transition frequency is the whole game — and the §6.3.1 cost model (lines 516-521) has no weight-state-transition term, no notion of weight state as schedulable state, and no quiesce-vs-queue policy, even though transition cost is exactly what A1 checkpoint-affinity routing must weigh. The §8.6 safe-point-swap pattern shows the doc knows the shape but never applies it here. §9.4 calls this "the one real new subsystem"; §12 carries no risk entry for it.
---
## Major — training/RL
### M12. "Step bodies are plain tensor programs, so autograd composes" is contradicted by the distillation code the substrate must absorb
**Where:** §8.2 lines 1016-1019; §6.2.2; §6.3.2.
Self-forcing does not "drive `DenoiseLoop.step`": its rollout samples per-block exit indices broadcast across ranks, runs no-grad steps to the exit, runs exactly *one* grad-enabled forward, then a separate no-grad `store_kv=True` context-caching pass with context noise, gated by `start_gradient_frame`. None of this fits `init/step/finalize` + `StepResult(done, emit)` without grad-gating flags, per-step cache-write control, and per-block exit policies — training-only surface in substrate code, or the method keeps its own loop and the "3 copies → 1" dedup claim dies for the hardest case. The KV path needs grad/AC-aware semantics the engine pool lacks: today's causal model snapshots KV indices whenever `torch.is_grad_enabled()` so activation-checkpoint recompute doesn't double-advance the cache (`wan_causal.py:119-120,405-431`), and never recycles blocks mid-rollout — while §6.3.2 specs vLLM-style out-of-window block recycling, and §8.5's own profile taxonomy says "training forward … *no caches*," showing the grad+KV case was never designed. §8.3 explicitly stakes the architecture on ChunkKVPool serving self-forcing training.
(Note: the related forward-context-backward attack was refuted — the Phase-1 retirement of the global plus explicit metadata passing *helps* autograd composition. The surviving residue is the grad-window/cache-mode design above.)
### M13. Behavior Record cost is understated ~1.5 orders of magnitude for its own flagship case (MoE diffusion)
**Where:** §8.5 lines 1156-1160; §5 miles row line 1088.
The miles ~60 MB/sample figure is per-token routing, one forward per generated token. Diffusion re-routes the *entire packed sequence at every denoise step, twice under CFG*: the record is steps × CFG × tokens × MoE-layers × top_k. For a Cosmos3-class request (Qwen3-VL-MoE config: 60 experts, top_k 4, ~24 sparse layers via `decoder_sparse_step=1`, ~50K packed tokens, 35-50 steps × 2 branches) that is ~1.3–1.9 GB/sample int32 — ~20–30 GB per 16-sample GRPO group, before latents. "Cheap because trajectory capture is already an OutputSpec feature" conflates plumbing cost with byte cost; at these sizes the Record forces a buffering/transport/storage design (GB-scale trajectories through connectors from disaggregated rollout fleets) that appears nowhere — not in §8.7's TrajectoryBuffer, not in §12, not in the known-gaps list. (The RNG-draws sub-claim was refuted: seeded generators in a single shared loop reproduce draws; uint8 expert IDs also cut 4×. The routing-record problem stands.)
### M14. The omni-RL pilot is a Phase-4 deliverable with no objective design
**Where:** §8.7 lines 1236-1240; §10 Phase 4.
The section establishes *expressibility* (one trajectory, two segment types — true, and a real structural advantage over engine-per-stage stacks) and quietly upgrades it to a deliverable without posing the algorithm problem:
- **Scale mismatch:** token log-probs are O(1–10) nats over 10²–10³ tokens; per-step diffusion SDE log-probs are Gaussian densities over 10⁶–10⁷ latent dims — any joint clipped-ratio objective needs principled per-segment normalization that none of the cited recipes (FlowGRPO/DanceGRPO/NFT/AIPO/GSPO) provides; get it wrong and one modality silently dominates the shared trunk.
- **Credit assignment:** the reasoner influences video reward only through *sampled discrete tokens* re-entering as conditioning — a non-differentiable boundary, so token segments get sparse trajectory-level REINFORCE signal while denoise segments get dense per-step ratios, both updating shared attention-trunk weights, with no interference analysis.
- **Reasoning regression:** RL-updating the und pathway on video-reward-correlated signal risks degrading its reasoning; reference-model KL anchoring for hybrid episodes is never mentioned.
The entire treatment is the phrase "optimized with mixed objectives," and §12's 15 open questions contain nothing on it — for the capability marketed as "the capability nobody else has." Either it gets an algorithm sketch and an open-question entry with an owner, or the Phase-4 item should be demoted from "pilot" to "trajectory capture demonstrated."
---
## Major — the migration plan (the weakest section)
### M15. The "frozen legacy stack" premise is empirically false in this very repo
**Where:** lines 5, 110, 1026; §11.4; risk 5.
The anti-third-stack defense is a declared freeze plus intent to delete — and this repo has already run that experiment and it failed within weeks. Verified from git: `fastvideo/train/` landed 2026-03-09 (#1159); since then **19 commits modified the "frozen" `fastvideo/training/`**, including a *brand-new* `cosmos2_5_training_pipeline.py` added to the legacy stack on 2026-05-11 (#1227) — **nine days after `training/AGENTS.md` explicitly forbade adding new models there**, and eleven days after the same model landed in `train/` (#1224). World-model training (#1179) and LongCat finetuning (#1244) also landed in the frozen stack in May; EMA bugfixes as recently as June 8-9; `AGENTS.md` still calls `training/` "authoritative for shipped models."
The doc invokes the training/-vs-train/ "lesson" but proposes nothing mechanically different from what was tried: no CI gate rejecting new files under legacy paths, no codeowner veto, no named owner per family, no calendar date for Phase 5. "Phase 5 is a scheduled deletion, not an aspiration" (risk 5) — but nothing in the document is scheduled. Under the same model-port pressure that broke the training/ freeze (measurably higher on the inference side), this freeze breaks the same way. Name the enforcement mechanism that did not exist last time, or the deprecation commitment is the prior failure restated with more confidence.
### M16. Phase dependency inversion: Phases 1–2 consume the substrate Phase 4 builds
**Where:** §10 lines 1406-1446 vs §6.3.1 lines 487-489, §6.3.2; three lenses converged on this.
Phase 1 migrates causal Wan ("exercises chunk-KV"); Phase 2 ships "AR continuous batching" — which §6.3.1 *constitutively defines* as "(continuous batching; paged KV; chunked prefill)"; the CacheManager owning both lands in Phase 4, and risk 3 even defers the StepScheduler+KVPool prototype to "early in Phase 4," contradicting Phase 2. Compounding it: **no AR-pathway model exists on the new runtime before the Phase-4 Cosmos3 re-port** (Wan-causal is chunked denoise, not token AR; thinkers/talkers are Phase 4), so Phase 2's headline deliverable has neither a cache backing nor a workload — and none of Phase 2's gates (lines 1427-1430) tests AR batching.
The Phase-1 half is softenable: an interim per-request chunk-KV behind the unchanged `KVHandle` seam, with a Phase-4 allocator swap, is normal incremental staging — but the doc never states this, and its own "no third stack / every phase deletes what it replaces" principle cuts against unstated throwaway implementations. Fix structurally: pull a CacheManager v0 (chunk-KV slabs + minimal paged text-KV) into Phases 1–2, or move AR batching to Phase 4 and rewrite the Phase-2 gate to what it actually exercises.
### M17. Phase 4 re-ports a baseline that is not on main, and the plan schedules neither its merge nor its rebase
**Where:** §10 Phase 0 line 1405, Phase 4 lines 1439-1446; §1 lines 42-49; Appendix.
`fastvideo/pipelines/basic/cosmos3/` on main contains only `__pycache__` — the design's forcing function exists solely as the unmerged 5-branch stacked chain (`feat/cosmos3-tier-a-port` → … → `feat/cosmos3-reasoning`). Phase 0's "Cosmos3 audio leaves `batch.extra`" cannot execute against main: it presupposes the chain is merged (a major-model review effort the plan never schedules) or means maintaining the migration on a side branch, continuously rebased across the most churn-heavy refactors in the repo's history (ForwardBatch→RequestState, loop inversion, executor→engine) — months of conflict-resolution work, unowned and unsized, on the artifact whose 150/150 bit-exactness is the design's proudest credential and whose parity suite the Phase-4 gate requires ("every phase ships green" cannot apply to a suite that is not in the tree). The plan sequences other in-flight work explicitly (`fastvideo/api/` in Phase 0, PR #1438 in Phase 1) but skips this. Needs an explicit merge milestone before Phase 0 touches the port.
### M18. G5's enforcement instrument has holes: ~6-7 shipped families have no SSIM test, and the CI-cost mitigation is incoherent for substrate PRs
**Where:** G5 lines 128-129; Phase 0 gate line 1405; risk 6.
`fastvideo/tests/ssim/` covers ~14 of 20+ families. Cosmos(2/2.5), Hunyuan, Hunyuan15(+SR), HYWorld, MagiHuman, Waypoint, and MatrixGame-v1 have no SSIM test — "all SSIM suites unchanged" passes *vacuously* for roughly a third of shipped pipelines, exactly the ones sitting on the shared loop being refactored. And risk 6's "gated to touched families" mitigation is designed for model-local PRs; Phases 0–2 are by construction not model-local — the ForwardBatch adapter, loop inversion, and executor replacement sit under every family, so "touched families" = all of them on precisely the riskiest PRs. Either substrate PRs run the full GPU matrix (a cost the plan should budget — SSIM runs on Modal L40S today) or gating quietly degrades to sampling, which is how regressions slip through. Needs: a reference-seeding work item before Phase 1, or G5 restated as "zero regression for the SSIM-covered subset," plus a stated per-phase GPU-CI budget.
### M19. Phase 5's deletion milestone breaks the "frozen and untouched" legacy training/ stack
**Where:** lines 5-6, 144, 1026-1027 vs Phase 5 line 1448.
The frozen stack is a live consumer of exactly the code Phase 5 deletes: `fastvideo/training/training_pipeline.py:39` imports `ComposedPipelineBase`/`ForwardBatch`/`LoRAPipeline`, holds `validation_pipeline: ComposedPipelineBase`, and its validation instantiates real legacy pipelines that run the legacy `DenoisingStage`; `distillation_pipeline.py:31` likewise. So Phase 5 cannot remove `ComposedPipelineBase` and `DenoisingStage` while leaving `training/` untouched — either the deletion milestone hollows to "delete except what legacy training/ needs" (the old path never dies — the very smell being fixed) or the scope statement is false and `training/` breaks on this plan's schedule. Relatedly, "loop inversion makes the step functions the single shared implementation" is arithmetically 3→2, not 3→1: the legacy inlined copies are out of scope forever. The doc needs an explicit answer: what happens to `fastvideo/training/` at Phase 5?
### M20. "Retire `fastvideo/forward_context.py` (Phase 1)" is infeasible as scheduled
**Where:** §6.3.3 lines 618-621; Phase 1 lines 1412-1414; vs N2/N4; Appendix line 1791.
194 references across ~50 files. The global is read inside `fastvideo/attention/layer.py` — the shared Attention module on *every* family's hot path — and set in 27 places inside the frozen `training/` stack (8 module-level imports). Phase 1 migrates only Wan+Flux2; the other ~16 families run "unmodified" behind the legacy adapter (N4) and still set the global. So in Phase 1 the file cannot be deleted (touches the frozen stack, violating N2; breaks every unmigrated family), and `attention/layer.py` must serve both worlds simultaneously — a dual-sourcing branch in the hottest shared layer, undesigned. The honest description: Phase 1 *adds a second context mechanism beside the global*, and the global survives until Phase 5 at the earliest — where the deliverables list never mentions it. Appendix A states "retired Phase 1" as accomplished fact. Rewrite as "new-path-only StageContext; `forward_context` frozen for legacy consumers; deletion gated on Phase 5," and design the dual-mechanism cost.
### M21. §10 is a dependency ordering, not a plan — no timeline, no staffing, no sizing, and no policy for the ~1-2 new model ports per month that arrive during the migration
**Where:** §10; N4 line 153; risk 1.
The scope — typed I/O, loop inversion + policies, extension system, async engine + StepScheduler + online-calibrated cost model, four-class CacheManager, PackedSeq/MoT layers, declarative parallelism compiler, workflow compiler, RL layer, Dynamo contract, config collapse — is plainly multi-engineer-years, with zero dates, headcount, per-phase sizing, or owners; "by Phase 2" decision deadlines (§11.1, risks 7/15) are unanchored because Phase 2 is not a date.
The sharper, unanswered problem is **inflow**: git shows ~1–2 new families landing per month (Flux2 Klein and Lucy Edit on 2026-06-09 alone; MatrixGame3 05-27; MagiHuman 05-12; Stable Audio 05-01; Gen3C 04-01…). Over multi-quarter Phases 0–4, another 10–15 models arrive, and the doc never says what they target: land them on legacy abstractions and the Phase-5 tail grows faster than phases retire it (negative net migration velocity); force them onto the new stack and every port blocks on machinery that doesn't exist until Phase 1/3/4. Either answer materially changes the plan; choosing neither means the terminal state recedes indefinitely. Minimum fix: per-phase engineer-month estimates, a named owner per phase, a calendar target for Phase 5, and an explicit "new ports target the new stack starting at Phase X" rule with its porting-velocity cost stated.
---
## Major — product/trust surfaces
### M22. Per-request plugin enablement is an unsandboxed third-party-code and noisy-neighbor surface; only workflow JSON is named untrusted
**Where:** §6.4 lines 859-861 vs §12 input-hardening gap lines 1693-1695.
Entry-point plugins execute arbitrary code inside the serving engine, and the doc makes their selection part of the *request* (`diffusion.plugins=[{"name": "cache_dit", "Fn": 8, "Bn": 8}]`) in the same engine pitched as a multi-tenant cloud — and since the OpenAI protocol is *generated from the request schema* (lines 907-908), the field derives into the public API with no carve-out. Consequences forcing a design change: (a) **correctness** — a caller can attach a distribution-altering interceptor to a request the product has labeled "exact mode" (the §9.4 trust claim), or pass unvalidated kwargs into third-party code; (b) **isolation** — a `needs_eager` observer on one request drops compile/cudagraph capture for scopes shared with co-scheduled tenants (line 809), a noisy-neighbor vector with no cost attribution anywhere in the metrics design; (c) **supply chain** — entry-point resolution imports whatever package claims the name. The needed contract: enablement/allowlisting at DeployConfig scope only; requests merely parameterize pre-enabled plugins against per-plugin validated schemas; plugin overhead attributed per-request in the cost model. §12's input-hardening gap names only workflow JSON — a categorically different surface.
### M23. No versioning or stability contract for the serialized schemas shipped to external consumers mid-migration
**Where:** §6.4 line 861; §6.6 lines 920-927; §10; open question 12.
By Phase 3 there are at least four externally consumed serialized surfaces: hub-published ModelSpec manifests (interchange with diffusers' `modular_model_index.json` — a format co-owned with an external party), compiled-workflow PipelineSpecs (content-hash-keyed in the weight-fleet cache — schema changes silently change hashes and invalidate fleet affinity), the OmniEvent streaming schema (Dreamverse's frontend; proposed as Dynamo ask A3's wire format), and per-model ModelOptions blocks. Phase 4 then lands PackedSeq, session-scoped inputs, and the Cosmos3 re-port — guaranteed churn after consumers exist. The migration plan gates *behavior* at every phase (SSIM, parity, load) and gates *interfaces* at none; the only versioning commitment in the document is hook-point names (open question 12 is scoped to hook points). Without per-surface decisions now — `schema_version` fields, frozen-vs-experimental tiers per phase, a deprecation window — Phase 4 either breaks published artifacts or gets paralyzed by accidental freezing. G5 protects only the Python `VideoGenerator` call.
---
## Minor (confirmed)
1. **ForwardBatch has 111 fields, not ~250** (AST-verified; stated twice, lines 33/188). P3 survives at 111, but the headline metric is inflated 2.3× in a doc that brands its pain points "evidence-backed" — it invites discounting of the numbers that *do* verify exactly (1381 lines and 35 probes both check out).
2. **"Prediction is a table lookup" vs the design's own flagship features** (§6.3.1 vs §6.4): DBCache/FBCache/TaylorSeer decide per step from runtime residual similarity — a stochastic per-step cost multiplier unknowable at admission; VSA tile selection is content-dependent; and AR decode lengths are unbounded (the doc concedes vLLM "must guess decode lengths," then silently exempts its own AR group).
3. **Worked example (g) is internally contradictory**: cache-dit + C1 + "identical trajectories" are pairwise incompatible under §8.5's own `distribution_altering` contract (§8.7 states the rule correctly: cache acceleration is C0). Matters because (g) is the template PR #1438 is told to target in Phase 1.
4. **The Phase-2 Dreamverse gate is untestable as written**: at ~4.55 s GPU-saturating per 5 s clip (line 1263), "≥2 concurrent sessions per GPU at unchanged segment latency" is only passable under an unstated think-time/collision-rate assumption — the gate can be passed or failed at will by choosing the test's session behavior. More broadly, no quantitative multiplexing target (sessions/GPU under a stated load profile, GPU-utilization, cost/clip) exists anywhere, so there is no way to conclude after Phase 2 whether step-level scheduling earned its complexity over the §11.6-rejected simpler design.
5. **The exec summary launders Dynamo contingencies into outcomes** (line 75: "each with a fallback — so Dynamo fronts both production serving and RL rollout fleets"): the body is honest (A1–A7 with fallbacks; §11.9; §12.15), but the asks are unfiled RFCs on an NVIDIA-governed roadmap; A5's own fallback "weakens fleet-scale async RL," and if A3 misses Phase 2, Dreamverse ships on the direct-WebSocket bypass and the production-hardened fallback becomes permanent — the exact "permanent workaround" dynamic §11.9 claims the direct relationship avoids. Ask-sequencing (§12.15) has no owner or decision dates.
6. **diffusers as "convergent validation" cuts both ways** (see M2): its four-wrappers-per-family shape is the subclass forest again; the citation supports the rejected alternative as well as the chosen one.
7. **Punica/ComfyUI LoRA semantics gap** — see M4.
---
## Fact-check corrections
70 concrete claims were checked; **none was fabricated**; 13 need correction. Everything else verified, including the claims most likely to be embellished: vLLM RFC #42770 (author/date/content/two-tier resolution), PR #42304 **merged** 2026-05-16 with `VLLM_USE_BREAKABLE_CUDAGRAPH`, vllm-omni RFC #4084, the Thinking Machines numbers (80/1000 unique outputs, divergence at token 103, 26s→42s, KL results), the Dynamo worker's `asyncio.Lock`, cache-dit, the cosmos-framework MoT details (PackedAttentionMoT, MoTDecoderLayer, ReasonerKVCache, MoE gen-MLP), miles/verl-omni/sglang-omni mechanics, sglang's cache-dit monkeypatch scars, and `enable_teacache` genuinely having no consumer.
| # | design.md says | Reality |
|---|---|---|
| 1 | "1381-line `DenoisingStage`" (lines 34, 201) | 1381 is the **file**; the class is ~670 lines (47–715) plus 6 subclasses in-file. The 35-probe count is exact for the file. |
| 2 | "~250-field ForwardBatch" (33, 188) | **111 fields** (whole file incl. TrainingBatch/PreprocessBatch: ~153). |
| 3 | "19 denoising-stage classes" (201) | **22** model/variant classes (+ base = 23); the list omits Magi-class and two other same-category stages predating the doc. |
| 4 | "Cosmos2.5 clamping … hardcoded in the shared loop" (201) | Clamping lives in the `Cosmos25DenoisingStage` **subclass**; the Wan2.2 expert switch (`denoising.py:229-235, 352-376`) and TI2V inline VAE encode (`:239-268, 399-404, 570-572`) are in the shared loop as claimed. |
| 5 | `SamplingParam` "~170 fields" (887) | **75**. The ~170 figure belongs to TrainingArgs (90 own + 81 inherited = 171). |
| 6 | `FastVideoArgs` "~96 fields" (885) | **81** (TrainingArgs subclassing claim correct). |
| 7 | "TP and SP (Ulysses/ring)" (192) | Main is **Ulysses-only** (`all_to_all_4D`); no ring-attention SP is wired into FastVideo. |
| 8 | CFG "3 copies: `stages/conditioning.py` vs …" (993) | Right count, wrong citation: the inference-stack copy is in `denoising.py`, not `conditioning.py`. |
| 9 | ComfyUI "~45 `comfy_extras` packs", "90+ blueprints" (1335-1339) | **117** packs (matching nodes.py's 117-entry registration list); **80** in-tree blueprints (the larger library ships via the registry). 64 core nodes, 39 API providers, GPL-3.0, FIFO-no-batching all verify. |
| 10 | kv-router events "`{sequence_hash, block_hash, removed}`" (707, A1 733-739) | Paraphrase: actual shape is `KvCacheEventData::Stored{parent_hash, blocks[{block_hash, tokens_hash}]}` / `Removed` / `Cleared` (`protocols.rs:627-646`). Token-prefix-derived keying verifies. |
| 11 | miles TIS clamp "to `[0.5, 2.0]`" (1086) | Configurable `[tis_clip_low, tis_clip]`, CLI defaults [0, 2.0]; the 0.5/2.0 pair comes from the MIS example config (`mis.yaml`). |
| 12 | sglang-omni "`DllmScheduler` for a DiT talker" (269) | DllmScheduler serves the **LLaDA2-Uni thinker** (diffusion-LLM); the DiT talker is Ming-Omni's, on a different scheduler. |
| 13 | `_iter_packed_batches` under `model/vfm/` (236); §11.3's claim that the port's "own status notes" list reasoning-KV/batching/streaming/prefix-reuse as "missing for production" | Lives at `cosmos_framework/inference/inference.py:66`. PORT_STATUS.md confirms 150/150 but contains no such missing-for-production list — that framing is the design doc's own and should not be attributed to the port's status notes. |
---
## Attacks that failed (the doc survives these)
The refute-by-default verifiers killed 36 findings, several of them attacks a hostile reviewer would lead with — worth knowing they don't land:
- **ChunkRollout/DenoiseLoop nesting is expressible** in the stated Stage/LoopStage/StepResult contracts ("one solver step / one token / one chunk" + composition).
- **N1 vs engine-internal pools** is consistent on a careful read (N1 is about datacenter orchestration; §6.3.5 states the reconciliation).
- **The trainer-scope line (N2 vs §8)** is drawn consistently — N2's own text enumerates exactly what §8 changes.
- **G6 vs the ≤2% Phase-2 gate** is goal-vs-acceptance-gate, not contradiction (Phase 1 is gated bit-identical).
- **The clean-room GPL posture holds**: sampler/scheduler math (DPM-Solver, Karras sigmas, flow-match shift) is published outside GPL sources.
- **C2 for the video denoise path is fine**: batch-1 fixed shapes are trivially batch-invariant — the doc's own analysis at lines 1145-1147 is correct; the AR/image/sharding exposures are correctly identified there too.
- **Self-forcing's cross-chunk gradients truncate by construction** (KV written under `no_grad` on detached context), so the engine KV pool is not blocked the way one might fear — the surviving residue is M12's grad-window cache mode.
- **"Every phase deletes or freezes something" survives audit** at the phase-deliverable level (the failures are the specific items in M19/M20).
- **The tier-1 ComfyUI vocabulary claim survives** blueprint-corpus measurement under the doc's actual claim (curated canonical workflows, not top-N node frequency).
- **The sglang reconvergence deferral** is substantively defended in §11.1 with reasons valid under either outcome.
- **WeightSyncPlan's "literal no-op"** is correctly scoped to colocated same-layout in the doc's own sentence; FSDP-vs-TP/SP is explicitly routed to in-place reshard.
---
## Ranked recommendations
1. **Design the abort/cancellation/OOM path with Phase 2** (C1) **and add memory as a budget axis with admission planning and preemption semantics** (M3). These two are the soundness conditions of the multiplexing bet; everything else in the execution plane sits on them.
2. **Re-derive §6.3.2 from the real vLLM constraint** (M1). The two-pool→one-pool reversal was made on a false premise; either accept uniform page bytes (and redesign the slab story) or bring back two pools with an explicit fragmentation/deadlock argument.
3. **Fix the migration plan's three structural defects**: CacheManager v0 into Phases 1–2 or AR batching out of Phase 2 (M16); a merge milestone for the cosmos3 chain before Phase 0 touches it (M17); a new-port inflow rule plus a freeze-enforcement mechanism that did not exist last time — CI path gate, codeowners, a date (M15, M21). Also reconcile Phase 5 with the frozen `training/` stack (M19) and restate the `forward_context` retirement honestly (M20).
4. **Specify the step skeleton and the policy contracts** — ordered, typed extension points; a policy state-scoping rule (state in LoopState, like plugins); a policy-observation channel — and work the mapping through Cosmos2.5 and LTX2 in the doc (M7, M8). Decide per-node request parameter binding in Phase 0 (M9).
5. **Give MoT a stated parallelism answer** (M5) and make the single-pool-spans-nodes decision explicit, including the fate of `RayDistributedExecutor` (M6).
6. **Close the workflow-cloud trust/correctness holes before Phase 3**: adapter-aware feature-cache keys (M10), weight-state transitions as a scheduled, costed operation (M11), DeployConfig-scoped plugin allowlisting (M22), per-surface schema stability tiers (M23), and an honest assessment of Punica's fit (M4).
7. **Right-size the RL claims**: design the grad+KV cache mode or scope self-forcing out of the shared loop (M12); budget the Behavior Record at real byte counts (M13); demote the omni-RL pilot or give it an objective sketch and an owner (M14); fix worked example (g).
8. **Reclassify loop inversion as unprecedented at scheduler granularity** in risk 3 and drop the diffusers "validation" (M2). The bet may still be right — but it should be made with open eyes, and the parity-gate plan is then carrying more weight than the doc admits.
9. **Correct the thirteen numbers above before circulating.** The doc's credibility rests on its "evidence-backed" brand; ~250-vs-111 is the kind of error that makes a reader re-check everything else — and most of everything else checks out.
@@ -0,0 +1,180 @@
#!/usr/bin/env bash
set -euo pipefail
# Build a text-only Parquet dataset from a one-prompt-per-line file. The script
# shards prompts across GPU_NUM single-GPU torchrun workers, runs
# v1_preprocess.py with --preprocess_task text_only, and writes prompt
# embeddings/captions under OUTPUT_DIR for DMD2/DiffusionNFT text-only runs.
INPUT_FILE="${1:-train.txt}"
OUTPUT_DIR="${OUTPUT_DIR:-data/train_text_only_dmd_preprocessed}"
MODEL_PATH="${MODEL_PATH:-Wan-AI/Wan2.1-T2V-1.3B-Diffusers}"
GPU_NUM="${GPU_NUM:-2}"
BATCH_SIZE="${BATCH_SIZE:-1}"
SAMPLES_PER_FILE="${SAMPLES_PER_FILE:-8}"
FLUSH_FREQUENCY="${FLUSH_FREQUENCY:-8}"
TEXT_MAX_LENGTH="${TEXT_MAX_LENGTH:-512}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
MIN_FREE_GPU_MB="${MIN_FREE_GPU_MB:-22000}"
if [[ ! -f "$INPUT_FILE" ]]; then
echo "Input text file not found: $INPUT_FILE" >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
visible_gpus=$(nvidia-smi --query-gpu=index --format=csv,noheader 2>/dev/null | wc -l | tr -d ' ')
if [[ "$visible_gpus" -lt "$GPU_NUM" ]]; then
echo "Expected at least $GPU_NUM GPUs, found $visible_gpus" >&2
exit 1
fi
for gpu_id in $(seq 0 $((GPU_NUM - 1))); do
free_mb=$(nvidia-smi --id="$gpu_id" --query-gpu=memory.free --format=csv,noheader,nounits | tr -d ' ')
if [[ "$free_mb" -lt "$MIN_FREE_GPU_MB" ]]; then
echo "GPU $gpu_id has only ${free_mb} MiB free; text-only Wan preprocessing needs" \
"about ${MIN_FREE_GPU_MB} MiB." >&2
echo "Free the GPU or lower MIN_FREE_GPU_MB if you know this run will fit." >&2
exit 1
fi
done
MARKER="$OUTPUT_DIR/.fastvideo_text_only_dmd_output"
if [[ -d "$OUTPUT_DIR" && ! -f "$MARKER" ]]; then
echo "Refusing to overwrite existing non-script output directory: $OUTPUT_DIR" >&2
echo "Set OUTPUT_DIR to a new path or remove the directory manually." >&2
exit 1
fi
rm -rf "$OUTPUT_DIR"
mkdir -p "$OUTPUT_DIR"
touch "$MARKER"
echo "Text-only DMD preprocessing config:"
echo " input: $INPUT_FILE"
echo " output: $OUTPUT_DIR"
echo " model: $MODEL_PATH"
echo " gpus: $GPU_NUM"
echo " batch size per GPU: $BATCH_SIZE"
echo " text max length: $TEXT_MAX_LENGTH"
echo " samples per parquet file: $SAMPLES_PER_FILE"
echo " flush frequency: $FLUSH_FREQUENCY"
SHARD_DIR="$OUTPUT_DIR/_text_shards"
mkdir -p "$SHARD_DIR"
python - "$INPUT_FILE" "$SHARD_DIR" "$GPU_NUM" <<'PY'
from pathlib import Path
import sys
input_path = Path(sys.argv[1])
shard_dir = Path(sys.argv[2])
num_shards = int(sys.argv[3])
prompts = [line.rstrip("\n") for line in input_path.read_text(encoding="utf-8").splitlines() if line.strip()]
if not prompts:
raise SystemExit(f"No non-empty prompts found in {input_path}")
for shard_idx in range(num_shards):
shard_prompts = prompts[shard_idx::num_shards]
shard_path = shard_dir / f"train_text_shard_{shard_idx}.txt"
shard_path.write_text("\n".join(shard_prompts) + "\n", encoding="utf-8")
print(f"Wrote {len(shard_prompts)} prompts to {shard_path}")
PY
run_preprocess_worker() {
local gpu_id="$1"
local shard_file="$2"
local shard_output="$3"
local log_file="$4"
local master_port="$5"
local -a cmd=(
torchrun
--nnodes=1
--nproc_per_node=1
--master_port "$master_port"
fastvideo/pipelines/preprocess/v1_preprocess.py
--model_path "$MODEL_PATH"
--data_merge_path "$shard_file"
--preprocess_video_batch_size "$BATCH_SIZE"
--seed 42
--max_height 448
--max_width 832
--num_frames 77
--dataloader_num_workers 0
--output_dir "$shard_output"
--train_fps 16
--samples_per_file "$SAMPLES_PER_FILE"
--flush_frequency "$FLUSH_FREQUENCY"
--text_max_length "$TEXT_MAX_LENGTH"
--video_length_tolerance_range 5
--preprocess_task text_only
)
{
echo "[gpu${gpu_id}] log file: $log_file"
echo "[gpu${gpu_id}] command: CUDA_VISIBLE_DEVICES=${gpu_id} ${cmd[*]}"
} | tee "$log_file"
CUDA_VISIBLE_DEVICES="$gpu_id" "${cmd[@]}" 2>&1 \
| sed -u "s/^/[gpu${gpu_id}] /" \
| tee -a "$log_file"
local status=${PIPESTATUS[0]}
if [[ "$status" -ne 0 ]]; then
echo "[gpu${gpu_id}] preprocessing failed with exit code $status" | tee -a "$log_file"
fi
return "$status"
}
pids=()
for gpu_id in $(seq 0 $((GPU_NUM - 1))); do
shard_file="$SHARD_DIR/train_text_shard_${gpu_id}.txt"
shard_output="$OUTPUT_DIR/shard_${gpu_id}"
mkdir -p "$shard_output"
log_file="$OUTPUT_DIR/preprocess_gpu_${gpu_id}.log"
echo "Launching text-only preprocessing on GPU ${gpu_id}: ${shard_file}"
run_preprocess_worker "$gpu_id" "$shard_file" "$shard_output" "$log_file" "$((29610 + gpu_id))" &
pids+=("$!")
done
failed=0
for pid in "${pids[@]}"; do
if ! wait "$pid"; then
failed=1
fi
done
if [[ "$failed" -ne 0 ]]; then
echo "One or more preprocessing workers failed. Check logs under $OUTPUT_DIR." >&2
exit 1
fi
num_parquet=$(find "$OUTPUT_DIR" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files were produced under $OUTPUT_DIR" >&2
exit 1
fi
echo "Text-only preprocessing complete."
echo "Parquet files: $num_parquet"
echo "Use this training data_path:"
echo "$OUTPUT_DIR"
+118
View File
@@ -0,0 +1,118 @@
#!/usr/bin/env bash
set -euo pipefail
CONFIG="${CONFIG:-examples/train/configs/rl/wan/diffusion_nft_pick_clip.yaml}"
DATA_PATH="${DATA_PATH:-data/pickscore_text_only_preprocessed}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/wan2.1_diffusion_nft_pick_clip}"
NUM_GPUS="${NUM_GPUS:-4}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29531}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$NUM_GPUS}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
NUM_FRAMES="${NUM_FRAMES:-1}"
NUM_LATENT_T="${NUM_LATENT_T:-1}"
PROJECT_NAME="${PROJECT_NAME:-diffusion_nft_wan}"
RUN_NAME="${RUN_NAME:-wan2.1_diffusion_nft_pick_clip}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
LOG_DIR="${LOG_DIR:-logs/train}"
if [[ ! -f "$CONFIG" ]]; then
echo "Training config not found: $CONFIG" >&2
exit 1
fi
if [[ ! -d "$DATA_PATH" ]]; then
echo "Preprocessed dataset directory not found: $DATA_PATH" >&2
echo "Run preprocessing first, for example:" >&2
echo " GPU_NUM=4 BATCH_SIZE=1 OUTPUT_DIR=$DATA_PATH \\" >&2
echo " bash scripts/preprocess/preprocess_train_text_only_dmd.sh DiffusionNFT/dataset/pickscore/train.txt" >&2
exit 1
fi
num_parquet=$(find "$DATA_PATH" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files found under $DATA_PATH" >&2
echo "Run preprocessing first, for example:" >&2
echo " GPU_NUM=4 BATCH_SIZE=1 OUTPUT_DIR=$DATA_PATH \\" >&2
echo " bash scripts/preprocess/preprocess_train_text_only_dmd.sh DiffusionNFT/dataset/pickscore/train.txt" >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-online}"
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_diffusion_nft_wan}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
mkdir -p "$LOG_DIR" "$OUTPUT_DIR"
timestamp="$(date +%Y%m%d_%H%M%S)"
log_file="$LOG_DIR/diffusion_nft_wan_pick_clip_${timestamp}.log"
cmd=(
torchrun
--nnodes "$NNODES"
--node_rank "$NODE_RANK"
--nproc_per_node "$NUM_GPUS"
--master_addr "$MASTER_ADDR"
--master_port "$MASTER_PORT"
-m fastvideo.train.entrypoint.train
--config "$CONFIG"
--training.data.data_path "$DATA_PATH"
--training.data.preprocessed_data_type text_only
--training.data.dataloader_num_workers "$DATALOADER_NUM_WORKERS"
--training.data.num_frames "$NUM_FRAMES"
--training.data.num_latent_t "$NUM_LATENT_T"
--training.distributed.num_gpus "$NUM_GPUS"
--training.distributed.sp_size "$SP_SIZE"
--training.distributed.tp_size "$TP_SIZE"
--training.distributed.hsdp_replicate_dim "$HSDP_REPLICATE_DIM"
--training.distributed.hsdp_shard_dim "$HSDP_SHARD_DIM"
--training.checkpoint.output_dir "$OUTPUT_DIR"
--training.tracker.project_name "$PROJECT_NAME"
--training.tracker.run_name "$RUN_NAME"
)
echo "DiffusionNFT Wan single-frame RL training config:"
echo " config: $CONFIG"
echo " data path: $DATA_PATH"
echo " parquet files: $num_parquet"
echo " output dir: $OUTPUT_DIR"
echo " frames / latent T: $NUM_FRAMES / $NUM_LATENT_T"
echo " rewards: pickscore + clipscore"
echo " learning rate: 3e-5"
echo " GPUs: $NUM_GPUS"
echo " SP/TP: $SP_SIZE/$TP_SIZE"
echo " HSDP replicate/shard: $HSDP_REPLICATE_DIM/$HSDP_SHARD_DIM"
echo " W&B mode: $WANDB_MODE"
echo " log file: $log_file"
echo "Command:"
printf ' %q' "${cmd[@]}" "$@"
echo
"${cmd[@]}" "$@" 2>&1 | tee "$log_file"
+170
View File
@@ -0,0 +1,170 @@
#!/usr/bin/env bash
set -euo pipefail
CONFIG="${CONFIG:-examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml}"
DATA_PATH="${DATA_PATH:-data/train_text_only_dmd_preprocessed}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/wan2.1_dmd2_text_only}"
NUM_GPUS="${NUM_GPUS:-2}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29521}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$NUM_GPUS}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
NUM_FRAMES="${NUM_FRAMES:-1}"
NUM_LATENT_T="${NUM_LATENT_T:-1}"
VALIDATION_NUM_FRAMES="${VALIDATION_NUM_FRAMES:-$NUM_FRAMES}"
VALIDATION_PROMPT_FILE="${VALIDATION_PROMPT_FILE:-}"
VALIDATION_FILE="${VALIDATION_FILE:-examples/train/configs/distribution_matching/wan/dmd2_text_only_validation.json}"
VALIDATION_OFFLOAD_TRAINING_STATE="${VALIDATION_OFFLOAD_TRAINING_STATE:-true}"
VALIDATION_UNLOAD_PIPELINE_AFTER="${VALIDATION_UNLOAD_PIPELINE_AFTER:-true}"
CFG_UNCOND_TEXT="${CFG_UNCOND_TEXT:-zero}"
CFG_UNCOND_ON_MISSING="${CFG_UNCOND_ON_MISSING:-ignore}"
PROJECT_NAME="${PROJECT_NAME:-distillation_wan_text_only}"
RUN_NAME="${RUN_NAME:-wan2.1_dmd2_text_only}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
LOG_DIR="${LOG_DIR:-logs/train}"
if [[ ! -f "$CONFIG" ]]; then
echo "Training config not found: $CONFIG" >&2
exit 1
fi
if [[ ! -d "$DATA_PATH" ]]; then
echo "Preprocessed dataset directory not found: $DATA_PATH" >&2
echo "Run scripts/preprocess/preprocess_train_text_only_dmd.sh first." >&2
exit 1
fi
num_parquet=$(find "$DATA_PATH" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files found under $DATA_PATH" >&2
echo "Run scripts/preprocess/preprocess_train_text_only_dmd.sh first." >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-offline}"
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_dmd2_text_only}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
mkdir -p "$LOG_DIR" "$OUTPUT_DIR"
if [[ -n "$VALIDATION_PROMPT_FILE" ]]; then
if [[ ! -f "$VALIDATION_PROMPT_FILE" ]]; then
echo "Validation prompt file not found: $VALIDATION_PROMPT_FILE" >&2
echo "Set VALIDATION_PROMPT_FILE to a text file, or leave it empty and set VALIDATION_FILE=<validation.json>." >&2
exit 1
fi
python - "$VALIDATION_PROMPT_FILE" "$VALIDATION_FILE" <<'PY'
import json
import os
import sys
prompt_file, validation_file = sys.argv[1:3]
with open(prompt_file, encoding="utf-8") as f:
prompts = [line.strip() for line in f if line.strip()]
if not prompts:
raise SystemExit(f"No validation prompts found in {prompt_file}")
validation_dir = os.path.dirname(os.path.abspath(validation_file))
if validation_dir:
os.makedirs(validation_dir, exist_ok=True)
with open(validation_file, "w", encoding="utf-8") as f:
json.dump(
{"data": [{"caption": prompt} for prompt in prompts]},
f,
indent=2,
ensure_ascii=False,
)
f.write("\n")
print(f"Wrote {len(prompts)} validation prompts to {validation_file}")
PY
elif [[ ! -f "$VALIDATION_FILE" ]]; then
echo "Validation dataset file not found: $VALIDATION_FILE" >&2
exit 1
fi
timestamp="$(date +%Y%m%d_%H%M%S)"
log_file="$LOG_DIR/dmd2_t2v_text_only_${timestamp}.log"
cmd=(
torchrun
--nnodes "$NNODES"
--node_rank "$NODE_RANK"
--nproc_per_node "$NUM_GPUS"
--master_addr "$MASTER_ADDR"
--master_port "$MASTER_PORT"
-m fastvideo.train.entrypoint.train
--config "$CONFIG"
--training.data.data_path "$DATA_PATH"
--training.data.preprocessed_data_type text_only
--training.data.dataloader_num_workers "$DATALOADER_NUM_WORKERS"
--training.data.num_frames "$NUM_FRAMES"
--training.data.num_latent_t "$NUM_LATENT_T"
--callbacks.validation.dataset_file "$VALIDATION_FILE"
--callbacks.validation.num_frames "$VALIDATION_NUM_FRAMES"
--callbacks.validation.offload_training_state "$VALIDATION_OFFLOAD_TRAINING_STATE"
--callbacks.validation.unload_pipeline_after_validation "$VALIDATION_UNLOAD_PIPELINE_AFTER"
--method.cfg_uncond.text "$CFG_UNCOND_TEXT"
--method.cfg_uncond.on_missing "$CFG_UNCOND_ON_MISSING"
--training.distributed.num_gpus "$NUM_GPUS"
--training.distributed.sp_size "$SP_SIZE"
--training.distributed.tp_size "$TP_SIZE"
--training.distributed.hsdp_replicate_dim "$HSDP_REPLICATE_DIM"
--training.distributed.hsdp_shard_dim "$HSDP_SHARD_DIM"
--training.checkpoint.output_dir "$OUTPUT_DIR"
--training.tracker.project_name "$PROJECT_NAME"
--training.tracker.run_name "$RUN_NAME"
)
echo "DMD2 T2V text-only training config:"
echo " config: $CONFIG"
echo " data path: $DATA_PATH"
echo " parquet files: $num_parquet"
echo " output dir: $OUTPUT_DIR"
echo " frames / latent T: $NUM_FRAMES / $NUM_LATENT_T"
echo " validation prompt file: ${VALIDATION_PROMPT_FILE:-<none>}"
echo " validation dataset: $VALIDATION_FILE"
echo " validation frames: $VALIDATION_NUM_FRAMES"
echo " validation offload training state: $VALIDATION_OFFLOAD_TRAINING_STATE"
echo " validation unload pipeline after: $VALIDATION_UNLOAD_PIPELINE_AFTER"
echo " GPUs: $NUM_GPUS"
echo " SP/TP: $SP_SIZE/$TP_SIZE"
echo " HSDP replicate/shard: $HSDP_REPLICATE_DIM/$HSDP_SHARD_DIM"
echo " CFG uncond text/on_missing: $CFG_UNCOND_TEXT/$CFG_UNCOND_ON_MISSING"
echo " W&B mode: $WANDB_MODE"
echo " log file: $log_file"
echo "Command:"
printf ' %q' "${cmd[@]}" "$@"
echo
"${cmd[@]}" "$@" 2>&1 | tee "$log_file"
@@ -0,0 +1,99 @@
from types import SimpleNamespace
import torch
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
class _FakeEMA:
def __init__(self):
self.updates = 0
def update(self, module):
del module
self.updates += 1
def test_reward_diagnostic_metrics_match_per_prompt_groups():
method = object.__new__(DiffusionNFTMethod)
method._trained_prompt_hashes = set()
sample_items = [{
"prompts": ["a", "a"],
}, {
"prompts": ["b", "b"],
}]
rewards = {"avg": torch.tensor([1.0, 3.0, 2.0, 6.0])}
metrics = method._reward_diagnostic_metrics(sample_items, rewards)
assert metrics["group_size"] == 2.0
assert metrics["trained_prompt_num"] == 2.0
assert torch.isclose(metrics["zero_std_ratio"], torch.tensor(0.0))
assert torch.isclose(metrics["reward_std_mean"], torch.tensor(1.5))
assert torch.isclose(metrics["mean_reward_100"], torch.tensor(3.0))
assert torch.isclose(metrics["mean_reward_50"], torch.tensor(4.5))
method._reward_diagnostic_metrics(sample_items, rewards)
assert len(method._trained_prompt_hashes) == 2
def test_update_ema_honors_update_after_step():
method = object.__new__(DiffusionNFTMethod)
method._ema_enabled = True
method._student_ema = _FakeEMA()
method._ema_update_count = 0
method._ema_update_after_step = 1
method.student = SimpleNamespace(transformer=object())
method._update_ema()
assert method._student_ema.updates == 0
assert method._ema_update_count == 1
method._update_ema()
assert method._student_ema.updates == 1
assert method._ema_update_count == 2
def test_num_train_timesteps_uses_explicit_schedule_length():
method = object.__new__(DiffusionNFTMethod)
method._sample_steps = 25
method._timestep_fraction = 0.5
method._sampling_config = SimpleNamespace(
timesteps=[900, 800, 700, 600, 500, 400, 300, 200, 100, 10],
sigmas=None,
)
assert method._num_train_timesteps() == 5
def test_checkpoint_state_saves_frozen_old_policy_weights():
student = torch.nn.Linear(2, 2)
old = torch.nn.Linear(2, 2)
for param in old.parameters():
param.requires_grad_(False)
method = object.__new__(DiffusionNFTMethod)
method._role_models = {
"student": SimpleNamespace(
transformer=student,
_trainable=True,
),
"old": SimpleNamespace(
transformer=old,
_trainable=False,
),
}
method.student = method._role_models["student"]
method.old = method._role_models["old"]
method._student_optimizer = None
method._student_lr_scheduler = None
method._ema_enabled = False
states = method.checkpoint_state()
old_state = states["roles.old.transformer"].state_dict()
assert "weight" in old_state
assert "bias" in old_state
assert torch.equal(old_state["weight"], old.weight)
+57
View File
@@ -0,0 +1,57 @@
import torch
import pytest
from fastvideo.train.methods.rl.rewards import MultiRewardScorer, select_first_frame
def test_select_first_frame_for_video_tensor():
video = torch.arange(2 * 3 * 4 * 5 * 6).reshape(2, 3, 4, 5, 6)
frame = select_first_frame(video)
assert frame.shape == (2, 3, 5, 6)
torch.testing.assert_close(frame, video[:, :, 0])
def test_select_first_frame_keeps_frame_tensor():
frame = torch.randn(2, 3, 5, 6)
selected = select_first_frame(frame)
assert selected is frame
def test_multi_reward_weighted_sum_with_injected_scorers():
def pickscore(media, prompts):
assert media.shape == (2, 3, 4, 5, 6)
assert prompts == ["a", "b"]
return torch.tensor([1.0, 2.0])
def clipscore(media, prompts):
assert media.shape == (2, 3, 4, 5, 6)
assert prompts == ["a", "b"]
return torch.tensor([0.5, 1.5])
scorer = MultiRewardScorer(
{"pickscore": 2.0, "clipscore": 3.0},
scorers={
"pickscore": pickscore,
"clipscore": clipscore,
},
)
scores = scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
torch.testing.assert_close(scores["pickscore"], torch.tensor([1.0, 2.0]))
torch.testing.assert_close(scores["clipscore"], torch.tensor([0.5, 1.5]))
torch.testing.assert_close(scores["avg"], torch.tensor([3.5, 8.5]))
def test_multi_reward_validates_score_shape():
scorer = MultiRewardScorer(
{"pickscore": 1.0},
scorers={"pickscore": lambda media, prompts: torch.tensor([[1.0], [2.0]])},
)
with pytest.raises(ValueError, match="must return shape"):
scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
+238
View File
@@ -0,0 +1,238 @@
import torch
import pytest
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.methods.rl.common import (
DiffusionSampler,
SamplingConfig,
distributed_k_repeat_indices,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
from fastvideo.train.utils.config import load_run_config
class _FakeScheduler:
def __init__(self):
self.num_train_timesteps = 1000
self.set_timesteps_calls = []
self.timesteps = torch.empty(0)
self.sigmas = torch.empty(0)
self.step_calls = 0
def set_timesteps(self, num_inference_steps=None, device=None, timesteps=None, sigmas=None):
self.set_timesteps_calls.append({
"num_inference_steps": num_inference_steps,
"timesteps": timesteps,
"sigmas": sigmas,
})
if timesteps is not None:
self.timesteps = torch.tensor(timesteps, device=device, dtype=torch.float32)
else:
self.timesteps = torch.linspace(1000, 0, int(num_inference_steps), device=device)
if sigmas is not None:
self.sigmas = torch.tensor(sigmas, device=device, dtype=torch.float32)
else:
self.sigmas = torch.cat([self.timesteps / 1000.0, torch.zeros(1, device=device)])
def step(self, model_output, timestep, sample, return_dict=False):
del timestep
self.step_calls += 1
prev = sample + model_output
return (prev, ) if not return_dict else {"prev_sample": prev}
class _FakeModel:
def __init__(self):
self.noise_scheduler = _FakeScheduler()
self.add_noise_calls = 0
self.timestep_shapes = []
def predict_noise(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return torch.zeros_like(noisy_latents)
def predict_x0(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return noisy_latents
def add_noise(self, clean_latents, noise, timestep):
del timestep
self.add_noise_calls += 1
return clean_latents + noise
def _batch():
batch = TrainingBatch()
batch.latents = torch.zeros(2, 1, 3, 4, 4)
return batch
def test_sampler_preserves_latent_dtype():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
batch = _batch()
batch.latents = batch.latents.to(torch.bfloat16)
result = sampler.sample(model, batch, generator=torch.Generator().manual_seed(0))
assert result.latents.dtype is torch.bfloat16
def test_sampler_uses_scheduler_generated_timesteps_by_default():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [1000.0, 666.6666259765625, 333.3333435058594, 0.0]
def test_sampler_honors_explicit_timestep_override():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=3, timesteps=[900, 300, 10]))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [900.0, 300.0, 10.0]
def test_sampler_honors_explicit_timesteps_without_matching_num_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(timesteps=[900, 300, 10]))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [900.0, 300.0, 10.0]
assert model.noise_scheduler.set_timesteps_calls == []
def test_sampling_config_rejects_unknown_keys():
with pytest.raises(ValueError, match="Unsupported method.sampling key"):
SamplingConfig.from_mapping({"solver": "dpm2"})
def test_sampler_restores_original_batch_timestep_after_sampling():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=2))
batch = _batch()
original_timesteps = torch.tensor([123.0])
batch.timesteps = original_timesteps
sampler.sample(batch=batch, model=model, generator=torch.Generator().manual_seed(0))
assert batch.timesteps is original_timesteps
def test_euler_sampler_does_not_renoise_between_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert model.add_noise_calls == 0
assert model.timestep_shapes == [(2,), (2,), (2,), (2,)]
def test_sde_reflow_sampler_renoises_between_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4, trajectory="sde_reflow"))
sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert model.add_noise_calls == 3
def test_diffusion_nft_config_uses_rl_sampler_not_dmd_pipeline():
config_path = "examples/train/configs/rl/wan/diffusion_nft_pick_clip.yaml"
cfg = load_run_config(config_path)
raw_text = open(config_path, encoding="utf-8").read()
assert cfg.method["_target_"] == "fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod"
assert cfg.training.optimizer.learning_rate == 3.0e-5
assert cfg.training.data.num_latent_t == 1
assert cfg.training.data.num_frames == 1
assert "sampling_timesteps" not in raw_text
assert "WanDMDPipeline" not in raw_text
assert "solver" not in cfg.method["sampling"]
assert cfg.method["sampling"]["scheduler"] == "flow_match_euler"
assert cfg.method["sampling"]["trajectory"] == "ode"
assert cfg.method["sampling"]["flow_shift"] == "inherit"
assert "deterministic" not in cfg.method["sampling"]
assert "noise_level" not in cfg.method["sampling"]
assert cfg.method["validation"]["every_steps"] == 10
assert cfg.method["validation"]["num_steps"] == 40
assert cfg.method["validation"]["num_prompts"] == 16
assert cfg.method["validation"]["log_samples"] is True
def test_validation_shard_indices_are_stable_and_padded():
rank0 = validation_shard_indices(5, rank=0, world_size=2)
rank1 = validation_shard_indices(5, rank=1, world_size=2)
assert rank0 == [(0, True), (2, True), (4, True)]
assert rank1 == [(1, True), (3, True), (0, False)]
def test_distributed_k_repeat_indices_repeats_prompts_globally():
rank0 = distributed_k_repeat_indices(
dataset_length=100,
batch_size=6,
repeats_per_prompt=24,
world_size=4,
rank=0,
seed=123,
)
all_indices = []
for rank in range(4):
sample = distributed_k_repeat_indices(
dataset_length=100,
batch_size=6,
repeats_per_prompt=24,
world_size=4,
rank=rank,
seed=123,
)
all_indices.extend(sample.local_indices)
assert rank0.unique_prompt_count == 1
assert len(all_indices) == 24
assert len(set(all_indices)) == 1
def test_validation_caption_puts_rewards_first():
caption = validation_caption(
"a small blue cube",
{
"avg": 0.75,
"pickscore": 0.5,
},
)
assert caption.startswith("avg: 0.7500 | pickscore: 0.5000 | ")
assert caption.endswith("a small blue cube")
def test_media_to_video_array_treats_frame_as_single_frame_video():
frame = torch.ones(3, 4, 5)
video = media_to_video_array(frame)
assert video.shape == (1, 3, 4, 5)
assert video.dtype.name == "uint8"
def test_media_to_video_array_preserves_video_frames():
media = torch.ones(3, 2, 4, 5)
video = media_to_video_array(media)
assert video.shape == (2, 3, 4, 5)
@@ -0,0 +1,157 @@
from dataclasses import dataclass, field
import torch
from fastvideo.train.trainer import Trainer
from fastvideo.train.utils.training_config import (
CheckpointConfig,
DistributedConfig,
ModelTrainingConfig,
OptimizerConfig,
TrackerConfig,
TrainingConfig,
TrainingLoopConfig,
)
class _Tracker:
def __init__(self):
self.logged = []
self.finished = False
def log(self, metrics, step):
self.logged.append((step, metrics))
def finish(self):
self.finished = True
class _Callbacks:
def __init__(self):
self.before_optimizer_steps = 0
self.training_step_ends = 0
def on_train_start(self, method, iteration=0):
pass
def on_before_optimizer_step(self, method, iteration=0):
self.before_optimizer_steps += 1
def on_training_step_end(self, method, metrics, iteration=0):
self.training_step_ends += 1
def on_validation_begin(self, method, iteration=0):
pass
def on_validation_end(self, method, iteration=0):
pass
def on_train_end(self, method, iteration=0):
pass
class _World:
rank = 0
local_rank = 0
class _Method(torch.nn.Module):
def __init__(self):
super().__init__()
self.calls = 0
self.backward_calls = 0
self.optimizer_steps = 0
self.tracker = None
def set_tracker(self, tracker):
self.tracker = tracker
def on_train_start(self):
pass
def manages_optimization(self):
return True
def managed_train_step(self, data_stream, iteration):
batch = next(data_stream)
self.calls += 1
return (
{"total_loss": torch.tensor(float(batch["x"]))},
{},
{"managed_metric": float(iteration)},
)
def backward(self, *args, **kwargs):
self.backward_calls += 1
def optimizers_schedulers_step(self, iteration):
self.optimizer_steps += 1
def optimizers_zero_grad(self, iteration):
pass
class _MethodWithValidation(_Method):
def __init__(self):
super().__init__()
self.validation_iterations = []
def on_validation_begin(self, iteration=0):
self.validation_iterations.append(iteration)
return {"validation/fake": float(iteration)}
def test_trainer_skips_default_optimizer_path_for_managed_methods(monkeypatch):
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.build_tracker", lambda *args, **kwargs: _Tracker())
cfg = TrainingConfig(
distributed=DistributedConfig(),
optimizer=OptimizerConfig(),
loop=TrainingLoopConfig(max_train_steps=1, gradient_accumulation_steps=3),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
)
trainer = Trainer(cfg)
trainer.callbacks = _Callbacks()
method = _Method()
dataloader = [{"x": 2}]
trainer.run(method, dataloader=dataloader, max_steps=1)
assert method.calls == 1
assert method.backward_calls == 0
assert method.optimizer_steps == 0
assert trainer.callbacks.before_optimizer_steps == 0
assert trainer.callbacks.training_step_ends == 1
assert trainer.tracker.logged[0][1]["total_loss"] == 2.0
def test_trainer_logs_method_validation_at_step_zero(monkeypatch):
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.build_tracker", lambda *args, **kwargs: _Tracker())
cfg = TrainingConfig(
distributed=DistributedConfig(),
optimizer=OptimizerConfig(),
loop=TrainingLoopConfig(max_train_steps=1, gradient_accumulation_steps=1),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
)
trainer = Trainer(cfg)
trainer.callbacks = _Callbacks()
method = _MethodWithValidation()
dataloader = [{"x": 2}]
trainer.run(method, dataloader=dataloader, max_steps=1)
assert method.validation_iterations == [0, 1]
assert trainer.tracker.logged[0] == (0, {"validation/fake": 0.0})
@@ -0,0 +1,62 @@
import torch
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.train.models.wan.wan import WanModel
class _CPUWanModel(WanModel):
@property
def device(self) -> torch.device:
return torch.device("cpu")
class _AutocastProbe(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.autocast_enabled: bool | None = None
self.autocast_dtype: torch.dtype | None = None
def forward(
self,
*,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor,
timestep: torch.Tensor,
return_dict: bool,
) -> torch.Tensor:
del encoder_hidden_states, encoder_attention_mask, timestep, return_dict
self.autocast_enabled = torch.is_autocast_enabled("cpu")
self.autocast_dtype = torch.get_autocast_dtype("cpu")
self.hidden_states_dtype = hidden_states.dtype
return hidden_states
def test_wan_predict_noise_uses_training_dtype_autocast_for_fp32_inputs():
model = object.__new__(_CPUWanModel)
model.transformer = _AutocastProbe()
batch = TrainingBatch()
batch.timesteps = torch.tensor([1], dtype=torch.long)
batch.conditional_dict = {
"encoder_hidden_states": torch.randn(1, 4, 8, dtype=torch.float32),
"encoder_attention_mask": torch.ones(1, 4, dtype=torch.float32),
}
noisy_latents = torch.randn(1, 1, 2, 4, 4, dtype=torch.float32)
timestep = torch.tensor([1], dtype=torch.long)
pred_noise = model.predict_noise(
noisy_latents,
timestep,
batch,
conditional=True,
)
assert pred_noise.shape == noisy_latents.shape
assert pred_noise.dtype is torch.bfloat16
assert model.transformer.hidden_states_dtype is torch.bfloat16
assert model.transformer.autocast_enabled is True
assert model.transformer.autocast_dtype is torch.bfloat16