Compare commits
93
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
67f5b53595 | ||
|
|
aa6bf5079b | ||
|
|
4ca56d8951 | ||
|
|
1e19650eb3 | ||
|
|
de354f804c | ||
|
|
6feee4e4fa | ||
|
|
82e00615e5 | ||
|
|
30cfbfc4f6 | ||
|
|
a4d8978c37 | ||
|
|
803a6c99ae | ||
|
|
8a637c3215 | ||
|
|
51bae33c7b | ||
|
|
54b7fe3a55 | ||
|
|
1fe50c0092 | ||
|
|
290795daf8 | ||
|
|
0047d54a1b | ||
|
|
b2d55a7ba0 | ||
|
|
24cfe281b3 | ||
|
|
dc086b207f | ||
|
|
b9db151658 | ||
|
|
9124963238 | ||
|
|
f90e8f3e76 | ||
|
|
4ff33a28b7 | ||
|
|
940f94f435 | ||
|
|
ac9dcb63ea | ||
|
|
861f87e843 | ||
|
|
7ef2d9083e | ||
|
|
0ac5367b54 | ||
|
|
56f84018df | ||
|
|
d04127fe6f | ||
|
|
ae0cf3e2d1 | ||
|
|
f4af3cf886 | ||
|
|
39580ad73a | ||
|
|
521c2845e6 | ||
|
|
0467edbd07 | ||
|
|
bfbd90ea3a | ||
|
|
5fd6e23e30 | ||
|
|
fc3332550c | ||
|
|
7464ef8308 | ||
|
|
4a56274d10 | ||
|
|
9ede9af123 | ||
|
|
3d405f6dfa | ||
|
|
d891771ba3 | ||
|
|
67dea39052 | ||
|
|
5f1d2ef7d2 | ||
|
|
440b99523e | ||
|
|
1cb2b4e84c | ||
|
|
51898f48e9 | ||
|
|
4e331eb7ff | ||
|
|
c6d2976fc2 | ||
|
|
6096b00aeb | ||
|
|
1d23399d81 | ||
|
|
fefcd415ff | ||
|
|
7ac2ff0d1c | ||
|
|
fa9c58b419 | ||
|
|
0662b42510 | ||
|
|
ae6d8085de | ||
|
|
d0648ba8d9 | ||
|
|
e10828346f | ||
|
|
655f362cf4 | ||
|
|
3541e81d66 | ||
|
|
f51497ee6d | ||
|
|
d8af2e60d2 | ||
|
|
f79919ba8a | ||
|
|
1594f8e6be | ||
|
|
64cadaa0bf | ||
|
|
750fc1245b | ||
|
|
7725998b0c | ||
|
|
4b61dedc43 | ||
|
|
1e58d90d02 | ||
|
|
47f7a04e09 | ||
|
|
9b8838834c | ||
|
|
634f0828a2 | ||
|
|
2f044c02dd | ||
|
|
431f4daddb | ||
|
|
6220d02746 | ||
|
|
9489c6c1dd | ||
|
|
69c9871154 | ||
|
|
18dd295e8d | ||
|
|
4c333e0509 | ||
|
|
32a7a6b87b | ||
|
|
58223c0c41 | ||
|
|
b254d1affe | ||
|
|
c7e0a8e894 | ||
|
|
b3ddf6014d | ||
|
|
a9e5f6ee7a | ||
|
|
7467076d72 | ||
|
|
01dc0c3377 | ||
|
|
f1dc587c74 | ||
|
|
098bcf014a | ||
|
|
270fae959d | ||
|
|
4a14c1afa3 | ||
|
|
f1c19050c3 |
@@ -0,0 +1,207 @@
|
||||
# v2 ← M\*: Architecture Gap-Analysis & Improvement Roadmap
|
||||
|
||||
**Status:** exploration, flagged for review. **Date:** 2026-06-19.
|
||||
**Source paper:** *M\*: A Modular, Extensible, Serving System for Multimodal Models* (arXiv 2606.12688,
|
||||
Stanford/UW/CMU; Jha, Sagan, Kamahori, …, Kasikci, S. Wang). It is a universal serving runtime for composite
|
||||
multimodal models built on the **Walk Graph** abstraction (a model is a dataflow graph `G`; a request is a
|
||||
*Walk* — a labeled subgraph — and the runtime executes walks). It beats vLLM-Omni (~20% lower T2I latency on
|
||||
**BAGEL**, up to 2.64× on I2I), SGLang-Omni (2.7× TTS throughput on **Qwen3-Omni**), and native V-JEPA2
|
||||
rollout (12.5×). It explicitly names **FastVideo's own** sparse/sliding-tile attention, xDiT/PipeFusion/USP,
|
||||
Inferix, and FlashDrive as techniques integratable into the graph runtime.
|
||||
|
||||
**Method:** a 28-agent workflow — 6 parallel v2-subsystem maps → 10 M\*-dimension analyses, each
|
||||
*adversarially verified against the actual v2 code* → synthesis + a completeness critic. The critic's
|
||||
corrections and three P0 claims were then **spot-verified by hand** (file:line below). This doc folds those
|
||||
corrections in; it is the corrected, authoritative synthesis.
|
||||
|
||||
---
|
||||
|
||||
## 1. Executive summary
|
||||
|
||||
v2 already implements the **harder half** of M\*'s thesis and in several axes **exceeds** it:
|
||||
|
||||
- v2's `Program` *is* M\*'s graph `G` (typed `ComponentNode`/`ModelLoopNode` + edges).
|
||||
- v2's `shared_weight_components` *is* M\*'s cross-Walk node sharing — BAGEL/Cosmos3/LTX2 each bind two
|
||||
`ModelLoopNode`s to **one resident transformer** (`instance.component()` returns the same live object). This
|
||||
is the exact MoT serving property the omni cards in this repo already express.
|
||||
- v2 adds three things M\* (serving-only) has **no equivalent for**: a required+validated per-loop **cost
|
||||
model**, a non-negotiable **interleave bit-parity gate**, and an **integrated training plane** (RL→distill
|
||||
flywheel driving the *same* serving Loop).
|
||||
- The `extend/` plugin seam (interceptors/observers/registry with capability negotiation) is precisely the
|
||||
hook M\*'s "extensible / integrate FastVideo-STA, xDiT, Inferix, FlashDrive" call-out asks for — **v2
|
||||
already has the seam M\* only gestures at.**
|
||||
|
||||
What v2 lacks is M\*'s **declarative authoring layer above the substrate**, and — the key insight — *much of
|
||||
that substrate is already authored but inert*: v2 has declared the metadata for "minimum components per
|
||||
request" (`required_for`/`optional_for` on every omni card) and "branch as a cache axis" (`guidance_sig`,
|
||||
`CacheKey`) but **never wired it to an executor**. The substrate is ~80% built and switched off.
|
||||
|
||||
**Highest-leverage cluster:** three small, parity-safe wires that turn on inert substrate and unblock the
|
||||
BAGEL/Qwen-Omni/Cosmos3 latency wins M\* measured **on the exact models this repo already runs** — plus one
|
||||
P1 that aligns v2 with the paper's headline "extensible" claim using a seam v2 already has.
|
||||
|
||||
### Verified P0 correctness findings (spot-checked by hand)
|
||||
1. **Runner divergence (real bug).** `v2/runtime/engine.py:88` → `nodes = self.program.nodes`;
|
||||
`v2/runtime/disaggregated.py:96` → `nodes = self.program.active_nodes(self.request)`. The inline and
|
||||
disaggregated runners execute *different node sets*. ✅ confirmed.
|
||||
2. **EOS is faked.** `v2/recipes/omni/ar_loop.py` docstring says "done on EOS/max_tokens"; `next()` (`:46-48`)
|
||||
checks **only** `max_tokens`. M\*'s marquee `DynamicLoop` use case (EOS) is unimplemented in the loop that
|
||||
serves the Qwen-Omni Thinker/Talker and Cosmos3 reasoner. ✅ confirmed.
|
||||
3. **`required_for`/`optional_for` have zero runtime consumers** (grep outside `specs.py`/recipes/tests is
|
||||
empty). The min-components metadata is declared on every card and never read. ✅ confirmed.
|
||||
|
||||
---
|
||||
|
||||
## 2. Dimension table (corrected)
|
||||
|
||||
| # | Dimension | v2 status | Gap | Priority | Effort | Payoff | Action |
|
||||
|---|---|---|---|---|---|---|---|
|
||||
| 1 | Min-components per request (`required_for` + `when_task`) | substrate built, **inert** | real, cheap | **P0** | S | Consume `required_for` in `active_nodes`; unify `engine.py:88` onto `active_nodes`; deliver via registry/card builder so all ~40 cards inherit it |
|
||||
| 2 | Real EOS + declarative `DynamicLoop` | early-exit emergent; **EOS faked** | real | **P0** | S | `ARDecodeLoop` honors `eos_id` + `req.sampling.stop`; add `LoopSpec.dynamic_stop` + `register_loop_stop`. **Training-enabling** (world-model rollout horizon) |
|
||||
| 3 | CFG/branch as label over one paged KV pool | absent (`PagedKVCache` is a counter) | real | **P1** | L | `(namespace,label)` paged store w/ one budget; reuse `guidance_sig` for hash (NOT `partition_field`); by-ref via existing `InProcKVConnector`. AR path only (diffusion has no KV) |
|
||||
| 4 | `extend/` plugin seam → integrate FastVideo-STA / Inferix | **seam exists, unused for attn** | real (paper headline) | **P1** | M | Expose FastVideo sparse/sliding-tile attention + Inferix block-diffusion as `Interceptor`/`EngineKind` plugins — the paper's named integration targets, on this repo's own code |
|
||||
| 5 | `ParitySpec.output_determinism` (C3 distributional) | C3 rung defined, **0 users** | real, dormant | **P1** | S | Add field; `compare_outputs` consults it. **Training-enabling** (SDE/FlowGRPO stochastic rollouts) |
|
||||
| 6 | Registry-driven delivery of #1 | present, not leveraged | integration | **P1** | S | Express `when_task`/min-components through `WorkflowRegistry`/card builders, not 3 bespoke recipe patches |
|
||||
| 7 | Serving conductor + pluggable data plane | conductor exists (`serving/http.py`); **single-process transport** | real | **P2** | L | v2 already has the step-scheduled worker surface; gap is ZeroMQ/Mooncake + direct worker→worker tensor routing (today `InProcKVConnector` only) |
|
||||
| 8 | Fleet/Dynamo placement + replicas | **live** (`deploy/fleet.py`,`dynamo.py`) | partial | **P2** | M | Fleet-level placement/affinity/replica is real & ≥M\*; missing piece is only the intra-engine `(node,Walk)→rank` map decoupled from model code |
|
||||
| 9 | Per-node TP / SP + cross-rank transport | axis vocab **exists** (`sp` incl.); not wired to runtime | partial | **P2** | XL | Wire declarative degrees into runtime; Wan/LTX are **SP-native** (TP is a no-op there); populate `parallel_plan_hash` on the serving cache path |
|
||||
| 10 | Named Walks + per-model state machine | `Program`=G, sharing real; no Walk/SM | real | **P2** | M | Defer until a *re-entrant* phase graph (Thinker↔Talker, rollout) needs it; #1 captures the min-components win without it |
|
||||
| 11 | Declarative `Parallel/Sequential/Loop` IR | imperative loop classes | real (authoring) | **P2** | M | Thin Section IR lowering to flat `Program`; scope to one AR recipe |
|
||||
| 12 | Streaming `ChunkPolicy` + `StreamBuffer` | causal-chunk emit **already ships** (`wan_causal`); `EdgeKind.STREAM` inert | real | **P2** | L | Declarative `ChunkPolicy` vocab over the existing chunk mechanism; needs concurrent producer/consumer runner (= pipelined scheduling). Inferix integration point |
|
||||
| 13 | Speculative deferred-termination; loop-spanning CUDA graphs; N+1 prefetch; attn double-buffer | absent / per-step capture (14 cards) | real | **P3** | L | Gate behind a real GPU executor; unobservable on CPU-toy CI; loop-span needs an `allows_interleaving=False` carve-out |
|
||||
| — | Cost model + interleave/consistency parity | **exceeds M\*** | none | **guard** | — | Do not regress; keep `step_cost_model` mandatory + `bit_identical` default |
|
||||
| — | Integrated training plane (flywheel, weight-sync) | **exceeds M\*** | none | **guard** | — | Protect train==serve loop identity with a toy fixture |
|
||||
|
||||
---
|
||||
|
||||
## 3. P0/P1 deep-dives (sequenced)
|
||||
|
||||
```
|
||||
PR-1 (P0) min-components ──┐
|
||||
PR-2 (P0) real EOS ─┼─► prereqs for honest "DynamicLoop" + min-component claims; both training-enabling
|
||||
PR-3 (P1) output_determinism (independent)
|
||||
PR-5 (P1) extend/ plugin: FastVideo-STA / Inferix as Interceptors (independent; highest paper-alignment)
|
||||
PR-4 (P1) CFG-as-label paged pool ──► depends on PR-2 (AR loop is the only KV consumer)
|
||||
```
|
||||
PR-1, PR-2, PR-3, PR-5 are mutually independent; PR-4 depends on PR-2.
|
||||
|
||||
### PR-1 (P0) — Turn on the inert min-components substrate + fix runner divergence
|
||||
- **Change.** Extend `Program.active_nodes(request)` (`v2/program/specs.py`) to also drop any node whose bound
|
||||
`ComponentSpec.required_for` (`v2/card/specs.py:144`) excludes `request.task` (and isn't in `optional_for`).
|
||||
**Fix the bug:** change `v2/runtime/engine.py:88` to `nodes = self.program.active_nodes(self.request)` so the
|
||||
inline `ProgramRunner` matches `DisaggregatedRunner` (`disaggregated.py:96`). Deliver the `when_task` gating
|
||||
through the **registry/card builder** (`recipes/__init__.py`, `program/workflow.py:WorkflowRegistry`) so all
|
||||
~40 cards inherit it uniformly — not three bespoke `program.py` patches.
|
||||
- **Why (this repo's models).** BAGEL T2I currently steps the AR-text loop and Cosmos3 t2v materializes the
|
||||
reasoner even though the cards declare `transformer required_for={'reason','t2i'}`, `vae required_for={'t2i'}`.
|
||||
On the GPU backend that is wasted resident-weight load + wasted steps on every single-modality request —
|
||||
exactly M\*'s "execute the MINIMUM components per request," delivered by consuming existing metadata.
|
||||
- **Risk/invariant.** Validate in `ModelCard.validate()` that every active node's `reads` are produced by an
|
||||
active node for each declared `TaskType` (avoid dropping a producer). Pure node-id filtering ⇒ serial and
|
||||
interleaved still walk the same filtered list ⇒ §9.3 interleave bit-parity holds by construction. CPU-toy clean.
|
||||
|
||||
### PR-2 (P0) — Real EOS + declarative `dynamic_stop` *(also training-enabling)*
|
||||
- **Change.** In `v2/recipes/omni/ar_loop.py`, `advance()` reads the emitted token; if it equals the model
|
||||
`eos_id` (toy backend exposes `EOS=0`) or matches `req.sampling.stop` (`params.py:21`, currently dead),
|
||||
register termination; `next()` returns `Done()` on stop OR `max_tokens`. Add `StopRegistry` to `LoopState` +
|
||||
`register_loop_stop(name)` to the `LoopContext` protocol (`contracts.py:204`) and to
|
||||
`DisaggregatedRunner`'s `RuntimeLoopContext`. Add `LoopSpec.dynamic_stop: bool=False`, opt the AR cards in.
|
||||
- **Why.** The docstring-vs-code lie sits in the loop serving Qwen-Omni Thinker/Talker and the Cosmos3 reasoner;
|
||||
M\*'s second named `DynamicLoop` use case (world-model **rollout horizon**) is exactly what `self_forcing` RL
|
||||
needs — so this is both a serving-credibility fix and a training enabler (raise its payoff accordingly).
|
||||
- **Risk/invariant.** `dynamic_stop=False` is byte-identical back-compat. Must pass **all three** parity gates:
|
||||
serial==interleaved AND disaggregated==inline. **Not** in this PR: speculative deferred-termination (unobservable
|
||||
on CPU-toy, fights the interleave invariant — P3, gated on GPU executor).
|
||||
|
||||
### PR-3 (P1) — `ParitySpec.output_determinism` (close the dormant C3 hole) *(training-enabling)*
|
||||
- **Change.** Add `output_determinism: str = "bit_identical"` to `ParitySpec` (`card/specs.py:88`); make
|
||||
`compare_outputs` (`parity/interleave_gate.py:54`) consult it (`bit_identical` → today's exact check;
|
||||
`distributional` → a moment/tolerance check — land a simple moment match first; a real KS test is new code).
|
||||
- **Why.** `ConsistencyLevel.C3` is defined and used by zero recipes; an SDE/FlowGRPO stochastic rollout cannot
|
||||
honestly declare its parity contract and would falsely fail the bit-identical gate. Additive; default unchanged.
|
||||
|
||||
### PR-5 (P1) — Expose FastVideo's own attention + Inferix as `extend/` plugins *(highest paper-alignment)*
|
||||
- **Change.** Use the existing `extend/{interceptors,observers,registry}.py` seam (capability-negotiated, with
|
||||
per-(request,branch) `plugin_state` that already passes the interleave gate) to register FastVideo's
|
||||
sparse/sliding-tile attention and Inferix-style block-diffusion as `Interceptor`s / an `EngineKind` plugin.
|
||||
- **Why.** M\*'s title is "Modular, **Extensible**" and it explicitly lists FastVideo-STA, xDiT/PipeFusion/USP,
|
||||
Inferix, FlashDrive as integratable. v2 already has the seam M\* only describes — this is where v2 most
|
||||
directly answers the paper, using this repo's own attention code. Low risk (the seam + capability negotiation
|
||||
already exist and are tested).
|
||||
|
||||
### PR-4 (P1) — CFG/branch as a LABEL over one paged KV pool
|
||||
- **Change.** Rewrite `PagedKVCache` (`cache/classes.py:155-172`) from a block *counter* into a real
|
||||
`(namespace,label)->[block-handle]` store with **one shared `total_blocks` budget** (M\*'s single-pool
|
||||
property). Reuse the existing-but-unpopulated `CacheKey.guidance_sig` (`keys.py:53`) for the hash. Thread the
|
||||
label through `ar_loop.py` (alloc/append/get per `(request_id, branch)`; prefill once per shared-prefix label;
|
||||
combine via `CFGPolicy.combine`). Wire `ResourceRequest.cache_blocks` (`contracts.py:64`, zero consumers) into
|
||||
admission per (class,label).
|
||||
- **Why.** The dossier-identified driver of M\*'s BAGEL win (3 CFG contexts as 3 labels over ONE pool vs dense
|
||||
per-context). Targets AR_DECODE (BAGEL `generate_text`, omni Thinker); **correctly excludes diffusion**
|
||||
(Wan/LTX are bidirectional, no KV — their CFG stays dense-but-batched).
|
||||
- **Corrections to bake in.** Do **NOT** add `branch_label` to `CacheKey.partition_field()` (CFG branches share
|
||||
embeddings; partitioning by branch is a semantic bug). Do **NOT** add a new by-ref type — reuse
|
||||
`InProcKVConnector` + `TransferManifest.cache_key`. Wiring `cache_blocks` admission is greenfield ⇒ effort **L**.
|
||||
CPU version proves label/sharing semantics; the real latency win needs a FlashInfer paged kernel (out of scope)
|
||||
— **merge** with a future "real KVCacheEngine" effort rather than landing isolated.
|
||||
|
||||
---
|
||||
|
||||
## 4. What v2 already does ≥ M\* — do NOT regress
|
||||
1. **Required+validated cost model** on every `LoopSpec` (13-kind `WorkUnitKind`) — typed, pre-GPU-validated.
|
||||
2. **Interleave bit-parity as a hard gate** (`parity.interleave_required=True` on 40+ cards). M\* has no such
|
||||
gate (its speculative scheduling deliberately wastes steps). Load-bearing invariant; every new primitive
|
||||
must pass it.
|
||||
3. **C0–C4 consistency ladder** wired into RL methods, with first-divergence tap reporting. No M\* equivalent.
|
||||
4. **Integrated training plane** — DiffusionNFT/DMD2/self_forcing, RL→distill flywheel, `WeightSyncController`
|
||||
hot weight-sync with drain-to-boundary + scoped cache invalidation, driving the **same** serving Loop.
|
||||
M\* is serving-only. Protect with a toy fixture asserting `rollout_loop` drives the served Loop object.
|
||||
5. **CPU-toy parity for the whole stack** — loops/CFG/caches/parity/RL run in CI without a GPU. Every new
|
||||
primitive must ship a toy exercise (this is what makes all PRs above testable without H100s).
|
||||
6. **Partition-not-flush cache invalidation** + four independent per-class pools.
|
||||
7. **`extend/` plugin seam** with capability negotiation (a 4-step distilled card *rejects* a residual-skip
|
||||
interceptor) — M\* describes extensibility; v2 has the mechanism.
|
||||
8. **Dynamo citizenship** (`deploy/dynamo.py`: one `DeploymentCard`+cost model, two consumers) — beyond M\*'s
|
||||
self-contained runtime.
|
||||
|
||||
---
|
||||
|
||||
## 5. Dropped / merged / deferred (and why)
|
||||
- **DROP declarative `Parallel` as a CFG-execution win.** The runner walks nodes linearly (ignores
|
||||
`Program.edges`), so `Parallel` lowers to sequential sugar and the CFG 3-pass braid is already one
|
||||
co-scheduled `WorkPlan.run`; splitting it risks the interleave gate. Salvage only the no-op refactor
|
||||
extracting `branch_forward` from `WanDenoiseLoop._velocity`. Reassign `Parallel` to the placement workstream.
|
||||
- **MERGE the full Walk/state-machine layer** into "defer until a re-entrant phase graph needs it" (PR-1 gets the
|
||||
min-components win with ~20 lines, no new abstraction). If built: the validator must check a walk's node-id
|
||||
order is a *subsequence* of `program.nodes` (not just membership) or the runner can reorder and break parity.
|
||||
- **MERGE `StreamBuffer`/`ChunkPolicy` into pipelined-scheduling.** Causal-chunk emit *already ships*
|
||||
(`wan_causal/loop.py` per-chunk `StepResult.emit` + slab-KV); the gap is the declarative `ChunkPolicy` vocab
|
||||
+ a concurrent producer/consumer runner. If built: keep all policies pure (per-request `StreamBuffer` history,
|
||||
not shared edge state) and restrict the bit-identical claim to the token-only handoff.
|
||||
- **MERGE CFG-fan-out exec + cross-rank transport + PD loop-splitting into a multi-GPU-runtime program.** These
|
||||
need real collectives (`v2/distributed/` is a stub) and KV-by-reference (KV lives in `CacheManager`, not the
|
||||
transferable `slots`). **Keep cheaply now:** the *declarative* halves — per-component degree, `(node,Walk)`
|
||||
placement key with node-only fallback, `ReplicaSet` under `LocalFleet`, and populate `parallel_plan_hash` on
|
||||
the **serving** cache path (it is already populated in `training/behavior.py:40` — the gap is serving-only).
|
||||
- **DEFER** speculative deferred-termination, loop-spanning CUDA graphs, N+1 prefetch, attention-plan
|
||||
double-buffer — all gated on a real GPU executor; benefit unobservable on CPU-toy CI. Keep the cheap
|
||||
`EngineKind` tag (`STATELESS|KV_CACHE|DIFFUSION`) now. Correct the stale `cudagraph.py:51-52` docstring
|
||||
(per-step capture ships in 14 cards, not just wan21).
|
||||
- **RESCOPE per-node TP.** Wan/LTX use `ReplicatedLinear` + **sequence parallelism** (`sp`), not TP; the `sp`
|
||||
axis already exists in `parallel/plan.py:AXIS_NAMES`. The work is wiring degrees into the runtime, not
|
||||
inventing vocabulary; a `tp_size=2` "one-line activation" is a no-op for the shipped models.
|
||||
|
||||
---
|
||||
|
||||
## 6. The first integration test, if/when multi-GPU placement work starts
|
||||
The **live Qwen-Omni 2-GPU bring-up** (Thinker on rank 0, Talker+Code2Wav on rank 1; see
|
||||
`v2_debug_videos/vlm.md` Session 4) is the natural first validation target for any `(node,Walk)→rank`
|
||||
placement work — it is the one place this repo already has real multi-rank composite-model execution.
|
||||
|
||||
---
|
||||
|
||||
## Anchor files for P0/P1
|
||||
`v2/program/specs.py`, `v2/runtime/engine.py` (**line 88 fix**), `v2/runtime/disaggregated.py`,
|
||||
`v2/recipes/omni/ar_loop.py`, `v2/loop/contracts.py`, `v2/card/specs.py`, `v2/cache/{classes.py,keys.py}`,
|
||||
`v2/parity/interleave_gate.py`, `v2/extend/{interceptors,registry}.py`, `recipes/__init__.py` +
|
||||
`v2/program/workflow.py` (registry-driven delivery).
|
||||
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
|
||||
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
@@ -148,7 +148,8 @@ For each metric with at least two non-null source values:
|
||||
4. Stop if any source record regresses against the batch median by more than
|
||||
`max_intra_batch_regression`.
|
||||
|
||||
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
|
||||
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
|
||||
otherwise `0.05`. Print a table with per-source values, batch median, and
|
||||
worst intra-batch regression.
|
||||
|
||||
This check prevents uploading a mixed batch where one JSON is materially
|
||||
@@ -182,7 +183,7 @@ present, that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records under /tmp
|
||||
|
||||
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
@@ -191,7 +192,7 @@ The sync command pattern is:
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory under `/tmp`:
|
||||
@@ -199,8 +200,8 @@ Then back up only the sanitized model directory under `/tmp`:
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.performance.hf_store import sanitize
|
||||
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
|
||||
from hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
@@ -234,7 +235,7 @@ first baseline seed. Continue, but report that baseline history was empty.
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from fastvideo.performance.hf_store import load_records_for_model
|
||||
from hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
@@ -371,7 +372,7 @@ prepared records plus backup on disk.
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```python
|
||||
from fastvideo.performance.hf_store import upload_record
|
||||
from hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
@@ -459,7 +460,7 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
|
||||
@@ -1,9 +1,5 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
+1
-28
@@ -114,17 +114,6 @@ steps:
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: LoRA Extraction Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
retry:
|
||||
automatic:
|
||||
- exit_status: 128
|
||||
limit: 3
|
||||
- exit_status: -1
|
||||
limit: 2
|
||||
agents:
|
||||
queue: "default"
|
||||
- label: ":test_tube: Training Tests"
|
||||
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
@@ -382,21 +371,6 @@ steps:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "scripts/lora_extraction/**"
|
||||
- "fastvideo/tests/lora_extraction/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/training/training_utils.py"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 90m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Extraction Tests"
|
||||
env:
|
||||
- TEST_TYPE=lora_extraction
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
@@ -436,7 +410,7 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile"
|
||||
config:
|
||||
command: "timeout 25m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: ":test_tube: LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
@@ -481,7 +455,6 @@ steps:
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/worker/**"
|
||||
- "fastvideo/entrypoints/**"
|
||||
- "fastvideo/performance/**"
|
||||
- "fastvideo/tests/performance/**"
|
||||
- ".buildkite/performance-benchmarks/**"
|
||||
- "pyproject.toml"
|
||||
|
||||
@@ -80,23 +80,6 @@ MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUI
|
||||
|
||||
POST_RUN_HOOK=""
|
||||
|
||||
is_truthy() {
|
||||
case "${1:-}" in
|
||||
1|true|TRUE|yes|YES|on|ON) return 0 ;;
|
||||
*) return 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
ssim_bootstrap_args() {
|
||||
local title="${PR_TITLE:-}"
|
||||
local message="${BUILDKITE_MESSAGE:-}"
|
||||
if is_truthy "${FASTVIDEO_SSIM_BOOTSTRAP_MODE:-}" \
|
||||
|| [[ "$title" == *"[new-model]"* ]] \
|
||||
|| [[ "$message" == *"[new-model]"* ]]; then
|
||||
printf ' --bootstrap-mode'
|
||||
fi
|
||||
}
|
||||
|
||||
upload_performance_artifacts() {
|
||||
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
|
||||
LOCAL_DIR="downloaded_reports"
|
||||
@@ -189,12 +172,7 @@ case "$TEST_TYPE" in
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
SSIM_BOOTSTRAP_ARGS=$(ssim_bootstrap_args)
|
||||
if [ -n "$SSIM_BOOTSTRAP_ARGS" ]; then
|
||||
log "SSIM bootstrap mode enabled for new-model reference draft generation"
|
||||
fi
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run "
|
||||
MODAL_COMMAND+="$MODAL_SSIM_TEST_FILE::run_ssim_tests$SSIM_BOOTSTRAP_ARGS"
|
||||
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_SSIM_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
|
||||
@@ -1,133 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# Gate the expensive Buildkite full suite on the cheap GitHub checks.
|
||||
#
|
||||
# Polls the workflow runs for the PR head commit and only exits 0 once the
|
||||
# watched cheap workflows (pre-commit, docs build) have succeeded, so the
|
||||
# 'ready' label cannot burn ~20 GPU lanes on a head that a cheap check has
|
||||
# already doomed.
|
||||
#
|
||||
# Semantics:
|
||||
# - watched run completed with a bad conclusion -> exit 1 (fail CLOSED:
|
||||
# no full suite; the next push re-arms via the 'synchronize' trigger)
|
||||
# - watched run cancelled -> still pending: the docs
|
||||
# workflow's repo-global 'pages' concurrency group cancels runs superseded
|
||||
# by unrelated pushes, so 'cancelled' is not a verdict on this PR
|
||||
# - watched runs pending -> poll until done
|
||||
# - docs run absent -> not applicable after a
|
||||
# short grace period ('Deploy Documentation' is path-filtered on PRs)
|
||||
# - pre-commit run absent -> keep polling: pre-commit
|
||||
# is never path-filtered, so its absence is always anomalous
|
||||
# - 'ready' label removed while waiting -> exit 1 (fail CLOSED:
|
||||
# un-labeling is a deliberate maintainer action)
|
||||
# - GitHub API unreachable or timeout -> exit 0 (fail OPEN,
|
||||
# loud warning: never brick CI on a GitHub outage)
|
||||
#
|
||||
# Required env: PR_SHA (PR head commit), PR_NUMBER, GITHUB_REPOSITORY, GH_TOKEN.
|
||||
set -euo pipefail
|
||||
|
||||
: "${PR_SHA:?PR_SHA (PR head commit) is required}"
|
||||
: "${PR_NUMBER:?PR_NUMBER (pull request number) is required}"
|
||||
: "${GITHUB_REPOSITORY:?GITHUB_REPOSITORY is required}"
|
||||
|
||||
# Workflow-level `name:` values that must be green before the full suite
|
||||
# may start. "Deploy Documentation" is path-filtered on PRs, so its run may
|
||||
# legitimately never exist; pre-commit always runs, so it must appear.
|
||||
WATCHED_NAMES='["pre-commit", "Deploy Documentation"]'
|
||||
WATCHED_REGEX='^(pre-commit|Deploy Documentation)$'
|
||||
POLL_SECS="${POLL_SECS:-20}"
|
||||
GRACE_SECS="${GRACE_SECS:-60}"
|
||||
MAX_WAIT_SECS="${MAX_WAIT_SECS:-1500}"
|
||||
|
||||
# Bound each API call so a hung connection hits the 3-strike fail-open path
|
||||
# instead of pinning the loop until the job timeout (which would fail closed
|
||||
# on exactly the GitHub-outage case this script is meant to survive).
|
||||
if command -v timeout >/dev/null 2>&1; then
|
||||
gh_api() { timeout 30 gh api "$@"; }
|
||||
else
|
||||
gh_api() { gh api "$@"; } # macOS dev boxes; CI always has coreutils timeout
|
||||
fi
|
||||
|
||||
# The workflow checked the label before starting the gate, but the wait can
|
||||
# last ~25 min: re-check once before any exit 0 and fail closed if 'ready'
|
||||
# was removed in the meantime. An API error here proceeds (the label was
|
||||
# present when the gate started; never brick CI on an outage).
|
||||
recheck_ready_label() {
|
||||
local pr_json
|
||||
if pr_json=$(gh_api "repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" 2>/dev/null); then
|
||||
if ! jq -e '[.labels[]?.name] | index("ready")' <<<"$pr_json" >/dev/null 2>&1; then
|
||||
echo "::error::PR #${PR_NUMBER} no longer has the 'ready' label —" \
|
||||
"NOT triggering the Buildkite full suite. Re-add the label to re-arm."
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
echo "::warning::Could not re-check the 'ready' label on PR #${PR_NUMBER}; proceeding (it was present when the gate started)."
|
||||
fi
|
||||
}
|
||||
|
||||
start=$(date +%s)
|
||||
api_fails=0
|
||||
missing=""
|
||||
|
||||
while true; do
|
||||
elapsed=$(( $(date +%s) - start ))
|
||||
|
||||
if runs_json=$(gh_api "repos/${GITHUB_REPOSITORY}/actions/runs?head_sha=${PR_SHA}&per_page=100" 2>/dev/null) \
|
||||
&& state=$(jq --arg re "$WATCHED_REGEX" '
|
||||
[.workflow_runs[]? | select(.name // "" | test($re))]
|
||||
| group_by(.name) | map(max_by(.id))
|
||||
| map({name, status, conclusion})' <<<"$runs_json" 2>/dev/null); then
|
||||
api_fails=0
|
||||
echo "t+${elapsed}s watched checks: $(jq -c . <<<"$state")"
|
||||
|
||||
failed=$(jq -r '[.[] | select(.status == "completed"
|
||||
and (.conclusion | IN("success", "skipped", "neutral", "cancelled") | not))]
|
||||
| map(.name) | join(", ")' <<<"$state")
|
||||
if [ -n "$failed" ]; then
|
||||
echo "::error::Cheap check(s) failed on ${PR_SHA}: ${failed}." \
|
||||
"NOT triggering the Buildkite full suite. Push a fix (the 'ready'" \
|
||||
"label re-arms on every push), or re-run the failed check and then" \
|
||||
"re-run this workflow."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 'cancelled' counts as pending: wait for a re-run to reach a real verdict
|
||||
# (bounded by MAX_WAIT, then the fail-open below).
|
||||
pending=$(jq '[.[] | select(.status != "completed" or .conclusion == "cancelled")] | length' <<<"$state")
|
||||
missing=$(jq -r --argjson watched "$WATCHED_NAMES" '($watched - map(.name)) | join(", ")' <<<"$state")
|
||||
if [ "$pending" -eq 0 ]; then
|
||||
if [ -z "$missing" ]; then
|
||||
recheck_ready_label
|
||||
echo "All watched cheap checks are green — full suite may proceed."
|
||||
exit 0
|
||||
fi
|
||||
case "$missing" in
|
||||
*pre-commit*)
|
||||
echo "pre-commit run not found for ${PR_SHA} yet; waiting (pre-commit is never path-filtered, so its absence is anomalous)."
|
||||
;;
|
||||
*)
|
||||
if [ "$elapsed" -ge "$GRACE_SECS" ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::Watched run(s) never appeared for ${PR_SHA}: ${missing} (path-filtered, likely not applicable). Proceeding on the checks that did run."
|
||||
exit 0
|
||||
fi
|
||||
echo "Waiting up to ${GRACE_SECS}s grace for path-filtered run(s) to appear: ${missing}."
|
||||
;;
|
||||
esac
|
||||
fi
|
||||
else
|
||||
api_fails=$(( api_fails + 1 ))
|
||||
echo "::warning::GitHub API error querying workflow runs for ${PR_SHA} (attempt ${api_fails}/3)."
|
||||
if [ "$api_fails" -ge 3 ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the full suite WITHOUT the cheap-check gate."
|
||||
exit 0
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$elapsed" -ge "$MAX_WAIT_SECS" ]; then
|
||||
recheck_ready_label
|
||||
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the full suite anyway."
|
||||
exit 0
|
||||
fi
|
||||
sleep "$POLL_SECS"
|
||||
done
|
||||
@@ -1,122 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
# Self-test for gate_full_suite.sh using a mocked `gh`. No network, runs on
|
||||
# any dev box: bash .github/scripts/test_gate_full_suite.sh
|
||||
set -u
|
||||
here=$(cd "$(dirname "$0")" && pwd)
|
||||
tmp=$(mktemp -d)
|
||||
trap 'rm -rf "$tmp"' EXIT
|
||||
|
||||
# Mock gh. Asserts the exact endpoint (including head_sha) it is called
|
||||
# with — an endpoint typo in the gate script fails the test rather than
|
||||
# silently serving canned data. On the runs endpoint it serves
|
||||
# $MOCK_DIR/response_<call#>.json, sticking on the highest existing file,
|
||||
# and exits 1 if none exist (simulates a GitHub API outage). On the pulls
|
||||
# endpoint it serves $MOCK_DIR/pr.json, defaulting to a 'ready'-labeled PR.
|
||||
cat > "$tmp/gh" <<'EOF'
|
||||
#!/usr/bin/env bash
|
||||
if [ "${1:-}" != "api" ]; then
|
||||
echo "unexpected gh invocation: $*" >> "$MOCK_DIR/endpoint_error"
|
||||
exit 2
|
||||
fi
|
||||
case "${2:-}" in
|
||||
"repos/o/r/actions/runs?head_sha=deadbeef&per_page=100")
|
||||
n=$(( $(cat "$MOCK_DIR/count" 2>/dev/null || echo 0) + 1 ))
|
||||
echo "$n" > "$MOCK_DIR/count"
|
||||
while [ "$n" -gt 0 ]; do
|
||||
if [ -f "$MOCK_DIR/response_$n.json" ]; then
|
||||
cat "$MOCK_DIR/response_$n.json"
|
||||
exit 0
|
||||
fi
|
||||
n=$(( n - 1 ))
|
||||
done
|
||||
echo "api outage" >&2
|
||||
exit 1
|
||||
;;
|
||||
"repos/o/r/pulls/42")
|
||||
if [ -f "$MOCK_DIR/pr.json" ]; then
|
||||
cat "$MOCK_DIR/pr.json"
|
||||
else
|
||||
echo '{"labels": [{"name": "ready"}]}'
|
||||
fi
|
||||
;;
|
||||
*)
|
||||
echo "unexpected gh endpoint: $2" >> "$MOCK_DIR/endpoint_error"
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
EOF
|
||||
chmod +x "$tmp/gh"
|
||||
|
||||
PC_OK='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "success"}'
|
||||
PC_BAD='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "failure"}'
|
||||
PC_PENDING='{"name": "pre-commit", "id": 1, "status": "in_progress", "conclusion": null}'
|
||||
DOCS_OK='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "success"}'
|
||||
DOCS_BAD='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "failure"}'
|
||||
DOCS_CANCELLED='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "cancelled"}'
|
||||
OTHER='{"name": "Trigger Full Suite", "id": 3, "status": "in_progress", "conclusion": null}'
|
||||
NULL_NAME='{"name": null, "id": 4, "status": "completed", "conclusion": "failure"}'
|
||||
PC_OK_RERUN='{"name": "pre-commit", "id": 5, "status": "completed", "conclusion": "success"}'
|
||||
|
||||
fails=0
|
||||
want_log="" # optional: expect() also greps out.log for this regex, then resets
|
||||
pr_json="" # optional: served for the pulls (label re-check) endpoint, then resets
|
||||
raw_body="" # optional: serve responses verbatim instead of wrapping in workflow_runs
|
||||
expect() { # <name> <expected-exit> <response json>...
|
||||
local name=$1 want=$2 dir i=1
|
||||
shift 2
|
||||
dir=$(mktemp -d "$tmp/test_XXXXXX")
|
||||
for body in "$@"; do
|
||||
if [ -n "$raw_body" ]; then
|
||||
printf '%s' "$body" > "$dir/response_$i.json"
|
||||
else
|
||||
printf '{"workflow_runs": [%s]}' "$body" > "$dir/response_$i.json"
|
||||
fi
|
||||
i=$(( i + 1 ))
|
||||
done
|
||||
[ -n "$pr_json" ] && printf '%s' "$pr_json" > "$dir/pr.json"
|
||||
( export PATH="$tmp:$PATH" MOCK_DIR="$dir" PR_SHA=deadbeef PR_NUMBER=42 \
|
||||
GITHUB_REPOSITORY=o/r POLL_SECS=0 GRACE_SECS=1 MAX_WAIT_SECS=3
|
||||
bash "$here/gate_full_suite.sh" > "$dir/out.log" 2>&1 )
|
||||
local rc=$?
|
||||
if [ "$rc" -ne "$want" ]; then
|
||||
echo "FAIL: $name (exit $rc, want $want)"
|
||||
cat "$dir/out.log"
|
||||
fails=1
|
||||
elif [ -f "$dir/endpoint_error" ]; then
|
||||
echo "FAIL: $name (mock gh got an unexpected call)"
|
||||
cat "$dir/endpoint_error"
|
||||
fails=1
|
||||
elif [ -n "$want_log" ] && ! grep -Eq "$want_log" "$dir/out.log"; then
|
||||
echo "FAIL: $name (log does not match: $want_log)"
|
||||
cat "$dir/out.log"
|
||||
fails=1
|
||||
else
|
||||
echo "ok: $name"
|
||||
fi
|
||||
want_log="" pr_json="" raw_body=""
|
||||
}
|
||||
|
||||
expect "both green -> proceed" 0 "$PC_OK, $DOCS_OK, $OTHER, $NULL_NAME"
|
||||
expect "docs build failed -> blocked" 1 "$PC_OK, $DOCS_BAD"
|
||||
expect "pre-commit failed -> blocked" 1 "$PC_BAD"
|
||||
expect "pending then green -> proceed" 0 "$PC_PENDING" "$PC_OK, $DOCS_OK"
|
||||
want_log="never appeared.*Deploy Documentation"
|
||||
expect "docs run absent (path-filtered) -> proceed after grace" 0 "$PC_OK"
|
||||
expect "API outage -> fail open" 0
|
||||
want_log="FAILING OPEN"
|
||||
expect "pending past MAX_WAIT -> fail open" 0 "$PC_PENDING"
|
||||
want_log="FAILING OPEN"
|
||||
expect "unrelated runs only -> no grace, fail open at MAX_WAIT" 0 "$OTHER"
|
||||
expect "cancelled docs then green -> proceed" 0 \
|
||||
"$PC_OK, $DOCS_CANCELLED" "$PC_OK, $DOCS_OK"
|
||||
want_log="FAILING OPEN"
|
||||
expect "cancelled docs forever -> fail open at MAX_WAIT" 0 "$PC_OK, $DOCS_CANCELLED"
|
||||
want_log="FAILING OPEN"
|
||||
expect "pre-commit absent -> no grace, fail open at MAX_WAIT" 0 "$DOCS_OK"
|
||||
expect "duplicate run names -> latest wins" 0 "$PC_BAD, $PC_OK_RERUN, $DOCS_OK"
|
||||
raw_body=1
|
||||
expect "garbage response body -> fail open" 0 "this is not json"
|
||||
pr_json='{"labels": [{"name": "other"}]}'
|
||||
expect "ready label removed mid-gate -> blocked" 1 "$PC_OK, $DOCS_OK"
|
||||
|
||||
exit "$fails"
|
||||
@@ -1,11 +1,7 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
# pull_request_target instead of pull_request: the workflow definition and
|
||||
# the hook config are always taken from the BASE branch, so fork /
|
||||
# first-time-contributor PRs run immediately without a maintainer clicking
|
||||
# "Approve and run". The PR head is checked out as data only.
|
||||
pull_request_target:
|
||||
pull_request:
|
||||
branches: [main]
|
||||
workflow_call:
|
||||
inputs:
|
||||
@@ -19,25 +15,12 @@ permissions:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
if: github.event.pull_request.draft != true
|
||||
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ inputs.ref || '' }}
|
||||
# For PR events, lint the PR head — but keep the hook definitions from
|
||||
# the base branch so an untrusted PR cannot alter what gets executed.
|
||||
- name: Save trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
|
||||
- uses: actions/checkout@v4
|
||||
if: github.event_name == 'pull_request_target'
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.head.sha }}
|
||||
persist-credentials: false
|
||||
- name: Restore trusted hook config
|
||||
if: github.event_name == 'pull_request_target'
|
||||
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.12"
|
||||
@@ -47,6 +30,3 @@ jobs:
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
# After pre-commit so a self-test failure cannot mask lint failures.
|
||||
- name: Full-suite gate self-test
|
||||
run: bash .github/scripts/test_gate_full_suite.sh
|
||||
|
||||
@@ -52,7 +52,6 @@ jobs:
|
||||
core.setOutput('pr_sha', pr.head.sha);
|
||||
core.setOutput('pr_branch', pr.head.ref);
|
||||
core.setOutput('pr_number', String(prNumber));
|
||||
core.setOutput('pr_title', pr.title);
|
||||
|
||||
- name: Trigger Full Suite
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
@@ -61,7 +60,6 @@ jobs:
|
||||
PR_SHA: ${{ steps.label.outputs.pr_sha }}
|
||||
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
|
||||
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
|
||||
PR_TITLE: ${{ steps.label.outputs.pr_title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -73,7 +71,6 @@ jobs:
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
@@ -83,12 +80,11 @@ jobs:
|
||||
pull_request_id: $pr_id,
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring),
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
}
|
||||
}')"
|
||||
|
||||
parse-command:
|
||||
if: >-
|
||||
@@ -129,7 +125,7 @@ jobs:
|
||||
set -euo pipefail
|
||||
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
|
||||
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
|
||||
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
|
||||
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
|
||||
exit 1
|
||||
@@ -140,7 +136,6 @@ jobs:
|
||||
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
|
||||
[ssim]=ssim [training]=training
|
||||
[lora-inference]=inference_lora [lora-training]=training_lora
|
||||
[lora-extraction]=lora_extraction
|
||||
[distillation]=distillation_dmd [self-forcing]=self_forcing
|
||||
[vsa]=training_vsa [vmoba]=inference_vmoba
|
||||
[performance]=performance [api]=api_server
|
||||
@@ -245,7 +240,6 @@ jobs:
|
||||
TEST_SCOPE: ${{ needs.parse-command.outputs.test_scope }}
|
||||
FULL_SUITE: ${{ needs.parse-command.outputs.full_suite }}
|
||||
TEST_TYPE: ${{ needs.parse-command.outputs.test_type }}
|
||||
PR_TITLE: ${{ github.event.issue.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -262,7 +256,6 @@ jobs:
|
||||
--arg full_suite "$FULL_SUITE" \
|
||||
--arg test_type "$TEST_TYPE" \
|
||||
--arg pr_number "$PR_NUMBER" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
'{
|
||||
commit: $commit,
|
||||
branch: $branch,
|
||||
@@ -272,9 +265,8 @@ jobs:
|
||||
pull_request_base_branch: "main",
|
||||
env: {
|
||||
TEST_SCOPE: $test_scope,
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number,
|
||||
PR_TITLE: $pr_title
|
||||
}
|
||||
}')"
|
||||
FULL_SUITE: $full_suite,
|
||||
TEST_TYPE: $test_type,
|
||||
PR_NUMBER: $pr_number
|
||||
}
|
||||
}')"
|
||||
|
||||
@@ -7,7 +7,6 @@ on:
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
actions: read
|
||||
|
||||
concurrency:
|
||||
group: full-suite-${{ github.event.pull_request.number }}
|
||||
@@ -19,8 +18,6 @@ jobs:
|
||||
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|
||||
|| github.event.action == 'synchronize'
|
||||
runs-on: ubuntu-latest
|
||||
# Gate below may wait for cheap checks (up to MAX_WAIT_SECS = 25 min).
|
||||
timeout-minutes: 35
|
||||
steps:
|
||||
- name: Check ready label
|
||||
id: check
|
||||
@@ -52,20 +49,6 @@ jobs:
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
|
||||
done
|
||||
|
||||
# Checks out the BASE branch (default for pull_request_target), so PR
|
||||
# authors cannot tamper with the gate script.
|
||||
- name: Checkout gate script
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
|
||||
|
||||
- name: Wait for pre-commit and docs build
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
run: bash .github/scripts/gate_full_suite.sh
|
||||
|
||||
- name: Trigger Buildkite Full Suite
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
@@ -73,7 +56,6 @@ jobs:
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_TITLE: ${{ github.event.pull_request.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
run: |
|
||||
@@ -85,7 +67,6 @@ jobs:
|
||||
--arg commit "$PR_SHA" \
|
||||
--arg branch "$PR_BRANCH" \
|
||||
--arg message "Full Suite for PR #${PR_NUMBER}" \
|
||||
--arg pr_title "$PR_TITLE" \
|
||||
--argjson pr_id "$PR_NUMBER" \
|
||||
'{
|
||||
commit: $commit,
|
||||
@@ -97,7 +78,6 @@ jobs:
|
||||
env: {
|
||||
TEST_SCOPE: "full",
|
||||
FULL_SUITE: "true",
|
||||
PR_NUMBER: ($pr_id | tostring),
|
||||
PR_TITLE: $pr_title
|
||||
PR_NUMBER: ($pr_id | tostring)
|
||||
}
|
||||
}')"
|
||||
|
||||
@@ -13,33 +13,12 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
|
||||
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
|
||||
# push trigger is a sufficient change detector on its own -- no separate
|
||||
# detect-changes/paths-filter job is needed now that there is a single
|
||||
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
|
||||
# rocm Dockerfile stay manual-dispatch only.
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'docker/Dockerfile'
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
# One static group, no cancellation: every run of this workflow writes the same
|
||||
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
|
||||
# concurrent push/dispatch runs would race on those tags, and cancelling a run
|
||||
# mid-publish can strand the cu126/cu130 tag families at different commits. An
|
||||
# in-flight superseded build wastes its runner time, but its tags are then
|
||||
# overwritten by the newer queued run. GitHub keeps a single pending run per
|
||||
# group: the newest queued run replaces any older queued one.
|
||||
concurrency:
|
||||
group: infra-build-image
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
|
||||
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
|
||||
@@ -49,11 +28,7 @@ jobs:
|
||||
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
|
||||
# 2.8.3 comes from the architecture-specific prebuilt releases.
|
||||
build-cuda-images:
|
||||
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
|
||||
# on a push that changed docker/Dockerfile (inputs are null on push). The
|
||||
# repository guard keeps fork syncs from auto-building; manual dispatch
|
||||
# still works in forks.
|
||||
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -100,10 +75,10 @@ jobs:
|
||||
secrets: inherit
|
||||
|
||||
publish-cuda-manifests:
|
||||
# !cancelled(): publish lanes whose digests exist even if a sibling build
|
||||
# leg failed (the digest-count check fails incomplete lanes); it also
|
||||
# bypasses skipped-needs propagation, hence the explicit skipped check.
|
||||
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
|
||||
# !cancelled(): a failed sibling build leg must not skip the manifests for a
|
||||
# CUDA lane whose own digests all exist; the digest-count check below fails
|
||||
# the incomplete lane loudly instead.
|
||||
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
|
||||
needs: build-cuda-images
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
|
||||
+2
-8
@@ -37,11 +37,6 @@ logs/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# Cosmos3 local parity assets (symlinked from main worktree)
|
||||
/official_weights/
|
||||
/converted_weights/
|
||||
/cosmos-framework
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
**/.cache/**
|
||||
@@ -77,7 +72,8 @@ docs/distillation/examples/
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos (negations must come after the catch-all on line below)
|
||||
# Reference videos
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/assets/images/**/*.png
|
||||
@@ -131,8 +127,6 @@ apps/dreamverse/web/.env.production.local
|
||||
.sisyphus/
|
||||
openspec/
|
||||
fastvideo/tests/ssim/reference_videos/**
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.png
|
||||
|
||||
# Editor logs and local Python version pins (accidentally committed)
|
||||
*.nvimlog
|
||||
|
||||
@@ -10,6 +10,8 @@ exclude: |
|
||||
scripts/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
v2/(layers|attention|platforms|configs|distributed|models|logging_utils|third_party|hooks|api)/.*|
|
||||
v2/(envs|logger|utils|version|forward_context|fastvideo_args)\.py|
|
||||
^apps/dreamverse/web/.*|
|
||||
examples/.*|
|
||||
\.agents/.*|
|
||||
|
||||
@@ -84,12 +84,9 @@ RUN source /opt/venv/bin/activate \
|
||||
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
|
||||
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
|
||||
|
||||
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
|
||||
# the dreamverse extra and is validated with it, so enable it here.
|
||||
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
|
||||
STREAM_MODE=av_fmp4 \
|
||||
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
|
||||
FASTVIDEO_FA4=1 \
|
||||
HF_HOME=/root/.cache/huggingface
|
||||
|
||||
RUN mkdir -p /var/lib/dreamverse
|
||||
|
||||
@@ -12,17 +12,13 @@ Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Records can include source metadata and rolling-baseline policy context:
|
||||
Records can include source metadata:
|
||||
|
||||
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
|
||||
- `baseline_eligible`: only successful scheduled-main records should be true
|
||||
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
|
||||
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
|
||||
floors used for recomputed status context
|
||||
|
||||
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
|
||||
crossings; `regressed` remains the gated CI-failure signal.
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
@@ -93,8 +89,7 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
|
||||
- PR number, branch, and Buildkite URL when present
|
||||
|
||||
The latest status table uses the stored JSON `success` value. Recomputed
|
||||
baseline context applies each metric's percent and absolute regression floors
|
||||
and does not override stored status.
|
||||
baseline context is shown separately and does not override stored status.
|
||||
|
||||
## API
|
||||
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
||||
import { fetchSummary, fetchTrends, refreshData } from "./api";
|
||||
import type { CohortValue, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
import { fetchSummary, fetchTrends, refreshData, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
|
||||
|
||||
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
|
||||
const RUN_SOURCES: Array<{ value: "" | RunSource; label: string }> = [
|
||||
@@ -110,61 +109,6 @@ function metricLabel(metricKey: string) {
|
||||
return METRIC_DEFINITIONS[metricKey]?.label ?? metricKey;
|
||||
}
|
||||
|
||||
type CohortFields = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
workload_id: CohortValue;
|
||||
variant_id: CohortValue;
|
||||
benchmark_version: CohortValue;
|
||||
recipe_fingerprint: CohortValue;
|
||||
hardware_profile_id: CohortValue;
|
||||
software_profile_id: CohortValue;
|
||||
};
|
||||
|
||||
function cohortValue(value: CohortValue) {
|
||||
if (value === null || value === undefined || value === "") {
|
||||
return "legacy";
|
||||
}
|
||||
return String(value);
|
||||
}
|
||||
|
||||
function shortCohortValue(value: CohortValue) {
|
||||
const text = cohortValue(value);
|
||||
if (text === "legacy" || text.length <= 14) {
|
||||
return text;
|
||||
}
|
||||
return text.slice(0, 12);
|
||||
}
|
||||
|
||||
function cohortKey(cohort: CohortFields) {
|
||||
return [
|
||||
cohort.model_id,
|
||||
cohort.gpu_type,
|
||||
cohortValue(cohort.workload_id),
|
||||
cohortValue(cohort.variant_id),
|
||||
cohortValue(cohort.benchmark_version),
|
||||
cohortValue(cohort.recipe_fingerprint),
|
||||
cohortValue(cohort.hardware_profile_id),
|
||||
cohortValue(cohort.software_profile_id)
|
||||
].join("|");
|
||||
}
|
||||
|
||||
function cohortTitle(cohort: CohortFields) {
|
||||
const workload = cohortValue(cohort.workload_id);
|
||||
const variant = cohortValue(cohort.variant_id);
|
||||
const version = cohortValue(cohort.benchmark_version);
|
||||
const versionLabel = version === "legacy" ? version : `v${version}`;
|
||||
return `${workload} / ${variant} / ${versionLabel}`;
|
||||
}
|
||||
|
||||
function cohortDetail(cohort: CohortFields) {
|
||||
return [
|
||||
`recipe ${shortCohortValue(cohort.recipe_fingerprint)}`,
|
||||
shortCohortValue(cohort.hardware_profile_id),
|
||||
shortCohortValue(cohort.software_profile_id)
|
||||
].join(" | ");
|
||||
}
|
||||
|
||||
function formatMetricValue(metricKey: string, value: number | null | undefined, tooltip = false) {
|
||||
const definition = METRIC_DEFINITIONS[metricKey];
|
||||
if (!definition) {
|
||||
@@ -227,9 +171,7 @@ function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string
|
||||
top: `${(activePoint.y / height) * 100}%`
|
||||
}
|
||||
: undefined;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}, ${cohortTitle(
|
||||
group
|
||||
)}`;
|
||||
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}`;
|
||||
|
||||
return (
|
||||
<div className="chart-shell">
|
||||
@@ -477,7 +419,7 @@ export default function App() {
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Latest Status</h2>
|
||||
<span>{latestRows.length} comparison cohorts</span>
|
||||
<span>{latestRows.length} model/GPU groups</span>
|
||||
</div>
|
||||
{latestRows.length === 0 ? (
|
||||
<div className="empty">No records match the selected filters.</div>
|
||||
@@ -490,7 +432,6 @@ export default function App() {
|
||||
<th>Recomputed</th>
|
||||
<th>Model</th>
|
||||
<th>GPU</th>
|
||||
<th>Cohort</th>
|
||||
<th>Commit</th>
|
||||
<th>Source</th>
|
||||
<th>Baseline</th>
|
||||
@@ -499,13 +440,11 @@ export default function App() {
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
<th>Exceeded</th>
|
||||
<th>Failing</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{latestRows.map((row) => (
|
||||
<tr key={cohortKey(row)}>
|
||||
<tr key={`${row.model_id}-${row.gpu_type}`}>
|
||||
<td>
|
||||
<span className={`badge ${row.status}`}>{row.status}</span>
|
||||
</td>
|
||||
@@ -516,12 +455,6 @@ export default function App() {
|
||||
</td>
|
||||
<td>{row.model_id}</td>
|
||||
<td>{row.gpu_type}</td>
|
||||
<td>
|
||||
<div className="cohort-cell">
|
||||
<strong>{cohortTitle(row)}</strong>
|
||||
<span>{cohortDetail(row)}</span>
|
||||
</div>
|
||||
</td>
|
||||
<td>{shortSha(row.commit_sha)}</td>
|
||||
<td>
|
||||
<span className={`source-badge source-${row.run_source}`}>{runSourceLabel(row.run_source)}</span>
|
||||
@@ -532,12 +465,6 @@ export default function App() {
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
<td>
|
||||
{row.threshold_exceeded_metrics.length
|
||||
? row.threshold_exceeded_metrics.join(", ")
|
||||
: "none"}
|
||||
</td>
|
||||
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
@@ -560,13 +487,11 @@ export default function App() {
|
||||
) : (
|
||||
trends.map((group) =>
|
||||
METRIC_KEYS.map((metricKey) => (
|
||||
<article className="trend-card" key={`${cohortKey(group)}-${metricKey}`}>
|
||||
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
|
||||
<div>
|
||||
<h3>{metricLabel(metricKey)}</h3>
|
||||
<p>
|
||||
{group.model_id} | {group.gpu_type}
|
||||
<span>{cohortTitle(group)}</span>
|
||||
<span>{cohortDetail(group)}</span>
|
||||
</p>
|
||||
</div>
|
||||
<TrendChart group={group} metricKey={metricKey} />
|
||||
|
||||
@@ -2,28 +2,11 @@ export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
absolute_delta: number | null;
|
||||
threshold_percent: number;
|
||||
threshold_absolute: number;
|
||||
gated: boolean;
|
||||
threshold_exceeded: boolean;
|
||||
regressed: boolean;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
};
|
||||
|
||||
export type CohortValue = string | number | null;
|
||||
|
||||
export type ComparisonCohort = {
|
||||
workload_id: CohortValue;
|
||||
variant_id: CohortValue;
|
||||
benchmark_version: CohortValue;
|
||||
recipe_fingerprint: CohortValue;
|
||||
hardware_profile_id: CohortValue;
|
||||
software_profile_id: CohortValue;
|
||||
};
|
||||
|
||||
export type SummaryRow = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
@@ -32,8 +15,7 @@ export type SummaryRow = {
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
threshold_exceeded_metrics: string[];
|
||||
failing_metrics: string[];
|
||||
regression_threshold_pct: number;
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
run_source: RunSource;
|
||||
@@ -45,7 +27,7 @@ export type SummaryRow = {
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, MetricValue>;
|
||||
} & ComparisonCohort;
|
||||
};
|
||||
|
||||
export type RunSource = "pr" | "local" | "scheduled_main" | "unknown";
|
||||
|
||||
@@ -79,13 +61,13 @@ export type TrendPoint = {
|
||||
build_id: string;
|
||||
job_id: string;
|
||||
metrics: Record<string, number | null>;
|
||||
} & ComparisonCohort;
|
||||
};
|
||||
|
||||
export type TrendGroup = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
points: TrendPoint[];
|
||||
} & ComparisonCohort;
|
||||
};
|
||||
|
||||
export type TrendsResponse = {
|
||||
groups: TrendGroup[];
|
||||
|
||||
@@ -149,11 +149,6 @@ h3 {
|
||||
font-size: 0.82rem;
|
||||
}
|
||||
|
||||
.trend-card p {
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.stat strong {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
@@ -191,7 +186,7 @@ h3 {
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
min-width: 1260px;
|
||||
min-width: 1120px;
|
||||
border-collapse: collapse;
|
||||
}
|
||||
|
||||
@@ -214,25 +209,6 @@ td {
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.cohort-cell {
|
||||
display: grid;
|
||||
gap: 2px;
|
||||
}
|
||||
|
||||
.cohort-cell strong,
|
||||
.trend-card p span {
|
||||
color: #1b2836;
|
||||
font-size: 0.78rem;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.cohort-cell span,
|
||||
.trend-card p span + span {
|
||||
color: #607080;
|
||||
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", monospace;
|
||||
font-size: 0.72rem;
|
||||
}
|
||||
|
||||
.badge {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
|
||||
@@ -1,22 +0,0 @@
|
||||
{
|
||||
"alpha_yaw": 0.08734091699186919,
|
||||
"alpha_pitch": 0.08169667696275307,
|
||||
"alpha_turn": 5.724587470723463e-17,
|
||||
"beta_fwd": 0.02842768078408099,
|
||||
"beta_strafe": 0.022531015077067108,
|
||||
"focal_length": 457.0,
|
||||
"frame_shape": [
|
||||
352,
|
||||
640
|
||||
],
|
||||
"calibrated_from": [
|
||||
"1_wasd_only",
|
||||
"camera",
|
||||
"camera4hold_alpha1",
|
||||
"fully_random",
|
||||
"wasdonly_alpha1",
|
||||
"wasd4holdrandview_simple_1key1mouse1"
|
||||
],
|
||||
"residual_rms": 15.890399609478676,
|
||||
"n_equations": 4125232
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
# FastVideo — Design Philosophy
|
||||
|
||||
One page on *why* FastVideo is built the way it is. The full architecture, the as-built status, and the
|
||||
forward roadmap live in **[`v2/README.md`](v2/README.md)** — this is the philosophy beneath it.
|
||||
|
||||
---
|
||||
|
||||
**A deployable model is a post-training artifact.** Unlike an LLM — where inference optimizes frozen weights
|
||||
after the fact — a *usable* video/omni model is *created* by training: step distillation for latency, QAT for
|
||||
precision, distillation + self-forcing for causal/world models. So every inference capability is a
|
||||
**(recipe, runtime) pair**: the weights and the loop that produced-and-assumes them are one versioned object.
|
||||
This is the source of the moat — whoever owns *both* sides of the pair owns the optimization frontier — and it
|
||||
is why training and serving cannot be two systems.
|
||||
|
||||
**The work is loops, not `forward()`.** Denoise timesteps, AR decode, chunked rollout, VAE tiles, encoder
|
||||
chunks, audio tokens, reward batches, optimizer steps, media chunks — video and omni inference is iteration. A
|
||||
runtime that collapses everything to a single `forward` can't schedule, batch, cancel, stream, reserve memory
|
||||
for, or capture the behavior of what actually runs. So loops are first-class, and they are **driven**: the
|
||||
model describes the next step it needs, the runtime decides when and with whom it runs, the model folds the
|
||||
result back. The model keeps content-adaptive control flow; the runtime keeps admission, batching, streaming,
|
||||
and behavior capture. Per-request state lives in typed `LoopState`, never in module globals — so interleaving
|
||||
requests through one model instance cannot smear state, by construction.
|
||||
|
||||
**The model is the center; everything else is a view over it.** A typed `ModelCard` owns components, loops,
|
||||
the recipe, and the parity contract. Programs compose a card's loops into a task; Workflows compose cards into
|
||||
pipelines; the scheduler runs the *steps* of all loops as `WorkUnit`s under one currency (predicted GPU-time,
|
||||
because a bidirectional denoise step and an AR token are ~1000× apart and incommensurable in counts);
|
||||
deployment places and routes; products stream artifacts. None of them define model semantics — they reference
|
||||
the Model Plane. One resident instance can run many loop types on shared weights, which is what makes omni/MoT
|
||||
native rather than a DAG that doubles weights.
|
||||
|
||||
**Correctness is a typed contract, not a hope.** Caches are correct by *key* — if a field can change output
|
||||
semantics it is in the key, so reuse is partitioned, never blindly flushed. Parity between the train-forward
|
||||
and the serve-forward is *measured* on a declared ladder (component → loop → behavioral → distribution →
|
||||
artifact-quality), never assumed. And the non-negotiable gate is **interleave bit-parity**: N requests
|
||||
interleaved at step granularity must be bit-identical to running them serially — the test the whole
|
||||
loop-inversion bet lives or dies on.
|
||||
|
||||
**One substrate for inference, training, and RL.** The rollout forward *is* the serve forward plus capture —
|
||||
same loop, same caches, same batcher, same numerics — so every serving optimization is automatically a rollout
|
||||
optimization, and there is one numerics surface the ladder measures rather than a correction layer papering
|
||||
over it. The engine doubles as the RL rollout engine under a strict rule: `training` consumes the engine; the
|
||||
**engine never imports `training`**.
|
||||
|
||||
**Borrow aggressively; copy nothing as the core.** vLLM/SGLang scheduling, vLLM-Omni/SGLang-Omni omni serving,
|
||||
Dynamo fleet orchestration, diffusers components, xDiT parallelism, TorchTitan mesh discipline,
|
||||
verl-omni/miles RL lessons, ComfyUI workflows, Dreamverse/LiveKit sessions — each contributes a take, none is
|
||||
the center. Deployment orchestration (Dynamo) sits *above* the engine, never inside it. Extensions are
|
||||
versioned hook points, never monkeypatching. New frontier capabilities arrive as a card, a method, a loop, a
|
||||
workflow, or a controller — **not a rewrite**.
|
||||
|
||||
> 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.
|
||||
+4
-6
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
|
||||
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
|
||||
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
|
||||
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
|
||||
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
|
||||
ARG FA4_CUTE_REF=940cd9680f3315f2f06b43ab5bea2c2cf2d96806
|
||||
|
||||
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
|
||||
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
|
||||
@@ -161,7 +161,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
uv pip install flash-attn==${FLASH_ATTN_VERSION} --no-build-isolation; \
|
||||
fi
|
||||
|
||||
# Overlay the CuTe-DSL-4.6-compatible upstream FA4 cute (FA4_CUTE_REF) over the
|
||||
# Overlay the cutlass-4.5-safe upstream FA4 cute (FA4_CUTE_REF) over the
|
||||
# wheel/source one so the image runs FA4, not the FA2 fallback. This pulls the FA4
|
||||
# runtime stack (cutlass-dsl, quack-kernels, apache-tvm-ffi, torch-c-dlpack-ext) --
|
||||
# the same deps the [dreamverse] extra already installs in CI; the installed torch
|
||||
@@ -170,14 +170,12 @@ RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
|
||||
# Then verify both survive so a broken overlay fails the build instead of shipping
|
||||
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
|
||||
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
|
||||
# arm64 and the image runs FA3/FA2 as usual.
|
||||
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
|
||||
RUN --mount=type=cache,target=/opt/uv/cache \
|
||||
source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
if [ "${TARGETARCH}" = "arm64" ]; then \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
|
||||
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
|
||||
else \
|
||||
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
|
||||
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
|
||||
|
||||
@@ -99,18 +99,10 @@ status.
|
||||
Full Suite is also path-filtered. It validates broader behavior before Mergify
|
||||
can merge a PR.
|
||||
|
||||
A `ready`-labeled PR does not hit Buildkite immediately:
|
||||
`ci-trigger-full-suite.yml` first runs `.github/scripts/gate_full_suite.sh`,
|
||||
which waits for the cheap Tier-1 checks (pre-commit, docs build) on the PR
|
||||
head. A red cheap check blocks the suite (fail closed; the next push re-arms
|
||||
it), while a GitHub outage or a >25 min wait lets it run anyway (fail open).
|
||||
`/test full` bypasses the gate.
|
||||
|
||||
| Buildkite label | `TEST_TYPE` | Main watched paths |
|
||||
|---|---|---|
|
||||
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
|
||||
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
|
||||
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
|
||||
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
|
||||
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
|
||||
@@ -152,7 +144,6 @@ Valid direct test names:
|
||||
| `/test training` | `training` |
|
||||
| `/test lora-inference` | `inference_lora` |
|
||||
| `/test lora-training` | `training_lora` |
|
||||
| `/test lora-extraction` | `lora_extraction` |
|
||||
| `/test distillation` | `distillation_dmd` |
|
||||
| `/test self-forcing` | `self_forcing` |
|
||||
| `/test vsa` | `training_vsa` |
|
||||
|
||||
@@ -72,18 +72,14 @@ fastvideo/tests/performance/
|
||||
│ writes Markdown summary + (optionally) uploads new records
|
||||
├── dashboard.py
|
||||
│ └── builds time-series Plotly HTML from HF history
|
||||
|
||||
fastvideo/performance/
|
||||
├── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
└── metric_policy.py # shared rolling-baseline threshold policy
|
||||
└── hf_store.py # shared HF I/O + DataFrame helpers
|
||||
```
|
||||
|
||||
The HF dataset (`FastVideo/performance-tracking` by default) holds one
|
||||
normalized JSON per run. For v2 records, the rolling baseline is the median of
|
||||
the last 5 successful, baseline-eligible records in the same comparison cohort:
|
||||
`model_id`, `gpu_type`, `workload_id`, `variant_id`, `benchmark_version`,
|
||||
`recipe_fingerprint`, `hardware_profile_id`, and `software_profile_id`. PR and
|
||||
local records are visible in the dashboard but are not baseline eligible.
|
||||
normalized JSON per `(model_id, gpu_type, run)` tuple. The rolling baseline is
|
||||
the median of the last 5 successful, baseline-eligible records for that
|
||||
model+GPU. PR and local records are visible in the dashboard but are not
|
||||
baseline eligible.
|
||||
|
||||
## Planned Coverage
|
||||
|
||||
@@ -96,28 +92,25 @@ and recipe changes instead of treating all records for a model as equivalent.
|
||||
|
||||
## Metrics
|
||||
|
||||
Each benchmark records six metrics. The rolling-baseline comparator also has a
|
||||
per-metric policy with direction, percent threshold, absolute threshold, and a
|
||||
`gated` flag.
|
||||
Each benchmark records six metrics:
|
||||
|
||||
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|
||||
|---|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
|
||||
| Metric | Raw key | Normalized key | Direction |
|
||||
|---|---|---|---|
|
||||
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
|
||||
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
|
||||
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
|
||||
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
|
||||
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
|
||||
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
|
||||
|
||||
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
|
||||
while it runs so pipeline stage execution times are available in
|
||||
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
|
||||
`prompt_encoding_stage` so duplicate stage classes do not collide. For
|
||||
`PipelineStage` entries, shared component stage bases emit a stable
|
||||
`component_metric`: text encoding stages map to `text_encoder_time_s`,
|
||||
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
|
||||
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
|
||||
older logs that do not include `component_metric` or that used the class name as
|
||||
`PipelineStage` entries, the extractor maps the `stage_class` field:
|
||||
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
|
||||
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
|
||||
`vae_decode_time_s`, with a fallback for older logs that used the class name as
|
||||
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
|
||||
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
|
||||
does not report one of the mapped stages, that component metric is stored as
|
||||
@@ -159,29 +152,13 @@ unrealistic memory growth, and optionally large component-specific slowdowns
|
||||
even when the rolling baseline is empty. They are hand-set with generous
|
||||
headroom and almost never need touching.
|
||||
|
||||
### Rolling baseline (per comparison cohort)
|
||||
### Rolling baseline (per `(model_id, gpu_type)`)
|
||||
|
||||
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
|
||||
for the same comparison cohort from the HF dataset, computes the median for
|
||||
each available metric, and evaluates the current run with the metric's
|
||||
rolling regression policy. For v2 records, that cohort is `model_id`,
|
||||
`gpu_type`, `workload_id`, `variant_id`, `benchmark_version`,
|
||||
`recipe_fingerprint`, `hardware_profile_id`, and `software_profile_id`. For
|
||||
latency, memory, and component times, higher values are regressions. For
|
||||
throughput, lower values are regressions.
|
||||
|
||||
A metric exceeds its rolling threshold when both of these are true:
|
||||
|
||||
```text
|
||||
percent_delta > threshold_percent
|
||||
absolute_delta > threshold_absolute
|
||||
```
|
||||
|
||||
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
|
||||
for metrics that should remain visible in reports and the dashboard without
|
||||
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
|
||||
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
|
||||
metrics are skipped.
|
||||
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
|
||||
for each available metric, and fails if the current run regresses by more than
|
||||
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
|
||||
higher values are regressions. For throughput, lower values are regressions.
|
||||
|
||||
This is the **drift detector** — it catches sub-threshold regressions that
|
||||
slowly add up. Only scheduled-main successful records are baseline eligible.
|
||||
@@ -195,51 +172,6 @@ agent skill to advance the rolling median.
|
||||
|
||||
## Schemas
|
||||
|
||||
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
|
||||
|
||||
Benchmark configs without `config_schema_version` are treated as legacy v1
|
||||
configs and remain loadable. New or migrated configs should use
|
||||
`config_schema_version: 2` and include explicit comparable identity fields:
|
||||
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2
|
||||
}
|
||||
```
|
||||
|
||||
`benchmark_id` is still required in this phase because raw artifact names,
|
||||
generated-video directories, normalized record paths, and the current rolling
|
||||
baseline comparator still depend on it. The v2 identity fields are config
|
||||
metadata that make the measured workload explicit:
|
||||
|
||||
| Field | Purpose |
|
||||
|---|---|
|
||||
| `workload_id` | Stable benchmark family, such as `wan-t2v`. |
|
||||
| `variant_id` | Intentional recipe family, including model size and parallelism config, such as `1.3b-sp2`. |
|
||||
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
|
||||
|
||||
If a config declares `config_schema_version: 2`, loading fails clearly when any
|
||||
required v2 identity field is missing. If v2 identity or metadata fields are
|
||||
added without `config_schema_version: 2`, loading also fails so partial
|
||||
migrations do not silently run as v1 configs. Optional v2 metadata fields
|
||||
reserved for follow-up work, such as `metric_threshold_policy` and
|
||||
`quality_metadata`, must be JSON objects when present. (`recipe` is emitted
|
||||
by the harness and is not config-declarable.)
|
||||
|
||||
Recipe fingerprinting, hardware/software profile IDs, exact-identity
|
||||
comparison, and dashboard cohort grouping land with this change: v2 records
|
||||
compare only within their identity cohort, and a record that opens a NEW
|
||||
cohort is marked `baseline_status: "initialized_new_cohort"` (regression
|
||||
gating starts once that cohort accumulates history). Legacy v1 configs still
|
||||
run and are normalized for reporting, but their records skip rolling-baseline
|
||||
comparison entirely (`baseline_status: "skipped_missing_identity"`, never
|
||||
baseline eligible); only static thresholds gate them. Metric-specific
|
||||
threshold policies and promoted baselines remain separate follow-ups.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
@@ -247,10 +179,6 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -268,62 +196,12 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
"max_dit_time_s": 10.0,
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.10,
|
||||
"threshold_absolute": 1.0,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"commit": "<full sha>",
|
||||
"run_source": "pr",
|
||||
"branch": "feature/perf-change",
|
||||
"pr_number": "1234",
|
||||
"test_scope": "direct",
|
||||
"build_url": "https://buildkite.example/build",
|
||||
"build_id": "<buildkite-build-id>",
|
||||
"job_id": "<buildkite-job-id>",
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"quality_metadata": { "quality_status": "canonical" },
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"recipe": {
|
||||
"recipe_schema_version": 1,
|
||||
"benchmark": {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2
|
||||
},
|
||||
"model": { "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" },
|
||||
"init_kwargs": { "num_gpus": 2, "sp_size": 2, "tp_size": 1 },
|
||||
"generation_kwargs": { "height": 480, "width": 832, "num_frames": 45 },
|
||||
"inputs": { "prompt_count": 1, "prompt_sha256": ["<measured-prompt-sha256>"] },
|
||||
"attention": { "requested_backend": "FLASH_ATTN", "resolved_backend": "FLASH_ATTN" }
|
||||
},
|
||||
"recipe_fingerprint": "<sha256>",
|
||||
"hardware_profile": {
|
||||
"device_type": "cuda",
|
||||
"gpu_count": 2,
|
||||
"gpus": [{ "name": "NVIDIA L40S", "memory_gb": 48, "compute_capability": "8.9" }],
|
||||
"interconnect": "none_or_partial"
|
||||
},
|
||||
"hardware_profile_id": "hw-<sha256-prefix>",
|
||||
"software_profile": {
|
||||
"python": "3.12",
|
||||
"pytorch": "2.12",
|
||||
"cuda": "13.0",
|
||||
"packages": {
|
||||
"fastvideo_kernel": "0.3.2",
|
||||
"flashinfer": "0.2.11",
|
||||
"nvidia_cutlass_dsl": "4.5.0",
|
||||
"triton": "3.4.1"
|
||||
}
|
||||
},
|
||||
"software_profile_id": "sw-<sha256-prefix>",
|
||||
"environment_metadata": { "env": { "IMAGE_VERSION": "py3.12-cuda13.0.0" } },
|
||||
"environment_fingerprint": "env-<sha256-prefix>"
|
||||
"vae_decode_time_s": 3.208
|
||||
}
|
||||
```
|
||||
|
||||
@@ -335,10 +213,6 @@ result, used as the rolling-baseline source of truth.
|
||||
```jsonc
|
||||
{
|
||||
"model_id": "wan-t2v-1.3b-2gpu",
|
||||
"result_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp2",
|
||||
"benchmark_version": 2,
|
||||
"timestamp": "2026-05-08T22:00:00+00:00",
|
||||
"commit_sha": "<full sha>",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
@@ -348,69 +222,34 @@ result, used as the rolling-baseline source of truth.
|
||||
"text_encoder_time_s": 2.141,
|
||||
"dit_time_s": 8.437,
|
||||
"vae_decode_time_s": 3.208,
|
||||
"regression_thresholds": {
|
||||
"latency": {
|
||||
"threshold_percent": 0.08,
|
||||
"threshold_absolute": 0.5,
|
||||
"gated": true
|
||||
}
|
||||
},
|
||||
"recipe_fingerprint": "<sha256>",
|
||||
"hardware_profile_id": "hw-<sha256-prefix>",
|
||||
"software_profile_id": "sw-<sha256-prefix>",
|
||||
"environment_fingerprint": "env-<sha256-prefix>",
|
||||
"run_source": "pr",
|
||||
"branch": "feature/perf-change",
|
||||
"pr_number": "1234",
|
||||
"test_scope": "direct",
|
||||
"build_url": "https://buildkite.example/build",
|
||||
"build_id": "<buildkite-build-id>",
|
||||
"job_id": "<buildkite-job-id>",
|
||||
"quality_metadata": { "quality_status": "canonical" },
|
||||
"success": true
|
||||
}
|
||||
```
|
||||
|
||||
### Compatibility with legacy records
|
||||
|
||||
Older records in the HF dataset may not have `result_schema_version`,
|
||||
component timing fields, or v2 identity/profile fields. Records without
|
||||
`result_schema_version` are treated as v1. The comparator ignores missing or
|
||||
`null` metrics when computing a median, and the dashboard lists skipped plots
|
||||
for metric series that have no non-null values. Records missing both
|
||||
`run_source` and `baseline_eligible` are treated as legacy successful
|
||||
main/full-suite uploads and remain eligible for rolling baselines.
|
||||
Current `perf_*.json` artifacts that lack the v2 comparison identity are
|
||||
normalized for reporting but skip rolling-baseline comparison and are not marked
|
||||
baseline eligible.
|
||||
|
||||
New records compare only against the same `model_id`, `gpu_type`,
|
||||
`workload_id`, `variant_id`, `benchmark_version`, `recipe_fingerprint`,
|
||||
`hardware_profile_id`, and `software_profile_id` cohort.
|
||||
`environment_metadata` and `environment_fingerprint` are audit data and are not
|
||||
part of the comparison key.
|
||||
The recipe prompt digests describe the prompts actually measured by the
|
||||
benchmark run; extra configured prompts are ignored unless the benchmark runner
|
||||
executes them.
|
||||
Software profile package cohorts keep exact versions for relevant
|
||||
attention/kernel packages, including FastVideo kernels, FlashAttention,
|
||||
FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
|
||||
Older records in the HF dataset may not have component timing fields. The
|
||||
comparator ignores missing or `null` metrics when computing a median, and the
|
||||
dashboard lists skipped plots for metric series that have no non-null values.
|
||||
Records missing both `run_source` and `baseline_eligible` are treated as legacy
|
||||
successful main/full-suite uploads and remain eligible for rolling baselines.
|
||||
|
||||
## Environment variable reference
|
||||
|
||||
| Variable | Default | Used by | Purpose |
|
||||
|---|---|---|---|
|
||||
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
|
||||
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
|
||||
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py`, `test_inference_performance.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
|
||||
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
|
||||
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
|
||||
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
|
||||
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
|
||||
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
|
||||
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
|
||||
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
|
||||
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
|
||||
|
||||
## CI integration
|
||||
@@ -421,22 +260,17 @@ point is `fastvideo/tests/modal/pr_test.py:run_performance_tests` and the
|
||||
Buildkite artifact upload is in
|
||||
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
|
||||
|
||||
Each performance build runs pytest first. PR and direct runs only continue to
|
||||
`compare_baseline.py` when that fixed-threshold phase passes; if pytest fails,
|
||||
Markdown summaries and normalized JSON artifacts are not emitted. Scheduled
|
||||
main runs set `PERF_UPLOAD_POLICY=always`, so they still run
|
||||
`compare_baseline.py` (with `PERF_PYTEST_RC` set) after a fixed-threshold
|
||||
failure. Those failed scheduled main runs emit summaries and normalized
|
||||
records, upload records with `success=false`, and are excluded from future
|
||||
rolling baselines. The dashboard still runs best-effort for observability.
|
||||
When the rolling-baseline phase runs, it emits:
|
||||
Each performance build runs pytest first. If that fixed-threshold phase fails,
|
||||
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
|
||||
artifacts are not emitted. The dashboard still runs best-effort for
|
||||
observability. When pytest passes, the rolling-baseline phase emits:
|
||||
|
||||
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
|
||||
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
|
||||
per-benchmark row with current vs. baseline values for latency, throughput,
|
||||
memory, text encoder time, DiT time, and VAE decode time.
|
||||
* **Plotly dashboard** — `dashboard_<sha>_<ts>.html` showing time-series for
|
||||
each metric grouped by comparison cohort.
|
||||
each metric grouped by `(model_id, gpu_type)`.
|
||||
* **Normalized records** — `normalized_perf_*.json`, one per benchmark.
|
||||
Useful as input to the
|
||||
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
|
||||
@@ -445,16 +279,11 @@ When the rolling-baseline phase runs, it emits:
|
||||
## Adding a new benchmark
|
||||
|
||||
1. Drop a new JSON config into
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
|
||||
use v2 identity fields:
|
||||
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
|
||||
|
||||
```json
|
||||
{
|
||||
"benchmark_id": "<unique-id>",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "<stable-workload-id>",
|
||||
"variant_id": "<variant, e.g. 1.3b-sp2>",
|
||||
"benchmark_version": 1,
|
||||
"model": { "model_path": "...", "model_short_name": "..." },
|
||||
"init_kwargs": { "num_gpus": 1, ... },
|
||||
"generation_kwargs": { "num_frames": 45, ... },
|
||||
@@ -470,19 +299,9 @@ When the rolling-baseline phase runs, it emits:
|
||||
"max_vae_decode_time_s": 10.0
|
||||
},
|
||||
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
|
||||
},
|
||||
"regression_thresholds": {
|
||||
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
|
||||
}
|
||||
}
|
||||
|
||||
```
|
||||
|
||||
Legacy v1 configs without `config_schema_version` still load, but should not
|
||||
gain v2 identity or metadata fields until they are migrated to
|
||||
`config_schema_version: 2`. For v2 configs, `workload_id`, `variant_id`,
|
||||
and `benchmark_version` are part of the comparison key; benchmark runs
|
||||
fail if any of these identity fields are missing.
|
||||
}
|
||||
```
|
||||
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
@@ -501,17 +320,10 @@ When the rolling-baseline phase runs, it emits:
|
||||
a useful fixed gate. The rolling baseline will still track component times
|
||||
when static component thresholds are omitted.
|
||||
|
||||
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
|
||||
include only benchmark-specific deviations. Tune these independently from
|
||||
the fixed thresholds when a metric is noisy or should be informational. The
|
||||
fixed `thresholds` block is an absolute pytest ceiling. The
|
||||
`regression_thresholds` block controls rolling-baseline comparisons against
|
||||
recent scheduled-main records.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**"No baseline for ... Initializing"** — first run for this comparison cohort.
|
||||
Run will pass and (if persisting) seed the first record.
|
||||
**"No baseline for ... Initializing"** — first run for this `(model_id,
|
||||
gpu_type)`. Run will pass and (if persisting) seed the first record.
|
||||
|
||||
**Persistent failure right after a torch / kernel / image upgrade** —
|
||||
genuine regression *or* baseline drift. Compare the failing normalized record
|
||||
@@ -524,5 +336,5 @@ pipelines that did not report a mapped component stage.
|
||||
|
||||
**Component timing is `null`** — the generated result did not include a mapped
|
||||
stage in `logging_info.stages`. Check that the pipeline emits stage logging
|
||||
and that the stage emits `component_metric` or is covered by the legacy
|
||||
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
|
||||
and that the stage name is listed in `STAGE_METRIC_MAP` in
|
||||
`test_inference_performance.py`.
|
||||
|
||||
@@ -180,30 +180,6 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--device-folder L40S_reference_videos
|
||||
```
|
||||
|
||||
### SSIM Bootstrap Mode
|
||||
|
||||
Normal SSIM runs are strict: if a reference video or latent is missing, the
|
||||
test fails. For new-model PRs, CI can run SSIM in bootstrap mode so missing
|
||||
references are uploaded as draft artifacts for review instead of immediately
|
||||
blocking on a missing canonical reference.
|
||||
|
||||
Buildkite enables SSIM bootstrap mode when either condition is true:
|
||||
|
||||
- the PR title or Buildkite message contains `[new-model]`;
|
||||
- `FASTVIDEO_SSIM_BOOTSTRAP_MODE=1` is set for the Buildkite job.
|
||||
|
||||
Bootstrap mode passes `--ssim-bootstrap-mode` to pytest. When a generated
|
||||
artifact is available, the test uploads it under the `drafts/...` namespace in
|
||||
the SSIM reference repo and marks that case as expected-failed. After reviewing
|
||||
the draft, promote it into the canonical reference layout:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py promote-draft \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id <model_id>
|
||||
```
|
||||
|
||||
## CI Integration
|
||||
|
||||
FastVideo CI tests are orchestrated by Buildkite and run on Modal GPU
|
||||
|
||||
@@ -191,9 +191,6 @@ surfaces:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
|
||||
color_correction_strength:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.dreamx_world.DreamXWorld5BARPipelineConfig
|
||||
default_camera_rotation:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
|
||||
@@ -458,8 +455,6 @@ surfaces:
|
||||
guidance_scale: request.sampling.guidance_scale
|
||||
guidance_scale_2: request.sampling.guidance_scale_2
|
||||
guidance_rescale: request.sampling.guidance_rescale
|
||||
use_embedded_guidance: request.sampling.use_embedded_guidance
|
||||
true_cfg_scale: request.sampling.true_cfg_scale
|
||||
boundary_ratio: request.sampling.boundary_ratio
|
||||
sigmas: request.sampling.sigmas
|
||||
enable_teacache: request.runtime.enable_teacache
|
||||
|
||||
@@ -1,116 +0,0 @@
|
||||
# 🌊 AnyFlow Any-Step Video Distillation
|
||||
|
||||
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
|
||||
|
||||
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
|
||||
|
||||
```
|
||||
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
|
||||
```
|
||||
|
||||
for any ``t > r``.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
|
||||
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
|
||||
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
|
||||
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
|
||||
|
||||
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
|
||||
|
||||
## ⚙️ Inference
|
||||
|
||||
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
|
||||
|
||||
## 🧠 Algorithm
|
||||
|
||||
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
|
||||
|
||||
### Stage 1 — Pretrain (flow-map central-difference)
|
||||
|
||||
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
|
||||
|
||||
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
|
||||
|
||||
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
|
||||
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
|
||||
- the remainder is free.
|
||||
|
||||
The student forward at ``(t, r)`` is trained against the central-difference target
|
||||
|
||||
```
|
||||
target = (eps - x_0) - (t - r) · dF/dt
|
||||
```
|
||||
|
||||
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
|
||||
|
||||
### Stage 2 — On-policy DMD
|
||||
|
||||
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
|
||||
|
||||
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
|
||||
|
||||
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
### Stage 1 — pretrain
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
|
||||
```
|
||||
|
||||
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
|
||||
|
||||
- Global batch size: 32 (8 GPUs × 4 per-GPU)
|
||||
- Learning rate: 5e-5
|
||||
- Flow shift: 5.0
|
||||
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
|
||||
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
|
||||
- ``weight_type``: ``beta08``
|
||||
- ``fuse_guidance_scale``: 3.0
|
||||
- Training steps: 6000
|
||||
|
||||
### Stage 2 — on-policy
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
|
||||
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
|
||||
```
|
||||
|
||||
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
|
||||
|
||||
**Key configuration**:
|
||||
|
||||
- Global batch size: 8 (8 GPUs × 1 per-GPU)
|
||||
- Learning rate: 2e-6
|
||||
- Flow shift: 5.0
|
||||
- ``student_sample_steps``: 4
|
||||
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
|
||||
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
|
||||
- ``real_score_guidance_scale``: 3.0
|
||||
- ``generator_update_interval``: 5 (DMD2 alternation)
|
||||
- Training steps: 4000
|
||||
|
||||
## 🔌 Loading published AnyFlow checkpoints
|
||||
|
||||
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
|
||||
|
||||
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
|
||||
|
||||
## 🧭 Note on ``fuse_guidance_scale``
|
||||
|
||||
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
|
||||
|
||||
```
|
||||
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
|
||||
```
|
||||
|
||||
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
|
||||
|
||||
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
|
||||
@@ -74,23 +74,6 @@ uv pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
### Flash Attention 4 (opt-in)
|
||||
|
||||
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
|
||||
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
|
||||
runtime on some GPU/shape combinations. To use FA4, install the pinned
|
||||
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
|
||||
|
||||
```bash
|
||||
export FASTVIDEO_FA4=1
|
||||
```
|
||||
|
||||
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
|
||||
cannot serve there: grad-enabled (training) attention (FA4's backward requires
|
||||
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
|
||||
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
|
||||
FastVideo fails loudly instead of silently falling back.
|
||||
|
||||
### FP4 Flash Attention 4 (Blackwell only)
|
||||
|
||||
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
|
||||
|
||||
@@ -58,8 +58,6 @@ pipeline initialization and sampling.
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
|
||||
@@ -1,91 +0,0 @@
|
||||
# Training Trackers
|
||||
|
||||
FastVideo can send training metrics and validation media to Weights & Biases
|
||||
or SwanLab. Tracking runs only on global rank 0, and local tracker files are
|
||||
stored under `<output_dir>/tracker`.
|
||||
|
||||
## Supported Trackers
|
||||
|
||||
| Value | Backend | Installation |
|
||||
|-------|---------|--------------|
|
||||
| `wandb` | Weights & Biases | Included with FastVideo |
|
||||
| `swanlab` | SwanLab | Install the optional `swanlab` dependency |
|
||||
| `none` | Disable external tracking | No additional package |
|
||||
|
||||
You can enable more than one backend, for example `trackers: [wandb, swanlab]`.
|
||||
Metrics and validation media are converted to the artifact type required by
|
||||
each backend.
|
||||
|
||||
## Install SwanLab
|
||||
|
||||
For a published FastVideo installation, install the SwanLab extra:
|
||||
|
||||
```bash
|
||||
uv pip install "fastvideo[swanlab]"
|
||||
```
|
||||
|
||||
For an editable source checkout, include the same extra during installation:
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[swanlab]"
|
||||
```
|
||||
|
||||
If FastVideo is already installed, you can install the compatible SDK directly:
|
||||
|
||||
```bash
|
||||
uv pip install "swanlab>=0.6.7"
|
||||
```
|
||||
|
||||
Authenticate once before starting a training run:
|
||||
|
||||
```bash
|
||||
swanlab login
|
||||
```
|
||||
|
||||
See the [SwanLab login documentation](https://docs.swanlab.cn/en/api/cli-swanlab-login.html)
|
||||
for non-interactive and self-hosted setups.
|
||||
|
||||
## Configure Tracking
|
||||
|
||||
Select SwanLab in the YAML config used by the modular training framework:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
checkpoint:
|
||||
output_dir: outputs/my_run
|
||||
tracker:
|
||||
trackers: [swanlab]
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
To log to both supported services:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
trackers: [wandb, swanlab]
|
||||
project_name: my_project
|
||||
run_name: my_run
|
||||
```
|
||||
|
||||
An empty or omitted `trackers` list selects W&B when `project_name` is set.
|
||||
Use an explicit `none` entry to disable external tracking:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
tracker:
|
||||
trackers: [none]
|
||||
```
|
||||
|
||||
## Validation Videos
|
||||
|
||||
SwanLab currently accepts GIF video artifacts. FastVideo converts validation
|
||||
MP4 files and in-memory video arrays to GIF automatically before logging them.
|
||||
For video files, FastVideo uses the sampling frame rate supplied by the caller,
|
||||
or the source file's frame rate when no value is supplied. In-memory arrays use
|
||||
the frame rate supplied by the caller. Both forms fall back to 16 FPS when no
|
||||
frame rate is available.
|
||||
|
||||
For details about configuring validation callbacks, see
|
||||
[Training Infrastructure](train_infra.md#callbacks-pluggable-hooks).
|
||||
@@ -161,21 +161,6 @@ training:
|
||||
decay_interval_steps: 0
|
||||
```
|
||||
|
||||
`training.data.data_path` can also mix multiple preprocessed datasets by using a mapping from dataset path to repeat count:
|
||||
|
||||
```yaml
|
||||
training:
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
data/multi3d_games: 2
|
||||
```
|
||||
|
||||
The repeat count duplicates that dataset's parquet file list before shuffling/sampling, so the example above trains with roughly twice as much `multi3d_games` exposure as `zeldam2-clean`. Paths are just suggested locations; use any local path that contains a FastVideo preprocessed parquet dataset.
|
||||
|
||||
See [Training Trackers](trackers.md) to configure Weights & Biases or SwanLab,
|
||||
including SwanLab installation and authentication.
|
||||
|
||||
### `callbacks` — Pluggable hooks
|
||||
|
||||
Callbacks run at specific points in the training loop (before/after optimizer
|
||||
@@ -338,40 +323,6 @@ Self-Forcing inherits all DMD2 parameters, plus:
|
||||
| `enable_gradient_in_rollout` | `true` | Enable backprop through rollout |
|
||||
| `start_gradient_frame` | `0` | Frame index where gradients begin |
|
||||
|
||||
### Streaming Long Tuning
|
||||
|
||||
`StreamingLongTuningMethod` extends Self-Forcing for LongLive-style rollouts. It
|
||||
keeps a streaming state, generates overlapping chunks, and trains only the new
|
||||
frames while preserving context from earlier chunks.
|
||||
|
||||
For the MatrixGame2/Zelda world-model example, self-forcing and long tuning are
|
||||
separate runs: first train or load the 1k-step self-forcing checkpoint using
|
||||
`examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml`,
|
||||
then run
|
||||
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
|
||||
from that checkpoint for the 3k-step streaming long-tuning stage.
|
||||
|
||||
```yaml
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
streaming_reencode_overlap_anchor: true
|
||||
streaming_anchor_inject_k: 1
|
||||
streaming_require_full_blocks: true
|
||||
multi_phased_distill_schedule:
|
||||
- stage: streaming_long
|
||||
start_step: 0
|
||||
end_step: 3000
|
||||
num_latent_t: 39
|
||||
streaming_training: true
|
||||
```
|
||||
|
||||
See
|
||||
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
|
||||
for a complete MatrixGame2/Zelda configuration.
|
||||
|
||||
---
|
||||
|
||||
## Callbacks
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
# v2 porting status — fastvideo models → the v2 (recipe, runtime) substrate
|
||||
|
||||
Goal: every model in fastvideo's registry resolves through the **v2 `VideoGenerator`** / `Engine`
|
||||
(typed `fastvideo.api` configs + the real torch backend) to a recipe that can construct and run it.
|
||||
|
||||
**Scope: ALL fastvideo models (achieved).** v2 now resolves **63/64** of fastvideo's registered HF ids
|
||||
by exact id (PRIMARY), plus the architecture fallback for local/unregistered checkpoints. The single
|
||||
remaining id — `FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers` — is **environment-blocked**: its VSA
|
||||
(Sparse-Linear Attention) kernels require `nvcc` (not built in this bring-up). It arch-resolves to the
|
||||
base Wan card but needs the VSA kernel build to run faithfully.
|
||||
|
||||
Dispatch is **architecture-driven** (`v2/registry.py`): exact HF id → short-name → architecture
|
||||
inference from the checkpoint (pipeline / transformer / VAE class names + `z_dim`, `transformer_2`,
|
||||
`spatial_upsampler`). Adding a model is one `_BUCKET_C` row (HF ids → builders + transformer class).
|
||||
|
||||
## The porting mechanism — self-contained recipe packages
|
||||
Every net-new arch is a **self-contained recipe package** (`v2/recipes/<arch>/` = `card.py` `loop.py`
|
||||
`program.py` [+ `sampler.py`] + an optional `v2/platform/backends/torch_<arch>.py` adapter). The card
|
||||
declares its torch adapter via **`ComponentSpec.adapter="module:Class"`** (the `_explicit_adapter` seam in
|
||||
`torch_backend.py`) instead of editing the shared `_make_dit`/`_make_vae`/`_make_text_encoder` dispatch —
|
||||
so a port adds **only new files**, never touching shared code, and parallel ports never conflict. New
|
||||
samplers/loops live in-package. Registration is one row in `v2/registry.py:_BUCKET_C`.
|
||||
|
||||
## Working today (GPU-verified, real video/audio) — committed on `v2`
|
||||
| Official example(s) | Model | v2 card |
|
||||
|---|---|---|
|
||||
| `basic.py`, `basic_mps.py`, `basic_ray.py` | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | wan21 |
|
||||
| `basic_self_forcing_causal.py` | `wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers` | wan_causal |
|
||||
| `basic_ltx2_distilled.py` | `FastVideo/LTX2-Distilled-Diffusers` (2-stage + spatial upsampler) | ltx2 |
|
||||
| `basic_wan2_2_ti2v.py` | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | wan2.2-ti2v |
|
||||
| `basic_wan2_2.py` | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` (MoE + CPU expert offload) | wan2.2-a14b |
|
||||
| `basic_ltx2.py` | `Davids048/LTX2-Base-Diffusers` | ltx2 base |
|
||||
| `basic_ltx2_3_distilled.py` | `FastVideo/LTX-2.3-Distilled-Diffusers` (joint T2VS, video+audio) | ltx2.3-distilled |
|
||||
|
||||
Plus the **Wan2.1 i2v cluster** (Fun-1.3B-InP GPU-verified; I2V-14B-480P/720P + Wan2.2-I2V-A14B MoE reuse
|
||||
the i2v card) — CLIP image-encoder + first-frame `[mask|cond]` → 36ch DiT.
|
||||
|
||||
## GPU bring-up results (real weights on H100 NVL, single-GPU, TORCH_SDPA)
|
||||
**20 models generate real video/audio on GPU** — the 7 above + **13 of the newly-ported** archs, each run
|
||||
end-to-end through the real `VideoGenerator` (resolve → stamp → CUDA load → generate). The rest are blocked
|
||||
by a **fastvideo-shared-code / missing-kernel / HF-access** wall, NOT a v2 recipe bug (the v2 recipes are
|
||||
faithful — e.g. cosmos25's DiT+VAE produced finite output; only its Qwen2.5-VL encoder hit a library
|
||||
incompat). All ports also resolve + run end-to-end on the CPU toy backend (`test_bucket_c_ports.py`).
|
||||
|
||||
| GPU status | Models |
|
||||
|---|---|
|
||||
| ✅ **Verified** (real GPU output) | stable_audio (audio), matrixgame2, matrixgame3, gen3c, wan_fun_control, lucy_edit, hunyuangamecraft, hunyuan_video, hunyuan_video15, longcat (13.58B), sfwan22 (2×14B MoE, expert offload), lingbotworld (2×14B, offload), fastwan (TI2V-5B-FullAttn DMD) |
|
||||
| 🚫 fastvideo/env-blocked | **cosmos25** (DiT+VAE ran; Qwen2.5-VL encoder → transformers 5.12.1 incompat in fastvideo); **kandinsky5** (fastvideo registry registers a bare `PipelineConfig`); **hyworld** (fastvideo DiT hardcodes `flash_attn`, not built); **turbowan** 1.3B/i2v + **fastwan** VSA-variants (SLA/VSA sparse-attn params + Triton kernels need nvcc) |
|
||||
| 🚫 access-blocked (HF-gated) | cosmos2, flux2, sd35 (no HF token in this env) |
|
||||
|
||||
To unblock the env-blocked: build `fastvideo-kernel` (SLA/VSA Triton, needs nvcc); pin a fastvideo-compatible
|
||||
`transformers` for the Qwen2.5-VL encoder; add a Kandinsky5 `PipelineConfig` + an SDPA fallback in the
|
||||
hyworld DiT (all fastvideo-side / environment, not v2 recipe work).
|
||||
|
||||
## Newly ported (recipe details)
|
||||
Each resolves through the registry AND runs end-to-end on the CPU toy backend via the public `Engine`
|
||||
path (the `v2/tests/test_bucket_c_ports.py` regression guard), emitting the correct modality artifact.
|
||||
|
||||
**15 net-new architectures** (each a new `TorchComponent` adapter + recipe):
|
||||
- **cosmos2** (Cosmos-Predict2-2B-Video2World) — EDM-Karras denoiser; new `CosmosDenoiseLoop` +
|
||||
`build_karras_sigmas` (the reference port). **cosmos25** (Cosmos-Predict2.5 2B/14B) — flow-match,
|
||||
per-frame plain-sigma timestep, Reason1/Qwen2.5-VL encoder. **gen3c** (GEN3C) — EDM + 82ch pose-buffer.
|
||||
- **hunyuan_video** (+FastHunyuan) — reuses WanDenoiseLoop, dual LLaMA+CLIP encoders, Hunyuan VAE.
|
||||
**hunyuan_video15** (480p/720p). **hunyuangamecraft**, **hyworld** — interactive (camera/action).
|
||||
- **longcat** (T2V/I2V/VC). **kandinsky5** (5.0 T2V Lite).
|
||||
- **sd35** (MMDiT, image, triple-encoder). **flux2** (dev/klein, MMDiT image). **stable_audio** (audio).
|
||||
- **lingbotworld** (camera/Plucker), **matrixgame2**, **matrixgame3** — interactive world models.
|
||||
|
||||
**5 Wan-family variants** (reuse the Wan/Causal arch, new in-package sampler/loop/conditioning):
|
||||
- **turbowan** — rCM few-step (faithful RCMScheduler port), 1.3B/14B T2V + I2V-A14B MoE.
|
||||
- **lucy_edit** — v2v editor (video-VAE-encode node → 96ch DiT input). **wan_fun_control** — control input.
|
||||
- **sfwan22** — Self-Forcing Wan2.2-A14B causal + MoE (i2v + t2v). **fastwan** — DMD 3-step (TI2V-5B-FullAttn
|
||||
loadable; VSA-trained variants + non-strict `to_gate_compress` load are BRINGUP).
|
||||
|
||||
BRINGUP scope per port (documented in each package): GPU load/run; for interactive/world-model archs the
|
||||
action/camera/memory conditioning needs a request-API extension (the t2v/degenerate path is what
|
||||
CPU-verifies); video2world/i2v frame-replace conditioning is threaded but inert without conditioning inputs.
|
||||
|
||||
## Environment
|
||||
v2 bring-up runs **single-GPU, resident, on the `TORCH_SDPA` backend** (no fastvideo-kernel / VSA / FP4).
|
||||
The box has been rescheduled across hosts/arches/python versions mid-session; rebuild the venv for the
|
||||
current arch when that happens: `uv venv --python 3.12 .venv`; comment out `fastvideo-kernel` in
|
||||
`pyproject.toml`; `uv pip install -e ".[dev]"`. Source `/home/scratch.willlin_ent/.bringup_env`
|
||||
(`HF_HOME=./.cache` on scratch, `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`). v2 CPU mini: 240 passed, 2 skipped.
|
||||
|
||||
## How to add a model to the v2 substrate
|
||||
1. `v2/recipes/<arch>/` — card (declare adapters via `ComponentSpec.adapter`; per-model `SamplingDefaults`),
|
||||
loop (reuse `WanDenoiseLoop`/`chunk_rollout` or a new in-package loop+sampler), program.
|
||||
2. `v2/platform/backends/torch_<arch>.py` — a `TorchComponent` subclass (only the forward semantics) if the
|
||||
arch is genuinely new; reuse `WanDiT`/`LTX2DiT`/`WanVAE`/`T5Encoder` via `load_id` when it isn't.
|
||||
3. One row in `v2/registry.py:_BUCKET_C` (HF ids → builders; `transformer_cls` for the arch fallback, or
|
||||
`""` for explicit-id-only capability variants of an existing arch).
|
||||
4. CPU-verify: it resolves + runs on the toy backend (auto-covered by `test_bucket_c_ports.py`). Then GPU
|
||||
bring-up (`stamp_*_checkpoints` → real weights) per BRINGUP notes.
|
||||
@@ -1,81 +0,0 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — image-to-video (I2V) path through
|
||||
# FastVideo's native Cosmos3 pipeline. The input image conditions latent frame 0
|
||||
# (kept clean during denoising); the rest of the clip is generated to follow it.
|
||||
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
|
||||
# ``official_weights/cosmos3``) to skip the Hugging Face download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3_i2v"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
image_path = os.environ.get("COSMOS3_IMAGE_PATH", "assets/images/cyclist.jpg")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A mountain biker rides forward along the sunlit forest trail, wheels "
|
||||
"kicking up dust as trees and dappled light sweep past, smooth cinematic "
|
||||
"tracking shot from behind."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
sampling=SamplingConfig(
|
||||
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
|
||||
# overridable via env for quick smoke runs.
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,77 +0,0 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — this example exercises the
|
||||
# text-to-video (T2V) path through FastVideo's native Cosmos3 pipeline.
|
||||
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
|
||||
# ``official_weights/cosmos3``) to skip the Hugging Face download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A golden retriever puppy runs across a sunlit meadow toward the camera, "
|
||||
"ears flopping and wildflowers swaying in the breeze. Shallow depth of "
|
||||
"field, warm afternoon light, smooth cinematic tracking shot."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
|
||||
# overridable via env for quick smoke runs.
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,78 +0,0 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — text-to-image (T2I) path through
|
||||
# FastVideo's native Cosmos3 pipeline. T2I is the single-frame case
|
||||
# (num_frames=1); the canonical Cosmos3 T2I resolution is 960x960 (the model's
|
||||
# "720" bucket, UniPC flow_shift=10.0). Point COSMOS3_MODEL_PATH at a local
|
||||
# diffusers checkpoint (e.g. ``official_weights/cosmos3``) to skip the download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3_t2i"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A photograph of a red panda sitting on a mossy log in a misty bamboo "
|
||||
"forest, soft golden morning light filtering through the leaves, shallow "
|
||||
"depth of field, crisp fur detail, serene atmosphere."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
# T2I is single-frame; canonical Cosmos3 T2I is 960x960. Overridable
|
||||
# via env for quick smoke runs.
|
||||
num_frames=1,
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "960")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "960")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate image: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,67 +0,0 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
# t2vs (text -> video + sound). The Cosmos3 denoise stage generates a joint
|
||||
# [vision | sound] latent and AVAE-decodes the sound to a waveform muxed into the
|
||||
# mp4. The joint-sound path is gated on COSMOS3_T2VS (set here for the example).
|
||||
os.environ.setdefault("COSMOS3_T2VS", "1")
|
||||
|
||||
from fastvideo import VideoGenerator # noqa: E402
|
||||
from fastvideo.api import ( # noqa: E402
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_cosmos3_t2vs"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(text_encoder=True, pin_cpu_memory=True, dit=False, vae=False),
|
||||
),
|
||||
)
|
||||
|
||||
load_start = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start
|
||||
|
||||
prompt = (
|
||||
"Ocean waves crash against a rocky shore at sunset, white foam spraying "
|
||||
"into the air as seagulls wheel overhead. Golden light, cinematic wide "
|
||||
"shot, the rhythmic roar of the surf."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True, return_frames=False),
|
||||
)
|
||||
|
||||
start = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video+sound: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,64 +0,0 @@
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
|
||||
|
||||
def _env_int(name: str, default: int) -> int:
|
||||
return int(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def _env_float(name: str, default: float) -> float:
|
||||
return float(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
|
||||
prompt = os.getenv(
|
||||
"DREAMX_WORLD_PROMPT",
|
||||
"A cinematic first-person drive through a futuristic coastal city at "
|
||||
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
|
||||
)
|
||||
image_path = os.getenv(
|
||||
"DREAMX_WORLD_IMAGE_PATH",
|
||||
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
|
||||
)
|
||||
|
||||
kwargs = {
|
||||
"output_path": OUTPUT_PATH,
|
||||
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
|
||||
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
|
||||
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
|
||||
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
|
||||
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
|
||||
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
|
||||
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
|
||||
"action_speed_list": [
|
||||
float(value)
|
||||
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
|
||||
],
|
||||
}
|
||||
if image_path:
|
||||
kwargs["image_path"] = image_path
|
||||
|
||||
try:
|
||||
generator.generate_video(prompt, **kwargs)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,140 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import os
|
||||
import re
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
(
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
|
||||
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
|
||||
"35mm, bokeh"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _safe_filename(text: str, max_len: int = 100) -> str:
|
||||
"""Make a stable, filesystem-friendly filename base."""
|
||||
s = text[:max_len].strip()
|
||||
s = s.replace(os.sep, "_")
|
||||
if os.altsep:
|
||||
s = s.replace(os.altsep, "_")
|
||||
s = re.sub(r"\s+", " ", s)
|
||||
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
|
||||
s = s.strip(" .")
|
||||
return s or "prompt"
|
||||
|
||||
|
||||
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
|
||||
if not os.path.isdir(out_dir):
|
||||
return
|
||||
|
||||
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
|
||||
for fn in os.listdir(out_dir):
|
||||
if pattern.match(fn):
|
||||
with contextlib.suppress(FileNotFoundError):
|
||||
os.remove(os.path.join(out_dir, fn))
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(
|
||||
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--model-path",
|
||||
default="official_weights/FLUX.1-dev",
|
||||
help="Local Diffusers checkpoint dir or HF repo id.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
default="outputs/flux_dev/samples",
|
||||
help="Directory for saved PNG outputs.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--prompt",
|
||||
action="append",
|
||||
default=None,
|
||||
help="Prompt. Repeat for multiple images.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
|
||||
)
|
||||
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
|
||||
p.add_argument("--height", type=int, default=1024, help="Output height.")
|
||||
p.add_argument("--width", type=int, default=1024, help="Output width.")
|
||||
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
|
||||
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
|
||||
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
|
||||
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": args.num_gpus,
|
||||
"workload_type": "t2i",
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"image_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": False,
|
||||
"use_fsdp_inference": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=args.model_path,
|
||||
**init_kwargs,
|
||||
)
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
filename_base = (
|
||||
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
)
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
generation_kwargs = {
|
||||
"output_path": output_path,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": args.steps,
|
||||
"guidance_scale": args.guidance,
|
||||
"use_embedded_guidance": True,
|
||||
"true_cfg_scale": 1.0,
|
||||
"seed": seed,
|
||||
"save_video": True,
|
||||
}
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
print(f"[flux] done. outputs written to: {args.out_dir}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,107 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image text-to-image generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
|
||||
text-to-image generation command, saved as a PNG."
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/landscape.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default=("A beautiful landscape photography with rolling hills, "
|
||||
"a winding river, and a vibrant sunset in the background. "
|
||||
"Warm golden light, photorealistic style."),
|
||||
help="Text prompt.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,37 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
|
||||
|
||||
IMAGE_PATH = "assets/girl.png"
|
||||
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A woman stands up and walks away"
|
||||
)
|
||||
_ = generator.generate_video(
|
||||
prompt,
|
||||
image_path=IMAGE_PATH,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
height=1024,
|
||||
width=1024,
|
||||
num_frames=121,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,37 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
|
||||
|
||||
def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
|
||||
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
|
||||
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,120 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
|
||||
|
||||
User story:
|
||||
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
|
||||
want a minimal edit command (text + image -> edited image), saved as a PNG."
|
||||
|
||||
GLM-Image is a single unified pipeline: passing a condition image switches it
|
||||
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
|
||||
write pass), so the generator config is identical to `basic_glm_image.py` — the
|
||||
`inputs.pil_image` on the request is what selects the edit mode.
|
||||
"""
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
default="zai-org/GLM-Image",
|
||||
help="HF id or local diffusers-format GLM-Image weights directory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image",
|
||||
default="assets/images/couple.jpg",
|
||||
help="Condition image to edit.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output",
|
||||
default="image_output/edited.png",
|
||||
help="Output PNG path.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
default="Change the background to a snowy mountain landscape at golden hour.",
|
||||
help="Edit instruction.",
|
||||
)
|
||||
parser.add_argument("--height", type=int, default=1024)
|
||||
parser.add_argument("--width", type=int, default=1024)
|
||||
parser.add_argument("--steps", type=int, default=50)
|
||||
parser.add_argument("--guidance-scale", type=float, default=1.5)
|
||||
parser.add_argument("--seed", type=int, default=1024)
|
||||
parser.add_argument("--num-gpus", type=int, default=1)
|
||||
parser.add_argument("--tp-size", type=int, default=None)
|
||||
parser.add_argument("--sp-size", type=int, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
|
||||
output = Path(args.output)
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
condition = Image.open(args.image).convert("RGB")
|
||||
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
|
||||
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
|
||||
|
||||
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
|
||||
# pipeline class come from the model's registered defaults — don't override.
|
||||
# The pipeline is registered as t2i; passing inputs.pil_image below switches
|
||||
# it to the edit path.
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
trust_remote_code=True,
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
|
||||
),
|
||||
pipeline=PipelineSelection(workload_type="t2i"),
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
inputs=InputConfig(pil_image=condition),
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=1,
|
||||
fps=1,
|
||||
num_inference_steps=args.steps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output.parent),
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
if isinstance(result, list):
|
||||
result = result[0]
|
||||
|
||||
frames = result.frames
|
||||
if frames is not None and len(frames):
|
||||
Image.fromarray(frames[0]).save(output)
|
||||
print(f"Saved image to {output}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,36 @@
|
||||
"""v2 port of basic.py — Wan2.1-T2V-1.3B through the v2 VideoGenerator.
|
||||
|
||||
Same convenience API as upstream (from_pretrained + generate_video); only delta is importing
|
||||
VideoGenerator from v2. v2 bring-up: single-GPU, resident, SDPA; modest res/frames for a quick run.
|
||||
"""
|
||||
from v2 import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "v2_video_samples"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
common = dict(output_path=OUTPUT_PATH, save_video=True,
|
||||
num_frames=25, height=480, width=832, num_inference_steps=30, guidance_scale=5.0)
|
||||
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with "
|
||||
"interest. The playful yet serene atmosphere is complemented by soft natural light "
|
||||
"filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_video_name="wan21_raccoon", **common)
|
||||
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame glistening under "
|
||||
"the warm afternoon sun. Low angle, steady tracking shot, cinematic.")
|
||||
video2 = generator.generate_video(prompt2, output_video_name="wan21_lion", **common)
|
||||
print(f"Outputs: {video.video_path} , {video2.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,29 @@
|
||||
"""v2 port of basic_ltx2.py — LTX-2 base (single-stage) through the v2 VideoGenerator.
|
||||
|
||||
Same convenience API as upstream; only delta is importing VideoGenerator from v2. LTX-2 base is the
|
||||
single-stage (non-distilled) model: the v2 single-stage card (build_ltx2_base_card) runs a request-driven
|
||||
many-step flow-match at FULL latent res (no distilled base/refine split, no spatial upsampler), reusing
|
||||
the LTX-2 DiT/VAE/Gemma adapters. The SAME single-stage card also serves LTX-2.3-Distilled (which is also
|
||||
single-stage) — just pass fewer num_inference_steps for the few-step distilled schedule.
|
||||
|
||||
NOTE: modest res/frames here — upstream defaults to 1088x1920x121, which on an 18.88B base is very slow;
|
||||
raise them for full quality. v2 bring-up: single-GPU, resident, SDPA.
|
||||
"""
|
||||
from v2 import VideoGenerator
|
||||
|
||||
PROMPT = ("A warm sunny backyard, cinematic close-up of two people talking; the camera slowly pans right "
|
||||
"to reveal a grandfather in the garden wearing enormous butterfly wings, flapping his arms like "
|
||||
"he is trying to take off. Deadpan, absurd, quietly tragic.")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained("Davids048/LTX2-Base-Diffusers", num_gpus=1)
|
||||
video = generator.generate_video(
|
||||
prompt=PROMPT, output_path="v2_video_samples_ltx2_base", output_video_name="ltx2_base_backyard",
|
||||
save_video=True, num_frames=25, height=512, width=768, num_inference_steps=30)
|
||||
print(f"Output: {video.video_path}")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,36 @@
|
||||
"""v2 port of basic_ltx2_3_distilled.py — LTX-2.3 Distilled (single-stage, joint A/V) through the v2
|
||||
VideoGenerator.
|
||||
|
||||
Unlike LTX-2.0 distilled (two-stage, video-only), LTX-2.3 is a single-stage *audio+video* model. The
|
||||
shared registry (v2/registry.py) maps ``FastVideo/LTX-2.3-Distilled-Diffusers`` to its OWN card,
|
||||
``build_ltx2_3_card`` — distinct from the LTX-2 base/2-stage cards — which wires the 2.3-specific path:
|
||||
* SEPARATE video + audio text connectors (the Gemma encoder projects the prompt to two embeddings,
|
||||
2048-dim for audio, 4096-dim for video) plus gated attention;
|
||||
* a JOINT DiT forward where video and audio latents cross-attend in a single denoise per step;
|
||||
* a video VAE decode + an AudioDecoder→Vocoder decode → video frames AND a stereo waveform @24kHz.
|
||||
|
||||
Because the model advertises TEXT_TO_VIDEO_SOUND, the VideoGenerator issues a T2VS request by default,
|
||||
so ``generate_video`` returns BOTH modalities: the mp4 plus a sibling ``.wav`` (and ``result.audio`` /
|
||||
``result.audio_sample_rate`` in memory). Being distilled, it wants FEW steps (8). GPU-verified on the
|
||||
rebuilt x86 stack: video (3,33,256,384) + stereo audio (2×61920 @ 24kHz).
|
||||
"""
|
||||
from v2 import VideoGenerator
|
||||
|
||||
PROMPT = "ocean waves crashing on rocks at sunset, seagulls calling in the distance, cinematic, highly detailed"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained("FastVideo/LTX-2.3-Distilled-Diffusers", num_gpus=1)
|
||||
# audio=None auto-enables sound for this A/V model (pass audio=False to force video-only).
|
||||
result = generator.generate_video(
|
||||
prompt=PROMPT, output_path="v2_video_samples_ltx2_3", output_video_name="ltx2_3_ocean",
|
||||
save_video=True, num_frames=33, height=512, width=768, num_inference_steps=8, seed=1)
|
||||
print(f"Video: {result.video_path}")
|
||||
audio_path = result.extra.get("audio_path")
|
||||
if audio_path:
|
||||
print(f"Audio: {audio_path} ({result.audio_sample_rate} Hz)")
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,87 @@
|
||||
"""v2 typed-API inference example — mirrors ``basic_dmd_new_api.py`` but drives the **v2
|
||||
(recipe, runtime) substrate + real torch backend** for the three models brought up on GPU
|
||||
(Wan2.1, SF-causal Wan, LTX-2).
|
||||
|
||||
The ONLY delta from the upstream example is importing ``VideoGenerator`` from ``v2`` instead of
|
||||
``fastvideo`` — the typed config classes are the SAME ``fastvideo.api`` dataclasses.
|
||||
|
||||
Run (on a GPU box, with the v2 venv active):
|
||||
python examples/inference/basic/v2_basic_new_api.py
|
||||
|
||||
Notes vs upstream: the v2 bring-up runs single-GPU, resident, on the TORCH_SDPA backend (no
|
||||
fastvideo-kernel / VSA), so resolutions/steps are modest here for a quick runnable demo. LTX-2 loads
|
||||
an 18.88B DiT (slow first load).
|
||||
"""
|
||||
import os
|
||||
import time
|
||||
|
||||
from v2 import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "v2_video_samples"
|
||||
|
||||
MODELS = [
|
||||
{
|
||||
"family": "wan21",
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"prompt": "a red panda surfing on ocean waves at sunset, cinematic, highly detailed",
|
||||
"sampling": SamplingConfig(num_frames=25, height=480, width=832,
|
||||
num_inference_steps=30, guidance_scale=5.0, seed=1, fps=16),
|
||||
},
|
||||
{
|
||||
"family": "wan_causal",
|
||||
"model_path": "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
|
||||
"prompt": "a cat walking through a sunlit garden, cinematic",
|
||||
"sampling": SamplingConfig(num_frames=25, height=480, width=832,
|
||||
num_inference_steps=4, guidance_scale=5.0, seed=1, fps=16),
|
||||
},
|
||||
{
|
||||
"family": "ltx2",
|
||||
"model_path": "FastVideo/LTX2-Distilled-Diffusers",
|
||||
"prompt": "surfers riding ocean waves at sunset, cinematic, highly detailed",
|
||||
"sampling": SamplingConfig(num_frames=9, height=512, width=768,
|
||||
num_inference_steps=8, guidance_scale=1.0, seed=1, fps=16),
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def run_one(m: dict) -> None:
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=m["model_path"],
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(text_encoder=False, dit=False, vae=False, pin_cpu_memory=False),
|
||||
),
|
||||
)
|
||||
load_start = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start
|
||||
|
||||
request = GenerationRequest(
|
||||
prompt=m["prompt"],
|
||||
sampling=m["sampling"],
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, output_video_name=f"v2_{m['family']}",
|
||||
save_video=True, return_frames=False),
|
||||
)
|
||||
gen_start = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - gen_start
|
||||
|
||||
print(f"[{m['family']:10s}] load={load_time:6.1f}s gen={gen_time:6.1f}s -> {result.video_path}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
for m in MODELS:
|
||||
run_one(m)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,30 @@
|
||||
"""v2 port of basic_self_forcing_causal.py — SF-causal Wan2.1 (CausalWanTransformer3DModel) through
|
||||
the v2 VideoGenerator (chunk_rollout loop).
|
||||
|
||||
Same convenience API as upstream; only delta is importing VideoGenerator from v2. NOTE: the v2 causal
|
||||
loop runs per-chunk few-step (not the upstream kv-cache streaming + SF schedule), so output is coherent
|
||||
but lower-fidelity (a documented gap). num_frames is set by the card's chunk schedule; height/width
|
||||
drive the latent geometry.
|
||||
"""
|
||||
from v2 import VideoGenerator
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "v2_video_samples_causal"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
model_name = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name, num_gpus=1, text_encoder_cpu_offload=False, dit_cpu_offload=False)
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with "
|
||||
"interest. The playful yet serene atmosphere is complemented by soft natural light "
|
||||
"filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, output_video_name="causal_raccoon",
|
||||
save_video=True, sampling_param=sampling_param, height=480, width=832)
|
||||
print(f"Output: {video.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,40 @@
|
||||
"""v2 port of basic_wan2_2.py — Wan2.2-T2V-A14B (MoE) through the v2 VideoGenerator.
|
||||
|
||||
Same convenience API as upstream; only delta is importing VideoGenerator from v2. A14B is a 2-expert
|
||||
MoE: WanTransformer3DModel x2 (in_ch=16, Wan2.1 geometry) with a boundary-timestep switch
|
||||
(boundary_ratio 0.875) — ported via build_wan22_a14b_card (BoundaryTimestepRouting: transformer =
|
||||
high-noise expert, transformer_2 = low-noise), reusing the Wan adapters for both experts.
|
||||
|
||||
NOTE: upstream runs A14B with num_gpus=2 + dit_cpu_offload=True ("DiT need to be offloaded for MoE").
|
||||
The v2 bring-up is single-GPU + resident (no offload), so the two 14B experts (~56GB bf16) + UMT5 are
|
||||
near an 80GB GPU's limit — this example uses reduced res/frames to fit. If it OOMs, the A14B card is
|
||||
still correct; it just needs the (not-yet-ported) MoE DiT CPU offload. See V2_PORTING_STATUS.md.
|
||||
"""
|
||||
from v2 import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "v2_video_samples_wan2_2_14B_t2v"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.2-T2V-A14B-Diffusers",
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
prompt = ("A majestic lion strides across the golden savanna, its powerful frame glistening under "
|
||||
"the warm afternoon sun. The tall grass ripples gently in the breeze. Low angle, steady "
|
||||
"tracking shot, cinematic.")
|
||||
# Reduced res/frames so the two resident 14B experts fit a single 80GB GPU (upstream: 720x1280x81).
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, output_video_name="wan22_a14b_lion",
|
||||
save_video=True, num_frames=17, height=480, width=832,
|
||||
num_inference_steps=20, guidance_scale=5.0)
|
||||
print(f"Output: {video.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,38 @@
|
||||
"""v2 port of basic_wan2_2_ti2v.py — Wan2.2-TI2V-5B (T2V mode) through the v2 VideoGenerator.
|
||||
|
||||
Same convenience API as upstream (from_pretrained + generate_video); only delta is importing
|
||||
VideoGenerator from v2. Wan2.2-TI2V-5B reuses the Wan adapter classes (WanTransformer3DModel /
|
||||
AutoencoderKLWan / UMT5) with the higher-compression VAE geometry (z_dim=48, 16x spatial, 4x temporal).
|
||||
|
||||
NOTE: upstream also runs I2V (image_path=...). The v2 program here is T2V-only (image conditioning is
|
||||
not yet ported), so this mirrors the upstream *T2V* branch (prompt2). Modest res/frames for a quick run.
|
||||
"""
|
||||
from v2 import VideoGenerator
|
||||
|
||||
OUTPUT_PATH = "v2_video_samples_wan2_2_5B_ti2v"
|
||||
|
||||
|
||||
def main() -> None:
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
# T2V mode (the v2 program is text-to-video; upstream's image_path I2V branch is not ported yet).
|
||||
prompt = ("A majestic lion strides across the golden savanna, its powerful frame glistening under "
|
||||
"the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's "
|
||||
"commanding presence. Low angle, steady tracking shot, cinematic.")
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, output_video_name="wan22_ti2v_lion",
|
||||
save_video=True, num_frames=25, height=448, width=768,
|
||||
num_inference_steps=20, guidance_scale=5.0)
|
||||
print(f"Output: {video.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,69 +0,0 @@
|
||||
# LTX-2.3 distilled inference configs
|
||||
|
||||
Ready-to-run `fastvideo generate` run configs for the LTX-2.3
|
||||
distilled model (`FastVideo/LTX-2.3-Distilled-Diffusers`), covering both
|
||||
workloads (t2v / i2v), both two-stage step schedules (`5+2`, `8+3` = denoise
|
||||
+ refine), and four resolutions.
|
||||
|
||||
```bash
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
|
||||
```
|
||||
|
||||
Each config is self-contained (no preset registry needed): the two-stage
|
||||
refine is wired via `generator.pipeline.preset_overrides.refine`, and the
|
||||
base sampling knobs live under `request.sampling`. The refine upsampler
|
||||
auto-resolves from the model's `spatial_upscaler`.
|
||||
|
||||
## Configs
|
||||
|
||||
| workload | schedule | resolution (HxW) | file |
|
||||
|---|---|---|---|
|
||||
| t2v | 5+2 | 1280x832 | `t2v_5s2_1280x832.yaml` |
|
||||
| t2v | 5+2 | 1024x1536 | `t2v_5s2_1024x1536.yaml` |
|
||||
| t2v | 5+2 | 768x1280 | `t2v_5s2_768x1280.yaml` |
|
||||
| t2v | 5+2 | 512x768 | `t2v_5s2_512x768.yaml` |
|
||||
| t2v | 8+3 | 1280x832 | `t2v_8s3_1280x832.yaml` |
|
||||
| t2v | 8+3 | 1024x1536 | `t2v_8s3_1024x1536.yaml` |
|
||||
| t2v | 8+3 | 768x1280 | `t2v_8s3_768x1280.yaml` |
|
||||
| t2v | 8+3 | 512x768 | `t2v_8s3_512x768.yaml` |
|
||||
| i2v | 5+2 | 1280x832 | `i2v_5s2_1280x832.yaml` |
|
||||
| i2v | 5+2 | 1024x1536 | `i2v_5s2_1024x1536.yaml` |
|
||||
| i2v | 5+2 | 768x1280 | `i2v_5s2_768x1280.yaml` |
|
||||
| i2v | 5+2 | 512x768 | `i2v_5s2_512x768.yaml` |
|
||||
| i2v | 8+3 | 1280x832 | `i2v_8s3_1280x832.yaml` |
|
||||
| i2v | 8+3 | 1024x1536 | `i2v_8s3_1024x1536.yaml` |
|
||||
| i2v | 8+3 | 768x1280 | `i2v_8s3_768x1280.yaml` |
|
||||
| i2v | 8+3 | 512x768 | `i2v_8s3_512x768.yaml` |
|
||||
|
||||
## Overriding without editing a file
|
||||
|
||||
Dotted overrides (prefixes `generator.` / `request.`) let you tweak any field:
|
||||
|
||||
```bash
|
||||
# swap prompt
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml \
|
||||
--request.prompt "a red fox running through fresh snow"
|
||||
|
||||
# change output path / gpu count
|
||||
fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml \
|
||||
--request.output.output_path outputs/preview.mp4 \
|
||||
--generator.engine.num_gpus 4
|
||||
```
|
||||
|
||||
## i2v
|
||||
|
||||
The `i2v_*` configs take a first-frame image via
|
||||
`request.extensions.ltx2_images` (`[[path, frame_offset, weight]]`). Edit the
|
||||
path in the file, or override it:
|
||||
|
||||
```bash
|
||||
fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml \
|
||||
--request.extensions.ltx2_images '[["/data/portrait.jpg", 0, 1.0]]'
|
||||
```
|
||||
|
||||
## Schedules
|
||||
|
||||
`5+2` is the fast preview schedule; `8+3` is the higher-quality distilled
|
||||
recipe. Refine (`preset_overrides.refine.num_inference_steps`) only accepts 2
|
||||
or 3 steps.
|
||||
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_512x768.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 5+2 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_5s2_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_512x768.mp4
|
||||
save_video: true
|
||||
@@ -1,38 +0,0 @@
|
||||
# LTX-2.3 distilled i2v — 8+3 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: i2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
# i2v conditioning — replace the path with your own first-frame image.
|
||||
extensions:
|
||||
ltx2_images:
|
||||
- ["/path/to/your/first_frame.jpg", 0, 1.0]
|
||||
ltx2_image_crf: 0.0
|
||||
output:
|
||||
output_path: outputs/ltx2_3_i2v_8s3_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_512x768.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 5+2 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 5 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 5
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_5s2_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 1024x1536.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1024x1536.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1024
|
||||
width: 1536
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_1024x1536.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 1280x832.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 1280
|
||||
width: 832
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_1280x832.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 512x768.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_512x768.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 512
|
||||
width: 768
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_512x768.mp4
|
||||
save_video: true
|
||||
@@ -1,33 +0,0 @@
|
||||
# LTX-2.3 distilled t2v — 8+3 two-stage at 768x1280.
|
||||
#
|
||||
# Run:
|
||||
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_768x1280.yaml
|
||||
#
|
||||
# Stage 1 denoises for 8 steps at half resolution; the latents are then
|
||||
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
|
||||
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
|
||||
generator:
|
||||
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
|
||||
engine:
|
||||
num_gpus: 1
|
||||
pipeline:
|
||||
workload_type: t2v
|
||||
preset_overrides:
|
||||
refine:
|
||||
enabled: true
|
||||
num_inference_steps: 3
|
||||
guidance_scale: 1.0
|
||||
add_noise: true
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
|
||||
sampling:
|
||||
height: 768
|
||||
width: 1280
|
||||
num_frames: 121
|
||||
fps: 24
|
||||
guidance_scale: 1.0
|
||||
num_inference_steps: 8
|
||||
output:
|
||||
output_path: outputs/ltx2_3_t2v_8s3_768x1280.mp4
|
||||
save_video: true
|
||||
@@ -1,78 +0,0 @@
|
||||
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
|
||||
# the student matches an EMA copy of itself at the next timestep, all under
|
||||
# clean-history teacher forcing.
|
||||
#
|
||||
# All three roles initialize from the SAME checkpoint (the teacher-forcing
|
||||
# AR-diffusion model). Point init_from at that checkpoint for a real run.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
ema:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: false
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
|
||||
discrete_cd_N: 48
|
||||
guidance_scale: 3.0
|
||||
ema_decay: 0.99
|
||||
ema_start_step: 200
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_cd
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_cd_shift5
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,105 +0,0 @@
|
||||
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
|
||||
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
|
||||
# rollout from pure noise. Teacher provides the real score, critic
|
||||
# learns the fake score; both inherited from DMD2Method.
|
||||
#
|
||||
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
|
||||
# or with the NVIDIA-released checkpoint
|
||||
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
|
||||
# the paper weights (the delta_embedder rename is handled by the
|
||||
# param_names_mapping in WanVideoArchConfig).
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: <PATH_TO_PRETRAIN_CKPT>
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
real_score_guidance_scale: 3.0
|
||||
dmd_denoising_steps: [999, 937, 833, 624]
|
||||
warp_denoising_step: false
|
||||
|
||||
# AnyFlow rollout knobs.
|
||||
student_sample_steps: 4
|
||||
use_mean_velocity: true
|
||||
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
|
||||
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
|
||||
|
||||
# Critic optimizer (DMD2 inherited).
|
||||
fake_score_learning_rate: 8.0e-6
|
||||
fake_score_betas: [0.0, 0.999]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
attn_kind: vsa
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.0, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_onpolicy
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_onpolicy
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -1,83 +0,0 @@
|
||||
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
|
||||
#
|
||||
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
|
||||
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
|
||||
# can be sampled at arbitrary NFE in the on-policy stage.
|
||||
#
|
||||
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
|
||||
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
|
||||
diffusion_ratio: 0.5
|
||||
consistency_ratio: 0.25
|
||||
epsilon: 5 # finite-difference step in absolute train-timestep units
|
||||
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
|
||||
fuse_guidance_scale: 3.0
|
||||
# shift is taken from pipeline.flow_shift below.
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/preprocessed
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 4
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 21
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 81
|
||||
|
||||
optimizer:
|
||||
learning_rate: 5.0e-5
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 6000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_anyflow_pretrain
|
||||
training_state_checkpointing_steps: 500
|
||||
checkpoints_total_limit: 3
|
||||
resume_from_checkpoint: latest
|
||||
|
||||
tracker:
|
||||
project_name: anyflow-wan
|
||||
run_name: wan2.1_t2v_anyflow_pretrain
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5.0
|
||||
dit_config:
|
||||
# Enable AnyFlow dual-timestep conditioning. The student loads from
|
||||
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
|
||||
# get initialized identically to time_embedder via deep-copy in
|
||||
# WanTimeTextImageEmbedding.__init__.
|
||||
r_embedder: true
|
||||
r_embedder_fusion: gated
|
||||
r_embedder_gate_value: 0.25
|
||||
r_embedder_deltatime_type: r
|
||||
@@ -100,6 +100,9 @@ callbacks:
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 81
|
||||
# Validation/inference uses standard CFG in both clean and Self-Forcing,
|
||||
# so this directly matches Self-Forcing guidance_scale=3.0.
|
||||
guidance_scale: 3.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
|
||||
@@ -1,73 +0,0 @@
|
||||
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model with a block size of 1 frame
|
||||
# - Training: each frame gets its own independent noise level (frame-wise
|
||||
# diffusion forcing), versus the chunk-wise variant that shares one noise
|
||||
# level across num_frames_per_block frames.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
num_frames_per_block: 1
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 1
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_dfsft_framewise
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,72 +0,0 @@
|
||||
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
|
||||
#
|
||||
# - Student: trainable causal Wan model
|
||||
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
|
||||
# denoises the current block while attending to *clean* history (clean_x),
|
||||
# not its own noisy rollout.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.wan.WanCausalModel
|
||||
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/Wan-Syn_77x448x832_600k
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1000
|
||||
num_latent_t: 18
|
||||
num_height: 448
|
||||
num_width: 832
|
||||
num_frames: 69
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-6
|
||||
betas: [0.9, 0.999]
|
||||
weight_decay: 0.01
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 4000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/wan2.1_causal_tfsft
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: distillation_wan_r
|
||||
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
|
||||
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
|
||||
every_steps: 50
|
||||
sampling_steps: [40]
|
||||
guidance_scale: 6.0
|
||||
num_frames: 69
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,119 +1,32 @@
|
||||
# World-Model: Matrix-Game 2.0 I2V
|
||||
|
||||
Training scenarios for the Matrix-Game 2.0 I2V world model on Solaris (Minecraft)
|
||||
data and Zelda data, using the new YAML-driven trainer
|
||||
(`fastvideo/train/entrypoint/train.py`).
|
||||
|
||||
## Solaris Configs
|
||||
Three training scenarios for the Matrix-Game 2.0 I2V world model on the
|
||||
new YAML-driven trainer (`fastvideo/train/entrypoint/train.py`).
|
||||
|
||||
| Config | Method | Student | Notes |
|
||||
|---|---|---|---|
|
||||
| `solaris/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
|
||||
| `solaris/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
|
||||
| `solaris/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Matrix-Game 2.0 DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
|
||||
|
||||
## Zelda Configs
|
||||
|
||||
| Config | Method | Student | Notes |
|
||||
|---|---|---|---|
|
||||
| `zelda/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Zelda bidirectional I2V finetuning from `FastVideo/Matrix-Game-2.0-Base-Diffusers`. Uses 33-frame clips and Zelda validation with action overlays. |
|
||||
| `zelda/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Zelda causal Diffusion-Forcing SFT from `mignonjia/mg_bidirectional_zelda`. Uses the same Zelda data, resolution, optimizer, and validation defaults as the Zelda finetune config. |
|
||||
| `zelda/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Zelda DMD/Self-Forcing distillation; student init = `mignonjia/mg_causal_zelda`, teacher = bidirectional (`mignonjia/mg_bidirectional_zelda`), critic = bidirectional. |
|
||||
| `zelda/streaming_long_tuning_causal_i2v.yaml` | `StreamingLongTuningMethod` | `MatrixGame2CausalModel` | LongLive-style streaming long tuning from the 1k-step Zelda self-forcing checkpoint. |
|
||||
|
||||
Zelda world-model distillation is a two-run workflow: first run
|
||||
`zelda/self_forcing_causal_i2v.yaml` to train or load the 1k-step
|
||||
self-forcing checkpoint (`mignonjia/mg_sf_distilled_zelda_1k_steps`), then run
|
||||
`zelda/streaming_long_tuning_causal_i2v.yaml` for the 3k-step streaming
|
||||
long-tuning stage. The long-tuning YAML starts from that 1k-step checkpoint; it
|
||||
does not run the short self-forcing stage inside the same config.
|
||||
|
||||
## Zelda Training Data
|
||||
|
||||
The Zelda training configs use `data/zeldam2-clean` as a suggested local path.
|
||||
Download the dataset from Hugging Face before running those configs:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id mignonjia/zeldam2-clean \
|
||||
--local_dir data/zeldam2-clean \
|
||||
--repo_type dataset
|
||||
```
|
||||
|
||||
You can store the dataset elsewhere; update `training.data.data_path` in the
|
||||
YAML to point at that location.
|
||||
|
||||
## Multi3D Training Data
|
||||
|
||||
`zelda/finetune_i2v.yaml` includes an optional, commented-out Multi3D entry.
|
||||
Enable it only when you want to mix Zelda with multi-game data from
|
||||
`data/multi3d_games`. You can store this dataset anywhere; before enabling it,
|
||||
update the matching commented `training.data.data_path` key in the YAML to the
|
||||
correct location.
|
||||
|
||||
To mix datasets in a training YAML, set `training.data.data_path` to a
|
||||
path-to-repeat-count mapping. For example, `zelda/finetune_i2v.yaml` can use
|
||||
`data/zeldam2-clean: 1` and `# data/multi3d_games: 10`; uncommenting the
|
||||
Multi3D entry repeats the multi-game parquet list ten times before training
|
||||
samples are shuffled.
|
||||
|
||||
## World Model Validation Data
|
||||
|
||||
The Zelda validation configs expect a small public validation bundle under
|
||||
`data/zelda_validation_data`.
|
||||
|
||||
Download it from Hugging Face before running the Zelda scenarios:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id mignonjia/zelda_validation_data \
|
||||
--local_dir data/zelda_validation_data \
|
||||
--repo_type dataset
|
||||
```
|
||||
|
||||
The bundle contains `validation_zelda.json`, `images/`, and `actions/`.
|
||||
The Zelda configs point
|
||||
`callbacks.validation.dataset_file` at
|
||||
`data/zelda_validation_data/validation_zelda.json`.
|
||||
| `finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
|
||||
| `dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
|
||||
| `self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
|
||||
|
||||
## Usage
|
||||
|
||||
### Solaris
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/solaris/finetune_i2v.yaml
|
||||
examples/train/scenario/worldmodel/finetune_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml
|
||||
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/solaris/self_forcing_causal_i2v.yaml
|
||||
```
|
||||
|
||||
### Zelda
|
||||
|
||||
```bash
|
||||
# Finetuning / DFSFT
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/finetune_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/dfsft_causal_i2v.yaml
|
||||
|
||||
# Distillation / long tuning
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml
|
||||
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml
|
||||
examples/train/scenario/worldmodel/self_forcing_causal_i2v.yaml
|
||||
```
|
||||
|
||||
Override any field on the command line:
|
||||
|
||||
```bash
|
||||
bash examples/train/run.sh \
|
||||
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml \
|
||||
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml \
|
||||
--training.distributed.num_gpus 8 \
|
||||
--training.optimizer.learning_rate 1e-5
|
||||
```
|
||||
|
||||
+1
-1
@@ -97,4 +97,4 @@ callbacks:
|
||||
guidance_scale: 6.0
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
flow_shift: 5
|
||||
@@ -1,94 +0,0 @@
|
||||
# Diffusion-Forcing SFT: Zelda world model I2V Causal
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
|
||||
chunk_size: 3
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 42
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-5
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_causal_dfsft
|
||||
training_state_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
entity: hapo-exp
|
||||
project_name: mg_1.3b_zelda
|
||||
run_name: zelda_causal_dfsft
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 200
|
||||
sampling_steps: [40]
|
||||
sampling_timesteps: [1000, 975, 950, 925, 900, 875, 850, 825, 800, 775,
|
||||
750, 725, 700, 675, 650, 625, 600, 575, 550, 525,
|
||||
500, 475, 450, 425, 400, 375, 350, 325, 300, 275,
|
||||
250, 225, 200, 175, 150, 125, 100, 75, 50, 25]
|
||||
num_frames: 33
|
||||
overlay_actions: true
|
||||
guidance_scale: 6.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -1,88 +0,0 @@
|
||||
# Matrix-Game 2.0 Zelda + multi-game I2V finetune.
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: FastVideo/Matrix-Game-2.0-Base-Diffusers
|
||||
trainable: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path:
|
||||
data/zeldam2-clean: 1
|
||||
# data/multi3d_games: 10
|
||||
dataloader_num_workers: 1
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0 # unused for MatrixGame2 I2V; no text_embedding CFG dropout
|
||||
seed: 42
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 2.0e-5
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 60000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_with_mg_init
|
||||
training_state_checkpointing_steps: 5000
|
||||
checkpoints_total_limit: 3
|
||||
|
||||
tracker:
|
||||
project_name: mg_1.3b_zelda
|
||||
run_name: zelda_with_mg_init
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
grad_clip:
|
||||
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
_target_: fastvideo.train.callbacks.validation.ValidationCallback
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_i2v_pipeline.MatrixGame2I2VPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 200
|
||||
sampling_steps: [40]
|
||||
num_frames: 33
|
||||
overlay_actions: true
|
||||
guidance_scale: 6.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
@@ -1,120 +0,0 @@
|
||||
# Self-forcing distillation: Zelda world model I2V Causal
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_causal_zelda
|
||||
trainable: true
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
same_step_across_blocks: true
|
||||
last_step_only: false
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 3.0e-7
|
||||
fake_score_betas: [0.9, 0.95]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/zeldam2-clean
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1001
|
||||
num_latent_t: 9
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 33
|
||||
|
||||
optimizer:
|
||||
learning_rate: 3.0e-6
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 1000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/zelda_causal_self_forcing
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 1
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: mg2_self_forcing_9_latents
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
ema:
|
||||
decay: 0.99
|
||||
start_iter: 200
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 100
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 153
|
||||
overlay_actions: true
|
||||
keyboard_value_scale: 1.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -1,140 +0,0 @@
|
||||
# MatrixGame2 I2V LongLive-style streaming distillation: Zelda world model I2V Causal
|
||||
# Student init from self forcing after 1k steps
|
||||
|
||||
models:
|
||||
student:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
|
||||
init_from: mignonjia/mg_sf_distilled_zelda_1k_steps
|
||||
trainable: true
|
||||
# transformer_override_safetensor: outputs/matrixgame_dmd/checkpoints/mg_zelda_sf_m2/checkpoint-1000_weight_only/ema/generator_ema.safetensors
|
||||
teacher:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: false
|
||||
disable_custom_init_weights: true
|
||||
critic:
|
||||
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
|
||||
init_from: mignonjia/mg_bidirectional_zelda
|
||||
trainable: true
|
||||
disable_custom_init_weights: true
|
||||
|
||||
method:
|
||||
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
|
||||
rollout_mode: simulate
|
||||
generator_update_interval: 5
|
||||
dmd_denoising_steps: [1000, 750, 500, 250]
|
||||
warp_denoising_step: true
|
||||
min_timestep_ratio: 0.02
|
||||
max_timestep_ratio: 0.98
|
||||
|
||||
chunk_size: 3
|
||||
student_sample_type: sde
|
||||
same_step_across_blocks: true
|
||||
last_step_only: false
|
||||
context_noise: 0.0
|
||||
enable_gradient_in_rollout: true
|
||||
start_gradient_frame: 0
|
||||
|
||||
streaming_training: true
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
streaming_reencode_overlap_anchor: true
|
||||
streaming_anchor_inject_k: 1
|
||||
streaming_require_full_blocks: true
|
||||
|
||||
multi_phased_distill_schedule:
|
||||
- stage: streaming_long
|
||||
start_step: 0
|
||||
end_step: 3000
|
||||
num_latent_t: 39
|
||||
streaming_training: true
|
||||
streaming_chunk_size: 9
|
||||
streaming_max_length: 39
|
||||
streaming_fixed_overlap_latents: 3
|
||||
|
||||
# Critic optimizer
|
||||
fake_score_learning_rate: 3.0e-7
|
||||
fake_score_betas: [0.9, 0.95]
|
||||
fake_score_lr_scheduler: constant
|
||||
|
||||
training:
|
||||
distributed:
|
||||
num_gpus: 8
|
||||
sp_size: 1
|
||||
tp_size: 1
|
||||
hsdp_replicate_dim: 1
|
||||
hsdp_shard_dim: 8
|
||||
|
||||
data:
|
||||
data_path: data/zeldam2-clean
|
||||
dataloader_num_workers: 4
|
||||
train_batch_size: 1
|
||||
training_cfg_rate: 0.0
|
||||
seed: 1001
|
||||
num_latent_t: 39
|
||||
num_height: 480
|
||||
num_width: 832
|
||||
num_frames: 153
|
||||
|
||||
optimizer:
|
||||
learning_rate: 3.0e-6
|
||||
betas: [0.9, 0.95]
|
||||
weight_decay: 0.0
|
||||
lr_scheduler: constant
|
||||
lr_warmup_steps: 0
|
||||
|
||||
loop:
|
||||
max_train_steps: 3000
|
||||
gradient_accumulation_steps: 1
|
||||
|
||||
checkpoint:
|
||||
output_dir: outputs/zelda_causal_long_tuning
|
||||
training_state_checkpointing_steps: 1000
|
||||
checkpoints_total_limit: 2
|
||||
|
||||
tracker:
|
||||
project_name: wangame_sf
|
||||
run_name: mg2_39only_streaming_long
|
||||
|
||||
model:
|
||||
enable_gradient_checkpointing_type: full
|
||||
|
||||
dit_precision: fp32
|
||||
|
||||
callbacks:
|
||||
ema:
|
||||
decay: 0.99
|
||||
start_iter: 200
|
||||
grad_clip:
|
||||
max_grad_norm: 1.0
|
||||
validation:
|
||||
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
|
||||
dataset_file: data/zelda_validation_data/validation_zelda.json
|
||||
every_steps: 100
|
||||
sampling_steps: [4]
|
||||
sampling_timesteps: [1000, 750, 500, 250]
|
||||
num_frames: 153
|
||||
overlay_actions: true
|
||||
keyboard_value_scale: 1.0
|
||||
metrics:
|
||||
enabled: true
|
||||
names:
|
||||
- vbench.imaging_quality
|
||||
- vbench.aesthetic_quality
|
||||
- vbench.temporal_flickering
|
||||
- vbench.motion_smoothness
|
||||
- vbench.subject_consistency
|
||||
- vbench.background_consistency
|
||||
- vbench.dynamic_degree
|
||||
- optical_flow.synthetic_optical_flow
|
||||
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
|
||||
skip_missing_deps: true
|
||||
strict: false
|
||||
unload_after_validation: true
|
||||
|
||||
pipeline:
|
||||
flow_shift: 5
|
||||
dit_config:
|
||||
local_attn_size: 6
|
||||
sink_size: 1
|
||||
@@ -12,9 +12,6 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
if(NOT GPU_BACKEND)
|
||||
set(GPU_BACKEND "CUDA")
|
||||
endif()
|
||||
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
enable_language(HIP)
|
||||
@@ -53,16 +50,7 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
|
||||
if(_FASTVIDEO_USER_CUDA_ARCH)
|
||||
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
|
||||
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
|
||||
# Only numeric spellings translate; keywords like "native"/"all" would
|
||||
# otherwise be mangled into nonsense ("nativ.e").
|
||||
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
|
||||
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
|
||||
message(FATAL_ERROR
|
||||
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
|
||||
"supported. Use a numeric arch (e.g. 90a, 121), set "
|
||||
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
|
||||
"auto-detect from the visible GPU.")
|
||||
endif()
|
||||
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
|
||||
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
|
||||
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
|
||||
@@ -185,14 +173,6 @@ else()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
|
||||
# there, and tk's base_types.cuh brace-initializes signed-char vector members
|
||||
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
|
||||
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
|
||||
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
endif()
|
||||
|
||||
if(ENABLE_TK_KERNELS)
|
||||
message(STATUS "ThunderKittens kernels: ENABLED")
|
||||
else()
|
||||
@@ -283,6 +263,11 @@ set(CUDA_FLAGS
|
||||
"--expt-relaxed-constexpr"
|
||||
"-Xcompiler=-fno-strict-aliasing"
|
||||
"-Xcompiler=-fPIC"
|
||||
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
|
||||
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
|
||||
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
|
||||
# is a no-op on x86_64, where char is already signed.
|
||||
"-Xcompiler=-fsigned-char"
|
||||
"-DTORCH_COMPILE"
|
||||
"-Xnvlink=--verbose"
|
||||
"-Xptxas=--verbose"
|
||||
@@ -441,14 +426,3 @@ if(ENABLE_ATTN_QAT_INFER)
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
# One-look answer to "what is this build producing?" — kept last so it is the
|
||||
# final thing configure prints. The per-kernel matrix lives in README.md.
|
||||
message(STATUS "============== fastvideo-kernel build summary ==============")
|
||||
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
|
||||
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
|
||||
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
|
||||
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
|
||||
message(STATUS "============================================================")
|
||||
|
||||
@@ -2,42 +2,6 @@
|
||||
|
||||
CUDA kernels for FastVideo video generation.
|
||||
|
||||
## Kernel inventory
|
||||
|
||||
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
|
||||
|
||||
| Extension | Kernels | Sources | GPU arch | Build gate |
|
||||
|---|---|---|---|---|
|
||||
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
|
||||
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
|
||||
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
|
||||
|
||||
Runtime-JIT kernels (no build step, ship in every wheel/image):
|
||||
|
||||
| Kernels | Where | Used when |
|
||||
|---|---|---|
|
||||
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
|
||||
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
|
||||
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
|
||||
|
||||
## What gets built where, and when
|
||||
|
||||
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|
||||
|---|---|---|---|---|---|
|
||||
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
|
||||
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
|
||||
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
|
||||
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
|
||||
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
|
||||
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
|
||||
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
|
||||
|
||||
Notes:
|
||||
|
||||
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
|
||||
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
|
||||
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
|
||||
|
||||
## Installation
|
||||
|
||||
### Standard Installation (Local Development)
|
||||
|
||||
@@ -46,23 +46,6 @@ fi
|
||||
if git rev-parse --git-dir >/dev/null 2>&1; then
|
||||
git submodule update --init --recursive include/cutlass include/tk
|
||||
fi
|
||||
# Fail fast with a clear message if the headers are still missing (e.g. a
|
||||
# Docker context that excluded .git AND the submodule contents) instead of
|
||||
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
|
||||
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
|
||||
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
|
||||
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
|
||||
if [ ! -d include/cutlass/include ]; then
|
||||
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
|
||||
echo " CUTLASS sources must already be present (run" >&2
|
||||
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
|
||||
echo " in the source checkout, or include them in the build context)." >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ ! -d include/tk/include ]; then
|
||||
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
|
||||
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
|
||||
fi
|
||||
|
||||
# Install build dependencies
|
||||
uv pip install scikit-build-core cmake ninja
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxSamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 1024
|
||||
width: int = 1024
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 3.5
|
||||
use_embedded_guidance: bool = True
|
||||
true_cfg_scale: float = 1.0
|
||||
@@ -90,10 +90,6 @@ class SamplingParam:
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
|
||||
use_embedded_guidance: bool = False
|
||||
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
|
||||
true_cfg_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
@@ -329,18 +325,6 @@ class SamplingParam:
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-embedded-guidance",
|
||||
action="store_true",
|
||||
default=SamplingParam.use_embedded_guidance,
|
||||
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--true-cfg-scale",
|
||||
type=float,
|
||||
default=SamplingParam.true_cfg_scale,
|
||||
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
|
||||
@@ -150,7 +150,6 @@ class SamplingConfig:
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
use_embedded_guidance: bool | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@@ -5,10 +5,99 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.attention.utils.flash_attn_default import (
|
||||
fa_version,
|
||||
flash_attn_func_compilable,
|
||||
)
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
|
||||
fa_version = "4"
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer have a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
|
||||
# breaks the graph at the call site (observed: wanvideo.py self-attn,
|
||||
# once per layer every step), which fragments the compiled region and
|
||||
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
|
||||
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
# opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
# (correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
|
||||
if fa_version in ("2", "3"):
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
|
||||
# call shape used by FlashAttentionImpl.forward's default branch
|
||||
# (see `flash_attn_func_compilable(...)` call site below). The
|
||||
# masked/no-pad and varlen / cross-attn paths use different
|
||||
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
|
||||
# are intentionally out of scope for this PR — wrapping them is a
|
||||
# natural follow-up. The wrapper's signature is the contract: any
|
||||
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
|
||||
# return_attn_probs, ...) raises TypeError at the call site, so
|
||||
# silent loss of kwargs is not a failure mode.
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
|
||||
# same dtype/device as q (head dim taken from v).
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
@@ -19,6 +108,8 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_WARNED_NON_FA_DTYPE = False
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
@@ -180,8 +271,12 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
|
||||
orig_dtype = query.dtype
|
||||
if orig_dtype not in (torch.float16, torch.bfloat16):
|
||||
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
|
||||
f"bfloat16 for the kernel and restoring on output.")
|
||||
global _WARNED_NON_FA_DTYPE
|
||||
if not _WARNED_NON_FA_DTYPE:
|
||||
_WARNED_NON_FA_DTYPE = True
|
||||
logger.warning(
|
||||
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
|
||||
"kernel and restoring on output.", orig_dtype)
|
||||
query = query.to(torch.bfloat16)
|
||||
key = key.to(torch.bfloat16)
|
||||
value = value.to(torch.bfloat16)
|
||||
@@ -198,17 +293,9 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
|
||||
# Route through the *_compilable wrappers so dynamo sees one
|
||||
# traceable node for each masked entry point (the unpad/pad
|
||||
# bookkeeping runs eager inside the custom op). On FA2 these
|
||||
# wrappers go through ops with full register_autograd, so
|
||||
# training also backprops through the op (no graph break on
|
||||
# the training path); on FA3/FA4 they carve out to the
|
||||
# autograd.Function for grad-enabled calls — see
|
||||
# fastvideo/attention/utils/flash_attn_no_pad.py.
|
||||
from fastvideo.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_no_pad_compilable as flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad_compilable as flash_attn_varlen_qk_no_pad,
|
||||
flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad,
|
||||
)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
@@ -235,11 +322,7 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
)
|
||||
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, attn_mask.shape[-1]).to(device=query.device)
|
||||
if key_padding_mask.shape[-1] > qkv.shape[1]:
|
||||
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
|
||||
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
|
||||
attn_mask_padded = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
|
||||
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
|
||||
elif self.nvfp4_fa4:
|
||||
output = self._forward_nvfp4(query, key, value)
|
||||
|
||||
@@ -1,147 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""NABLA block-sparse flex-attention backend (Kandinsky5 "nabla" checkpoints).
|
||||
|
||||
The block mask is data-dependent: nablaT_v2 mean-pools 64-token blocks of the
|
||||
fractal-ordered sequence, thresholds the softmaxed block map, and ORs it with a
|
||||
precomputed spatio-temporal-window (STA) mask carried on the attention
|
||||
metadata. The mask spans the full sequence, so this backend does not support
|
||||
sequence parallelism — use it via LocalAttention only.
|
||||
"""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
from torch.nn.attention.flex_attention import BlockMask, flex_attention
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
|
||||
CAN_USE_FLEX_ATTN = True
|
||||
except ImportError:
|
||||
CAN_USE_FLEX_ATTN = False
|
||||
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
|
||||
|
||||
def nablaT_v2(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
sta: torch.Tensor,
|
||||
thr: float = 0.9,
|
||||
) -> "BlockMask":
|
||||
q = q.transpose(1, 2).contiguous()
|
||||
k = k.transpose(1, 2).contiguous()
|
||||
|
||||
# Map estimation
|
||||
B, h, S, D = q.shape
|
||||
s1 = S // 64
|
||||
qa = q.reshape(B, h, s1, 64, D).mean(-2)
|
||||
ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1)
|
||||
map = qa @ ka
|
||||
|
||||
map = torch.softmax(map / math.sqrt(D), dim=-1)
|
||||
# Map binarization
|
||||
vals, inds = map.sort(-1)
|
||||
cvals = vals.cumsum_(-1)
|
||||
mask = (cvals >= 1 - thr).int()
|
||||
mask = mask.gather(-1, inds.argsort(-1))
|
||||
|
||||
mask = torch.logical_or(mask, sta)
|
||||
|
||||
# BlockMask creation
|
||||
kv_nb = mask.sum(-1).to(torch.int32)
|
||||
kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32)
|
||||
return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None)
|
||||
|
||||
|
||||
class NablaAttentionBackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "NABLA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["NablaAttentionImpl"]:
|
||||
return NablaAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["NablaAttentionMetadata"]:
|
||||
return NablaAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["NablaAttentionMetadataBuilder"]:
|
||||
return NablaAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class NablaAttentionMetadata(AttentionMetadata):
|
||||
# Block-level STA window mask [1, 1, S/64, S/64], precomputed once per run.
|
||||
sta_mask: torch.Tensor = None # type: ignore[assignment]
|
||||
# Cumulative-probability threshold for block-map binarization.
|
||||
P: float = 0.9
|
||||
visual_shape: tuple[int, int, int] = (0, 0, 0)
|
||||
|
||||
|
||||
class NablaAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
||||
def prepare(self) -> None:
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
sta_mask: torch.Tensor,
|
||||
P: float,
|
||||
visual_shape: tuple[int, int, int],
|
||||
**kwargs: Any,
|
||||
) -> NablaAttentionMetadata:
|
||||
return NablaAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
sta_mask=sta_mask,
|
||||
P=P,
|
||||
visual_shape=visual_shape,
|
||||
)
|
||||
|
||||
|
||||
class NablaAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
causal: bool = False,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
if not CAN_USE_FLEX_ATTN:
|
||||
raise RuntimeError("NABLA attention requires torch.nn.attention.flex_attention, "
|
||||
"which is unavailable in this PyTorch build.")
|
||||
if causal:
|
||||
raise ValueError("NABLA attention does not support causal masking.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: NablaAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
# q/k/v: [B, S, heads, head_dim], fractal-ordered by the model; S % 64 == 0.
|
||||
block_mask = nablaT_v2(query, key, attn_metadata.sta_mask, thr=attn_metadata.P)
|
||||
return flex_attention(
|
||||
query=query.transpose(1, 2),
|
||||
key=key.transpose(1, 2),
|
||||
value=value.transpose(1, 2),
|
||||
block_mask=block_mask,
|
||||
).transpose(1, 2)
|
||||
@@ -2,7 +2,6 @@
|
||||
|
||||
import torch
|
||||
from dataclasses import dataclass
|
||||
from torch.nn import functional as F
|
||||
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from fastvideo.logger import init_logger
|
||||
@@ -20,7 +19,7 @@ class SDPABackend(AttentionBackend):
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "TORCH_SDPA"
|
||||
return "SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SDPAImpl"]:
|
||||
@@ -50,51 +49,9 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> SDPAMetadata:
|
||||
# Store the mask exactly as passed. The metadata is cross-backend:
|
||||
# call sites (HYWorld, HunyuanVideo15) build SDPAMetadata while the
|
||||
# layer's selector may pick FLASH_ATTN, and the shared convention for
|
||||
# padding masks is the tokenizer-style 2D [batch, key_len]. Any
|
||||
# reshaping for torch.sdpa happens inside the SDPA impl
|
||||
# (_normalize_attn_mask_for_sdpa).
|
||||
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
|
||||
|
||||
def _normalize_attn_mask_for_sdpa(
|
||||
attn_mask: torch.Tensor | None,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
) -> torch.Tensor | None:
|
||||
if attn_mask is None:
|
||||
return None
|
||||
|
||||
attn_mask = attn_mask.to(device=query.device)
|
||||
# F.scaled_dot_product_attention only accepts bool or float masks;
|
||||
# tokenizers commonly produce int64 0/1 padding masks.
|
||||
if attn_mask.dtype != torch.bool and not attn_mask.dtype.is_floating_point:
|
||||
attn_mask = attn_mask != 0
|
||||
|
||||
key_len = key.shape[-2]
|
||||
if attn_mask.shape[-1] > key_len:
|
||||
raise ValueError("Invalid attention mask length for SDPA: "
|
||||
f"expected at most {key_len}, got {attn_mask.shape[-1]}")
|
||||
if attn_mask.shape[-1] < key_len:
|
||||
# Front-pad as "attend": double-stream layouts (HYWorld) prepend
|
||||
# non-text tokens the tokenizer mask does not cover.
|
||||
valid_value = True if attn_mask.dtype == torch.bool else 0.0
|
||||
attn_mask = F.pad(attn_mask, (key_len - attn_mask.shape[-1], 0), value=valid_value)
|
||||
|
||||
if attn_mask.dim() == 2:
|
||||
# In-tree producers pass 2D [batch, key_len] padding masks; lift to a
|
||||
# broadcastable [batch, 1, 1, key_len] here so torch.sdpa does not
|
||||
# reinterpret 2D as its documented [query_len, key_len] broadcast.
|
||||
return attn_mask[:, None, None, :]
|
||||
if attn_mask.dim() == 3:
|
||||
return attn_mask[:, None, :, :]
|
||||
if attn_mask.dim() == 4:
|
||||
return attn_mask
|
||||
raise ValueError(f"Unsupported attention mask shape for SDPA: {attn_mask.shape}")
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -125,7 +82,6 @@ class SDPAImpl(AttentionImpl):
|
||||
|
||||
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
|
||||
and hasattr(attn_metadata, "attn_mask")) else None
|
||||
attn_mask = _normalize_attn_mask_for_sdpa(attn_mask, query, key)
|
||||
attn_kwargs = {
|
||||
"attn_mask": attn_mask,
|
||||
"dropout_p": self.dropout,
|
||||
|
||||
@@ -252,7 +252,6 @@ class LocalAttention(nn.Module):
|
||||
causal: bool = False,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -263,10 +262,7 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = get_compute_dtype()
|
||||
attn_backend = get_attn_backend(head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
default_backend=default_backend)
|
||||
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
|
||||
@@ -84,9 +84,8 @@ def get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
|
||||
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
|
||||
|
||||
|
||||
@cache
|
||||
@@ -95,7 +94,6 @@ def _cached_get_attn_backend(
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
default_backend: AttentionBackendEnum | None = None,
|
||||
) -> type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -114,12 +112,6 @@ def _cached_get_attn_backend(
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
# Layer-level default (e.g. a checkpoint that requires a specific sparse
|
||||
# backend). Lower precedence than the global force and the env var, so
|
||||
# users can still override it.
|
||||
if selected_backend is None and default_backend is not None:
|
||||
selected_backend = default_backend
|
||||
|
||||
# get device-specific attn_backend
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
|
||||
@@ -4,9 +4,10 @@ import functools
|
||||
from collections.abc import Callable
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -15,8 +16,7 @@ if torch.cuda.is_available():
|
||||
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
|
||||
except ImportError:
|
||||
# flash_attn.cute (FA4) is simply not installed -- expected on builds
|
||||
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
|
||||
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
|
||||
# without it; callers fall back to FA3/FA2 quietly.
|
||||
raise
|
||||
except Exception as e:
|
||||
# flash_attn.cute IS installed but failed to import -- almost always an
|
||||
@@ -24,57 +24,23 @@ if torch.cuda.is_available():
|
||||
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
|
||||
# ImportError). This is fixable by pinning a compatible
|
||||
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
|
||||
# callers can handle it uniformly.
|
||||
# callers fall back to FA3/FA2 instead of crashing worker init.
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r). "
|
||||
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
|
||||
"compatible nvidia-cutlass-dsl to restore FA4.", e)
|
||||
"flash_attn.cute (FA4) is installed but failed to import (%r); "
|
||||
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
|
||||
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
|
||||
"restore FA4.", e)
|
||||
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
|
||||
else:
|
||||
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
|
||||
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
|
||||
|
||||
try:
|
||||
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
|
||||
# Optional so FA4-only installs can still import this module.
|
||||
from flash_attn import flash_attn_func as _flash_attn_2_func
|
||||
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
|
||||
except ImportError:
|
||||
_flash_attn_2_func = None
|
||||
_flash_attn_2_varlen_func = None
|
||||
|
||||
|
||||
def _check_dropout(dropout_p: float) -> None:
|
||||
if dropout_p != 0.0:
|
||||
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _sm90_or_newer() -> bool:
|
||||
return current_platform.has_device_capability(90)
|
||||
|
||||
|
||||
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
|
||||
if _sm90_or_newer():
|
||||
return False
|
||||
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
|
||||
# capability gate, not a runtime fallback):
|
||||
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
|
||||
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
|
||||
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
|
||||
if q.shape[-2] != k.shape[-2]:
|
||||
return True
|
||||
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
|
||||
|
||||
|
||||
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
|
||||
if fa2_func is None:
|
||||
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
|
||||
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
|
||||
"not installed.")
|
||||
return fa2_func
|
||||
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_cute_forward",
|
||||
mutates_args=(),
|
||||
@@ -277,6 +243,70 @@ torch.library.register_autograd(
|
||||
)
|
||||
|
||||
|
||||
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
|
||||
# fail MLIR op creation at runtime even though the import succeeded (observed:
|
||||
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
|
||||
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
|
||||
class _FA4Policy:
|
||||
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
|
||||
|
||||
FA4 is skipped when:
|
||||
* a previous call failed at runtime -- CuTeDSL JIT compilation is
|
||||
shape-dependent, so the first failure disables FA4 for the rest of
|
||||
the process instead of retrying a broken JIT on every call; or
|
||||
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
|
||||
dies on its arch check) and is unvalidated for training in this repo
|
||||
(its lse is not even allocated through our inference-shaped custom
|
||||
op), so training keeps the pre-FA4 behavior: FA2 on every device.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.broken = False
|
||||
|
||||
def use_fa4(self, *tensors: torch.Tensor) -> bool:
|
||||
if self.broken:
|
||||
return False
|
||||
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
|
||||
|
||||
def mark_broken(self, error: Exception) -> None:
|
||||
if not self.broken:
|
||||
self.broken = True
|
||||
logger.warning(
|
||||
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
|
||||
"to FA2 for the rest of this process.", error)
|
||||
|
||||
|
||||
_FA4 = _FA4Policy()
|
||||
|
||||
|
||||
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
|
||||
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
|
||||
|
||||
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
|
||||
the first FA4 runtime failure) the call is served by ``fa2_func``.
|
||||
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
|
||||
failure, so it propagates without disabling FA4.
|
||||
"""
|
||||
|
||||
def decorator(fa4_func: Callable) -> Callable:
|
||||
|
||||
@functools.wraps(fa4_func)
|
||||
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
if _FA4.use_fa4(q, k, v):
|
||||
try:
|
||||
return fa4_func(q, k, v, *args, **kwargs)
|
||||
except NotImplementedError:
|
||||
raise
|
||||
except Exception as e: # CuTeDSL compile errors surface as ValueError
|
||||
_FA4.mark_broken(e)
|
||||
return fa2_func(q, k, v, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_func)
|
||||
def flash_attn_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -287,16 +317,6 @@ def flash_attn_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
|
||||
return out
|
||||
@@ -372,6 +392,7 @@ def flash_attn_fp4_func(
|
||||
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
|
||||
|
||||
|
||||
@_with_fa2_fallback(_flash_attn_2_varlen_func)
|
||||
def flash_attn_varlen_func(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
@@ -386,20 +407,6 @@ def flash_attn_varlen_func(
|
||||
deterministic: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Only returns the output, not the lse."""
|
||||
if _use_fa2(q, k, v):
|
||||
return _fa2_or_raise(_flash_attn_2_varlen_func)(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
)
|
||||
_check_dropout(dropout_p)
|
||||
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
|
||||
q,
|
||||
|
||||
@@ -1,251 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""torch.compile-traceable wrapper for the FA2/FA3/FA4 default attention path.
|
||||
|
||||
The FA4/cute path (`fa_version == "4"`) is already a registered
|
||||
`torch.library.custom_op` in `fastvideo.attention.utils.flash_attn_cute`, so
|
||||
dynamo treats it as a graph node. The external FA2/FA3 ``flash_attn_func`` is
|
||||
NOT — dynamo breaks the graph at the call site (observed: wanvideo.py
|
||||
self-attn, once per layer every step), which fragments the compiled region
|
||||
and blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom op
|
||||
(mirrors the FP4 `flash_attn_cute` template) so it becomes an
|
||||
opaque-but-traceable node. The kernel still runs eager inside the op
|
||||
(correct — flash-attn must run eager); only dynamo's treatment of the
|
||||
boundary changes, so numerics are unchanged (SSIM-gated).
|
||||
|
||||
Autograd: FA2 has full ``register_autograd`` parity — the custom op's
|
||||
backward calls flash_attn's ``_flash_attn_backward`` directly, so training
|
||||
backprops *through* the op (no graph break on the training path either).
|
||||
FA3 currently keeps the no-backward + carve-out pattern from PR #1373
|
||||
because FA3's private backward signature wants validation on a real Hopper
|
||||
box (gated on Kuan-Hao's Modal FA3 setup PR). Once that lands the FA3 path
|
||||
can mirror FA2.
|
||||
|
||||
Lives in `attention/utils/` (sibling of `flash_attn_cute.py` and
|
||||
`flash_attn_no_pad.py`) so it can be imported by any backend that wants the
|
||||
traceable FA default call without pulling in backend dispatch logic. The
|
||||
backend (`attention/backends/flash_attn.py`) just imports
|
||||
`flash_attn_func_compilable` and `fa_version` from here.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Pick the same backend the rest of FastVideo picked for `flash_attn_func`
|
||||
# (FA4/cute → FA3 → FA2). Mirror the precedence used in
|
||||
# `attention/utils/flash_attn_no_pad.py` so the two probes always agree.
|
||||
#
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1: its CuTeDSL
|
||||
# kernels JIT-compile per shape family and can fail at runtime on some
|
||||
# arch/shape combinations, so it is never auto-selected just because it is
|
||||
# installed.
|
||||
if envs.FASTVIDEO_FA4:
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
fa_version = "4"
|
||||
else:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 no longer has a different API, see following commit:
|
||||
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
try:
|
||||
if importlib.util.find_spec("flash_attn.cute") is not None:
|
||||
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
|
||||
"set FASTVIDEO_FA4=1 to use it for inference.")
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
if fa_version == "2":
|
||||
# Scope: this op covers exactly the q/k/v + softmax_scale + causal call
|
||||
# shape used by FlashAttentionImpl.forward's default branch (see
|
||||
# `flash_attn_func_compilable(...)` call site in
|
||||
# `attention/backends/flash_attn.py`). The masked/no-pad and varlen /
|
||||
# cross-attn paths use different entry points
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_*`) which live in
|
||||
# `attention/utils/flash_attn_no_pad.py`. The wrapper's signature is the
|
||||
# contract: any extra kwarg (dropout_p, window_size, alibi_slopes,
|
||||
# deterministic, return_attn_probs, ...) raises TypeError at the call
|
||||
# site, so silent loss of kwargs is not a failure mode.
|
||||
from flash_attn.flash_attn_interface import _flash_attn_backward as _fa2_backward
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
# `return_attn_probs=True` asks FA2 to also return softmax_lse +
|
||||
# S_dmask. We need softmax_lse to feed the backward; S_dmask is the
|
||||
# dropout mask (always None here since dropout_p is fixed at 0).
|
||||
out, softmax_lse, _ = _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal, return_attn_probs=True)
|
||||
return out, softmax_lse
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del softmax_scale, causal
|
||||
# FA2 default path: out = [batch, seqlen_q, nheads, head_dim_v],
|
||||
# softmax_lse = [batch, nheads, seqlen_q], fp32 regardless of q dtype.
|
||||
b, sq, hq = q.shape[0], q.shape[1], q.shape[2]
|
||||
out = q.new_empty(b, sq, hq, v.shape[-1])
|
||||
lse = q.new_empty(b, hq, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_default_setup_context(ctx, inputs, output):
|
||||
q, k, v, softmax_scale, causal = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(q, k, v, out, lse)
|
||||
# `lse` is an auxiliary output we save to feed FA2's backward; nobody
|
||||
# should differentiate through it. Mark it non-differentiable so
|
||||
# autograd errors loudly if a caller wires it into a loss, rather
|
||||
# than silently producing zero/None grads through the `del grad_lse`
|
||||
# in our backward.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's *forward* substitutes `1 / sqrt(head_dim)` for `softmax_scale=None`
|
||||
# internally; FA2's *backward* (`_flash_attn_backward`) demands a concrete
|
||||
# float in its C++ schema and rejects None at the binding boundary. Resolve
|
||||
# the default here so the value saved on ctx (and passed to backward) is
|
||||
# always a real float — matches what FA2's own autograd.Function does.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = q.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
|
||||
def _flash_attn_default_backward(ctx, grad_out, grad_lse):
|
||||
# We only differentiate `out`; softmax_lse is saved-for-backward, not
|
||||
# a real differentiable output. (Mirrors the FP4 cute template.)
|
||||
del grad_lse
|
||||
q, k, v, out, lse = ctx.saved_tensors
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
# FA2's `_flash_attn_backward` writes into dq/dk/dv in place. The
|
||||
# extra kwargs (window_size_*, softcap, alibi_slopes, deterministic,
|
||||
# rng_state) are pinned to the same defaults the forward wrapper
|
||||
# uses — flash-attn==2.8.1 (the version FastVideo pins) requires
|
||||
# all of them explicitly. `rng_state=None` is correct for our
|
||||
# `dropout_p=0` configuration.
|
||||
_fa2_backward(
|
||||
grad_out,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
out,
|
||||
lse,
|
||||
dq,
|
||||
dk,
|
||||
dv,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=False,
|
||||
rng_state=None,
|
||||
)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
_flash_attn_default_backward,
|
||||
setup_context=_flash_attn_default_setup_context,
|
||||
)
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Backward is registered: autograd flows through the op (training
|
||||
# path is also traceable; no carve-out needed). Public API matches
|
||||
# `flash_attn_func` — returns just `out`; we drop the saved-for-
|
||||
# backward `lse` here so callers see the original single-tensor
|
||||
# contract.
|
||||
out, _ = torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
return out
|
||||
elif fa_version == "3":
|
||||
# FA3 path: same forward+fake custom op as the original PR #1373, with
|
||||
# the autograd carve-out kept. The full backward (mirroring the FA2 leg
|
||||
# above) wants a Hopper box for grad-check validation, which we don't
|
||||
# have until Kuan-Hao's Modal FA3 setup PR lands. Until then this keeps
|
||||
# inference traceable + training correct (via the original
|
||||
# autograd.Function path + a pre-PR-style graph break on training).
|
||||
_fa_default = flash_attn_func
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_default_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_default_forward(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
|
||||
def _flash_attn_default_forward_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float | None,
|
||||
causal: bool,
|
||||
) -> torch.Tensor:
|
||||
del softmax_scale, causal
|
||||
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
|
||||
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
# Autograd carve-out. The custom op above registers a forward + fake
|
||||
# kernel but NO backward (register_autograd), so it is opaque to
|
||||
# autograd. Inference runs under no_grad / inference_mode and routes
|
||||
# through the traceable custom op — that is the torch.compile win, and
|
||||
# the only path this PR claims. Training backprops through attention,
|
||||
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
|
||||
# (itself an autograd.Function, so backward is correct) at the cost of a
|
||||
# dynamo graph break on the training path — i.e. pre-PR behavior, no
|
||||
# regression. Full autograd parity for the custom op (mirroring the FP4
|
||||
# cute template) is a tracked follow-up.
|
||||
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
|
||||
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
|
||||
elif fa_version == "4":
|
||||
# FA4 path: `flash_attn_func` is already a torch.library custom op
|
||||
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
|
||||
# passthrough is enough — no extra registration needed.
|
||||
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
|
||||
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
|
||||
else:
|
||||
# Defensive: the probe above only ever sets fa_version to "2", "3",
|
||||
# or "4"; an unexpected value means an import/probe regression and
|
||||
# we want a loud error at import, not a silent NameError later.
|
||||
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
|
||||
f"'2', '3', or '4' from the import probe above.")
|
||||
@@ -21,46 +21,27 @@ from einops import rearrange
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
|
||||
from fastvideo import envs
|
||||
|
||||
|
||||
def _resolve_flash_attn_varlen_func() -> tuple[Any, str]:
|
||||
if envs.FASTVIDEO_FA4:
|
||||
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
|
||||
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
|
||||
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
|
||||
# not ImportError: importers like bsa_attn.py treat ImportError as
|
||||
# "flash-attn not installed" and silently degrade to reference kernels.
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
except ImportError as e:
|
||||
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
|
||||
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
|
||||
"or unset FASTVIDEO_FA4.") from e
|
||||
|
||||
return flash_attn_varlen_func_cute, "4"
|
||||
def _resolve_flash_attn_varlen_func() -> Any:
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
|
||||
return flash_attn_varlen_func_interface, "3"
|
||||
return flash_attn_varlen_func_cute
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_flash, "2"
|
||||
return flash_attn_varlen_func_interface
|
||||
except ImportError:
|
||||
from flash_attn import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
|
||||
|
||||
return flash_attn_varlen_func_flash
|
||||
|
||||
|
||||
flash_attn_varlen_func_impl, _FA_VARLEN_VERSION = _resolve_flash_attn_varlen_func()
|
||||
|
||||
# FA2-only: the private varlen backward we register against the custom ops
|
||||
# below. FA3 / FA4 have different private signatures and validation paths
|
||||
# (Hopper / Blackwell boxes) — those legs keep the autograd carve-out
|
||||
# pattern from PR #1373 until their setup PRs land.
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
from flash_attn.flash_attn_interface import (
|
||||
_flash_attn_varlen_backward as _fa2_varlen_backward, )
|
||||
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
|
||||
|
||||
|
||||
def flash_attn_no_pad(
|
||||
@@ -199,473 +180,3 @@ def flash_attn_varlen_qk_no_pad(
|
||||
h=nheads,
|
||||
)
|
||||
return output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# torch.compile traceability + register_autograd parity for the masked /
|
||||
# varlen attention paths.
|
||||
#
|
||||
# Wraps the two entry points `FlashAttentionImpl.forward` calls
|
||||
# (`flash_attn_no_pad`, `flash_attn_varlen_qk_no_pad`) as
|
||||
# `torch.library.custom_op`s so dynamo sees one traceable node — the
|
||||
# internal unpad / pad bookkeeping (data-dependent `nnz` shapes) runs
|
||||
# eager inside the op, and the op's outputs are the statically-shaped
|
||||
# padded tensors. This mirrors the FA2 default-path wrapper in
|
||||
# `fastvideo/attention/backends/flash_attn.py`.
|
||||
#
|
||||
# Autograd: on FA2 we register a real backward (`register_autograd`)
|
||||
# that calls FA2's `_flash_attn_varlen_backward` on the unpadded form
|
||||
# — re-unpadding the saved padded tensors using the saved mask. The
|
||||
# `softmax_lse` from the varlen forward is naturally unpadded
|
||||
# (`[nheads, total_q]`); we pad it to `[batch, nheads, seqlen]` on
|
||||
# the way out (statically shaped) and re-unpad in backward. So
|
||||
# training backprops *through* the op (no graph break on the training
|
||||
# path either).
|
||||
#
|
||||
# FA3 / FA4 keep the autograd carve-out pattern from PR #1373: the
|
||||
# custom op has forward + fake only, and `*_compilable` falls back to
|
||||
# the original autograd.Function for grad-enabled calls. Those legs
|
||||
# are gated on Hopper-class / Blackwell-class boxes for backward
|
||||
# validation and ship as separate follow-ups.
|
||||
|
||||
if _FA_VARLEN_VERSION == "2":
|
||||
# ---------- masked self-attention: flash_attn_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, s, _three, h, d = qkv.shape
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_qkvpacked_func(x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz, h, d] -> [b, s, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), indices, b, s),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: FA2 varlen returns [nheads, total_q]. Transpose to [total_q,
|
||||
# nheads], pad to [b, s, nheads], permute to [b, nheads, s] — statically
|
||||
# shaped so register_fake matches.
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), indices, b, s).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
out = qkv.new_empty(b, s, h, d)
|
||||
lse = qkv.new_empty(b, h, s, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_no_pad_setup_context(ctx, inputs, output):
|
||||
qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(qkv, out, lse, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
# FA2's varlen backward requires a concrete float for softmax_scale.
|
||||
if softmax_scale is None:
|
||||
softmax_scale = qkv.shape[-1]**-0.5 # head_dim from qkv's last dim
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
# lse is saved-for-backward, not differentiated.
|
||||
del grad_lse
|
||||
qkv, out_padded, lse_padded, key_padding_mask = ctx.saved_tensors
|
||||
b, s, _three, h, d = qkv.shape
|
||||
|
||||
# One `unpad_input` call (on qkv) gives us indices + cu_seqlens + max_s;
|
||||
# reuse those for out / dout / lse below via direct indexing instead
|
||||
# of redundant `unpad_input` calls (each of which would re-run
|
||||
# `nonzero` + `cumsum` + a `.max().item()` GPU→CPU sync).
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
|
||||
q_unpad, k_unpad, v_unpad = (t.contiguous() for t in x_unpad.unbind(dim=1))
|
||||
|
||||
# Direct-index variants reuse `indices` (computed above).
|
||||
out_unpad = out_padded.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[indices].view(-1, h, d).contiguous()
|
||||
# lse_padded [b, h, s] -> [b, s, h] -> [nnz, h] -> [h, nnz].
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens,
|
||||
cu_seqlens_k=cu_seqlens,
|
||||
max_seqlen_q=max_s,
|
||||
max_seqlen_k=max_s,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# Re-pad each grad and stack into dqkv.
|
||||
def _repad(dt_unpad: torch.Tensor) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, b, s)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dqkv = torch.stack([_repad(dq_unpad), _repad(dk_unpad), _repad(dv_unpad)], dim=2)
|
||||
# 6 inputs total: qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic.
|
||||
return dqkv, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
_flash_attn_no_pad_backward,
|
||||
setup_context=_flash_attn_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- cross-attention: flash_attn_varlen_qk_no_pad (FA2) ----------
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
b, sq, h, d = query.shape
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
v_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
v_unpad = rearrange(v_unpad, "nnz (h d) -> nnz h d", h=h)
|
||||
out_unpad, lse_unpad, _ = flash_attn_varlen_func_impl(q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
cu_seqlens_q,
|
||||
cu_seqlens_k,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
deterministic=deterministic,
|
||||
return_attn_probs=True)
|
||||
# Pad out: [nnz_q, h, d] -> [b, sq, h, d]
|
||||
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), q_indices, b, sq),
|
||||
"b s (h d) -> b s h d",
|
||||
h=h)
|
||||
# Pad lse: [h, nnz_q] -> [b, h, sq]
|
||||
lse_padded = pad_input(lse_unpad.t().contiguous(), q_indices, b, sq).permute(0, 2, 1).contiguous()
|
||||
return out_padded, lse_padded
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# out_padded ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
out = query.new_empty(b, sq, h, value.shape[-1])
|
||||
lse = query.new_empty(b, h, sq, dtype=torch.float32)
|
||||
return out, lse
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_setup_context(ctx, inputs, output):
|
||||
(query, key, value, query_padding_mask, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic) = inputs
|
||||
out, lse = output
|
||||
ctx.save_for_backward(query, key, value, out, lse, query_padding_mask, key_padding_mask)
|
||||
# Auxiliary output, not differentiable — see default-path note.
|
||||
ctx.mark_non_differentiable(lse)
|
||||
if softmax_scale is None:
|
||||
softmax_scale = query.shape[-1]**-0.5
|
||||
ctx.softmax_scale = softmax_scale
|
||||
ctx.causal = causal
|
||||
ctx.dropout_p = dropout_p
|
||||
ctx.deterministic = deterministic
|
||||
|
||||
def _flash_attn_varlen_qk_no_pad_backward(ctx, grad_out, grad_lse):
|
||||
del grad_lse
|
||||
(query, key, value, out_padded, lse_padded, query_padding_mask, key_padding_mask) = ctx.saved_tensors
|
||||
b, sq, h, d = query.shape
|
||||
sk = key.shape[1]
|
||||
|
||||
# One `unpad_input` call per distinct mask; reuse the returned
|
||||
# indices via direct indexing for everything else that shares
|
||||
# the same mask (v with k_mask; out/dout/lse with q_mask; the
|
||||
# final repad of dk/dv also reuses k_indices). Avoids ~4
|
||||
# redundant `unpad_input` calls + their GPU→CPU `.max().item()`
|
||||
# syncs.
|
||||
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
|
||||
query_padding_mask)
|
||||
k_unpad, k_indices, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
|
||||
key_padding_mask)
|
||||
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
|
||||
v_unpad = value.flatten(0, 1)[k_indices].view(-1, h, d).contiguous()
|
||||
|
||||
# out / dout / lse follow q's shape, so index with q_indices.
|
||||
out_unpad = out_padded.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
dout_unpad = grad_out.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
|
||||
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[q_indices].t().contiguous()
|
||||
|
||||
dq_unpad = torch.empty_like(q_unpad)
|
||||
dk_unpad = torch.empty_like(k_unpad)
|
||||
dv_unpad = torch.empty_like(v_unpad)
|
||||
_fa2_varlen_backward(
|
||||
dout_unpad,
|
||||
q_unpad,
|
||||
k_unpad,
|
||||
v_unpad,
|
||||
out_unpad,
|
||||
lse_unpad,
|
||||
dq_unpad,
|
||||
dk_unpad,
|
||||
dv_unpad,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
cu_seqlens_k=cu_seqlens_k,
|
||||
max_seqlen_q=max_seqlen_q,
|
||||
max_seqlen_k=max_seqlen_k,
|
||||
dropout_p=ctx.dropout_p,
|
||||
softmax_scale=ctx.softmax_scale,
|
||||
causal=ctx.causal,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=ctx.deterministic,
|
||||
rng_state=None,
|
||||
)
|
||||
|
||||
# k_indices is already available from the unpad_input above —
|
||||
# no need to recompute it for the dk/dv repad.
|
||||
def _repad(dt_unpad: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int) -> torch.Tensor:
|
||||
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, batch, seqlen)
|
||||
return rearrange(padded, "b s (h d) -> b s h d", h=h)
|
||||
|
||||
dq_padded = _repad(dq_unpad, q_indices, b, sq)
|
||||
dk_padded = _repad(dk_unpad, k_indices, b, sk)
|
||||
dv_padded = _repad(dv_unpad, k_indices, b, sk)
|
||||
# 9 inputs total: query, key, value, q_mask, k_mask, causal, dropout_p,
|
||||
# softmax_scale, deterministic.
|
||||
return dq_padded, dk_padded, dv_padded, None, None, None, None, None, None
|
||||
|
||||
torch.library.register_autograd(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
_flash_attn_varlen_qk_no_pad_backward,
|
||||
setup_context=_flash_attn_varlen_qk_no_pad_setup_context,
|
||||
)
|
||||
|
||||
# ---------- public dispatchers (FA2: autograd flows through the op) -----
|
||||
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_no_pad`` (registered op,
|
||||
full register_autograd on FA2 — both inference and training go through
|
||||
the op, no graph break on either)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
return out
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
"""dynamo-traceable wrapper around ``flash_attn_varlen_qk_no_pad`` (registered
|
||||
op, full register_autograd on FA2)."""
|
||||
out, _ = torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
return out
|
||||
|
||||
else:
|
||||
# ---------- FA3 / FA4: carve-out (forward+fake only, no real backward) ---
|
||||
# Same pattern as the parked varlen-extension and the FA3 default leg in
|
||||
# `fastvideo/attention/backends/flash_attn.py`. Real backward for these
|
||||
# versions is a follow-up gated on Hopper / Blackwell box validation.
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_no_pad_forward(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_no_pad( # type: ignore[no-untyped-call]
|
||||
qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
|
||||
def _flash_attn_no_pad_forward_fake(
|
||||
qkv: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
|
||||
b, s, _three, h, d = qkv.shape
|
||||
return qkv.new_empty(b, s, h, d)
|
||||
|
||||
@torch.library.custom_op(
|
||||
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
|
||||
mutates_args=(),
|
||||
device_types="cuda",
|
||||
)
|
||||
def _flash_attn_varlen_qk_no_pad_forward(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
return flash_attn_varlen_qk_no_pad( # type: ignore[no-untyped-call]
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
|
||||
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
|
||||
def _flash_attn_varlen_qk_no_pad_forward_fake(
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
query_padding_mask: torch.Tensor,
|
||||
key_padding_mask: torch.Tensor,
|
||||
causal: bool,
|
||||
dropout_p: float,
|
||||
softmax_scale: float | None,
|
||||
deterministic: bool,
|
||||
) -> torch.Tensor:
|
||||
del key, query_padding_mask, key_padding_mask
|
||||
del causal, dropout_p, softmax_scale, deterministic
|
||||
b, sq, h, _ = query.shape
|
||||
# `out`'s head_dim comes from value (d_v), matching the real forward's
|
||||
# output ([b, sq, h, d_v]); it can differ from query's d_q.
|
||||
return query.new_empty(b, sq, h, value.shape[-1])
|
||||
|
||||
def flash_attn_no_pad_compilable(qkv,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and qkv.requires_grad:
|
||||
return flash_attn_no_pad(qkv,
|
||||
key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
|
||||
deterministic)
|
||||
|
||||
def flash_attn_varlen_qk_no_pad_compilable(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask,
|
||||
key_padding_mask,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=None,
|
||||
deterministic=False):
|
||||
if torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad):
|
||||
return flash_attn_varlen_qk_no_pad(query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=causal,
|
||||
dropout_p=dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
deterministic=deterministic)
|
||||
return torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
|
||||
key_padding_mask, causal, dropout_p,
|
||||
softmax_scale, deterministic)
|
||||
|
||||
@@ -1,9 +1,6 @@
|
||||
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
|
||||
from fastvideo.configs.models.dits.flux import FluxDiTConfig
|
||||
from fastvideo.configs.models.dits.flux_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
|
||||
@@ -16,8 +13,7 @@ from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
|
||||
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
|
||||
"StableAudioConfig", "GlmImageDiTConfig"
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
]
|
||||
|
||||
@@ -1,119 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 VFM Transformer FastVideo dataclass configs.
|
||||
|
||||
Architecture is 1:1 with the published ``nvidia/Cosmos3-Nano`` checkpoint
|
||||
(``transformer/config.json``; class ``Cosmos3OmniTransformer`` / framework
|
||||
``Cosmos3VFMNetwork``). Field values match that config so the FastVideo native
|
||||
DiT builds a parameter tree matching the checkpoint's state-dict surface
|
||||
(814 tensors / 44 patterns, validated 2026-06-06).
|
||||
|
||||
Reference of record: ``cosmos-framework`` (NVIDIA). The checkpoint is a single
|
||||
``layers`` ModuleList of dual-pathway (understanding/text + generation/vision)
|
||||
decoder blocks; per layer: ``self_attn`` with und (``to_{q,k,v}``/``to_out``)
|
||||
and gen (``add_{q,k,v}_proj``/``to_add_out``) projections + QK-norms, plus
|
||||
``mlp`` (und) and ``mlp_moe_gen`` (gen), and four RMSNorms. Top level adds
|
||||
``embed_tokens``/``norm``/``norm_moe_gen``/``lm_head``/``proj_in``/``proj_out``/
|
||||
``time_embedder`` and dormant ``action_*``/``audio_*`` heads. The checkpoint
|
||||
remap lives in ``scripts/checkpoint_conversion/cosmos3_convert.py``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_cosmos3_transformer_block(name: str, module) -> bool:
|
||||
"""FSDP shard boundary: the dual-pathway decoder blocks ``layers.{i}``."""
|
||||
del module
|
||||
parts = name.split(".")
|
||||
return "layers" in parts and parts[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3ArchConfig(DiTArchConfig):
|
||||
"""Architecture config for the Cosmos3 omni DiT (Cosmos3-Nano).
|
||||
|
||||
1:1 with ``transformer/config.json``. The action/sound heads ship in the
|
||||
checkpoint, so they are constructed for strict-load parity even though the
|
||||
PR1 video path (T2V/I2V/T2I) leaves them dormant.
|
||||
"""
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_cosmos3_transformer_block])
|
||||
|
||||
# Conversion is owned by scripts/checkpoint_conversion/cosmos3_convert.py;
|
||||
# the native module tree is the source of truth for parameter names.
|
||||
param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
# ---- Backbone (Qwen3-VL-text) ----
|
||||
hidden_size: int = 4096
|
||||
num_hidden_layers: int = 36
|
||||
num_attention_heads: int = 32
|
||||
num_key_value_heads: int = 8 # GQA (4 query groups)
|
||||
head_dim: int = 128
|
||||
intermediate_size: int = 12288
|
||||
hidden_act: str = "silu"
|
||||
vocab_size: int = 151936
|
||||
rms_norm_eps: float = 1e-6
|
||||
attention_bias: bool = False
|
||||
qk_norm_for_diffusion: bool = True
|
||||
qk_norm_for_text: bool = True
|
||||
use_moe: bool = True # dual-pathway weights; sparse routing unused
|
||||
joint_attn_implementation: str = "two_way"
|
||||
freeze_und: bool = False
|
||||
|
||||
# ---- Position embedding (unified 3D MRoPE) ----
|
||||
position_embedding_type: str = "unified_3d_mrope"
|
||||
rope_theta: float = 5_000_000.0
|
||||
max_position_embeddings: int = 262144
|
||||
mrope_section: list[int] = field(default_factory=lambda: [24, 20, 20])
|
||||
mrope_interleaved: bool = True
|
||||
unified_3d_mrope_reset_spatial_ids: bool = True
|
||||
temporal_modality_margin: int = 15000 # unified_3d_mrope_temporal_modality_margin
|
||||
|
||||
# ---- VAE / patch geometry ----
|
||||
latent_patch_size: int = 2
|
||||
latent_channel: int = 48
|
||||
patch_latent_dim: int = 192 # latent_patch_size**2 * latent_channel
|
||||
|
||||
# ---- Diffusion conditioning ----
|
||||
timestep_scale: float = 0.001
|
||||
|
||||
# ---- Temporal / FPS modulation ----
|
||||
base_fps: float = 24.0
|
||||
temporal_compression_factor: int = 4
|
||||
enable_fps_modulation: bool = True
|
||||
video_temporal_causal: bool = False
|
||||
|
||||
# ---- Action generation head (dormant in PR1 video path) ----
|
||||
action_gen: bool = True
|
||||
action_dim: int = 64
|
||||
max_action_dim: int = 64
|
||||
num_embodiment_domains: int = 32
|
||||
|
||||
# ---- Sound generation head (dormant in PR1 video path) ----
|
||||
sound_gen: bool = True
|
||||
sound_dim: int = 64
|
||||
sound_latent_fps: float = 25.0
|
||||
temporal_compression_factor_sound: int = 1
|
||||
|
||||
# ---- BaseDiT bookkeeping ----
|
||||
in_channels: int = 48
|
||||
out_channels: int = 48
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
# Video DiT contract: latent channels == VAE z_dim.
|
||||
self.num_channels_latents = self.latent_channel
|
||||
if not self.out_channels:
|
||||
self.out_channels = self.in_channels
|
||||
# Derived: patchify packs latent_patch_size**2 spatial patches * channels.
|
||||
self.patch_latent_dim = self.latent_patch_size**2 * self.latent_channel
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3VideoConfig(DiTConfig):
|
||||
"""Pipeline-level Cosmos3 DiT config (T2V / I2V / T2I share this surface)."""
|
||||
|
||||
arch_config: DiTArchConfig = field(default_factory=Cosmos3ArchConfig)
|
||||
prefix: str = "Cosmos3"
|
||||
@@ -1,71 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldArchConfig(WanVideoArchConfig):
|
||||
"""DreamX-World DiT config with camera PRoPE control fields."""
|
||||
|
||||
add_control_adapter: bool = True
|
||||
cam_method: str | None = "prope"
|
||||
attn_compress: int = 1
|
||||
cam_self_attn_layers: tuple[int, ...] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DreamXWorldArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldARArchConfig(DreamXWorldArchConfig):
|
||||
"""DreamX-World-5B autoregressive causal DiT config."""
|
||||
|
||||
model_type: str = "ti2v"
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
text_len: int = 512
|
||||
text_dim: int = 4096
|
||||
freq_dim: int = 256
|
||||
attn_compress: int = 4
|
||||
cam_self_attn_layers: tuple[int, ...] | None = tuple(range(30))
|
||||
local_attn_size: int = 12
|
||||
sink_size: int = 3
|
||||
num_frames_per_block: int = 3
|
||||
rope_cache_policy: str = "block_relativistic"
|
||||
# The official AR checkpoint (AMAP-ML/DreamX-World ``model.safetensors``)
|
||||
# already uses FastVideo's native key names and the converter copies the
|
||||
# tensors verbatim, so every rule is an identity. The rules enumerate the
|
||||
# full state-dict surface of ``DreamXWorldARTransformer3DModel`` (norm1 /
|
||||
# norm2 / head.norm are affine-free and have no parameters).
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^patch_embedding\.(.*)$": r"patch_embedding.\1",
|
||||
r"^text_embedding\.([02])\.(.*)$": r"text_embedding.\1.\2",
|
||||
r"^time_embedding\.([02])\.(.*)$": r"time_embedding.\1.\2",
|
||||
r"^time_projection\.1\.(.*)$": r"time_projection.1.\1",
|
||||
r"^blocks\.(\d+)\.self_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.self_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_(q|k)\.weight$": r"blocks.\1.self_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.cross_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.cross_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_(q|k)\.weight$": r"blocks.\1.cross_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.cam_self_attn\.(q_proj|k_proj|v_proj|out_proj)\.(.*)$": r"blocks.\1.cam_self_attn.\2.\3",
|
||||
r"^blocks\.(\d+)\.cam_self_attn\.norm_(q|k)\.weight$": r"blocks.\1.cam_self_attn.norm_\2.weight",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm3.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.([02])\.(.*)$": r"blocks.\1.ffn.\2.\3",
|
||||
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.modulation",
|
||||
r"^head\.head\.(.*)$": r"head.head.\1",
|
||||
r"^head\.modulation$": r"head.modulation",
|
||||
})
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorldARConfig(DreamXWorldConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=DreamXWorldARArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -1,27 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxTransformer2DArchConfig(DiTArchConfig):
|
||||
|
||||
patch_size: int = 1
|
||||
in_channels: int = 64
|
||||
out_channels: int | None = None
|
||||
num_layers: int = 19
|
||||
num_single_layers: int = 38
|
||||
attention_head_dim: int = 128
|
||||
num_attention_heads: int = 24
|
||||
joint_attention_dim: int = 4096
|
||||
pooled_projection_dim: int = 768
|
||||
guidance_embeds: bool = True
|
||||
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FluxDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=FluxTransformer2DArchConfig)
|
||||
prefix: str = "flux"
|
||||
@@ -1,61 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTArchConfig(DiTArchConfig):
|
||||
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
|
||||
hidden_size: int = 4096
|
||||
num_attention_heads: int = 32
|
||||
attention_head_dim: int = 128
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_layers: int = 30
|
||||
|
||||
text_embed_dim: int = 1472
|
||||
time_embed_dim: int = 512
|
||||
condition_dim: int = 256
|
||||
|
||||
prior_vq_quantizer_codebook_size: int = 16384
|
||||
|
||||
patch_size: int = 2
|
||||
|
||||
max_height: int = 2048
|
||||
max_width: int = 2048
|
||||
|
||||
qk_norm: str = "layer_norm"
|
||||
eps: float = 1e-5
|
||||
|
||||
exclude_lora_layers: list[str] = field(
|
||||
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
|
||||
|
||||
param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
|
||||
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
|
||||
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
|
||||
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
|
||||
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
|
||||
})
|
||||
|
||||
reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.num_channels_latents = self.out_channels
|
||||
|
||||
|
||||
@dataclass
|
||||
class GlmImageDiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
|
||||
prefix: str = "GlmImage"
|
||||
@@ -2,24 +2,14 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
def _is_kandinsky5_transformer_block(n: str, m) -> bool:
|
||||
return ("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
|
||||
|
||||
@dataclass
|
||||
class Kandinsky5ArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_kandinsky5_transformer_block])
|
||||
|
||||
# NABLA block-sparse attention for attention_type="nabla" checkpoints, plus
|
||||
# the dense backends every DiT supports.
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.NABLA_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
)
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
lambda n, m:
|
||||
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
|
||||
])
|
||||
|
||||
# Native FastVideo implementation uses the same parameter names as diffusers
|
||||
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Literal
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -20,11 +19,6 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
# AnyFlow dual-timestep checkpoints expose delta_embedder weights with the
|
||||
# same internal layout as time_embedder. The regex is harmless on plain
|
||||
# Wan checkpoints (no delta_embedder keys to match).
|
||||
r"^condition_embedder\.delta_embedder\.linear_1\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.delta_embedder\.linear_2\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
@@ -92,14 +86,6 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
|
||||
rope_cache_policy: str = "absolute"
|
||||
|
||||
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
|
||||
# the legacy single-timestep forward (no delta_embedder allocated, no
|
||||
# extra computation on the embedder forward path).
|
||||
r_embedder: bool = False
|
||||
r_embedder_fusion: Literal["additive", "gated"] = "additive"
|
||||
r_embedder_gate_value: float = 0.25
|
||||
r_embedder_deltatime_type: Literal["r", "t-r"] = "r"
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
self.out_channels = self.out_channels or self.in_channels
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user