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).
|
||||
@@ -1,9 +1,5 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"description": "Wan2.1 T2V 1.3B inference performance",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -125,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
|
||||
@@ -136,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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
+2
-4
@@ -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" && \
|
||||
|
||||
@@ -103,7 +103,6 @@ can merge a PR.
|
||||
|---|---|---|
|
||||
| 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 |
|
||||
@@ -145,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` |
|
||||
|
||||
@@ -172,45 +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-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1
|
||||
}
|
||||
```
|
||||
|
||||
`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-1.3b`. |
|
||||
| `variant_id` | Intentional recipe family, such as `canonical`. |
|
||||
| `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 `recipe`, `metric_threshold_policy`, and
|
||||
`quality_metadata`, must be JSON objects when present.
|
||||
|
||||
Recipe fingerprinting, hardware/software profile IDs, exact-identity
|
||||
comparison, metric-specific threshold policy behavior, promoted baselines, and
|
||||
dashboard regrouping are separate follow-up changes. Until those land, rolling
|
||||
baseline comparison remains keyed by `(model_id, gpu_type)`.
|
||||
|
||||
### Raw record (`results/perf_*.json`)
|
||||
|
||||
Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
@@ -218,10 +179,6 @@ Written by `test_inference_performance.py`. One file per benchmark run.
|
||||
```jsonc
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
|
||||
"device": "NVIDIA L40S",
|
||||
"num_gpus": 2,
|
||||
@@ -322,16 +279,11 @@ observability. When pytest passes, the rolling-baseline phase 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": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"model": { "model_path": "...", "model_short_name": "..." },
|
||||
"init_kwargs": { "num_gpus": 1, ... },
|
||||
"generation_kwargs": { "num_frames": 45, ... },
|
||||
@@ -351,10 +303,6 @@ observability. When pytest passes, the rolling-baseline phase emits:
|
||||
}
|
||||
```
|
||||
|
||||
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`.
|
||||
|
||||
2. The pytest test auto-discovers all configs — no test code needed. CI
|
||||
picks it up on the next `/test performance` run.
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 | `GD-ML/DreamX-World-5B-Cam` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B AR | `GD-ML/DreamX-World-5B` | 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 | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
|
||||
@@ -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,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", "GD-ML/DreamX-World-5B-Cam")
|
||||
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()
|
||||
@@ -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()
|
||||
@@ -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,39 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib.util
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
|
||||
# kernel package's FASTVIDEO_VSA_CUTEDSL: 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. Below sm90 a capability
|
||||
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
|
||||
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
|
||||
# JIT, observed on sm_89).
|
||||
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:
|
||||
except ImportError:
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
@@ -45,12 +21,6 @@ else:
|
||||
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
|
||||
|
||||
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
|
||||
# already a registered torch.library custom op, so dynamo treats it as a
|
||||
@@ -117,10 +87,9 @@ if fa_version in ("2", "3"):
|
||||
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` (from `flash_attn_cute`) goes through a
|
||||
# registered torch.library custom op (with an FA4 backward on sm90+;
|
||||
# grad-enabled and GQA calls below sm90 route to FA2), so a passthrough
|
||||
# is enough — no extra registration needed.
|
||||
# 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:
|
||||
@@ -130,6 +99,17 @@ else:
|
||||
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,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
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.
|
||||
@@ -291,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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -21,35 +21,24 @@ 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() -> Any:
|
||||
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
|
||||
try:
|
||||
from fastvideo.attention.utils.flash_attn_cute import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
|
||||
|
||||
return flash_attn_varlen_func_cute
|
||||
try:
|
||||
from flash_attn_interface import (
|
||||
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
|
||||
|
||||
return flash_attn_varlen_func_interface
|
||||
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
|
||||
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 = _resolve_flash_attn_varlen_func()
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
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_2 import Flux2Config
|
||||
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
|
||||
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
|
||||
@@ -14,7 +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", "LongCatVideoConfig", "LTX2VideoConfig",
|
||||
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
|
||||
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
|
||||
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
|
||||
]
|
||||
|
||||
@@ -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,7 +1,6 @@
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.cosmos import CosmosConfig
|
||||
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
@@ -17,6 +16,5 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
|
||||
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
|
||||
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
|
||||
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "MatrixGame2I2V480PConfig",
|
||||
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -1,127 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B-Cam FastVideo model configuration helpers."""
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.dreamx_world import (DreamXWorldARArchConfig, DreamXWorldARConfig,
|
||||
DreamXWorldArchConfig, DreamXWorldConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
|
||||
from fastvideo.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.configs.pipelines.wan import LucyEditDevConfig, t5_postprocess_text
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_dit_config() -> DreamXWorldConfig:
|
||||
"""Return the DreamX-World DiT config matching DreamX-World-5B-Cam."""
|
||||
return DreamXWorldConfig(arch_config=DreamXWorldArchConfig(
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
in_channels=48,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
cross_attn_norm=True,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=None,
|
||||
))
|
||||
|
||||
|
||||
def make_dreamx_world_5b_ar_dit_config() -> DreamXWorldARConfig:
|
||||
"""Return the DreamX-World-5B autoregressive causal DiT config."""
|
||||
return DreamXWorldARConfig(arch_config=DreamXWorldARArchConfig(
|
||||
model_type="ti2v",
|
||||
num_attention_heads=24,
|
||||
attention_head_dim=128,
|
||||
in_channels=48,
|
||||
out_channels=48,
|
||||
ffn_dim=14336,
|
||||
num_layers=30,
|
||||
cross_attn_norm=True,
|
||||
qk_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=4,
|
||||
cam_self_attn_layers=tuple(range(30)),
|
||||
local_attn_size=12,
|
||||
sink_size=3,
|
||||
num_frames_per_block=3,
|
||||
))
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_vae_config() -> WanVAEConfig:
|
||||
"""Return the Wan2.2 48-channel VAE config used by DreamX-World-5B-Cam."""
|
||||
return LucyEditDevConfig().vae_config
|
||||
|
||||
|
||||
def make_dreamx_world_5b_cam_text_encoder_config() -> T5Config:
|
||||
"""Return the UMT5-XXL text encoder config used by DreamX-World-5B-Cam."""
|
||||
return T5Config(
|
||||
arch_config=T5ArchConfig(
|
||||
vocab_size=256384,
|
||||
d_model=4096,
|
||||
d_kv=64,
|
||||
d_ff=10240,
|
||||
num_layers=24,
|
||||
num_decoder_layers=None,
|
||||
num_heads=64,
|
||||
relative_attention_num_buckets=32,
|
||||
dropout_rate=0.0,
|
||||
text_len=512,
|
||||
feed_forward_proj="gelu",
|
||||
is_encoder_decoder=False,
|
||||
),
|
||||
prefix="umt5",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorld5BCamPipelineConfig(PipelineConfig):
|
||||
"""Pipeline config for the first-scope DreamX-World-5B-Cam mode."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_cam_dit_config)
|
||||
vae_config: VAEConfig = field(default_factory=make_dreamx_world_5b_cam_vae_config)
|
||||
text_encoder_configs: tuple[EncoderConfig,
|
||||
...] = field(default_factory=lambda: (make_dreamx_world_5b_cam_text_encoder_config(), ))
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda: (t5_postprocess_text, ))
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
|
||||
flow_shift: float | None = 3.0
|
||||
ti2v_task: bool = True
|
||||
expand_timesteps: bool = True
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
vae_precision: str = "fp32"
|
||||
vae_decode_precision: str | None = "bf16"
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
self.dit_config.expand_timesteps = self.expand_timesteps
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXWorld5BARPipelineConfig(DreamXWorld5BCamPipelineConfig):
|
||||
"""Pipeline config for DreamX-World-5B autoregressive forcing."""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_ar_dit_config)
|
||||
flow_shift: float | None = 5.0
|
||||
ti2v_task: bool = True
|
||||
is_causal: bool = True
|
||||
dmd_denoising_steps: tuple[int, ...] = (1000, 750, 500, 250)
|
||||
warp_denoising_step: bool = True
|
||||
context_noise: float = 0.1
|
||||
num_frames_per_block: int = 3
|
||||
color_correction_strength: float = 1.0
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
super().__post_init__()
|
||||
self.dit_config.expand_timesteps = True
|
||||
@@ -20,7 +20,6 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: str | None = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_FA4: bool = False
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
@@ -208,18 +207,9 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
# - "VIDEO_SPARSE_ATTN": use Video Sparse Attention
|
||||
# - "SAGE_ATTN": use Sage Attention
|
||||
# - "SAGE_ATTN_THREE": use Sage Attention 3
|
||||
# FLASH_ATTN uses FlashAttention-3/2; to run FlashAttention-4 set
|
||||
# FASTVIDEO_FA4=1 as well (see below).
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# If set (=1), the FLASH_ATTN backend uses FlashAttention-4
|
||||
# (flash_attn.cute). FA4 is opt-in and never auto-selected just because it
|
||||
# is installed. Below sm90, grad-enabled and GQA calls are routed to FA2
|
||||
# (FA4's backward asserts sm90+ and its pack_gqa fails to JIT there).
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
@@ -323,9 +323,9 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
|
||||
@@ -1,511 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldConfig
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather_with_unpad,
|
||||
sequence_model_parallel_shard)
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
|
||||
from fastvideo.models.dits.wanvideo import (LayerNormScaleShift,
|
||||
PatchEmbed,
|
||||
WanTimeTextImageEmbedding,
|
||||
WanTransformer3DModel,
|
||||
WanTransformerBlock)
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
|
||||
|
||||
def _dreamx_invert_se3(transforms: torch.Tensor) -> torch.Tensor:
|
||||
assert transforms.shape[-2:] == (4, 4)
|
||||
rot_inv = transforms[..., :3, :3].transpose(-1, -2)
|
||||
out = torch.zeros_like(transforms)
|
||||
out[..., :3, :3] = rot_inv
|
||||
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", rot_inv,
|
||||
transforms[..., :3, 3])
|
||||
out[..., 3, 3] = 1.0
|
||||
return out.to(dtype=transforms.dtype)
|
||||
|
||||
|
||||
def _dreamx_lift_k(intrinsics: torch.Tensor) -> torch.Tensor:
|
||||
assert intrinsics.shape[-2:] == (3, 3)
|
||||
out = torch.zeros(intrinsics.shape[:-2] + (4, 4),
|
||||
device=intrinsics.device,
|
||||
dtype=intrinsics.dtype)
|
||||
out[..., :3, :3] = intrinsics
|
||||
out[..., 3, 3] = 1.0
|
||||
return out
|
||||
|
||||
|
||||
def _dreamx_invert_k(intrinsics: torch.Tensor) -> torch.Tensor:
|
||||
assert intrinsics.shape[-2:] == (3, 3)
|
||||
out = torch.zeros_like(intrinsics)
|
||||
out[..., 0, 0] = 1.0 / intrinsics[..., 0, 0]
|
||||
out[..., 1, 1] = 1.0 / intrinsics[..., 1, 1]
|
||||
out[..., 0, 2] = -intrinsics[..., 0, 2] / intrinsics[..., 0, 0]
|
||||
out[..., 1, 2] = -intrinsics[..., 1, 2] / intrinsics[..., 1, 1]
|
||||
out[..., 2, 2] = 1.0
|
||||
return out.to(dtype=intrinsics.dtype)
|
||||
|
||||
|
||||
def _dreamx_apply_tiled_projmat(feats: torch.Tensor,
|
||||
matrix: torch.Tensor) -> torch.Tensor:
|
||||
batch, num_heads, seq_len, feat_dim = feats.shape
|
||||
proj_dim = matrix.shape[-1]
|
||||
assert feat_dim % proj_dim == 0
|
||||
|
||||
if matrix.shape[1] == seq_len:
|
||||
feats = feats.view(batch, num_heads, seq_len, feat_dim // proj_dim,
|
||||
proj_dim)
|
||||
out = torch.einsum("btij,bntpj->bntpi", matrix, feats)
|
||||
return out.reshape(batch, num_heads, seq_len, feat_dim)
|
||||
|
||||
cameras = matrix.shape[1]
|
||||
assert seq_len > cameras and seq_len % cameras == 0
|
||||
feats = feats.reshape(batch, num_heads, cameras, -1,
|
||||
feat_dim // proj_dim, proj_dim)
|
||||
out = torch.einsum("bcij,bncpkj->bncpki", matrix, feats)
|
||||
return out.reshape(batch, num_heads, seq_len, feat_dim)
|
||||
|
||||
|
||||
def _dreamx_prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
viewmats: torch.Tensor, intrinsics: torch.Tensor):
|
||||
batch, num_heads, seq_len, head_dim = q.shape
|
||||
cameras = viewmats.shape[1]
|
||||
assert q.shape == k.shape == v.shape
|
||||
assert viewmats.shape == (batch, cameras, 4, 4)
|
||||
assert intrinsics.shape == (batch, cameras, 3, 3)
|
||||
assert head_dim % 4 == 0
|
||||
|
||||
intrinsics_norm = torch.zeros_like(intrinsics)
|
||||
intrinsics_norm[..., 0, 0] = intrinsics[..., 0, 0]
|
||||
intrinsics_norm[..., 1, 1] = intrinsics[..., 1, 1]
|
||||
intrinsics_norm[..., 2, 2] = 1.0
|
||||
|
||||
proj = torch.einsum("...ij,...jk->...ik",
|
||||
_dreamx_lift_k(intrinsics_norm), viewmats)
|
||||
proj_t = proj.transpose(-1, -2).to(dtype=viewmats.dtype)
|
||||
proj_inv = torch.einsum(
|
||||
"...ij,...jk->...ik",
|
||||
_dreamx_invert_se3(viewmats),
|
||||
_dreamx_lift_k(_dreamx_invert_k(intrinsics_norm)),
|
||||
).to(dtype=viewmats.dtype)
|
||||
|
||||
q = _dreamx_apply_tiled_projmat(q, proj_t)
|
||||
k = _dreamx_apply_tiled_projmat(k, proj_inv)
|
||||
v = _dreamx_apply_tiled_projmat(v, proj_inv)
|
||||
return q, k, v, proj
|
||||
|
||||
|
||||
class DreamXPropeSelfAttention(nn.Module):
|
||||
"""DreamX-World parallel PRoPE camera self-attention branch."""
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
attn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str | bool = True,
|
||||
eps: float = 1e-6,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
assert attn_dim % num_heads == 0
|
||||
self.attn_dim = attn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = attn_dim // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
|
||||
self.q_proj = ReplicatedLinear(dim,
|
||||
attn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.q_proj")
|
||||
self.k_proj = ReplicatedLinear(dim,
|
||||
attn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.k_proj")
|
||||
self.v_proj = ReplicatedLinear(dim,
|
||||
attn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.v_proj")
|
||||
self.out_proj = ReplicatedLinear(attn_dim,
|
||||
dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.out_proj")
|
||||
|
||||
if qk_norm == "rms_norm":
|
||||
self.norm_q = RMSNorm(self.head_dim, eps=eps)
|
||||
self.norm_k = RMSNorm(self.head_dim, eps=eps)
|
||||
elif qk_norm in (True, "rms_norm_across_heads"):
|
||||
self.norm_q = RMSNorm(attn_dim, eps=eps)
|
||||
self.norm_k = RMSNorm(attn_dim, eps=eps)
|
||||
elif qk_norm is False:
|
||||
self.norm_q = nn.Identity()
|
||||
self.norm_k = nn.Identity()
|
||||
else:
|
||||
raise ValueError(f"Unsupported qk_norm for DreamX PRoPE: {qk_norm}")
|
||||
|
||||
nn.init.zeros_(self.out_proj.weight)
|
||||
if self.out_proj.bias is not None:
|
||||
nn.init.zeros_(self.out_proj.bias)
|
||||
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor,
|
||||
y_camera: dict[str, torch.Tensor]) -> torch.Tensor:
|
||||
if get_sp_world_size() > 1:
|
||||
# The transformer shards the sequence before the block loop and
|
||||
# this branch uses LocalAttention (no all-to-all): under
|
||||
# sequence parallelism each rank would attend only within its
|
||||
# own shard — silently wrong output. Fail loudly until this
|
||||
# path is ported to DistributedAttention and validated.
|
||||
raise NotImplementedError(
|
||||
"DreamXPropeSelfAttention does not support sequence "
|
||||
"parallelism yet (LocalAttention on a sharded sequence "
|
||||
"corrupts output). Run with sp_size=1.")
|
||||
batch_size, seq_len, _ = hidden_states.shape
|
||||
|
||||
query, _ = self.q_proj(hidden_states)
|
||||
key, _ = self.k_proj(hidden_states)
|
||||
value, _ = self.v_proj(hidden_states)
|
||||
|
||||
if self.qk_norm == "rms_norm":
|
||||
query = query.view(batch_size, seq_len, self.num_heads,
|
||||
self.head_dim)
|
||||
key = key.view(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
else:
|
||||
query = self.norm_q(query).view(batch_size, seq_len,
|
||||
self.num_heads, self.head_dim)
|
||||
key = self.norm_k(key).view(batch_size, seq_len, self.num_heads,
|
||||
self.head_dim)
|
||||
|
||||
value = value.view(batch_size, seq_len, self.num_heads, self.head_dim)
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
|
||||
query, key, value, output_projection = _dreamx_prope_qkv(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
viewmats=y_camera["viewmats"],
|
||||
intrinsics=y_camera["K"],
|
||||
)
|
||||
|
||||
out = self.attn(query.transpose(1, 2), key.transpose(1, 2),
|
||||
value.transpose(1, 2))
|
||||
out = _dreamx_apply_tiled_projmat(out.transpose(1, 2),
|
||||
output_projection).transpose(1, 2)
|
||||
out = out.flatten(2)
|
||||
out, _ = self.out_proj(out)
|
||||
return out
|
||||
|
||||
|
||||
class DreamXWorldTransformerBlock(WanTransformerBlock):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: int | None = None,
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...]
|
||||
| None = None,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
add_control_adapter: bool = True,
|
||||
cam_method: str | None = "prope",
|
||||
attn_compress: int = 1,
|
||||
cam_self_attn_layers: tuple[int, ...] | None = None,
|
||||
layer_idx: int | None = None):
|
||||
super().__init__(dim, ffn_dim, num_heads, qk_norm, cross_attn_norm,
|
||||
eps, added_kv_proj_dim,
|
||||
supported_attention_backends, quant_config, prefix)
|
||||
self.cam_self_attn = None
|
||||
add_cam_attn = add_control_adapter and cam_method == "prope"
|
||||
if add_cam_attn and cam_self_attn_layers is not None:
|
||||
add_cam_attn = layer_idx in cam_self_attn_layers
|
||||
if add_cam_attn:
|
||||
if num_heads % attn_compress != 0 or dim % attn_compress != 0:
|
||||
raise ValueError("DreamX attn_compress must divide dim and num_heads")
|
||||
self.cam_self_attn = DreamXPropeSelfAttention(
|
||||
dim,
|
||||
dim // attn_compress,
|
||||
num_heads // attn_compress,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.cam_self_attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
original_seq_len: int,
|
||||
y_camera: dict[str, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
orig_dtype = hidden_states.dtype
|
||||
|
||||
if temb.dim() == 4:
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb.float()).chunk(
|
||||
6, dim=2)
|
||||
shift_msa = shift_msa.squeeze(2)
|
||||
scale_msa = scale_msa.squeeze(2)
|
||||
gate_msa = gate_msa.squeeze(2)
|
||||
c_shift_msa = c_shift_msa.squeeze(2)
|
||||
c_scale_msa = c_scale_msa.squeeze(2)
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
|
||||
attn_output, _ = self.attn1(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
original_seq_len,
|
||||
freqs_cis=freqs_cis,
|
||||
)
|
||||
attn_output = attn_output.flatten(2)
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
if self.cam_self_attn is not None and y_camera is not None:
|
||||
attn_output = attn_output + self.cam_self_attn(
|
||||
norm_hidden_states, y_camera)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DreamXWorldTransformer3DModel(WanTransformer3DModel):
|
||||
_fsdp_shard_conditions = DreamXWorldConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = DreamXWorldConfig()._compile_conditions
|
||||
_supported_attention_backends = DreamXWorldConfig(
|
||||
)._supported_attention_backends
|
||||
param_names_mapping = DreamXWorldConfig().param_names_mapping
|
||||
reverse_param_names_mapping = DreamXWorldConfig().reverse_param_names_mapping
|
||||
lora_param_names_mapping = DreamXWorldConfig().lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: DreamXWorldConfig, hf_config: dict[str,
|
||||
Any]) -> None:
|
||||
BaseDiT.__init__(self, config=config, hf_config=hf_config)
|
||||
self.quant_config = config.quant_config
|
||||
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.in_channels = config.in_channels
|
||||
self.out_channels = config.out_channels
|
||||
self.num_channels_latents = config.num_channels_latents
|
||||
self.patch_size = config.patch_size
|
||||
self.text_len = config.text_len
|
||||
|
||||
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
|
||||
|
||||
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
|
||||
embed_dim=inner_dim,
|
||||
patch_size=config.patch_size,
|
||||
flatten=False)
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=config.freq_dim,
|
||||
text_embed_dim=config.text_dim,
|
||||
image_embed_dim=config.image_dim,
|
||||
)
|
||||
self.blocks = nn.ModuleList([
|
||||
DreamXWorldTransformerBlock(
|
||||
inner_dim,
|
||||
config.ffn_dim,
|
||||
config.num_attention_heads,
|
||||
config.qk_norm,
|
||||
config.cross_attn_norm,
|
||||
config.eps,
|
||||
config.added_kv_proj_dim,
|
||||
self._supported_attention_backends,
|
||||
quant_config=config.quant_config,
|
||||
prefix=f"{config.prefix}.blocks.{i}",
|
||||
add_control_adapter=config.add_control_adapter,
|
||||
cam_method=config.cam_method,
|
||||
attn_compress=config.attn_compress,
|
||||
cam_self_attn_layers=config.cam_self_attn_layers,
|
||||
layer_idx=i)
|
||||
for i in range(config.num_layers)
|
||||
])
|
||||
self.norm_out = LayerNormScaleShift(inner_dim,
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
self.gradient_checkpointing = False
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
|
||||
| None = None,
|
||||
guidance=None,
|
||||
y_camera: dict[str, torch.Tensor] | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
orig_dtype = hidden_states.dtype
|
||||
if encoder_hidden_states is not None and not isinstance(
|
||||
encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
list) and len(encoder_hidden_states_image) > 0:
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
else:
|
||||
encoder_hidden_states_image = None
|
||||
|
||||
batch_size, _, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
d = self.hidden_size // self.num_attention_heads
|
||||
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||
(post_patch_num_frames, post_patch_height, post_patch_width),
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
|
||||
rope_theta=10000)
|
||||
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
|
||||
freqs_sin.to(hidden_states.device).float())
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
hidden_states, original_seq_len = sequence_model_parallel_shard(
|
||||
hidden_states, dim=1)
|
||||
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
timestep = timestep.flatten()
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep,
|
||||
encoder_hidden_states,
|
||||
encoder_hidden_states_image,
|
||||
timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
timestep_proj = timestep_proj.unflatten(2, (6, -1))
|
||||
else:
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
else:
|
||||
encoder_hidden_states = encoder_hidden_states_image
|
||||
|
||||
if current_platform.is_mps() or current_platform.is_npu():
|
||||
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, original_seq_len, y_camera)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, original_seq_len,
|
||||
y_camera=y_camera)
|
||||
|
||||
if temb.dim() == 3:
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(0) +
|
||||
temb.unsqueeze(2)).chunk(2, dim=2)
|
||||
shift = shift.squeeze(2)
|
||||
scale = scale.squeeze(2)
|
||||
else:
|
||||
shift, scale = (self.scale_shift_table +
|
||||
temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = sequence_model_parallel_all_gather_with_unpad(
|
||||
hidden_states, original_seq_len, dim=1)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
EntryClass = DreamXWorldTransformer3DModel
|
||||
@@ -1,920 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World autoregressive causal DiT.
|
||||
|
||||
Adapted from DreamX-World's Apache-2.0
|
||||
``wan/modules/causal_camera_model_2_2_prope_infinity.py``. The implementation is
|
||||
kept native to FastVideo: no production import from DreamX, Diffusers, or
|
||||
Transformers is required.
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.dreamx_world import (_dreamx_apply_tiled_projmat,
|
||||
_dreamx_prope_qkv)
|
||||
|
||||
|
||||
def attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
|
||||
# Deliberately raw SDPA rather than fastvideo.attention.LocalAttention:
|
||||
# (1) LocalAttention dispatches through the attention-backend registry, so
|
||||
# FLASH_ATTN could be selected and its kernel is not bit-identical to
|
||||
# torch SDPA — the AR KV-cache rollout must stay numerically frozen;
|
||||
# (2) LocalAttention requires an active ForwardContext, which direct
|
||||
# transformer invocations (parity tests) do not set;
|
||||
# (3) the sibling causal model keeps raw SDPA in the same KV-cache window
|
||||
# path (matrixgame2/causal_model.py).
|
||||
# Sequence-parallel gap: this model never shards the sequence; run with
|
||||
# sp_size=1 (see fastvideo/layers/AGENTS.md on documenting raw SDPA).
|
||||
q_bhld = q.transpose(1, 2)
|
||||
k_bhld = k.transpose(1, 2)
|
||||
v_bhld = v.transpose(1, 2)
|
||||
out = F.scaled_dot_product_attention(q_bhld, k_bhld, v_bhld, dropout_p=0.0)
|
||||
return out.transpose(1, 2)
|
||||
|
||||
|
||||
def prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
|
||||
viewmats: torch.Tensor, Ks: torch.Tensor):
|
||||
q, k, v, output_projection = _dreamx_prope_qkv(q, k, v, viewmats, Ks)
|
||||
|
||||
def apply_fn_o(x: torch.Tensor) -> torch.Tensor:
|
||||
return _dreamx_apply_tiled_projmat(x, output_projection)
|
||||
|
||||
return q, k, v, apply_fn_o
|
||||
|
||||
|
||||
def sinusoidal_embedding_1d(dim, position):
|
||||
assert dim % 2 == 0
|
||||
half = dim // 2
|
||||
position = position.type(torch.float64)
|
||||
sinusoid = torch.outer(
|
||||
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
|
||||
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||
|
||||
|
||||
def rope_params(max_seq_len, dim, theta=10000):
|
||||
assert dim % 2 == 0
|
||||
freqs = torch.outer(
|
||||
torch.arange(max_seq_len),
|
||||
1.0 / torch.pow(theta,
|
||||
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
|
||||
return torch.polar(torch.ones_like(freqs), freqs)
|
||||
|
||||
|
||||
class WanRMSNorm(nn.Module):
|
||||
"""Kept private instead of fastvideo.layers.layernorm.RMSNorm.
|
||||
|
||||
The official DreamX-World ``model_2_2.py`` computes the RMS statistics in
|
||||
the *input* dtype — the upstream code has the fp32 upcast explicitly
|
||||
commented out (``# return self._norm(x.float())...``). FastVideo's RMSNorm
|
||||
always normalizes in fp32, which is not bit-identical under bf16, so the
|
||||
verbatim implementation stays.
|
||||
"""
|
||||
|
||||
def __init__(self, dim, eps=1e-5):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.eps = eps
|
||||
self.weight = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return self._norm(x).type_as(x) * self.weight
|
||||
|
||||
def _norm(self, x):
|
||||
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||
|
||||
|
||||
class WanLayerNorm(nn.LayerNorm):
|
||||
"""Kept private instead of fastvideo.layers.layernorm.FP32LayerNorm.
|
||||
|
||||
The official DreamX-World ``model_2_2.py`` normalizes in the *input* dtype
|
||||
(no ``x.float()`` upcast, unlike Wan2.1). FP32LayerNorm casts input and
|
||||
affine params to fp32, which is not bit-identical under bf16, so the
|
||||
verbatim implementation stays.
|
||||
"""
|
||||
|
||||
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||
|
||||
def forward(self, x):
|
||||
return super().forward(x).type_as(x)
|
||||
|
||||
|
||||
class WanCrossAttention(nn.Module):
|
||||
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6):
|
||||
super().__init__()
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.q = ReplicatedLinear(dim, dim)
|
||||
self.k = ReplicatedLinear(dim, dim)
|
||||
self.v = ReplicatedLinear(dim, dim)
|
||||
self.o = ReplicatedLinear(dim, dim)
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def _kv(self, context, b, n, d):
|
||||
k, _ = self.k(context)
|
||||
k = self.norm_k(k).view(b, -1, n, d)
|
||||
v, _ = self.v(context)
|
||||
v = v.view(b, -1, n, d)
|
||||
return k, v
|
||||
|
||||
def forward(self, x, context, context_lens, crossattn_cache=None):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
q, _ = self.q(x)
|
||||
q = self.norm_q(q).view(b, -1, n, d)
|
||||
|
||||
if crossattn_cache is not None:
|
||||
if not crossattn_cache["is_init"]:
|
||||
crossattn_cache["is_init"] = True
|
||||
k, v = self._kv(context, b, n, d)
|
||||
crossattn_cache["k"] = k
|
||||
crossattn_cache["v"] = v
|
||||
else:
|
||||
k = crossattn_cache["k"]
|
||||
v = crossattn_cache["v"]
|
||||
else:
|
||||
k, v = self._kv(context, b, n, d)
|
||||
|
||||
x = attention(q, k, v)
|
||||
x = x.flatten(2)
|
||||
out, _ = self.o(x)
|
||||
return out
|
||||
|
||||
|
||||
def block_relativistic_rope(x, grid_sizes, freqs, start_frame=0, relative_frame_indices=None):
|
||||
"""
|
||||
Apply Block-Relativistic RoPE to input tensor.
|
||||
Adapted from Infinity-RoPE (https://arxiv.org/abs/2511.20649).
|
||||
|
||||
Args:
|
||||
x: Input tensor [B, L, num_heads, head_dim]
|
||||
grid_sizes: Tensor [B, 3] containing (F, H, W)
|
||||
freqs: RoPE frequencies
|
||||
start_frame: Starting frame index for sequential RoPE
|
||||
relative_frame_indices: Optional tensor [F] specifying explicit frame indices
|
||||
for Block-Relativistic RoPE. Overrides start_frame if provided.
|
||||
"""
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
|
||||
seq_len, n, -1, 2))
|
||||
|
||||
if relative_frame_indices is not None:
|
||||
frame_indices = relative_frame_indices.long()
|
||||
freqs_temporal = freqs[0][frame_indices].view(f, 1, 1, -1).expand(f, h, w, -1)
|
||||
else:
|
||||
freqs_temporal = freqs[0][start_frame:start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1)
|
||||
|
||||
freqs_i = torch.cat([
|
||||
freqs_temporal,
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||
], dim=-1).reshape(seq_len, 1, -1)
|
||||
|
||||
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]])
|
||||
output.append(x_i)
|
||||
|
||||
return torch.stack(output).type_as(x)
|
||||
|
||||
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
"""Self-attention with KV cache and Block-Relativistic RoPE for causal inference."""
|
||||
|
||||
def __init__(self, dim, num_heads, local_attn_size=6, sink_size=1,
|
||||
qk_norm=True, eps=1e-6):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
|
||||
|
||||
self.q = ReplicatedLinear(dim, dim)
|
||||
self.k = ReplicatedLinear(dim, dim)
|
||||
self.v = ReplicatedLinear(dim, dim)
|
||||
self.o = ReplicatedLinear(dim, dim)
|
||||
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs, kv_cache,
|
||||
current_start=0, cache_start=None, sink_recache_after_switch=False):
|
||||
"""
|
||||
Args:
|
||||
x: Shape [B, L, C]
|
||||
seq_lens: Shape [B]
|
||||
grid_sizes: Shape [B, 3] containing (F, H, W)
|
||||
freqs: RoPE frequencies [1024, head_dim / 2]
|
||||
kv_cache: Dict with 'k', 'v', 'global_end_index', 'local_end_index'
|
||||
current_start: Current position in the global token sequence
|
||||
"""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
q, _ = self.q(x)
|
||||
q = self.norm_q(q).view(b, s, n, d)
|
||||
k, _ = self.k(x)
|
||||
k = self.norm_k(k).view(b, s, n, d)
|
||||
v, _ = self.v(x)
|
||||
v = v.view(b, s, n, d)
|
||||
|
||||
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
|
||||
num_new_frames = grid_sizes[0][0].item()
|
||||
current_end = current_start + q.shape[1]
|
||||
sink_tokens = self.sink_size * frame_seqlen
|
||||
kv_cache_size = kv_cache["k"].shape[1]
|
||||
num_new_tokens = q.shape[1]
|
||||
|
||||
cache_update_info = None
|
||||
is_recompute = current_end <= kv_cache["global_end_index"].item() and current_start > 0
|
||||
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
# === ROLLING MODE: cache full, evict oldest non-sink tokens ===
|
||||
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
|
||||
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
temp_k = kv_cache["k"].detach().clone()
|
||||
temp_v = kv_cache["v"].detach().clone()
|
||||
temp_k[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
temp_k[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
temp_v[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
temp_v[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
|
||||
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
|
||||
|
||||
# Block-Relativistic RoPE: query uses window-relative indices
|
||||
query_relative_indices = torch.arange(
|
||||
self.local_attn_size - num_new_frames, self.local_attn_size, device=q.device)
|
||||
roped_query = block_relativistic_rope(
|
||||
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
|
||||
|
||||
# Block-Relativistic RoPE: cached K uses position-in-window indices
|
||||
num_cache_frames = local_end_index // frame_seqlen
|
||||
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
|
||||
cache_grid_sizes = grid_sizes.clone()
|
||||
cache_grid_sizes[0, 0] = num_cache_frames
|
||||
roped_temp_k = block_relativistic_rope(
|
||||
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
|
||||
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
|
||||
|
||||
cache_update_info = {
|
||||
"action": "roll_and_insert",
|
||||
"sink_tokens": sink_tokens,
|
||||
"num_rolled_tokens": num_rolled_tokens,
|
||||
"num_evicted_tokens": num_evicted_tokens,
|
||||
"local_start_index": local_start_index,
|
||||
"local_end_index": local_end_index,
|
||||
"write_start_index": write_start_index,
|
||||
"write_end_index": local_end_index,
|
||||
"new_k": k[:, roped_offset:roped_offset + write_len],
|
||||
"new_v": v[:, roped_offset:roped_offset + write_len],
|
||||
"current_end": current_end,
|
||||
"is_recompute": is_recompute
|
||||
}
|
||||
else:
|
||||
# === DIRECT INSERT MODE: cache not yet full ===
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
temp_k = kv_cache["k"].detach().clone()
|
||||
temp_v = kv_cache["v"].detach().clone()
|
||||
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
if sink_recache_after_switch:
|
||||
write_start_index = local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
|
||||
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
|
||||
|
||||
# RoPE with relative indices (growing sequentially before cache fills)
|
||||
current_frame_in_window = local_start_index // frame_seqlen
|
||||
query_relative_indices = torch.arange(
|
||||
current_frame_in_window, current_frame_in_window + num_new_frames, device=q.device)
|
||||
roped_query = block_relativistic_rope(
|
||||
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
|
||||
|
||||
num_cache_frames = local_end_index // frame_seqlen
|
||||
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
|
||||
cache_grid_sizes = grid_sizes.clone()
|
||||
cache_grid_sizes[0, 0] = num_cache_frames
|
||||
roped_temp_k = block_relativistic_rope(
|
||||
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
|
||||
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
|
||||
|
||||
cache_update_info = {
|
||||
"action": "direct_insert",
|
||||
"local_start_index": local_start_index,
|
||||
"local_end_index": local_end_index,
|
||||
"write_start_index": write_start_index,
|
||||
"write_end_index": local_end_index,
|
||||
"new_k": k[:, roped_offset:roped_offset + write_len],
|
||||
"new_v": v[:, roped_offset:roped_offset + write_len],
|
||||
"current_end": current_end,
|
||||
"is_recompute": is_recompute
|
||||
}
|
||||
|
||||
# Attention: sink tokens + local window
|
||||
if sink_tokens > 0:
|
||||
local_budget = self.max_attention_size - sink_tokens
|
||||
k_sink = roped_temp_k[:, :sink_tokens]
|
||||
v_sink = temp_v[:, :sink_tokens]
|
||||
if local_budget > 0:
|
||||
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
|
||||
k_local = roped_temp_k[:, local_start_for_window:local_end_index]
|
||||
v_local = temp_v[:, local_start_for_window:local_end_index]
|
||||
k_cat = torch.cat([k_sink, k_local], dim=1)
|
||||
v_cat = torch.cat([v_sink, v_local], dim=1)
|
||||
else:
|
||||
k_cat = k_sink
|
||||
v_cat = v_sink
|
||||
x = attention(roped_query, k_cat, v_cat)
|
||||
else:
|
||||
window_start = max(0, local_end_index - self.max_attention_size)
|
||||
x = attention(
|
||||
roped_query,
|
||||
roped_temp_k[:, window_start:local_end_index],
|
||||
temp_v[:, window_start:local_end_index])
|
||||
|
||||
x = x.flatten(2)
|
||||
x, _ = self.o(x)
|
||||
return x, (current_end, local_end_index, cache_update_info)
|
||||
|
||||
|
||||
class CausalPropeSelfAttention(nn.Module):
|
||||
"""PRoPE self-attention with optional KV cache for camera-controlled inference."""
|
||||
|
||||
def __init__(self, dim, attn_dim, num_heads, window_size=(-1, -1),
|
||||
local_attn_size=-1, sink_size=0, qk_norm=True, eps=1e-6):
|
||||
assert dim % num_heads == 0
|
||||
assert attn_dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.attn_dim = attn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = attn_dim // num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.window_size = window_size
|
||||
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
|
||||
|
||||
self.q_proj = ReplicatedLinear(dim, attn_dim)
|
||||
self.k_proj = ReplicatedLinear(dim, attn_dim)
|
||||
self.v_proj = ReplicatedLinear(dim, attn_dim)
|
||||
self.out_proj = ReplicatedLinear(attn_dim, dim)
|
||||
|
||||
self.norm_q = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
nn.init.zeros_(self.out_proj.weight)
|
||||
nn.init.zeros_(self.out_proj.bias)
|
||||
|
||||
def forward(self, x, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
|
||||
kv_cache=None, current_start=0, cache_start=None,
|
||||
sink_recache_after_switch=False, cache_update_policy="commit_detached"):
|
||||
"""
|
||||
Args:
|
||||
x: Shape [B, L, C]
|
||||
cam_viewmats: Camera view matrices
|
||||
cam_K: Camera intrinsics
|
||||
kv_cache: Optional KV cache dict. When None, runs full attention over current chunk.
|
||||
"""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
q, _ = self.q_proj(x)
|
||||
q = self.norm_q(q).view(b, s, n, d)
|
||||
k, _ = self.k_proj(x)
|
||||
k = self.norm_k(k).view(b, s, n, d)
|
||||
v, _ = self.v_proj(x)
|
||||
v = v.view(b, s, n, d)
|
||||
|
||||
# Apply PRoPE (Positional Rotary Position Embedding from camera parameters)
|
||||
q_t, k_t, v_t, apply_fn_o = prope_qkv(
|
||||
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2),
|
||||
viewmats=cam_viewmats, Ks=cam_K)
|
||||
proped_q = q_t.transpose(1, 2)
|
||||
proped_k = k_t.transpose(1, 2)
|
||||
proped_v = v_t.transpose(1, 2)
|
||||
|
||||
if kv_cache is None:
|
||||
# No cache: full attention over current chunk
|
||||
x_out = attention(proped_q, proped_k, proped_v)
|
||||
else:
|
||||
# KV cache mode with rolling cache support
|
||||
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
|
||||
num_new_tokens = s
|
||||
current_end = current_start + num_new_tokens
|
||||
sink_tokens = self.sink_size * frame_seqlen
|
||||
kv_cache_size = kv_cache["k"].shape[1]
|
||||
is_recompute = (current_end <= kv_cache["global_end_index"].item()) and (current_start > 0)
|
||||
|
||||
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
|
||||
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
|
||||
# === ROLLING MODE ===
|
||||
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
|
||||
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - \
|
||||
kv_cache["global_end_index"].item() - num_evicted_tokens
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
if cache_update_policy != "none":
|
||||
with torch.no_grad():
|
||||
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
with torch.no_grad():
|
||||
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
|
||||
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
|
||||
else:
|
||||
# === DIRECT INSERT MODE ===
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
if cache_update_policy != "none":
|
||||
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
|
||||
if sink_recache_after_switch:
|
||||
write_start_index = local_start_index
|
||||
roped_offset = max(0, write_start_index - local_start_index)
|
||||
write_len = max(0, local_end_index - write_start_index)
|
||||
if write_len > 0:
|
||||
with torch.no_grad():
|
||||
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
|
||||
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
|
||||
|
||||
# Attention: sink tokens + local window
|
||||
if sink_tokens > 0:
|
||||
local_budget = self.max_attention_size - sink_tokens
|
||||
k_sink = kv_cache["k"][:, :sink_tokens].detach()
|
||||
v_sink = kv_cache["v"][:, :sink_tokens].detach()
|
||||
if local_budget > 0:
|
||||
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
|
||||
k_local = kv_cache["k"][:, local_start_for_window:local_end_index].detach()
|
||||
v_local = kv_cache["v"][:, local_start_for_window:local_end_index].detach()
|
||||
k_cat = torch.cat([k_sink, k_local], dim=1)
|
||||
v_cat = torch.cat([v_sink, v_local], dim=1)
|
||||
else:
|
||||
k_cat = k_sink
|
||||
v_cat = v_sink
|
||||
x_out = attention(proped_q, k_cat, v_cat)
|
||||
else:
|
||||
window_start = max(0, local_end_index - self.max_attention_size)
|
||||
x_out = attention(
|
||||
proped_q,
|
||||
kv_cache["k"][:, window_start:local_end_index].detach(),
|
||||
kv_cache["v"][:, window_start:local_end_index].detach())
|
||||
|
||||
if not is_recompute and cache_update_policy != "none":
|
||||
kv_cache["global_end_index"].fill_(current_end)
|
||||
kv_cache["local_end_index"].fill_(local_end_index)
|
||||
|
||||
# Apply inverse PRoPE
|
||||
x = apply_fn_o(x_out.transpose(1, 2)).transpose(1, 2)
|
||||
x = x.flatten(2)
|
||||
x, _ = self.out_proj(x)
|
||||
return x
|
||||
|
||||
|
||||
class CausalWanAttentionBlock(nn.Module):
|
||||
|
||||
def __init__(self, dim, ffn_dim, num_heads, local_attn_size=-1, sink_size=0,
|
||||
qk_norm=True, cross_attn_norm=False, eps=1e-6, **kwargs):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
self.add_control_adapter = kwargs.get('add_control_adapter', False)
|
||||
self.cam_method = kwargs.get('cam_method')
|
||||
self.attn_compress = kwargs.get('attn_compress', 1)
|
||||
self.layer_idx = kwargs.get('layer_idx')
|
||||
cam_self_attn_layers = kwargs.get('cam_self_attn_layers')
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(dim, eps)
|
||||
self.self_attn = CausalWanSelfAttention(
|
||||
dim, num_heads, local_attn_size, sink_size, qk_norm, eps)
|
||||
self.norm3 = WanLayerNorm(
|
||||
dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.cross_attn = WanCrossAttention(dim, num_heads, (-1, -1), qk_norm, eps)
|
||||
self.norm2 = WanLayerNorm(dim, eps)
|
||||
# nn.Linear (not ReplicatedLinear) on purpose: the official checkpoint
|
||||
# stores these as positional Sequential keys (ffn.0 / ffn.2) that the
|
||||
# copy-only converter and the strict-load tests require verbatim, and
|
||||
# ReplicatedLinear's (out, bias) tuple return cannot compose inside
|
||||
# nn.Sequential without renaming the state-dict surface.
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn_dim, dim))
|
||||
|
||||
# PRoPE self-attention branch for camera control
|
||||
add_cam_attn = self.add_control_adapter and self.cam_method == 'prope'
|
||||
if add_cam_attn and cam_self_attn_layers is not None:
|
||||
add_cam_attn = self.layer_idx in cam_self_attn_layers
|
||||
if add_cam_attn:
|
||||
self.cam_self_attn = CausalPropeSelfAttention(
|
||||
dim, dim // self.attn_compress, num_heads,
|
||||
local_attn_size=local_attn_size, sink_size=sink_size,
|
||||
qk_norm=qk_norm, eps=eps)
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(self, x, e, seq_lens, grid_sizes, freqs, context, context_lens,
|
||||
kv_cache, crossattn_cache=None, current_start=0, cache_start=None,
|
||||
cam_viewmats=None, cam_K=None, sink_recache_after_switch=False,
|
||||
cache_update_policy="commit_detached"):
|
||||
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
|
||||
e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)
|
||||
|
||||
# self-attention
|
||||
attn_input = (self.norm1(x).unflatten(
|
||||
dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2)
|
||||
y, cache_update_info = self.self_attn(
|
||||
attn_input, seq_lens, grid_sizes, freqs, kv_cache,
|
||||
current_start, cache_start, sink_recache_after_switch)
|
||||
|
||||
# PRoPE camera attention (parallel branch)
|
||||
if hasattr(self, 'cam_self_attn') and cam_viewmats is not None and cam_K is not None:
|
||||
prope_kv_cache = None
|
||||
if kv_cache is not None and "prope_k" in kv_cache:
|
||||
prope_kv_cache = {
|
||||
"k": kv_cache["prope_k"],
|
||||
"v": kv_cache["prope_v"],
|
||||
"global_end_index": kv_cache["prope_global_end_index"],
|
||||
"local_end_index": kv_cache["prope_local_end_index"],
|
||||
}
|
||||
y = y + self.cam_self_attn(
|
||||
attn_input, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
|
||||
kv_cache=prope_kv_cache, current_start=current_start,
|
||||
cache_start=cache_start, cache_update_policy=cache_update_policy)
|
||||
|
||||
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(1, 2)
|
||||
|
||||
# cross-attention & FFN
|
||||
x = x + self.cross_attn(self.norm3(x), context, context_lens,
|
||||
crossattn_cache=crossattn_cache)
|
||||
y = self.ffn(
|
||||
(self.norm2(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
|
||||
* (1 + e[4]) + e[3]).flatten(1, 2))
|
||||
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[5]).flatten(1, 2)
|
||||
|
||||
return x, cache_update_info
|
||||
|
||||
|
||||
class CausalHead(nn.Module):
|
||||
|
||||
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.out_dim = out_dim
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
|
||||
out_dim = math.prod(patch_size) * out_dim
|
||||
self.norm = WanLayerNorm(dim, eps)
|
||||
self.head = ReplicatedLinear(dim, out_dim)
|
||||
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
|
||||
|
||||
def forward(self, x, e):
|
||||
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
|
||||
e = (self.modulation.unsqueeze(1) + e).chunk(2, dim=2)
|
||||
x, _ = self.head(
|
||||
self.norm(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
|
||||
* (1 + e[1]) + e[0])
|
||||
return x
|
||||
|
||||
|
||||
class DreamXWorldARTransformer3DModel(BaseDiT):
|
||||
"""DreamX-World-5B autoregressive causal transformer."""
|
||||
|
||||
_fsdp_shard_conditions = DreamXWorldARConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = DreamXWorldARConfig()._compile_conditions
|
||||
_supported_attention_backends = DreamXWorldARConfig()._supported_attention_backends
|
||||
param_names_mapping = DreamXWorldARConfig().param_names_mapping
|
||||
reverse_param_names_mapping = DreamXWorldARConfig().reverse_param_names_mapping
|
||||
lora_param_names_mapping = DreamXWorldARConfig().lora_param_names_mapping
|
||||
_no_split_modules = ["CausalWanAttentionBlock"]
|
||||
|
||||
def __init__(self, config: DreamXWorldARConfig, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
model_type = config.model_type
|
||||
patch_size = config.patch_size
|
||||
text_len = config.text_len
|
||||
in_dim = config.in_channels
|
||||
dim = config.hidden_size
|
||||
ffn_dim = config.ffn_dim
|
||||
freq_dim = config.freq_dim
|
||||
text_dim = config.text_dim
|
||||
out_dim = config.out_channels
|
||||
num_heads = config.num_attention_heads
|
||||
num_layers = config.num_layers
|
||||
local_attn_size = config.local_attn_size
|
||||
sink_size = config.sink_size
|
||||
qk_norm = bool(config.qk_norm)
|
||||
cross_attn_norm = config.cross_attn_norm
|
||||
eps = config.eps
|
||||
add_control_adapter = config.add_control_adapter
|
||||
cam_method = config.cam_method
|
||||
attn_compress = config.attn_compress
|
||||
cam_self_attn_layers = config.cam_self_attn_layers
|
||||
|
||||
assert model_type in ['t2v', 'i2v', 'ti2v']
|
||||
self.model_type = model_type
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
self.in_dim = in_dim
|
||||
self.dim = dim
|
||||
self.ffn_dim = ffn_dim
|
||||
self.freq_dim = freq_dim
|
||||
self.text_dim = text_dim
|
||||
self.out_dim = out_dim
|
||||
self.num_heads = num_heads
|
||||
self.num_layers = num_layers
|
||||
self.local_attn_size = local_attn_size
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
|
||||
# embeddings — nn.Linear inside nn.Sequential on purpose: the official
|
||||
# checkpoint keys are positional (text_embedding.0/.2, time_embedding.0/.2,
|
||||
# time_projection.1) and must load verbatim (see ffn comment above).
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
self.text_embedding = nn.Sequential(
|
||||
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(dim, dim))
|
||||
self.time_embedding = nn.Sequential(
|
||||
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(
|
||||
nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
|
||||
# transformer blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
CausalWanAttentionBlock(
|
||||
dim, ffn_dim, num_heads, local_attn_size, sink_size,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
add_control_adapter=add_control_adapter,
|
||||
cam_method=cam_method,
|
||||
attn_compress=attn_compress,
|
||||
layer_idx=layer_idx,
|
||||
cam_self_attn_layers=cam_self_attn_layers)
|
||||
for layer_idx in range(num_layers)
|
||||
])
|
||||
for layer_idx, block in enumerate(self.blocks):
|
||||
block.self_attn.layer_idx = layer_idx
|
||||
block.self_attn.num_layers = self.num_layers
|
||||
|
||||
# head
|
||||
self.head = CausalHead(dim, out_dim, patch_size, eps)
|
||||
|
||||
# RoPE frequencies
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
d = dim // num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6))
|
||||
], dim=1)
|
||||
|
||||
self.num_attention_heads = num_heads
|
||||
self.attention_head_dim = dim // num_heads
|
||||
self.hidden_size = dim
|
||||
self.in_channels = in_dim
|
||||
self.out_channels = out_dim
|
||||
self.num_channels_latents = out_dim
|
||||
self.init_weights()
|
||||
self.num_frame_per_block = config.arch_config.num_frames_per_block
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self, x=None, t=None, context=None, seq_len=None, y=None, y_camera=None,
|
||||
kv_cache=None, crossattn_cache=None, current_start=0,
|
||||
cache_start=0, cache_update_policy="commit_detached",
|
||||
hidden_states=None, encoder_hidden_states=None, timestep=None, **kwargs):
|
||||
"""
|
||||
Causal inference with KV caching.
|
||||
See Algorithm 2 of CausVid (https://arxiv.org/abs/2412.07772).
|
||||
|
||||
Args:
|
||||
x: List of input video tensors [C_in, F, H, W]
|
||||
t: Timestep tensor [B, L]
|
||||
context: List of text embeddings [L, C]
|
||||
seq_len: Maximum sequence length for positional encoding
|
||||
y: Optional conditional video inputs (I2V mode)
|
||||
y_camera: Camera parameters dict {'viewmats': ..., 'K': ...}
|
||||
kv_cache: List of KV cache dicts per transformer block
|
||||
crossattn_cache: List of cross-attention cache dicts
|
||||
current_start: Current position in global token sequence
|
||||
cache_start: Cache start position
|
||||
cache_update_policy: Cache update strategy ('commit_detached' or 'none')
|
||||
|
||||
Returns:
|
||||
Stacked output tensors [B, C_out, F, H/8, W/8]
|
||||
"""
|
||||
if x is None and hidden_states is not None:
|
||||
x = [sample for sample in hidden_states]
|
||||
if t is None and timestep is not None:
|
||||
t = timestep
|
||||
if context is None and encoder_hidden_states is not None:
|
||||
if isinstance(encoder_hidden_states, torch.Tensor):
|
||||
context = [sample for sample in encoder_hidden_states]
|
||||
else:
|
||||
context = encoder_hidden_states
|
||||
if seq_len is None:
|
||||
if torch.is_tensor(t):
|
||||
seq_len = int(t.shape[1]) if t.dim() > 1 else int(t.numel())
|
||||
elif x is not None:
|
||||
sample = x[0]
|
||||
seq_len = (sample.shape[1] // self.patch_size[0]) * (sample.shape[2] // self.patch_size[1]) * (sample.shape[3] // self.patch_size[2])
|
||||
if x is None or t is None or context is None or seq_len is None:
|
||||
raise ValueError("DreamXWorldARTransformer3DModel requires x/t/context/seq_len or FastVideo aliases")
|
||||
|
||||
device = self.patch_embedding.weight.device
|
||||
if self.freqs.is_meta or self.freqs.device != device:
|
||||
d = self.dim // self.num_heads
|
||||
self.freqs = torch.cat([
|
||||
rope_params(1024, d - 4 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
rope_params(1024, 2 * (d // 6)),
|
||||
], dim=1).to(device)
|
||||
|
||||
if y is not None:
|
||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y, strict=True)]
|
||||
|
||||
# patch embedding
|
||||
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
|
||||
assert seq_lens.max() <= seq_len
|
||||
x = torch.cat(x)
|
||||
|
||||
# time embedding
|
||||
e = self.time_embedding(
|
||||
sinusoidal_embedding_1d(self.freq_dim, t.flatten()).type_as(x))
|
||||
e0 = self.time_projection(e).unflatten(
|
||||
1, (6, self.dim)).unflatten(dim=0, sizes=t.shape)
|
||||
|
||||
# text embedding
|
||||
context_lens = None
|
||||
context = self.text_embedding(
|
||||
torch.stack([
|
||||
torch.cat(
|
||||
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||
for u in context
|
||||
]))
|
||||
|
||||
# camera parameters
|
||||
if y_camera is not None and isinstance(y_camera, dict):
|
||||
cam_viewmats = y_camera['viewmats']
|
||||
cam_K = y_camera['K']
|
||||
else:
|
||||
cam_viewmats = None
|
||||
cam_K = None
|
||||
|
||||
block_kwargs = dict(
|
||||
e=e0, seq_lens=seq_lens, grid_sizes=grid_sizes, freqs=self.freqs,
|
||||
context=context, context_lens=context_lens,
|
||||
cam_viewmats=cam_viewmats, cam_K=cam_K,
|
||||
cache_update_policy=cache_update_policy,
|
||||
)
|
||||
|
||||
cache_update_infos = []
|
||||
for block_index, block in enumerate(self.blocks):
|
||||
block_kwargs.update({
|
||||
"kv_cache": kv_cache[block_index] if kv_cache is not None else None,
|
||||
"crossattn_cache": crossattn_cache[block_index] if crossattn_cache is not None else None,
|
||||
"current_start": current_start,
|
||||
"cache_start": cache_start,
|
||||
})
|
||||
x, block_cache_update_info = block(x, **block_kwargs)
|
||||
if kv_cache is not None:
|
||||
cache_update_infos.append((block_index, block_cache_update_info))
|
||||
|
||||
# Apply deferred cache updates
|
||||
if kv_cache is not None and cache_update_infos and cache_update_policy != "none":
|
||||
self._apply_cache_updates(kv_cache, cache_update_infos)
|
||||
|
||||
# head & unpatchify
|
||||
x = self.head(x, e.unflatten(dim=0, sizes=t.shape).unsqueeze(2))
|
||||
x = self.unpatchify(x, grid_sizes)
|
||||
return torch.stack(x)
|
||||
|
||||
def _apply_cache_updates(self, kv_cache, cache_update_infos):
|
||||
"""Apply deferred cache updates collected from all transformer blocks.
|
||||
|
||||
For Block-Relativistic RoPE, this stores un-roped K values in the cache.
|
||||
RoPE is applied dynamically during attention based on each token's current
|
||||
relative position in the sliding window.
|
||||
"""
|
||||
with torch.no_grad():
|
||||
for block_index, (current_end, local_end_index, update_info) in cache_update_infos:
|
||||
if update_info is not None:
|
||||
cache = kv_cache[block_index]
|
||||
|
||||
if update_info["action"] == "roll_and_insert":
|
||||
sink_tokens = update_info["sink_tokens"]
|
||||
num_rolled_tokens = update_info["num_rolled_tokens"]
|
||||
num_evicted_tokens = update_info["num_evicted_tokens"]
|
||||
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
|
||||
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
|
||||
new_k = update_info["new_k"].detach()
|
||||
new_v = update_info["new_v"].detach()
|
||||
|
||||
cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
|
||||
cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
|
||||
|
||||
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
|
||||
cache["k"][:, write_start_index:write_end_index] = new_k
|
||||
cache["v"][:, write_start_index:write_end_index] = new_v
|
||||
|
||||
elif update_info["action"] == "direct_insert":
|
||||
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
|
||||
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
|
||||
new_k = update_info["new_k"].detach()
|
||||
new_v = update_info["new_v"].detach()
|
||||
|
||||
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
|
||||
cache["k"][:, write_start_index:write_end_index] = new_k
|
||||
cache["v"][:, write_start_index:write_end_index] = new_v
|
||||
|
||||
is_recompute = False if update_info is None else update_info.get("is_recompute", False)
|
||||
if not is_recompute:
|
||||
kv_cache[block_index]["global_end_index"].fill_(current_end)
|
||||
kv_cache[block_index]["local_end_index"].fill_(local_end_index)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
"""Reconstruct video tensors from patch embeddings."""
|
||||
c = self.out_dim
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist(), strict=True):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = torch.einsum('fhwpqrc->cfphqwr', u)
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size, strict=True)])
|
||||
out.append(u)
|
||||
return out
|
||||
|
||||
def init_weights(self):
|
||||
"""Initialize model parameters using Xavier initialization."""
|
||||
for m in self.modules():
|
||||
if isinstance(m, (nn.Linear, ReplicatedLinear)):
|
||||
nn.init.xavier_uniform_(m.weight)
|
||||
if m.bias is not None:
|
||||
nn.init.zeros_(m.bias)
|
||||
|
||||
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
|
||||
for m in self.text_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=.02)
|
||||
for m in self.time_embedding.modules():
|
||||
if isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, std=.02)
|
||||
|
||||
nn.init.zeros_(self.head.head.weight)
|
||||
|
||||
|
||||
EntryClass = DreamXWorldARTransformer3DModel
|
||||
@@ -537,9 +537,9 @@ class CausalMatrixGame2TransformerBlock(nn.Module):
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
|
||||
if self.norm_q is not None:
|
||||
query = self.norm_q(query)
|
||||
query = self.norm_q.forward_native(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(key)
|
||||
key = self.norm_k.forward_native(key)
|
||||
|
||||
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
|
||||
@@ -11,7 +11,7 @@ import tempfile
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable, Set
|
||||
from dataclasses import dataclass, field
|
||||
from functools import cache, lru_cache
|
||||
from functools import lru_cache
|
||||
from typing import NoReturn, TypeVar, cast
|
||||
|
||||
import cloudpickle
|
||||
@@ -32,8 +32,6 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"HYWorldTransformer3DModel":
|
||||
("dits", "hyworld", "HYWorldTransformer3DModel"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
|
||||
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
|
||||
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
|
||||
@@ -50,8 +48,6 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
|
||||
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"MatrixGame2WanModel": ("dits", "matrixgame2", "MatrixGame2WanModel"),
|
||||
"CausalMatrixGame2WanModel": ("dits", "matrixgame2", "CausalMatrixGame2WanModel"),
|
||||
@@ -145,7 +141,7 @@ _LEGACY_FAST_VIDEO_MODELS = {
|
||||
MODELS_PATH = os.path.dirname(__file__)
|
||||
|
||||
|
||||
@cache
|
||||
@lru_cache(maxsize=None)
|
||||
def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
|
||||
discovered_models: dict[str, tuple[str, str, str]] = {}
|
||||
for root, dirs, files in os.walk(MODELS_PATH):
|
||||
@@ -160,7 +156,7 @@ def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
|
||||
|
||||
filepath = os.path.join(root, filename)
|
||||
try:
|
||||
with open(filepath, encoding="utf-8") as f:
|
||||
with open(filepath, "r", encoding="utf-8") as f:
|
||||
source = f.read()
|
||||
tree = ast.parse(source, filename=filename)
|
||||
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.configs.pipelines.dreamx_world import (
|
||||
DreamXWorld5BARPipelineConfig,
|
||||
DreamXWorld5BCamPipelineConfig,
|
||||
make_dreamx_world_5b_ar_dit_config,
|
||||
make_dreamx_world_5b_cam_dit_config,
|
||||
make_dreamx_world_5b_cam_text_encoder_config,
|
||||
make_dreamx_world_5b_cam_vae_config,
|
||||
)
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import (
|
||||
DREAMX_Y_CAMERA_KEY,
|
||||
DreamXWorldCameraConditioningStage,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DREAMX_Y_CAMERA_KEY",
|
||||
"DreamXWorld5BARPipelineConfig",
|
||||
"DreamXWorld5BCamPipelineConfig",
|
||||
"DreamXWorldCameraConditioningStage",
|
||||
"DreamXWorldARPipeline",
|
||||
"DreamXWorldPipeline",
|
||||
"make_dreamx_world_5b_ar_dit_config",
|
||||
"make_dreamx_world_5b_cam_dit_config",
|
||||
"make_dreamx_world_5b_cam_text_encoder_config",
|
||||
"make_dreamx_world_5b_cam_vae_config",
|
||||
]
|
||||
@@ -1,219 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World autoregressive causal denoising stage."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import DREAMX_Y_CAMERA_KEY
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
|
||||
|
||||
class DreamXWorldARCausalDenoisingStage(DenoisingStage):
|
||||
"""Official DreamX AR-forcing denoising loop with KV cache."""
|
||||
|
||||
_AR_NOISE_SEED_OFFSET = 1_000_003
|
||||
|
||||
def __init__(self, transformer, scheduler, pipeline=None, vae=None) -> None:
|
||||
super().__init__(transformer=transformer, scheduler=scheduler, pipeline=pipeline, vae=vae)
|
||||
self.num_transformer_blocks = len(self.transformer.blocks)
|
||||
self.num_frame_per_block = int(getattr(self.transformer, "num_frame_per_block", 3))
|
||||
self.local_attn_size = int(getattr(self.transformer, "local_attn_size", 12))
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
assert batch.latents is not None, "latents must be prepared before DreamX AR denoising"
|
||||
assert batch.prompt_embeds, "prompt embeds must be prepared before DreamX AR denoising"
|
||||
latents = batch.latents
|
||||
device = latents.device
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = device.type == "cuda" and not fastvideo_args.disable_autocast
|
||||
|
||||
frame_seq_length = (latents.shape[-2] // self.transformer.patch_size[1]) * (latents.shape[-1] //
|
||||
self.transformer.patch_size[2])
|
||||
timesteps = torch.tensor(
|
||||
tuple(getattr(fastvideo_args.pipeline_config, "dmd_denoising_steps", (1000, 750, 500, 250))),
|
||||
dtype=torch.long,
|
||||
).cpu()
|
||||
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step", True):
|
||||
self.scheduler.set_timesteps(1000)
|
||||
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
|
||||
timesteps = scheduler_timesteps[1000 - timesteps]
|
||||
timesteps = timesteps.to(device)
|
||||
|
||||
if latents.shape[2] % self.num_frame_per_block != 0:
|
||||
raise ValueError("DreamX AR latent frames must be divisible by num_frame_per_block")
|
||||
|
||||
y_camera = batch.extra.get(DREAMX_Y_CAMERA_KEY, batch.extra.get("y_camera"))
|
||||
if isinstance(y_camera, dict):
|
||||
y_camera = {
|
||||
k: v.to(device=device, dtype=target_dtype) if torch.is_tensor(v) else v
|
||||
for k, v in y_camera.items()
|
||||
}
|
||||
|
||||
if batch.image_latent is not None and batch.image_latent.shape[1] == latents.shape[1]:
|
||||
latents[:, :, :batch.image_latent.shape[2]] = batch.image_latent.to(device=device, dtype=latents.dtype)
|
||||
|
||||
kv_cache = self._initialize_kv_cache(latents.shape[0], target_dtype, device, frame_seq_length)
|
||||
crossattn_cache = self._initialize_crossattn_cache(latents.shape[0], target_dtype, device)
|
||||
prompt = batch.prompt_embeds[0]
|
||||
if torch.is_tensor(prompt):
|
||||
prompt = prompt.to(device=device, dtype=target_dtype)
|
||||
context = [sample for sample in prompt]
|
||||
else:
|
||||
context = prompt
|
||||
|
||||
num_blocks = latents.shape[2] // self.num_frame_per_block
|
||||
start = 0
|
||||
first_frame_mask = torch.ones_like(latents)
|
||||
first_frame_mask[:, :, 0] = 0
|
||||
base_generator = batch.generator[0] if isinstance(batch.generator, list) else batch.generator
|
||||
noise_generator = self._make_noise_generator(base_generator, device)
|
||||
|
||||
with tqdm(total=num_blocks * len(timesteps), desc="DreamX AR denoising", leave=False) as progress:
|
||||
for _ in range(num_blocks):
|
||||
current_num_frames = self.num_frame_per_block
|
||||
block_latents = latents[:, :, start:start + current_num_frames]
|
||||
noisy_input = block_latents.clone()
|
||||
mask_block = first_frame_mask[:, :, start:start + current_num_frames]
|
||||
camera_block = self._slice_camera(y_camera, start, current_num_frames)
|
||||
|
||||
for idx, current_timestep in enumerate(timesteps):
|
||||
timestep = torch.full(
|
||||
(latents.shape[0], current_num_frames * frame_seq_length),
|
||||
int(current_timestep.item()),
|
||||
device=device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
if start == 0:
|
||||
timestep[:, :frame_seq_length] = 0
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
denoised = self.transformer(
|
||||
hidden_states=block_latents.to(target_dtype),
|
||||
encoder_hidden_states=torch.stack(context).to(target_dtype),
|
||||
timestep=timestep,
|
||||
y_camera=camera_block,
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=start * frame_seq_length,
|
||||
)
|
||||
denoised = denoised.to(latents.dtype)
|
||||
if idx < len(timesteps) - 1:
|
||||
next_timestep = torch.full((latents.shape[0], current_num_frames),
|
||||
int(timesteps[idx + 1].item()),
|
||||
device=device,
|
||||
dtype=torch.long)
|
||||
noise_kwargs = {"device": device, "dtype": denoised.dtype}
|
||||
if noise_generator is not None:
|
||||
noise_kwargs["generator"] = noise_generator
|
||||
noise = torch.randn(denoised.permute(0, 2, 1, 3, 4).shape, **noise_kwargs)
|
||||
block_btchw = self.scheduler.add_noise(
|
||||
denoised.permute(0, 2, 1, 3, 4).flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
next_timestep.flatten(),
|
||||
).unflatten(0, (latents.shape[0], current_num_frames))
|
||||
block_latents = block_btchw.permute(0, 2, 1, 3, 4)
|
||||
block_latents = block_latents * mask_block + noisy_input * (1 - mask_block)
|
||||
else:
|
||||
block_latents = denoised * mask_block + noisy_input * (1 - mask_block)
|
||||
progress.update()
|
||||
|
||||
latents[:, :, start:start + current_num_frames] = block_latents
|
||||
self._update_context_cache(block_latents, context, camera_block, kv_cache, crossattn_cache, start,
|
||||
frame_seq_length, target_dtype, autocast_enabled,
|
||||
float(getattr(fastvideo_args.pipeline_config, "context_noise", 0.1)))
|
||||
start += current_num_frames
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
def _make_noise_generator(self, generator: torch.Generator | None, device: torch.device) -> torch.Generator | None:
|
||||
if generator is None:
|
||||
return None
|
||||
if getattr(generator, "device", None) == device:
|
||||
return generator
|
||||
seed = int(generator.initial_seed()) + self._AR_NOISE_SEED_OFFSET
|
||||
return torch.Generator(device=device).manual_seed(seed)
|
||||
|
||||
@staticmethod
|
||||
def _context_noise_timestep(context_noise: float) -> int:
|
||||
if 0.0 < context_noise <= 1.0:
|
||||
return int(context_noise * 1000)
|
||||
return int(context_noise)
|
||||
|
||||
def _slice_camera(self, y_camera: Any, start: int, num_frames: int):
|
||||
if not isinstance(y_camera, dict):
|
||||
return y_camera
|
||||
return {
|
||||
"viewmats": y_camera["viewmats"][:, start:start + num_frames],
|
||||
"K": y_camera["K"][:, start:start + num_frames],
|
||||
}
|
||||
|
||||
def _initialize_kv_cache(self, batch_size: int, dtype: torch.dtype, device: torch.device,
|
||||
frame_seq_length: int) -> list[dict[str, Any]]:
|
||||
size = self.local_attn_size * frame_seq_length if self.local_attn_size != -1 else 18480
|
||||
heads = self.transformer.num_attention_heads
|
||||
head_dim = self.transformer.attention_head_dim
|
||||
cam_self_attn = next(
|
||||
(getattr(block, "cam_self_attn", None)
|
||||
for block in self.transformer.blocks if getattr(block, "cam_self_attn", None) is not None),
|
||||
None,
|
||||
)
|
||||
|
||||
caches = []
|
||||
for _ in range(self.num_transformer_blocks):
|
||||
cache = {
|
||||
"k": torch.zeros(batch_size, size, heads, head_dim, dtype=dtype, device=device),
|
||||
"v": torch.zeros(batch_size, size, heads, head_dim, dtype=dtype, device=device),
|
||||
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
|
||||
"local_end_index": torch.tensor([0], dtype=torch.long, device=device),
|
||||
}
|
||||
if cam_self_attn is not None:
|
||||
cam_heads = int(cam_self_attn.num_heads)
|
||||
cam_head_dim = int(cam_self_attn.head_dim)
|
||||
cache.update({
|
||||
"prope_k":
|
||||
torch.zeros(batch_size, size, cam_heads, cam_head_dim, dtype=dtype, device=device),
|
||||
"prope_v":
|
||||
torch.zeros(batch_size, size, cam_heads, cam_head_dim, dtype=dtype, device=device),
|
||||
"prope_global_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
"prope_local_end_index":
|
||||
torch.tensor([0], dtype=torch.long, device=device),
|
||||
})
|
||||
caches.append(cache)
|
||||
return caches
|
||||
|
||||
def _initialize_crossattn_cache(self, batch_size: int, dtype: torch.dtype, device: torch.device):
|
||||
heads = self.transformer.num_attention_heads
|
||||
head_dim = self.transformer.attention_head_dim
|
||||
return [{
|
||||
"k": torch.zeros(batch_size, 512, heads, head_dim, dtype=dtype, device=device),
|
||||
"v": torch.zeros(batch_size, 512, heads, head_dim, dtype=dtype, device=device),
|
||||
"is_init": False,
|
||||
} for _ in range(self.num_transformer_blocks)]
|
||||
|
||||
def _update_context_cache(self, block_latents: torch.Tensor, context: Any, camera_block: Any,
|
||||
kv_cache: list[dict[str, Any]], crossattn_cache: list[dict[str, Any]], start: int,
|
||||
frame_seq_length: int, target_dtype: torch.dtype, autocast_enabled: bool,
|
||||
context_noise: float) -> None:
|
||||
timestep = torch.full(
|
||||
(block_latents.shape[0], block_latents.shape[2] * frame_seq_length),
|
||||
self._context_noise_timestep(context_noise),
|
||||
device=block_latents.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
self.transformer(
|
||||
hidden_states=block_latents.to(target_dtype),
|
||||
encoder_hidden_states=torch.stack(context).to(target_dtype),
|
||||
timestep=timestep,
|
||||
y_camera=camera_block,
|
||||
kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=start * frame_seq_length,
|
||||
)
|
||||
@@ -1,228 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from scipy.interpolate import interp1d
|
||||
from scipy.spatial.transform import Rotation, Slerp
|
||||
|
||||
_ACTION_TO_MOTION = {
|
||||
"w": "forward",
|
||||
"a": "left",
|
||||
"d": "right",
|
||||
"s": "backward",
|
||||
"j": "left_rot",
|
||||
"l": "right_rot",
|
||||
"i": "up_rot",
|
||||
"k": "down_rot",
|
||||
}
|
||||
_TRANSLATION_BASE_UNIT = 1.0
|
||||
_ROTATION_BASE_UNIT = 10.0
|
||||
_INTRINSIC_ROW = [0.8, 0.5, 0.5, 0.5]
|
||||
|
||||
|
||||
@dataclass
|
||||
class DreamXCamera:
|
||||
fx: float
|
||||
fy: float
|
||||
cx: float
|
||||
cy: float
|
||||
w2c_mat: np.ndarray
|
||||
|
||||
@property
|
||||
def c2w_mat(self) -> np.ndarray:
|
||||
return np.linalg.inv(self.w2c_mat)
|
||||
|
||||
@classmethod
|
||||
def from_pose_row(cls, row: list[float]) -> DreamXCamera:
|
||||
w2c_mat = np.eye(4, dtype=np.float64)
|
||||
w2c_mat[:3, :] = np.asarray(row[7:], dtype=np.float64).reshape(3, 4)
|
||||
return cls(
|
||||
fx=float(row[1]),
|
||||
fy=float(row[2]),
|
||||
cx=float(row[3]),
|
||||
cy=float(row[4]),
|
||||
w2c_mat=w2c_mat,
|
||||
)
|
||||
|
||||
|
||||
def _translation_step(motion_type: str, current_pose: dict[str, np.ndarray], value: float, duration: int) -> np.ndarray:
|
||||
if motion_type in ("forward", "backward"):
|
||||
yaw = np.radians(current_pose["rotation"][1])
|
||||
pitch = np.radians(current_pose["rotation"][0])
|
||||
forward = np.array([-math.sin(yaw) * math.cos(pitch), math.sin(pitch), math.cos(yaw) * math.cos(pitch)])
|
||||
direction = 1 if motion_type == "forward" else -1
|
||||
return forward * value * direction / duration
|
||||
if motion_type in ("left", "right"):
|
||||
yaw = np.radians(current_pose["rotation"][1])
|
||||
right = np.array([math.cos(yaw), 0.0, math.sin(yaw)])
|
||||
direction = -1 if motion_type == "left" else 1
|
||||
return right * value * direction / duration
|
||||
return np.zeros(3)
|
||||
|
||||
|
||||
def _rotation_step(motion_type: str, value: float, duration: int) -> np.ndarray:
|
||||
if not motion_type.endswith("rot"):
|
||||
return np.zeros(3)
|
||||
axis = motion_type.split("_")[0]
|
||||
rotation = np.zeros(3)
|
||||
if axis == "left":
|
||||
rotation[1] = value
|
||||
elif axis == "right":
|
||||
rotation[1] = -value
|
||||
elif axis == "up":
|
||||
rotation[0] = -value
|
||||
elif axis == "down":
|
||||
rotation[0] = value
|
||||
return rotation / duration
|
||||
|
||||
|
||||
def _euler_to_quaternion(angles: np.ndarray) -> list[float]:
|
||||
pitch, yaw, roll = np.radians(angles)
|
||||
cy = math.cos(yaw * 0.5)
|
||||
sy = math.sin(yaw * 0.5)
|
||||
cp = math.cos(pitch * 0.5)
|
||||
sp = math.sin(pitch * 0.5)
|
||||
cr = math.cos(roll * 0.5)
|
||||
sr = math.sin(roll * 0.5)
|
||||
return [
|
||||
cy * cp * cr + sy * sp * sr,
|
||||
cy * sp * cr + sy * cp * sr,
|
||||
sy * cp * cr - cy * sp * sr,
|
||||
cy * cp * sr - sy * sp * cr,
|
||||
]
|
||||
|
||||
|
||||
def _quaternion_to_rotation_matrix(quaternion: list[float]) -> np.ndarray:
|
||||
qw, qx, qy, qz = quaternion
|
||||
return np.array([
|
||||
[1 - 2 * (qy**2 + qz**2), 2 * (qx * qy - qw * qz), 2 * (qx * qz + qw * qy)],
|
||||
[2 * (qx * qy + qw * qz), 1 - 2 * (qx**2 + qz**2), 2 * (qy * qz - qw * qx)],
|
||||
[2 * (qx * qz - qw * qy), 2 * (qy * qz + qw * qx), 1 - 2 * (qx**2 + qy**2)],
|
||||
])
|
||||
|
||||
|
||||
def _pose_rows_from_actions(action_seq: list[str], action_speed_list: list[float], duration: int) -> list[list[float]]:
|
||||
if len(action_seq) != len(action_speed_list):
|
||||
raise ValueError("action_seq and action_speed_list must have the same length")
|
||||
|
||||
positions: list[np.ndarray] = []
|
||||
rotations: list[np.ndarray] = []
|
||||
current_pose = {
|
||||
"position": np.array([0.0, 0.0, 0.0]),
|
||||
"rotation": np.array([0.0, 0.0, 0.0]),
|
||||
}
|
||||
|
||||
for action_id, speed in zip(action_seq, action_speed_list, strict=True):
|
||||
motion_types = [_ACTION_TO_MOTION[key] for key in list(action_id)]
|
||||
translation_step = np.zeros(3)
|
||||
rotation_step = np.zeros(3)
|
||||
for motion_type in motion_types:
|
||||
translation_step += _translation_step(motion_type, current_pose,
|
||||
float(speed) * _TRANSLATION_BASE_UNIT, duration)
|
||||
rotation_step += _rotation_step(motion_type, float(speed) * _ROTATION_BASE_UNIT, duration)
|
||||
|
||||
segment_positions = []
|
||||
segment_rotations = []
|
||||
for index in range(1, duration + 1):
|
||||
segment_positions.append(current_pose["position"] + translation_step * index)
|
||||
segment_rotations.append(current_pose["rotation"] + rotation_step * index)
|
||||
current_pose["position"] = segment_positions[-1].copy()
|
||||
current_pose["rotation"] = segment_rotations[-1].copy()
|
||||
positions.extend(segment_positions)
|
||||
rotations.extend(segment_rotations)
|
||||
|
||||
rows: list[list[float]] = [[0.0] + _INTRINSIC_ROW + [0.0, 0.0] +
|
||||
[1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0]]
|
||||
for index, (position, rotation) in enumerate(zip(positions, rotations, strict=False)):
|
||||
rotation_matrix = _quaternion_to_rotation_matrix(_euler_to_quaternion(rotation))
|
||||
translation = -rotation_matrix @ position
|
||||
extrinsic = np.hstack([rotation_matrix, translation.reshape(3, 1)])
|
||||
rows.append([float(index)] + _INTRINSIC_ROW + [0.0, 0.0] + extrinsic.flatten().tolist())
|
||||
return rows
|
||||
|
||||
|
||||
def _interpolate_camera_poses(
|
||||
cameras: list[DreamXCamera],
|
||||
src_indices: np.ndarray,
|
||||
tgt_indices: np.ndarray,
|
||||
) -> list[DreamXCamera]:
|
||||
if len(cameras) <= 1:
|
||||
return [cameras[0]] * len(tgt_indices) if cameras else []
|
||||
src_rot_mat = np.array([camera.w2c_mat[:3, :3] for camera in cameras])
|
||||
src_trans_vec = np.array([camera.w2c_mat[:3, 3] for camera in cameras])
|
||||
|
||||
dets = np.linalg.det(src_rot_mat)
|
||||
flip_handedness = dets.size > 0 and np.median(dets) < 0.0
|
||||
if flip_handedness:
|
||||
flip_mat = np.diag([1.0, 1.0, -1.0]).astype(src_rot_mat.dtype)
|
||||
src_rot_mat = src_rot_mat @ flip_mat
|
||||
|
||||
trans = interp1d(src_indices, src_trans_vec, axis=0, kind="linear", bounds_error=False,
|
||||
fill_value="extrapolate")(tgt_indices)
|
||||
quats = Rotation.from_matrix(src_rot_mat).as_quat().copy()
|
||||
for index in range(1, len(quats)):
|
||||
if np.dot(quats[index], quats[index - 1]) < 0:
|
||||
quats[index] = -quats[index]
|
||||
rot = Slerp(src_indices, Rotation.from_quat(quats))(tgt_indices).as_matrix()
|
||||
if flip_handedness:
|
||||
rot = rot @ flip_mat
|
||||
|
||||
ref = cameras[0]
|
||||
result = []
|
||||
for index in range(len(tgt_indices)):
|
||||
w2c_mat = np.eye(4, dtype=np.float64)
|
||||
w2c_mat[:3, :] = np.hstack([rot[index], trans[index].reshape(3, 1)])
|
||||
result.append(DreamXCamera(ref.fx, ref.fy, ref.cx, ref.cy, w2c_mat))
|
||||
return result
|
||||
|
||||
|
||||
def _relative_c2w_poses(cameras: list[DreamXCamera]) -> np.ndarray:
|
||||
abs_w2cs = [camera.w2c_mat for camera in cameras]
|
||||
abs_c2ws = [camera.c2w_mat for camera in cameras]
|
||||
target_cam_c2w = np.eye(4, dtype=np.float64)
|
||||
abs2rel = target_cam_c2w @ abs_w2cs[0]
|
||||
poses = [target_cam_c2w] + [abs2rel @ c2w for c2w in abs_c2ws[1:]]
|
||||
return np.asarray(poses, dtype=np.float32)
|
||||
|
||||
|
||||
def _invert_se3(transforms: torch.Tensor) -> torch.Tensor:
|
||||
rotation_inv = transforms[..., :3, :3].transpose(-1, -2)
|
||||
output = torch.zeros_like(transforms)
|
||||
output[..., :3, :3] = rotation_inv
|
||||
output[..., :3, 3] = -torch.einsum("...ij,...j->...i", rotation_inv, transforms[..., :3, 3])
|
||||
output[..., 3, 3] = 1.0
|
||||
return output
|
||||
|
||||
|
||||
def build_dreamx_camera_condition(
|
||||
action_seq: list[str],
|
||||
action_speed_list: list[float],
|
||||
*,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
device: torch.device | str = "cpu",
|
||||
) -> dict[str, torch.Tensor]:
|
||||
del height, width # DreamX-World-5B-Cam uses fixed normalized intrinsics.
|
||||
duration = math.ceil(num_frames / len(action_seq))
|
||||
rows = _pose_rows_from_actions(action_seq, action_speed_list, duration)[:num_frames]
|
||||
cameras = [DreamXCamera.from_pose_row(row) for row in rows]
|
||||
|
||||
latent_frame_count = 1 + (len(cameras) - 1) // 4
|
||||
src_indices = np.arange(len(cameras), dtype=np.float64)
|
||||
tgt_indices = np.linspace(0, len(cameras) - 1, latent_frame_count)
|
||||
cameras = _interpolate_camera_poses(cameras, src_indices, tgt_indices)
|
||||
|
||||
c2ws = torch.as_tensor(_relative_c2w_poses(cameras), dtype=dtype, device=device)
|
||||
viewmats = _invert_se3(c2ws)
|
||||
|
||||
intrinsics = torch.zeros((latent_frame_count, 3, 3), dtype=dtype, device=device)
|
||||
intrinsics[:, 0, 0] = 969.6969696969696 / (960.0 * 2)
|
||||
intrinsics[:, 1, 1] = 969.6969696969696 / (540.0 * 2)
|
||||
intrinsics[:, 2, 2] = 1.0
|
||||
return {"viewmats": viewmats, "K": intrinsics}
|
||||
@@ -1,20 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Compatibility exports for DreamX-World pipeline configs."""
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import (
|
||||
DreamXWorld5BARPipelineConfig,
|
||||
DreamXWorld5BCamPipelineConfig,
|
||||
make_dreamx_world_5b_ar_dit_config,
|
||||
make_dreamx_world_5b_cam_dit_config,
|
||||
make_dreamx_world_5b_cam_text_encoder_config,
|
||||
make_dreamx_world_5b_cam_vae_config,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"DreamXWorld5BARPipelineConfig",
|
||||
"DreamXWorld5BCamPipelineConfig",
|
||||
"make_dreamx_world_5b_ar_dit_config",
|
||||
"make_dreamx_world_5b_cam_dit_config",
|
||||
"make_dreamx_world_5b_cam_text_encoder_config",
|
||||
"make_dreamx_world_5b_cam_vae_config",
|
||||
]
|
||||
@@ -1,67 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B autoregressive pipeline entrypoint."""
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import (
|
||||
DreamXWorldCameraConditioningStage,
|
||||
DreamXWorldImageVAEEncodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
DecodingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DreamXWorldARPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""DreamX-World-5B autoregressive causal camera pipeline."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
|
||||
pipeline_config_cls = DreamXWorld5BARPipelineConfig
|
||||
sampling_params_cls = SamplingParam
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
self.modules["scheduler"].set_timesteps(1000)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")))
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None),
|
||||
))
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=DreamXWorldImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
self.add_stage(stage_name="dreamx_camera_conditioning_stage", stage=DreamXWorldCameraConditioningStage())
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DreamXWorldARCausalDenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
))
|
||||
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"), pipeline=self))
|
||||
logger.info("DreamXWorldARPipeline initialized with autoregressive causal denoising")
|
||||
|
||||
|
||||
EntryClass = DreamXWorldARPipeline
|
||||
@@ -1,78 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World video pipeline entrypoint."""
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import DreamXWorldCameraConditioningStage
|
||||
from fastvideo.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
DecodingStage,
|
||||
DenoisingStage,
|
||||
InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
TextEncodingStage,
|
||||
TimestepPreparationStage,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DreamXWorldPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""DreamX-World-5B-Cam pipeline with native FastVideo camera conditioning."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
|
||||
pipeline_config_cls = DreamXWorld5BCamPipelineConfig
|
||||
sampling_params_cls = SamplingParam
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer", None),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(stage_name="dreamx_camera_conditioning_stage", stage=DreamXWorldCameraConditioningStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
transformer_2=self.get_module("transformer_2", None),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"), pipeline=self))
|
||||
|
||||
logger.info("DreamXWorldPipeline initialized with native camera conditioning")
|
||||
|
||||
|
||||
EntryClass = DreamXWorldPipeline
|
||||
@@ -1,58 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World model family pipeline presets."""
|
||||
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_NEGATIVE_PROMPT_CN = ("色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,"
|
||||
"静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,"
|
||||
"多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,"
|
||||
"形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,"
|
||||
"背景人很多,倒着走")
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="DreamX-World camera-conditioned denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
DREAMX_WORLD_5B_CAM = InferencePreset(
|
||||
name="dreamx_world_5b_cam",
|
||||
version=1,
|
||||
model_family="dreamx_world",
|
||||
description="DreamX-World 5B camera-control video generation",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 161,
|
||||
"fps": 16,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 30,
|
||||
"negative_prompt": _NEGATIVE_PROMPT_CN,
|
||||
},
|
||||
)
|
||||
|
||||
DREAMX_WORLD_5B_AR = InferencePreset(
|
||||
name="dreamx_world_5b_ar",
|
||||
version=1,
|
||||
model_family="dreamx_world",
|
||||
description="DreamX-World 5B autoregressive camera-control generation",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 1005,
|
||||
"fps": 16,
|
||||
"guidance_scale": 1.0,
|
||||
"num_inference_steps": 4,
|
||||
"negative_prompt": _NEGATIVE_PROMPT_CN,
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (DREAMX_WORLD_5B_CAM, DREAMX_WORLD_5B_AR)
|
||||
@@ -1,156 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World pipeline stages."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
|
||||
build_dreamx_camera_condition, )
|
||||
|
||||
DREAMX_Y_CAMERA_KEY = "dreamx_y_camera"
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class DreamXWorldCameraConditioningStage(PipelineStage):
|
||||
"""Build PRoPE camera conditioning for DreamX-World-5B-Cam."""
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
del fastvideo_args
|
||||
if DREAMX_Y_CAMERA_KEY in batch.extra:
|
||||
return batch
|
||||
|
||||
action_seq = batch.extra.get("dreamx_action_seq", batch.action_list)
|
||||
action_speed_list = batch.extra.get("dreamx_action_speed_list", batch.action_speed_list)
|
||||
if action_seq is None:
|
||||
action_seq = ["w"]
|
||||
if action_speed_list is None:
|
||||
action_speed_list = [4]
|
||||
|
||||
if isinstance(action_seq, str):
|
||||
action_seq = [action_seq]
|
||||
if isinstance(action_speed_list, int | float):
|
||||
action_speed_list = [action_speed_list]
|
||||
if len(action_speed_list) == 1 and len(action_seq) > 1:
|
||||
action_speed_list = list(action_speed_list) * len(action_seq)
|
||||
action_speed_list = [float(speed) for speed in action_speed_list]
|
||||
|
||||
height = int(batch.height) if batch.height is not None else 704
|
||||
width = int(batch.width) if batch.width is not None else 1280
|
||||
num_frames = int(batch.num_frames)
|
||||
dtype = batch.latents.dtype if torch.is_tensor(batch.latents) else torch.float32
|
||||
device = batch.latents.device if torch.is_tensor(batch.latents) else "cpu"
|
||||
|
||||
y_camera = build_dreamx_camera_condition(
|
||||
list(action_seq),
|
||||
action_speed_list,
|
||||
num_frames=num_frames,
|
||||
height=height,
|
||||
width=width,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
)
|
||||
batch.extra[DREAMX_Y_CAMERA_KEY] = {key: value.unsqueeze(0) for key, value in y_camera.items()}
|
||||
return batch
|
||||
|
||||
def verify_output(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> VerificationResult:
|
||||
del fastvideo_args
|
||||
result = VerificationResult()
|
||||
y_camera = batch.extra.get(DREAMX_Y_CAMERA_KEY)
|
||||
result.add_check("dreamx_y_camera", y_camera, lambda value: isinstance(value, dict))
|
||||
if isinstance(y_camera, dict):
|
||||
result.add_check("dreamx_y_camera.viewmats", y_camera.get("viewmats"), torch.is_tensor)
|
||||
result.add_check("dreamx_y_camera.K", y_camera.get("K"), torch.is_tensor)
|
||||
return result
|
||||
|
||||
|
||||
class DreamXWorldImageVAEEncodingStage(PipelineStage):
|
||||
"""Encode the conditioning image into the first-frame latent.
|
||||
|
||||
Official AR-forcing flow (AMAP-ML/DreamX-World inference_ar_forcing.py):
|
||||
the input image is resized, normalized to [-1, 1], VAE-encoded
|
||||
deterministically, and written into frame 0 of the noise — the causal
|
||||
denoiser then treats frame 0 as clean context. This stage produces
|
||||
``batch.image_latent`` ([B, C, 1, H_lat, W_lat]); the injection into
|
||||
the latents happens in DreamXWorldARCausalDenoisingStage.
|
||||
"""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
if batch.pil_image is None:
|
||||
# No conditioning image: the causal denoiser falls back to
|
||||
# running from pure noise (frame 0 uninitialized). Warn loudly —
|
||||
# this pipeline is registered I2V and the official flow always
|
||||
# forces from a frame.
|
||||
logger.warning("DreamXWorldARPipeline called without an input image; "
|
||||
"first-frame context will be noise (T2V-style). Pass an "
|
||||
"image for the official AR-forcing behavior.")
|
||||
return batch
|
||||
|
||||
from fastvideo.platforms import get_local_torch_device
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
device = get_local_torch_device()
|
||||
image = batch.pil_image
|
||||
if not isinstance(image, torch.Tensor):
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
assert isinstance(image, PIL.Image.Image)
|
||||
width = batch.width if isinstance(batch.width, int) else batch.width[0]
|
||||
height = batch.height if isinstance(batch.height, int) else batch.height[0]
|
||||
image = image.convert("RGB").resize((width, height), PIL.Image.Resampling.LANCZOS)
|
||||
arr = torch.from_numpy(np.asarray(image)).float().permute(2, 0, 1) / 255.0
|
||||
image = (arr - 0.5) / 0.5 # official Normalize([0.5], [0.5])
|
||||
image = image.unsqueeze(0) # [1, C, H, W]
|
||||
if image.dim() == 4:
|
||||
image = image.unsqueeze(2) # [B, C, 1, H, W]
|
||||
elif image.dim() == 5:
|
||||
image = image[:, :, :1]
|
||||
image = image.to(device=device, dtype=torch.float32)
|
||||
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
self.vae = self.vae.to(device)
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
if not vae_autocast_enabled:
|
||||
image = image.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(image)
|
||||
|
||||
# Official encode_to_latent is deterministic ((mean - mu) / sigma per
|
||||
# channel); the posterior mean + shift/scale is the FastVideo
|
||||
# equivalent of that normalization.
|
||||
latent = encoder_output.mean
|
||||
if getattr(self.vae, "shift_factor", None) is not None:
|
||||
shift = self.vae.shift_factor
|
||||
latent = latent - (shift.to(latent.device, latent.dtype) if isinstance(shift, torch.Tensor) else shift)
|
||||
scale = self.vae.scaling_factor
|
||||
latent = latent * (scale.to(latent.device, latent.dtype) if isinstance(scale, torch.Tensor) else scale)
|
||||
|
||||
batch.image_latent = latent
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
result = VerificationResult()
|
||||
return result
|
||||
@@ -190,19 +190,6 @@ class DenoisingStage(PipelineStage):
|
||||
},
|
||||
)
|
||||
|
||||
dreamx_y_camera = batch.extra.get("dreamx_y_camera", batch.extra.get("y_camera"))
|
||||
if isinstance(dreamx_y_camera, dict):
|
||||
dreamx_y_camera = {
|
||||
key: value.to(device=local_device, dtype=target_dtype) if torch.is_tensor(value) else value
|
||||
for key, value in dreamx_y_camera.items()
|
||||
}
|
||||
dreamx_camera_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"y_camera": dreamx_y_camera,
|
||||
},
|
||||
)
|
||||
|
||||
for key in ("flux2_txt_ids", "flux2_img_ids"):
|
||||
value = batch.extra.get(key)
|
||||
if torch.is_tensor(value):
|
||||
@@ -254,11 +241,7 @@ class DenoisingStage(PipelineStage):
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
assert self.vae is not None, "VAE is not provided for TI2V task"
|
||||
vae_device = next(self.vae.parameters()).device
|
||||
self.vae = self.vae.to(local_device)
|
||||
z = self.vae.encode(batch.pil_image).mean.float()
|
||||
if getattr(fastvideo_args, "vae_cpu_offload", False):
|
||||
self.vae = self.vae.to(vae_device)
|
||||
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
z -= self.vae.shift_factor.to(z.device, z.dtype)
|
||||
@@ -511,7 +494,6 @@ class DenoisingStage(PipelineStage):
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**camera_kwargs,
|
||||
**dreamx_camera_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
**flux2_id_kwargs,
|
||||
)
|
||||
@@ -554,7 +536,6 @@ class DenoisingStage(PipelineStage):
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**camera_kwargs,
|
||||
**dreamx_camera_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
**flux2_id_kwargs,
|
||||
)
|
||||
|
||||
@@ -20,7 +20,6 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
|
||||
Cosmos25Config,
|
||||
Cosmos25_14BConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.gen3c import Gen3CConfig
|
||||
@@ -774,40 +773,6 @@ def _register_configs() -> None:
|
||||
model_family="wan",
|
||||
default_preset="wan_2_2_ti2v_5b",
|
||||
)
|
||||
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=DreamXWorld5BCamPipelineConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"GD-ML/DreamX-World-5B-Cam",
|
||||
],
|
||||
model_detectors=[
|
||||
# Mutually exclusive with the AR detector below: Cam requires an
|
||||
# explicit "cam" marker so hyphenated AR local paths (e.g.
|
||||
# /ckpts/dreamx-world-5b-converted) don't first-match here —
|
||||
# detector resolution is first-match in registration order.
|
||||
lambda path:
|
||||
("dreamx-world" in path.lower() and "cam" in path.lower()) or "dreamxworldpipeline" in path.lower()
|
||||
],
|
||||
model_family="dreamx_world",
|
||||
default_preset="dreamx_world_5b_cam",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=DreamXWorld5BARPipelineConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"GD-ML/DreamX-World-5B",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path:
|
||||
("dreamx-world-5b" in path.lower() and "cam" not in path.lower()) or "dreamxworldarpipeline" in path.lower(
|
||||
)
|
||||
],
|
||||
model_family="dreamx_world",
|
||||
default_preset="dreamx_world_5b_ar",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
|
||||
@@ -986,8 +951,6 @@ def _register_presets() -> None:
|
||||
from fastvideo.api.presets import register_preset
|
||||
from fastvideo.pipelines.basic.cosmos.presets import (
|
||||
ALL_PRESETS as COSMOS_PRESETS, )
|
||||
from fastvideo.pipelines.basic.dreamx_world.presets import (
|
||||
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.gamecraft.presets import (
|
||||
ALL_PRESETS as GAMECRAFT_PRESETS, )
|
||||
from fastvideo.pipelines.basic.gen3c.presets import (
|
||||
@@ -1021,7 +984,6 @@ def _register_presets() -> None:
|
||||
|
||||
all_preset_groups = (
|
||||
COSMOS_PRESETS,
|
||||
DREAMX_WORLD_PRESETS,
|
||||
FLUX2_PRESETS,
|
||||
GAMECRAFT_PRESETS,
|
||||
GEN3C_PRESETS,
|
||||
|
||||
@@ -3,16 +3,15 @@
|
||||
|
||||
Landed in PR #1225 slice 5 (Attn-QAT 5/12). The resolver centralises the
|
||||
varlen-flash-attn import-fallback logic that several backends
|
||||
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The resolution
|
||||
(``bsa_attn.py``, ``video_sparse_attn.py``) used to duplicate. The fallback
|
||||
order is:
|
||||
|
||||
1. ``fastvideo.attention.utils.flash_attn_cute`` -- only when
|
||||
``FASTVIDEO_FA4=1`` (explicit opt-in), and then it must import or the
|
||||
resolver raises RuntimeError instead of falling through
|
||||
1. ``fastvideo.attention.utils.flash_attn_cute``
|
||||
2. ``flash_attn_interface``
|
||||
3. ``flash_attn``
|
||||
|
||||
These tests verify the opt-in gate and the FA3/FA2 fallthrough. CPU-only, no
|
||||
These tests verify that the resolver picks the highest-priority impl
|
||||
available and falls through cleanly on ``ImportError``. CPU-only, no
|
||||
flash-attn install required.
|
||||
"""
|
||||
|
||||
@@ -36,36 +35,8 @@ def _reload_resolver_module():
|
||||
return importlib.import_module("fastvideo.attention.utils.flash_attn_no_pad")
|
||||
|
||||
|
||||
def test_resolver_skips_cute_without_opt_in(monkeypatch) -> None:
|
||||
"""Without ``FASTVIDEO_FA4=1`` the resolver must not even attempt the cute
|
||||
import."""
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
attempted: list[str] = []
|
||||
real_import = builtins.__import__
|
||||
|
||||
def spying_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
attempted.append(name)
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", spying_import)
|
||||
|
||||
mod = _reload_resolver_module()
|
||||
resolved = mod._resolve_flash_attn_varlen_func()
|
||||
assert resolved is not None
|
||||
assert resolved.__name__ == "flash_attn_varlen_func"
|
||||
assert "fastvideo.attention.utils.flash_attn_cute" not in attempted
|
||||
|
||||
|
||||
def test_resolver_raises_when_opted_in_but_cute_unavailable(monkeypatch) -> None:
|
||||
"""With ``FASTVIDEO_FA4=1`` an unimportable cute build fails loudly instead
|
||||
of silently falling through to FA3/FA2.
|
||||
|
||||
The resolver runs at module import time, so the reload itself must raise.
|
||||
It raises RuntimeError (not ImportError) so importers that treat
|
||||
ImportError as "flash-attn not installed" (``bsa_attn.py``) cannot swallow
|
||||
the opted-in failure.
|
||||
"""
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
def test_resolver_falls_back_when_cute_unavailable(monkeypatch) -> None:
|
||||
"""When ``flash_attn_cute`` is unimportable, resolver tries the next impl."""
|
||||
real_import = builtins.__import__
|
||||
|
||||
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
@@ -75,17 +46,21 @@ def test_resolver_raises_when_opted_in_but_cute_unavailable(monkeypatch) -> None
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", patched_import)
|
||||
|
||||
with pytest.raises(RuntimeError, match="cute disabled for test"):
|
||||
_reload_resolver_module()
|
||||
mod = _reload_resolver_module()
|
||||
resolved = mod._resolve_flash_attn_varlen_func()
|
||||
assert resolved is not None
|
||||
assert resolved.__name__ == "flash_attn_varlen_func"
|
||||
|
||||
|
||||
def test_resolver_returns_flash_attn_when_interface_unavailable(monkeypatch) -> None:
|
||||
def test_resolver_returns_flash_attn_when_cute_and_interface_unavailable(monkeypatch) -> None:
|
||||
"""The terminal fallback is the plain ``flash_attn`` import."""
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
real_import = builtins.__import__
|
||||
|
||||
def patched_import(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
if name == "flash_attn_interface":
|
||||
if name in {
|
||||
"fastvideo.attention.utils.flash_attn_cute",
|
||||
"flash_attn_interface",
|
||||
}:
|
||||
raise ImportError(f"{name} disabled for test")
|
||||
return real_import(name, globals, locals, fromlist, level)
|
||||
|
||||
|
||||
@@ -1,203 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import init_device_mesh
|
||||
from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
|
||||
WORLD_SIZE = 2
|
||||
HIDDEN_SIZE = 8
|
||||
SEED = 1379
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
|
||||
|
||||
def _run_torchrun(script_path: Path, mode: str, output_path: Path) -> None:
|
||||
# --standalone binds the rendezvous port atomically, avoiding the
|
||||
# free-port-probe race a hand-picked --master_port would have.
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--standalone",
|
||||
"--nproc_per_node",
|
||||
str(WORLD_SIZE),
|
||||
str(script_path),
|
||||
"--rmsnorm-fsdp-worker",
|
||||
"--mode",
|
||||
mode,
|
||||
"--output",
|
||||
str(output_path),
|
||||
]
|
||||
env = os.environ.copy()
|
||||
env.setdefault("TORCHDYNAMO_DISABLE", "1")
|
||||
try:
|
||||
process = subprocess.run(
|
||||
cmd,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=env,
|
||||
timeout=120,
|
||||
)
|
||||
except subprocess.TimeoutExpired as error:
|
||||
raise RuntimeError(
|
||||
f"{mode} worker timed out after 120 seconds\n"
|
||||
f"STDOUT:\n{error.stdout}\n"
|
||||
f"STDERR:\n{error.stderr}"
|
||||
) from error
|
||||
if process.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"{mode} worker failed with code {process.returncode}\n"
|
||||
f"STDOUT:\n{process.stdout}\n"
|
||||
f"STDERR:\n{process.stderr}"
|
||||
)
|
||||
|
||||
|
||||
def _summarize_tensor(tensor: torch.Tensor | Any) -> dict[str, Any]:
|
||||
return {
|
||||
"type": type(tensor).__name__,
|
||||
"is_dtensor": isinstance(tensor, DTensor),
|
||||
"shape": list(tensor.shape) if hasattr(tensor, "shape") else None,
|
||||
"device": str(tensor.device) if hasattr(tensor, "device") else None,
|
||||
"dtype": str(tensor.dtype) if hasattr(tensor, "dtype") else None,
|
||||
}
|
||||
|
||||
|
||||
def _run_worker(mode: str, output_path: Path) -> None:
|
||||
if mode not in {
|
||||
"module_no_offload",
|
||||
"direct_no_offload",
|
||||
"module_cpu_offload",
|
||||
"direct_cpu_offload",
|
||||
}:
|
||||
raise ValueError(f"Unsupported mode: {mode}")
|
||||
|
||||
dist.init_process_group("nccl")
|
||||
rank = dist.get_rank()
|
||||
world_size = dist.get_world_size()
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
torch.cuda.set_device(device)
|
||||
torch.manual_seed(SEED + rank)
|
||||
|
||||
try:
|
||||
mesh = init_device_mesh("cuda", (world_size,))
|
||||
norm = RMSNorm(HIDDEN_SIZE, eps=1e-6, has_weight=True).to(device)
|
||||
with torch.no_grad():
|
||||
norm.weight.fill_(1.0)
|
||||
|
||||
fsdp_kwargs: dict[str, Any] = {"mesh": mesh}
|
||||
if mode.endswith("cpu_offload"):
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=False)
|
||||
# fully_shard is applied to the bare RMSNorm to make the hook bypass
|
||||
# observable. Production sharding (fsdp_load.shard_model) only wraps
|
||||
# whole transformer blocks, whose pre-forward all-gather localizes norm
|
||||
# weights before the qk-norm call sites run, so this pins the dispatch
|
||||
# invariant rather than reproducing a production topology.
|
||||
fully_shard(norm, **fsdp_kwargs)
|
||||
|
||||
x = torch.randn(2, 3, HIDDEN_SIZE, device=device, dtype=torch.bfloat16)
|
||||
call_kind = "direct" if mode.startswith("direct") else "module"
|
||||
|
||||
try:
|
||||
if call_kind == "direct":
|
||||
output = norm.forward_native(x)
|
||||
else:
|
||||
output = norm(x)
|
||||
torch.cuda.synchronize(device)
|
||||
result = {
|
||||
"rank": rank,
|
||||
"ok": True,
|
||||
"mode": mode,
|
||||
"weight": _summarize_tensor(norm.weight),
|
||||
"output": _summarize_tensor(output),
|
||||
}
|
||||
except Exception as exc:
|
||||
result = {
|
||||
"rank": rank,
|
||||
"ok": False,
|
||||
"mode": mode,
|
||||
"error_type": type(exc).__name__,
|
||||
"error": str(exc),
|
||||
"weight": _summarize_tensor(norm.weight),
|
||||
}
|
||||
|
||||
gathered = [None for _ in range(world_size)] if rank == 0 else None
|
||||
dist.gather_object(result, object_gather_list=gathered, dst=0)
|
||||
if rank == 0:
|
||||
output_path.write_text(json.dumps(gathered, indent=2), encoding="utf-8")
|
||||
dist.barrier()
|
||||
finally:
|
||||
dist.destroy_process_group()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mode", "expect_ok"),
|
||||
[
|
||||
("module_no_offload", True),
|
||||
("direct_no_offload", False),
|
||||
("module_cpu_offload", True),
|
||||
("direct_cpu_offload", False),
|
||||
],
|
||||
)
|
||||
def test_rmsnorm_forward_native_bypasses_fsdp_hooks(mode: str, expect_ok: bool, tmp_path: Path) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("This test requires CUDA.")
|
||||
if torch.cuda.device_count() < WORLD_SIZE:
|
||||
pytest.skip(f"This test requires at least {WORLD_SIZE} CUDA devices.")
|
||||
|
||||
output_path = tmp_path / f"{mode}.json"
|
||||
_run_torchrun(Path(__file__).resolve(), mode, output_path)
|
||||
results = json.loads(output_path.read_text(encoding="utf-8"))
|
||||
print(f"\n{mode} results:\n{json.dumps(results, indent=2)}")
|
||||
|
||||
if expect_ok:
|
||||
failures = [result for result in results if not result["ok"]]
|
||||
assert not failures, json.dumps(results, indent=2)
|
||||
return
|
||||
|
||||
successes = [result for result in results if result["ok"]]
|
||||
assert not successes, json.dumps(results, indent=2)
|
||||
error_text = "\n".join(result.get("error", "") for result in results)
|
||||
# Pin the specific bypassed-hook failure: "got mixed torch.Tensor and
|
||||
# DTensor" ("Tensor" alone is a substring of "DTensor", so it adds nothing).
|
||||
assert "mixed" in error_text and "DTensor" in error_text, json.dumps(results, indent=2)
|
||||
|
||||
|
||||
def test_no_direct_forward_native_calls_in_models() -> None:
|
||||
"""Direct .forward_native(...) calls bypass nn.Module.__call__ and FSDP
|
||||
hooks (issue #1379); model code must use module dispatch instead."""
|
||||
models_dir = REPO_ROOT / "fastvideo" / "models"
|
||||
offenders = [
|
||||
str(path.relative_to(REPO_ROOT))
|
||||
for path in sorted(models_dir.rglob("*.py"))
|
||||
if ".forward_native(" in path.read_text(encoding="utf-8")
|
||||
]
|
||||
assert not offenders, f"Replace .forward_native(...) with module dispatch in: {offenders}"
|
||||
|
||||
|
||||
def _parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--rmsnorm-fsdp-worker", action="store_true")
|
||||
parser.add_argument("--mode", type=str, default=None)
|
||||
parser.add_argument("--output", type=str, default=None)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = _parse_args()
|
||||
if not args.rmsnorm_fsdp_worker:
|
||||
raise SystemExit("This module is intended to be run by pytest.")
|
||||
if args.mode is None or args.output is None:
|
||||
raise SystemExit("--mode and --output are required in worker mode.")
|
||||
_run_worker(mode=args.mode, output_path=Path(args.output))
|
||||
@@ -32,8 +32,7 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -41,9 +40,6 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App("fastvideo-gpu-job")
|
||||
|
||||
REPO_DIR = "/FastVideo"
|
||||
@@ -76,7 +72,12 @@ local_secrets = modal.Secret.from_dict({
|
||||
# Mutable tags inherit the registry image's baked backend, including custom
|
||||
# FASTVIDEO_MODAL_IMAGE overrides. Explicit CUDA tags also work with older
|
||||
# images that predate the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(IMAGE_TAG)
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in IMAGE_TAG.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(IMAGE_REF, add_python="3.12")
|
||||
@@ -97,9 +98,6 @@ image = (
|
||||
"TOKENIZERS_PARALLELISM": "false",
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
|
||||
# references. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
})
|
||||
)
|
||||
|
||||
|
||||
@@ -5,7 +5,6 @@ through the ``fastvideo`` package, so they keep working on thin CI hosts that
|
||||
have ``modal`` but not torch.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import urllib.request
|
||||
|
||||
_REGISTRY = "ghcr.io"
|
||||
@@ -56,23 +55,3 @@ def resolve_image_ref(image_ref: str) -> str:
|
||||
print(f"WARNING: could not resolve {image_ref} to a digest ({error}); "
|
||||
"Modal may reuse a stale cached image for this tag.")
|
||||
return image_ref
|
||||
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
"""UV_TORCH_BACKEND for a launcher image tag.
|
||||
|
||||
A caller-set UV_TORCH_BACKEND always wins. Otherwise sniff the CUDA
|
||||
version from an explicit image tag (cuda13 -> cu130, cuda12.6 -> cu126)
|
||||
so uv resolves torch against the image's toolkit. Mutable tags (e.g.
|
||||
py3.12-latest) return None and inherit the registry image's baked
|
||||
backend, which keeps a latest-tag CUDA transition safe.
|
||||
"""
|
||||
override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if override:
|
||||
return override
|
||||
tag = image_tag.lower()
|
||||
if "cuda13" in tag:
|
||||
return "cu130"
|
||||
if "cuda12.6" in tag:
|
||||
return "cu126"
|
||||
return None
|
||||
|
||||
@@ -5,8 +5,7 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -14,9 +13,6 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -28,7 +24,12 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
|
||||
image = (modal.Image.from_registry(
|
||||
image_ref, add_python="3.12"
|
||||
@@ -60,10 +61,6 @@ image = (modal.Image.from_registry(
|
||||
**({
|
||||
"UV_TORCH_BACKEND": uv_torch_backend_override
|
||||
} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
|
||||
# SSIM/perf baselines. Caller override wins.
|
||||
"FASTVIDEO_FA4":
|
||||
os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
"HF_REPO_ID":
|
||||
"FastVideo/performance-tracking",
|
||||
}))
|
||||
|
||||
@@ -13,8 +13,7 @@ import modal
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
try:
|
||||
from modal_image_utils import ( # noqa: E402
|
||||
resolve_image_ref, resolve_uv_torch_backend)
|
||||
from modal_image_utils import resolve_image_ref # noqa: E402
|
||||
except ModuleNotFoundError:
|
||||
# Remote Modal containers re-import this module but mount only the
|
||||
# entrypoint file; the digest resolution already happened at local
|
||||
@@ -22,9 +21,6 @@ except ModuleNotFoundError:
|
||||
def resolve_image_ref(image_ref: str) -> str:
|
||||
return image_ref
|
||||
|
||||
def resolve_uv_torch_backend(image_tag: str) -> str | None:
|
||||
return os.environ.get("UV_TORCH_BACKEND")
|
||||
|
||||
app = modal.App()
|
||||
|
||||
model_vol = modal.Volume.from_name("hf-model-weights")
|
||||
@@ -36,7 +32,12 @@ print(f"Using image: {image_ref}")
|
||||
# Mutable tags inherit the registry image's baked backend, keeping a latest-tag
|
||||
# transition safe. Explicit CUDA tags also work with older images that predate
|
||||
# the baked setting, and a caller override always wins.
|
||||
uv_torch_backend_override = resolve_uv_torch_backend(image_tag)
|
||||
uv_torch_backend_override = os.environ.get("UV_TORCH_BACKEND")
|
||||
if not uv_torch_backend_override:
|
||||
if "cuda13" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu130"
|
||||
elif "cuda12.6" in image_tag.lower():
|
||||
uv_torch_backend_override = "cu126"
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(image_ref, add_python="3.12")
|
||||
@@ -63,9 +64,6 @@ image = (
|
||||
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"IMAGE_VERSION": image_version,
|
||||
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
|
||||
# FA4 is opt-in (FASTVIDEO_FA4); the SSIM references were seeded
|
||||
# with FA4 inference, so keep it enabled in CI. Caller override wins.
|
||||
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
@@ -1,116 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.tests.performance.test_inference_performance import (
|
||||
_benchmark_display_id,
|
||||
_config_identity_metadata,
|
||||
_is_v2_config,
|
||||
_validate_benchmark_config,
|
||||
)
|
||||
|
||||
|
||||
def _v2_config():
|
||||
return {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
|
||||
def test_v1_benchmark_config_without_schema_version_validates():
|
||||
cfg = {
|
||||
"benchmark_id": "legacy-benchmark",
|
||||
}
|
||||
|
||||
_validate_benchmark_config(cfg, "legacy.json")
|
||||
|
||||
assert _is_v2_config(cfg) is False
|
||||
assert _config_identity_metadata(cfg) == {}
|
||||
assert _benchmark_display_id(cfg) == "legacy-benchmark"
|
||||
|
||||
|
||||
def test_v2_benchmark_config_identity_validates_and_is_preserved():
|
||||
cfg = _v2_config()
|
||||
cfg["quality_metadata"] = {"some": "data"}
|
||||
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
assert _is_v2_config(cfg) is True
|
||||
assert _config_identity_metadata(cfg) == {
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": {"some": "data"},
|
||||
}
|
||||
|
||||
|
||||
def test_v2_benchmark_config_missing_identity_fields_fails_clearly():
|
||||
cfg = _v2_config()
|
||||
del cfg["variant_id"]
|
||||
del cfg["benchmark_version"]
|
||||
|
||||
expected = "wan.json: missing required v2 identity fields: variant_id, benchmark_version"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("field", "value"),
|
||||
[
|
||||
("workload_id", {}),
|
||||
("workload_id", ""),
|
||||
("workload_id", " "),
|
||||
("variant_id", []),
|
||||
("variant_id", ""),
|
||||
("variant_id", " "),
|
||||
],
|
||||
)
|
||||
def test_v2_benchmark_config_rejects_invalid_string_identity_values(field, value):
|
||||
cfg = _v2_config()
|
||||
cfg[field] = value
|
||||
|
||||
expected = f"wan.json: v2 identity field {field!r} must be a non-empty string"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, "1", 1.5, True])
|
||||
def test_v2_benchmark_config_rejects_invalid_benchmark_version_values(value):
|
||||
cfg = _v2_config()
|
||||
cfg["benchmark_version"] = value
|
||||
|
||||
expected = "wan.json: v2 identity field 'benchmark_version' must be an integer"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_partial_v2_identity_requires_schema_version():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
}
|
||||
|
||||
expected = "wan.json: v2 benchmark identity fields require config_schema_version=2"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
|
||||
|
||||
def test_optional_v2_metadata_fields_must_be_objects():
|
||||
cfg = {
|
||||
"benchmark_id": "wan-t2v-1.3b-2gpu",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v-1.3b",
|
||||
"variant_id": "canonical",
|
||||
"benchmark_version": 1,
|
||||
"quality_metadata": ["not", "an", "object"],
|
||||
}
|
||||
|
||||
expected = "wan.json: optional v2 metadata field 'quality_metadata' must be an object"
|
||||
with pytest.raises(ValueError, match=expected):
|
||||
_validate_benchmark_config(cfg, "wan.json")
|
||||
@@ -28,17 +28,6 @@ STAGE_METRIC_MAP: dict[str, str] = {
|
||||
"DmdDenoisingStage": "dit_time_s",
|
||||
"DecodingStage": "vae_decode_time_s",
|
||||
}
|
||||
V2_CONFIG_SCHEMA_VERSION = 2
|
||||
V2_REQUIRED_IDENTITY_FIELDS = (
|
||||
"workload_id",
|
||||
"variant_id",
|
||||
"benchmark_version",
|
||||
)
|
||||
V2_OPTIONAL_METADATA_FIELDS = (
|
||||
"recipe",
|
||||
"metric_threshold_policy",
|
||||
"quality_metadata",
|
||||
)
|
||||
|
||||
# -- Config discovery -------------------------------------------------------
|
||||
|
||||
@@ -53,71 +42,6 @@ _BENCHMARKS_DIR = os.path.join(
|
||||
)
|
||||
|
||||
|
||||
def _has_v2_fields(cfg):
|
||||
v2_fields = V2_REQUIRED_IDENTITY_FIELDS + V2_OPTIONAL_METADATA_FIELDS
|
||||
return any(field in cfg for field in v2_fields)
|
||||
|
||||
|
||||
def _is_v2_config(cfg):
|
||||
return cfg.get("config_schema_version") == V2_CONFIG_SCHEMA_VERSION
|
||||
|
||||
|
||||
def _validate_non_empty_string(value, field, path):
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be a non-empty string")
|
||||
|
||||
|
||||
def _validate_integer(value, field, path):
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{path}: v2 identity field {field!r} must be an integer")
|
||||
|
||||
|
||||
def _validate_benchmark_config(cfg, path="<memory>"):
|
||||
missing_common = [field for field in ("benchmark_id",) if field not in cfg]
|
||||
if missing_common:
|
||||
raise ValueError(f"{path}: missing required benchmark config fields: {', '.join(missing_common)}")
|
||||
|
||||
schema_version = cfg.get("config_schema_version")
|
||||
if schema_version is None:
|
||||
if _has_v2_fields(cfg):
|
||||
raise ValueError(f"{path}: v2 benchmark identity fields require config_schema_version=2")
|
||||
return
|
||||
|
||||
if schema_version != V2_CONFIG_SCHEMA_VERSION:
|
||||
raise ValueError(f"{path}: unsupported benchmark config_schema_version={schema_version!r}")
|
||||
|
||||
missing_v2 = [field for field in V2_REQUIRED_IDENTITY_FIELDS if field not in cfg]
|
||||
if missing_v2:
|
||||
raise ValueError(f"{path}: missing required v2 identity fields: {', '.join(missing_v2)}")
|
||||
|
||||
_validate_non_empty_string(cfg["workload_id"], "workload_id", path)
|
||||
_validate_non_empty_string(cfg["variant_id"], "variant_id", path)
|
||||
_validate_integer(cfg["benchmark_version"], "benchmark_version", path)
|
||||
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg and not isinstance(cfg[field], Mapping):
|
||||
raise ValueError(f"{path}: optional v2 metadata field {field!r} must be an object")
|
||||
|
||||
|
||||
def _config_identity_metadata(cfg):
|
||||
if not _is_v2_config(cfg):
|
||||
return {}
|
||||
metadata = {
|
||||
"config_schema_version": cfg["config_schema_version"],
|
||||
"workload_id": cfg["workload_id"],
|
||||
"variant_id": cfg["variant_id"],
|
||||
"benchmark_version": cfg["benchmark_version"],
|
||||
}
|
||||
for field in V2_OPTIONAL_METADATA_FIELDS:
|
||||
if field in cfg:
|
||||
metadata[field] = cfg[field]
|
||||
return metadata
|
||||
|
||||
|
||||
def _benchmark_display_id(cfg):
|
||||
return cfg["benchmark_id"]
|
||||
|
||||
|
||||
def _discover_benchmarks():
|
||||
"""Glob benchmark JSON configs and return list of (id, config) tuples."""
|
||||
pattern = os.path.join(_BENCHMARKS_DIR, "*.json")
|
||||
@@ -125,7 +49,6 @@ def _discover_benchmarks():
|
||||
for path in sorted(glob.glob(pattern)):
|
||||
with open(path) as f:
|
||||
cfg = json.load(f)
|
||||
_validate_benchmark_config(cfg, path)
|
||||
configs.append(cfg)
|
||||
return configs
|
||||
|
||||
@@ -296,7 +219,6 @@ def _run_benchmark(cfg):
|
||||
|
||||
results = {
|
||||
"benchmark_id": cfg["benchmark_id"],
|
||||
**_config_identity_metadata(cfg),
|
||||
"model_short_name": model_info.get("model_short_name", ""),
|
||||
"device": device_name,
|
||||
"num_gpus": init_kwargs.get("num_gpus", 1),
|
||||
@@ -353,7 +275,7 @@ def _run_benchmark(cfg):
|
||||
@pytest.mark.parametrize(
|
||||
"cfg",
|
||||
_BENCHMARK_CONFIGS,
|
||||
ids=[_benchmark_display_id(c) for c in _BENCHMARK_CONFIGS],
|
||||
ids=[c["benchmark_id"] for c in _BENCHMARK_CONFIGS],
|
||||
)
|
||||
def test_inference_performance(cfg):
|
||||
"""Measure generation latency, peak GPU memory, and component-level timings
|
||||
|
||||
@@ -3,9 +3,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from logging import Logger
|
||||
from typing import Iterator
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
@@ -67,12 +67,6 @@ def _find_reference_video(reference_folder: str, prompt: str) -> str:
|
||||
raise FileNotFoundError("Reference video missing")
|
||||
|
||||
|
||||
def _remove_stale_generated_video(output_dir: str, output_video_name: str) -> None:
|
||||
stale_path = os.path.join(output_dir, output_video_name)
|
||||
if os.path.exists(stale_path):
|
||||
os.remove(stale_path)
|
||||
|
||||
|
||||
def _assert_similarity(
|
||||
*,
|
||||
logger: Logger,
|
||||
@@ -220,7 +214,6 @@ def run_text_to_video_similarity_test(
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
_remove_stale_generated_video(output_dir, output_video_name)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
@@ -296,7 +289,6 @@ def run_image_to_video_similarity_test(
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
_remove_stale_generated_video(output_dir, output_video_name)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
|
||||
@@ -1,214 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
resolve_inference_device_reference_folder,
|
||||
run_image_to_video_similarity_test,
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
device_reference_folder = resolve_inference_device_reference_folder(logger)
|
||||
|
||||
_LOCAL_CONVERTED_MODEL = Path("converted_weights/dreamx_world")
|
||||
_MODEL_PATH = os.getenv(
|
||||
"DREAMX_WORLD_SSIM_MODEL_PATH",
|
||||
str(_LOCAL_CONVERTED_MODEL),
|
||||
)
|
||||
_LOCAL_AR_CANDIDATES = (
|
||||
Path("/tmp/converted_dreamx_world_ar"),
|
||||
Path("/root/data/dreamx_world_ar_converted"),
|
||||
)
|
||||
_DEFAULT_AR_MODEL_PATH = next(
|
||||
(str(path) for path in _LOCAL_AR_CANDIDATES if path.exists()),
|
||||
str(_LOCAL_AR_CANDIDATES[0]),
|
||||
)
|
||||
_AR_MODEL_PATH = os.getenv(
|
||||
"DREAMX_WORLD_AR_SSIM_MODEL_PATH",
|
||||
_DEFAULT_AR_MODEL_PATH,
|
||||
)
|
||||
|
||||
DREAMX_WORLD_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": _MODEL_PATH,
|
||||
"height": 64,
|
||||
"width": 64,
|
||||
"num_frames": 9,
|
||||
"num_inference_steps": 1,
|
||||
"guidance_scale": 1.0,
|
||||
"seed": 1024,
|
||||
"fps": 16,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_FULL_QUALITY_PARAMS = {
|
||||
**DREAMX_WORLD_PARAMS,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 161,
|
||||
"num_inference_steps": 30,
|
||||
"guidance_scale": 5.0,
|
||||
}
|
||||
|
||||
|
||||
DREAMX_WORLD_AR_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": _AR_MODEL_PATH,
|
||||
"height": 192,
|
||||
"width": 192,
|
||||
"num_frames": 81,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 1.0,
|
||||
"seed": 2048,
|
||||
"fps": 16,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_AR_FULL_QUALITY_PARAMS = {
|
||||
**DREAMX_WORLD_AR_PARAMS,
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 1005,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B-Cam": DREAMX_WORLD_PARAMS,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B": DREAMX_WORLD_AR_PARAMS,
|
||||
}
|
||||
|
||||
FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B-Cam": DREAMX_WORLD_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B": DREAMX_WORLD_AR_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_TEST_CASES = [
|
||||
(
|
||||
"A cinematic first-person drive through a futuristic coastal city at sunrise, "
|
||||
"reflective glass towers, clean streets, soft volumetric light.",
|
||||
("w", "d", "w"),
|
||||
(4.0, 2.0, 4.0),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
DREAMX_WORLD_AR_TEST_CASES = [
|
||||
(
|
||||
"A long autonomous drive through a futuristic coastal city at sunrise, "
|
||||
"smooth forward camera motion, reflective glass towers, clean streets.",
|
||||
("w", "d", "w", "a"),
|
||||
(2.0, 1.0, 2.0, 1.0),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _write_deterministic_reference_image(path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
|
||||
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
|
||||
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
|
||||
draw.line((0, 80, 96, 66), fill=(238, 209, 124), width=3)
|
||||
image.save(path)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_TEST_CASES)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_MODEL_TO_PARAMS.keys()))
|
||||
def test_dreamx_world_inference_similarity(
|
||||
prompt: str,
|
||||
action_list: tuple[str, ...],
|
||||
action_speed_list: tuple[float, ...],
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
model_path = Path(str(DREAMX_WORLD_MODEL_TO_PARAMS[model_id]["model_path"]))
|
||||
if not model_path.exists():
|
||||
pytest.skip(
|
||||
f"DreamX-World converted model path is missing: {model_path}. "
|
||||
"Set DREAMX_WORLD_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
|
||||
)
|
||||
|
||||
image_path = tmp_path / "dreamx_world_ssim_input.png"
|
||||
_write_deterministic_reference_image(image_path)
|
||||
|
||||
run_image_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
image_path=str(image_path),
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=DREAMX_WORLD_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.98,
|
||||
init_kwargs_override={
|
||||
"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",
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"action_list": list(action_list),
|
||||
"action_speed_list": list(action_speed_list),
|
||||
},
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_AR_TEST_CASES)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_AR_MODEL_TO_PARAMS.keys()))
|
||||
def test_dreamx_world_ar_inference_similarity(
|
||||
prompt: str,
|
||||
action_list: tuple[str, ...],
|
||||
action_speed_list: tuple[float, ...],
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
model_path = Path(str(DREAMX_WORLD_AR_MODEL_TO_PARAMS[model_id]["model_path"]))
|
||||
if not model_path.exists():
|
||||
pytest.skip(
|
||||
f"DreamX-World AR converted model path is missing: {model_path}. "
|
||||
"Set DREAMX_WORLD_AR_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
|
||||
)
|
||||
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=DREAMX_WORLD_AR_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.98,
|
||||
init_kwargs_override={
|
||||
"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": "DreamXWorldARPipeline",
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"action_list": list(action_list),
|
||||
"action_speed_list": list(action_speed_list),
|
||||
},
|
||||
)
|
||||
@@ -127,12 +127,4 @@ def test_wan_causal_dfsft_single_train_step(
|
||||
|
||||
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
|
||||
# Skips when the current GPU has no seeded reference.
|
||||
# rtol above the harness default: the causal model compiles flex_attention
|
||||
# with max-autotune (required for Wan 1.3B's head config), and the
|
||||
# timing-based kernel selection is bimodal across L40S containers —
|
||||
# observed 3.2562 vs 3.5860 (10.13% apart) with identical code, straddling
|
||||
# the default 10%. 12% covers both winners; real wiring breakage (dead
|
||||
# grads, scale bugs) still lands far outside it.
|
||||
check_grad_norm_regression("test_wan_causal_dfsft",
|
||||
model.transformer,
|
||||
rtol=0.12)
|
||||
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
# FastVideo Runtime — Aggressive Implementation Plan
|
||||
|
||||
**Companion to** `design.md` (v19) and `design_summary.md` · **Stance:** this plan trades interface stability for
|
||||
speed. Where it deviates from design.md's conservative migration (§10), the deviation is flagged with **⚡**.
|
||||
design.md remains the architectural authority; this is the execution order.
|
||||
|
||||
---
|
||||
|
||||
## 1. Rules of engagement
|
||||
|
||||
**We break, freely and early:**
|
||||
|
||||
- The public Python API: `generate_video(**kwargs)` and `SamplingParam` are **deleted**, not deprecated.
|
||||
- Config schemas: `FastVideoArgs` (1,272 lines, 81 fields) stops being a public or threaded surface.
|
||||
- `fastvideo/api/compat.py` (651 lines): **deleted in M1** ⚡ (design.md §6.6 shrinks it monotonically to Phase 5 —
|
||||
that policy existed only to honor signatures we are now licensed to break).
|
||||
- CLI flags, YAML schemas, package layout, `fastvideo.api` exports, ComfyUI node params, every example.
|
||||
- In-repo dependents (`apps/dreamverse`, `comfyui/`, `examples/`, `scripts/`) get **fixed in the same PR train** —
|
||||
we own them; no deprecation period, no shims.
|
||||
|
||||
**We never break, at any speed:**
|
||||
|
||||
- **Numerics.** Bit-identical loop parity and SSIM gates are not "conservative" — they are the definition of
|
||||
correct. Aggression applies to interfaces, never to outputs.
|
||||
- **Model coverage** for the families that matter (tier list in §6 — the tail is a decision, not a casualty).
|
||||
- The frozen legacy `fastvideo/training/` stack (N2) and the bit-exact porting methodology (N3).
|
||||
- External users get **batched breakage**: all user-visible breaks land in at most two releases (R1 = request/config
|
||||
cut, R2 = engine default), each with a migration guide and a `fastvideo migrate` codemod — never a drip.
|
||||
|
||||
## 2. The sequencing argument (answering "fix omni request first, then separate the planes?")
|
||||
|
||||
**Yes to the first half. The second half should not be a project.** The three planes are not separated by moving
|
||||
code into plane-named directories — today's monolithic stages would just get reshuffled and then rewritten when
|
||||
loops invert. The planes are *born* from two cuts, and a third that is really a config change:
|
||||
|
||||
1. **The request-plane cut (M1)** — `OmniRequest` becomes the only currency crossing the boundary. Everything
|
||||
behind it is implementation. This is your "fix the omni input and request first," and it goes first because it
|
||||
is low-risk, it defines the vocabulary every later stage consumes, and it gets the user-facing pain over with
|
||||
while the codebase is still familiar.
|
||||
2. **Loop inversion (M2)** — this *is* the pipeline/execution plane separation. Once families expose
|
||||
`init/step/finalize` step bodies, something other than the family must own iteration; that owner is the
|
||||
executor, and the execution plane exists by construction. Before inversion there is nothing for an execution
|
||||
plane to schedule — "separating" it would be an empty directory.
|
||||
3. **The config cut (inside M1)** — the real coupling between planes today is `FastVideoArgs`: one 81-field object
|
||||
threaded through entrypoints, pipelines, stages, and executors, mixing deploy-time, model-time, and
|
||||
request-time concerns. Splitting it into `DeployConfig` / `ModelSpec` / `OmniRequest` (design.md §6.6's four
|
||||
layers) is the single highest-leverage "separation" action, and it's schema work, not architecture work.
|
||||
|
||||
So the order is: **M1 request+config cut → M2 loop inversion (planes now exist) → M3 engine on top.** Plane
|
||||
separation is the *outcome* of M1+M2, not a milestone.
|
||||
|
||||
## 3. Milestones
|
||||
|
||||
Timeline assumes 3–4 engineers on the runtime critical path. Overlap is deliberate; gates are not ⚡-able.
|
||||
|
||||
### M0 — Baselines, harness, enforcement (weeks 0–3, overlaps M1)
|
||||
|
||||
The license for everything aggressive afterward. Not skippable, not shrinkable. M0 does not block M1 (which
|
||||
changes no numerics) — the only hard rule is **no family's M2 migration starts before its baseline exists**.
|
||||
|
||||
- Merge the `feat/cosmos3-reasoning` chain (design.md sizes this alone at 2–3 engineer-months — it runs as its own
|
||||
track); seed SSIM references for the ~7 uncovered families.
|
||||
- ParityAligner v0: record/compare named taps on *current* pipelines (it must exist before anything changes).
|
||||
- **The enforcement package, on day one** (design.md §10 — the prior freeze was broken 19× for lack of exactly
|
||||
this): CI path gates (reject new `fastvideo/training/` files now; reject new `DenoisingStage` subclasses once the
|
||||
first M2 family lands), CODEOWNERS on the frozen and migrating paths, a named owner per milestone, and the
|
||||
inflow rule — new model families land on the new abstractions from the first Wan/Flux2 landing onward.
|
||||
- Announce the M1 freeze window for in-flight PRs touching `fastvideo/api/`, `fastvideo_args.py`, entrypoints.
|
||||
|
||||
*Gate: every tier-A family has a recorded SSIM + activation baseline; CI gates live.*
|
||||
|
||||
### M1 — The request-plane cut (weeks 1–4) → **breaking release R1**
|
||||
|
||||
The typed API is partway there: `VideoGenerator.generate(GenerationRequest)` is already the documented primary
|
||||
entrypoint (`generate_video` carries a deprecation warning), and `fastvideo/entrypoints/openai/` already serves
|
||||
`POST /v1/videos` and `POST /v1/images`. But the legacy path is still what's *used*: Dreamverse calls
|
||||
`generate_video(**kwargs)` (`apps/dreamverse/dreamverse/video_generation.py:508`), as do ComfyUI and most
|
||||
examples. M1 finishes the cut instead of bridging it:
|
||||
|
||||
- **`OmniRequest` / `OmniOutput` / `OmniEvent`**: evolve `api/schema.py`'s `GenerationRequest` in place — typed
|
||||
modality parts, `TaskType`, per-model `ModelOptions` registered blocks (formalizing the `api/matrixgame2.py`
|
||||
pattern), seeds/priority/streaming flags. `api/results.py`'s `Video*Event` types become `OmniEvent`.
|
||||
- **Config: four layers, one owner each** (§6.6): extract `DeployConfig` (placement, parallelism axes, memory/
|
||||
offload, compile, plugins) from the runtime third of `FastVideoArgs` + `EngineConfig`/`ParallelismConfig`;
|
||||
`ModelSpec` manifest v0 (manifest-first component resolution; today's name-detectors as fallback);
|
||||
`OmniRequest` absorbs every per-call field. CLI flags, OpenAI protocol models, and presets are **generated**
|
||||
from the schema.
|
||||
- **Delete** ⚡: `compat.py` (651), `sampling_param.py` (411), `generate_video()`, the `fastvideo.api` legacy
|
||||
exports, `FastVideoArgs` as a *public* type. Internally it survives as a boundary-constructed shim for as long
|
||||
as anything still receives it: migrated families drop it per-family in M2, but unmigrated tier-B stages
|
||||
(`LegacyPipelineNode`) and the frozen `training/` stack (whose `TrainingArgs` subclasses it) carry it until M6 —
|
||||
it dies as a type with the tail, not before.
|
||||
- **Fix in-train**: ComfyUI nodes (legacy-API callers), all `examples/` (~75 files, mostly mechanical),
|
||||
`scripts/`, docs. Ship `fastvideo migrate` (codemod: old kwargs/YAML → `OmniRequest`/`DeployConfig`).
|
||||
- Internals unchanged: `ForwardBatch` is built *from* `OmniRequest` at the boundary; the executor and stages are
|
||||
untouched in M1.
|
||||
|
||||
*Gate: all SSIM suites unchanged; Dreamverse + ComfyUI + examples green on the new surface; R1 notes + codemod
|
||||
published.*
|
||||
|
||||
### M2 — Loop inversion (weeks 4–10): the pipeline plane is born
|
||||
|
||||
- `DenoiseLoop` / `ARDecodeLoop` with `init/step/finalize`; runtime owns iteration; custom-step escape hatch from
|
||||
day one (the Cosmos3-port and self-forcing pattern is legitimate, §6.2.3).
|
||||
- **Family order** (each lands step body + policies, and **deletes its legacy stage code in the same PR** ⚡ —
|
||||
continuous deletion, no end-of-plan cliff): **Wan 2.1/2.2 + Flux2 first, jointly** — design.md's rationale
|
||||
stands: together they exercise CFG variants, expert routing, chunk-KV, and the image path, so the step-body
|
||||
contract freezes only after all four are exercised → Wan-causal (self-forcing student) → LTX-2 → HunyuanVideo →
|
||||
Stable Audio → remaining image families. Unmigrated families keep running via `LegacyPipelineNode`.
|
||||
- Policies: `CFGPolicy` (absorbs the 3 CFG copies), `AttnMetadataProvider`, `FlowShiftPolicy`, `PrecisionPolicy`.
|
||||
- Extension core lands with the loop (it's why the loop is being rebuilt): observer bus, ParityAligner promoted to
|
||||
per-request observer, Profiler/NaNWatch, and **cache-dit as the first interceptor** (retiring `enable_teacache`).
|
||||
- `forward_context.py` off the *migrated* inference path (194 references across ~68 files today: ~8 importer files
|
||||
in frozen `training/`, the rest spread across train/ models, tier-B inference stages, quantization, and tests);
|
||||
the module survives as a shim for frozen `training/` **and unmigrated tier-B stages** until M6 — what M2
|
||||
guarantees is that no migrated family and no new code touches it.
|
||||
- **`train/` migrates per-family, immediately behind inference**: DMD2 and the landed DiffusionNFT (#1450) adopt
|
||||
the shared step functions as each family's body lands — `rl/common/sampling.py`'s loop is deleted, #1396
|
||||
grad-norm refs extended to the migrated methods (RL included).
|
||||
|
||||
*Gate, per family: old-vs-new loop bit-identical (ParityAligner) + SSIM + a recorded loop-overhead / batch-of-1
|
||||
latency measurement (the baseline M3 gates against); for train/: seeded rollout latents identical, reward metrics
|
||||
- grad-norms neutral. No family is ever dual-maintained.*
|
||||
|
||||
### M3 — Execution plane: engine + scheduler (weeks 8–14, overlaps M2) → **breaking release R2**
|
||||
|
||||
- `AsyncEngine` (queue, admission, cancellation-as-common-path, failure isolation); offline `VideoGenerator` keeps
|
||||
its name, becomes a thin sync wrapper that can bypass the queue.
|
||||
- `StepScheduler` v0: multiplexes denoise steps across requests in a pool; budget currency = **predicted GPU-time**
|
||||
from a calibrated per-(model, phase, shape) cost table (the cost *model* matures later; the currency is right
|
||||
from day one). Carries the `ARDecodeLoop` contract; AR batching itself waits for its workload (N5).
|
||||
- CacheManager v0: per-request chunk-KV slabs behind `KVHandle`; CFG-parallel axis (2-branch in practice).
|
||||
- **Dynamo stock worker** (registration, health/drain, cost metrics), retiring the locked
|
||||
`dynamo/examples/diffusers/worker.py` pattern.
|
||||
- **Dreamverse hard-cut** (per design.md Phase 2; the aggressive delta is doing it in one PR): `gpu_pool.py`,
|
||||
queue, warmup, and stream relay deleted and replaced by engine-client calls; the duty-cycle concurrency study
|
||||
runs on the result.
|
||||
- Colocated weight-sync RPC + component-granular sleep/wake + `RolloutClient` (engine-client RL mode for #1450).
|
||||
|
||||
*Gate: serving load tests; batch-of-1 latency regression ≤ 2% vs the M2-recorded measurement; Dreamverse
|
||||
single-session parity; RL engine-client seeded final-latent parity vs in-process; deploys under stock Dynamo.*
|
||||
|
||||
### M4 — Graphs, parallelism, multi-session (weeks 14–20)
|
||||
|
||||
- `PipelineSpec` graph IR: per-family pipeline classes shrink to **spec + step body + policies**
|
||||
(`create_pipeline_stages()` retires); LTX-2 and Hunyuan15+SR land as real fan-out graphs.
|
||||
- Role pools + connectors (port `multimodal_gen`'s disagg state machine); declarative stacked-parallelism axes
|
||||
compiled to DeviceMesh; general cross-mesh `WeightSyncPlan`.
|
||||
- ComfyUI workflow→spec compiler MVP (tier-1 ~20-node vocabulary) + weight/adapter fleet cache.
|
||||
|
||||
*Gate (design.md Phase 3's, in full): ≥2 Dreamverse sessions/GPU on the recorded duty-cycle trace, p95 within SLO
|
||||
— this is also where the loop-inversion **falsifier** is evaluated (see §7); LTX-2 A/V full-fan-out end-to-end;
|
||||
disaggregated-vs-colocated throughput benchmark; CPU-only topology validation suite; ComfyUI tier-1 workflows
|
||||
compile and run with equivalence reports; spec-built pipelines SSIM-identical to M2 loop versions.*
|
||||
|
||||
### M5 — Omni/MoT native + RL hardening (weeks 20–30)
|
||||
|
||||
- Cosmos3 re-port onto specs: packed factored sequences, dual-pathway attention, reasoner paged KV, joint denoise,
|
||||
world-model `ChunkRollout`; `/v1/chat/completions`; AR continuous batching arrives **with** this workload (N5).
|
||||
- Consistency ladder enforced end-to-end: C1 default in CI, C2 bitwise mode for goldens, Behavior Record opt-in;
|
||||
first GRPO-class method lands on the engine-client rollout path (log-prob drift becomes the gated metric).
|
||||
|
||||
*Gate: Cosmos3 150-test parity suite on the new runtime; reasoner pool efficiency — tokens/s/GPU at target
|
||||
concurrent denoise throughput, with the ≥10×-vs-re-prefill sanity floor; C1 drift ≈ 0 on a Wan RL run with the
|
||||
drift dashboard live.*
|
||||
|
||||
### M6 — The tail and the precondition (week 30+)
|
||||
|
||||
Continuous deletion (M1/M2) shrinks the final phase but does not eliminate it: what remains by M5 is the tier-B
|
||||
tail on `LegacyPipelineNode` and the frozen `training/` stack — which is a *live consumer* of
|
||||
`ComposedPipelineBase` and `forward_context`, so its retirement is the precondition, exactly as design.md Phase 5
|
||||
states. M6 = execute the §6 tail decision (migrate or deprecate each tier-B family), retire `training/` per the
|
||||
checklist, then delete `ComposedPipelineBase`, the legacy `DenoisingStage`, `forward_context.py`,
|
||||
`FastVideoArgs`/`TrainingArgs`, and `RayDistributedExecutor` together. **4 loop copies → 1.**
|
||||
|
||||
## 4. Breakage manifest (user-visible)
|
||||
|
||||
| Release | What breaks | Replacement | Aid |
|
||||
|---|---|---|---|
|
||||
| **R1** (M1) | `generate_video(prompt, **kwargs)`, `SamplingParam`, `FastVideoArgs` as public type, `fastvideo.api` legacy exports, CLI flag names, YAML config schema, streaming event types (`Video*Event` → `OmniEvent`, `schema_version`'d from day one) | `VideoGenerator.generate(OmniRequest)`, `DeployConfig`, generated CLI/protocol, `OmniEvent` | `fastvideo migrate` codemod, migration guide, R0 pinned |
|
||||
| **R2** (M3) | Default execution path becomes the engine (offline bypass preserved); server lifecycle (queue/admission semantics, job states) | `AsyncEngine` | guide; `OmniEvent` schema unchanged from R1 |
|
||||
| after R2 | nothing user-visible — M4/M5 are additive | — | — |
|
||||
|
||||
## 5. Deviations from design.md §10, stated honestly
|
||||
|
||||
| design.md | this plan | why it's safe now |
|
||||
|---|---|---|
|
||||
| Phase 0 keeps `VideoGenerator`/CLI signatures; `compat.py` shrinks to Phase 5 | M1 breaks signatures, deletes `compat.py` ⚡ | the only argument for the shim was signature stability — explicitly revoked |
|
||||
| Legacy code deleted at Phase 5 | per-family deletion at parity, M2 onward ⚡ | parity gate is per-family anyway; carrying dead code to a final phase only invites the 19×-broken-freeze failure mode |
|
||||
| Phases strictly sequential | M2/M3 overlap ⚡ | the engine consumes step bodies, not finished families; the step-body contract freezes at the Wan+Flux2 landing |
|
||||
| Phases −1 through 4 sized at 36–54 engineer-months | ~21–28 engineer-months (3–4 eng × 30 wks) ⚡ | the delta is real deleted work — no compat maintenance, no adapter upkeep, no dual-stack carry — plus M2/M3 overlap; treat 30 weeks as the aggressive case and 36–40 as the planning case |
|
||||
| Unchanged | parity/SSIM gates (restored in full at every milestone), enforcement package (CI path gates, CODEOWNERS, inflow rule — now at M0), train/RL migration timing (design.md Phase 1 already migrates NFT), Dreamverse hard-cut (Phase 2 already prescribes it), N2/N3/N5, cost-model currency, Dynamo asks + fallbacks, schema versioning | aggression budget is spent on interfaces only |
|
||||
|
||||
## 6. Decisions needed before M0
|
||||
|
||||
1. **Tier the model zoo.** Tier A (migrated, coverage guaranteed): Wan 2.1/2.2, Wan-causal/self-forcing, LTX-2,
|
||||
Flux2, HunyuanVideo, Stable Audio, Cosmos3 (contingent on the M0 merge — it is not on `main` today), image
|
||||
families. Tier B (runs on `LegacyPipelineNode` until someone claims it, candidate for deprecation at M6):
|
||||
gen3c, matrixgame2/3, longcat, the rest. **Approve or edit the split** — it bounds M2.
|
||||
2. **Release framing.** R1 as `v0.3.0` (pre-1.0 semantics, loud notes) vs holding breaks for a `v1.0` story.
|
||||
Recommendation: `v0.3.0` now — waiting taxes every milestone.
|
||||
3. **Freeze windows.** M1 freezes `api/`/args/entrypoints PRs ~2 weeks; M2 freezes per-family stage PRs while that
|
||||
family migrates (days each). Needs maintainer sign-off.
|
||||
4. **Staffing.** Critical path is M2's per-family step bodies — parallelizable per family after the Wan+Flux2
|
||||
reference lands. 3–4 engineers ≈ 30 weeks to M5 in the aggressive case (design.md's own sizing implies 36–40
|
||||
weeks at the same staffing — see §5); 2 engineers ≈ stretch ~1.5×. The Cosmos3-chain merge (M0) is its own
|
||||
2–3 engineer-month track and should be staffed separately from the runtime critical path.
|
||||
|
||||
## 7. Risks specific to the aggressive posture
|
||||
|
||||
- **In-flight PR collisions** with layout/schema moves → freeze windows (above) + landing schema cuts at
|
||||
milestone *starts*, not ends.
|
||||
- **Community churn at R1** (ComfyUI users, script users) → codemod covers the mechanical 90%; the 10% that isn't
|
||||
mechanical (kwargs with changed semantics) is enumerated in the guide; previous version stays pinned and
|
||||
installable.
|
||||
- **Parity harness becomes the bottleneck** — every aggressive deletion is licensed by it. Mitigation: it is the
|
||||
*first* deliverable (M0), and per-family migration PRs are template-driven (record → port → compare → delete).
|
||||
- **Overlap risk (M2/M3)**: the engine team building against a moving step-body contract → the contract
|
||||
(`init/step/finalize` + `StepResult`) freezes at the *first* family (Wan), enforced by the same schema-version
|
||||
discipline as external surfaces.
|
||||
- **The known unknown**: loop inversion at scheduler granularity has no production precedent (design.md §1). The
|
||||
falsifier stands, on design.md §11.6's schedule: the M3 duty-cycle study *publishes the targets*; the falsifier
|
||||
is **evaluated at the M4 gate** — if step-level multiplexing can't beat request-level serialization on real
|
||||
Dreamverse traces, StepScheduler retreats to request-level dispatch and the loop contract keeps only its
|
||||
streaming/preemption seams, with no family code changing — step bodies and the M1/M2 cuts retain full value.
|
||||
+3
-2
@@ -223,8 +223,9 @@ follow_imports = "silent"
|
||||
skip = "./data,./wandb,apps/fastvideo_studio/package-lock.json,apps/performance_dashboard/frontend/package-lock.json,*/_vendored/*"
|
||||
# "tread" matches daVinci-MagiHuman's acronym "TReAD" (Token Routing and
|
||||
# Early Drop). codespell lowercases ignore-words entries, so the single
|
||||
# lowercase form silences all case variants.
|
||||
ignore-words-list = "tread,passt"
|
||||
# lowercase form silences all case variants. "mot" = Mixture-of-Transformers
|
||||
# (MoT); "clen" = a Content-Length local; "te" = a text-embeds local.
|
||||
ignore-words-list = "tread,passt,mot,clen,te"
|
||||
|
||||
[tool.ruff]
|
||||
# Allow lines to be as long as 120.
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
# Adversarial Review of `design.md` (v12) — FastVideo Next-Generation Inference Runtime
|
||||
|
||||
**Date:** 2026-06-11
|
||||
**Method:** Multi-agent adversarial review. 9 fact-check agents verified 70 concrete claims against the repo, the local reference checkouts (`cosmos-framework/`, `dynamo/`, `vllm-omni/`, `~/sglang`, `~/vllm`, `~/miles`, `~/verl-omni`, `~/diffusers`, `~/torchtitan`, `~/xDiT`, `~/ComfyUI`, `~/sglang-omni`, `~/cosmos-rl`), and GitHub. 9 attack lenses (abstractions, scheduler/perf, memory/cache, training/RL, strategy, migration, internal consistency, omissions, external borrowings) plus a completeness critic raised 70 findings; every finding went to a refute-by-default verifier. 36 findings were refuted; this document contains only the 34 that survived (1 critical, 26 major — consolidated below where lenses converged — 7 minor), plus fact-check corrections.
|
||||
|
||||
---
|
||||
|
||||
## Verdict
|
||||
|
||||
The architecture survives its strongest attacks — loop inversion's expressibility, the typed-state hybrid, the N1/N5 scope discipline, the clean-room GPL posture, and the C2-for-batch-1-video argument all held under refutation attempts. What does not survive is:
|
||||
|
||||
1. **The migration plan**, which consumes its own substrate two phases before building it and rests on a "frozen legacy stack" premise this repo has already empirically falsified.
|
||||
2. **Two load-bearing factual errors** about reference systems (vLLM's BlockPool page sizes, diffusers' loop ownership) that each drove a recorded design decision.
|
||||
3. **A family of undesigned failure/memory/trust paths** that the multiplexing bet itself creates. One is critical.
|
||||
|
||||
---
|
||||
|
||||
## Critical
|
||||
|
||||
### C1. No failure-isolation or cancellation semantics for the multiplexed pool — the blast-radius problem the architecture itself creates
|
||||
**Where:** §6.3.1; absent from §12.
|
||||
|
||||
Today one request per pool means one request's CUDA error is its own problem. Step-multiplexing changes the failure class categorically: a mid-step OOM/illegal-access/NaN from one request poisons the CUDA context and desyncs in-flight NCCL collectives for *every* co-scheduled tenant on the pool, including resident Dreamverse session caches. The doc designs none of the machinery: no SPMD-consistent abort broadcast (the dual of its scheduling broadcast), no request-fatal vs pool-fatal classification, no pool re-init + cache-invalidation policy, no partial-artifact semantics for fan-out graphs. "OOM" and request cancellation appear nowhere in 1799 lines; "abort" appears once (RL stragglers).
|
||||
|
||||
Ordinary cancellation is also missing — and vibe directing makes abandoning in-flight generations the *common* path. Worse, Phase 2 retires Dreamverse's `gpu_pool.py`, which today has a working sentinel-fd worker-death watch (`gpu_pool.py:542-586`), into engine-client calls — a reliability regression for the flagship customer if the gate ships as written. vLLM v1, the doc's own scheduler template, needed first-class machinery for exactly this (`ENGINE_CORE_DEAD`, `EngineDeadError`, `abort_requests`). Risk 4 covers only scheduling-decision divergence; the long-job-resilience known-gap is single-job-framed.
|
||||
|
||||
The abort path shapes the StepScheduler loop, the worker RPC surface, and CacheManager handle lifetimes — it must be designed *with* Phase 2, and by the doc's own standard ("absence reads as a decision"), this absence is an oversight.
|
||||
|
||||
---
|
||||
|
||||
## Major — reference-system misreads that drove recorded decisions
|
||||
|
||||
### M1. The single-BlockPool CacheManager rests on a property vLLM explicitly does not have: per-group page sizes
|
||||
**Where:** §6.3.2 lines 555-559.
|
||||
|
||||
The sentence asserts two mutually exclusive properties. vLLM's one-pool/no-fragmentation guarantee exists *only because* physical bytes-per-block are uniform across all groups: `kv_cache_utils.py` asserts a single page size (`get_uniform_page_size`), and its docstring says verbatim that breaking this "is non-trivial due to memory fragmentation concerns." Groups differ only in tokens-per-block at equal byte size; the unification mechanism inflates the smaller group's `block_size`.
|
||||
|
||||
Apply that to FastVideo's groups: a text-KV page (~64 KB/layer) vs a latent-frame slab (9.6–32 MB/layer for 1.3B/14B causal Wan) is a 150–500× ratio — unification means a 500-token reasoner prompt strands a multi-MB slab per layer-group. The one vLLM path with multiple page sizes (DeepseekV4) statically partitions capacity at startup over a single global block-id free list, which is harmless when group demand is token-coupled (every token passes through all layer groups) but wasteful exactly when demand is workload-decoupled — FastVideo's regime, where text-KV and chunk-KV demand vary independently with request mix.
|
||||
|
||||
Since this misread is what reversed the two-pool sketch (recorded at line 280), the decision rests on a false premise: either chunk-KV stays uniformly fine-paged (losing the slab semantics the MoT "falls out naturally" story depends on), or the two-pool design returns and needs its own fragmentation/deadlock argument.
|
||||
|
||||
### M2. diffusers Modular is not loop inversion — the "strongest external validation" of the keystone doesn't validate it
|
||||
**Where:** §5 line 277, §6.2.3 lines 426-428.
|
||||
|
||||
`LoopSequentialPipelineBlocks.__call__` raises `NotImplementedError`; every concrete family hand-writes `for i, t in enumerate(timesteps)` inside its own blocking wrapper (`wan/denoise.py:434`, `stable_diffusion_xl/denoise.py:701` — SDXL ships four such wrappers, the subclass forest again). The iteration is block-owned, invisible to any runtime — no init/step/finalize, no external driver, none of the properties §6.2.2 says inversion exists for (scheduling, interleaving, preemption, streaming, fair sharing). In scheduling terms it is the current `DenoisingStage` with a refactored body — i.e., it validates the Guiders/policy pillar but as evidence for inversion it is *equally consistent with the alternative the design rejects* ("keep loops in stages, make bodies pluggable"). The class also carries an explicit experimental warning.
|
||||
|
||||
Consequence: no surveyed system — vLLM, sglang, multimodal_gen, diffusers — implements runtime-owned diffusion iteration at scheduler granularity. Loop inversion is the design's most novel element with zero production precedent, and risk 3 (which admits novelty only for the hybrid AR+denoise slice) should say so instead of borrowing validation the reference doesn't provide.
|
||||
|
||||
### M3. Cost-currency scheduling drops the memory half of vLLM's admission — and memory is never a scheduling resource anywhere in the design
|
||||
**Where:** §6.3.1 (lines 476-547), §6.3.2; two lenses converged here.
|
||||
|
||||
vLLM's token budget is not a prediction — it is an exact cap checked *in the same loop as memory admission* (`allocate_slots` per request, preempt on allocation failure; activation memory separately bounded by a profiled worst case). The design takes the accounting structure, swaps the currency for a *forecast* (predicted GPU-time), and drops the memory dimension entirely: latents, conditioning sets, CFG duplicates, and activation peaks live in `RequestState`, explicitly outside the CacheManager, and nothing bounds how many concurrent LoopStates a pool admits — for a workload the doc itself calls memory-bound (line 499). Two items that each fit alone can jointly OOM, and a GPU-seconds currency cannot see it; combined with C1, that OOM is a pool-wide event. "Preemption only at step boundaries" never defines what happens to a preempted request's multi-GB resident state (offload? drop-and-resume-from-LoopState? — different economics from KV recompute).
|
||||
|
||||
Related internal contradiction, verified: cost is "static and known at admission... a table lookup" (line 538), but the same cost model is cache-dit-aware (line 520) — DBCache skip decisions are runtime data-dependent residual comparisons, unknowable at admission.
|
||||
|
||||
**Fix:** the budget needs a memory axis (resident-state + peak-activation per schedulable item), admission needs a memory planner over RequestState, and preemption semantics must be specified. The Phase-2 "≥2 sessions per GPU" gate rests on unaccounted memory until then.
|
||||
|
||||
### M4. Punica cannot express ComfyUI LoRA semantics
|
||||
**Where:** §9.4 lines 1376-1380 (also §6.3.2 lines 575-579). *Verifier rated minor-to-major; grouped here with the borrowings cluster.*
|
||||
|
||||
vLLM's `LoRARequest` carries one `lora_int_id` and no strength field; scaling is baked into `lora_b` at registration; the Punica wrapper maps one adapter index per token. ComfyUI traffic — the workload §9.4 names — is N stacked LoRAs per request with continuous user-set `strength_model` *and* `strength_clip`, routinely tweaked per generation. Pushing that through Punica means registering each (ordered-set, strengths) tuple as a synthetic concatenated adapter: near-zero cache-hit rate across strength tweaks, registration churn in the stacked GPU weight slots, and concatenated ranks colliding with `max_lora_rank`. "Strictly better than hot-swap-only" is unsupported without a composition layer that doesn't exist anywhere, including in vLLM.
|
||||
|
||||
---
|
||||
|
||||
## Major — execution-plane gaps
|
||||
|
||||
### M5. MoT mode multiplexing has no parallelism answer
|
||||
**Where:** §6.3.1 lines 502-503 vs §6.3.4; Phase 4 gate.
|
||||
|
||||
The "mode multiplexer" claim assumes both loop types share one static pool layout (`parallel: [dp, cfg, sp, tp]`), but their optimal layouts are disjoint: denoise wants SP+CFG; AR decode is sequence-length-1 — SP has nothing to shard and CFG doesn't exist. On a `[cfg(2), sp(4)]` 8-GPU pool the reasoner either runs replicated (1/8 useful work, paged KV duplicated 8×) or needs TP — and TP-everywhere regresses the bread-and-butter denoise workload on the flagship pool. Per-phase re-layout of the same resident weights is not expressible in the §6.3.4 spec (one static stack per pool), and resharding machinery exists only for train↔rollout weight sync (§8.6). §6.3.1's own jumbo-step mitigation (split cost classes across pools) is structurally unavailable for MoT — AR steps and denoise steps are the same weights — so concurrent reasoner token latency is gated by indivisible 50–500 ms denoise steps.
|
||||
|
||||
A workable resolution exists (AR continuous batching data-parallel across the cfg×sp weight-replica axes onto TP subgroups, plus §6.3.3 per-pathway TP, plus routing pure-REASON traffic to differently-shaped pools), but the doc never states one, and the Phase-4 gate ("reasoner ≥10× faster than re-prefill") is measured against an O(n²) strawman baseline that certifies nothing about pool efficiency. Risk 3's "prototype early in Phase 4" defers a *design contradiction*, not an implementation unknown.
|
||||
|
||||
### M6. The engine's own multi-node story is unstated, and the Ray executor silently disappears
|
||||
**Where:** N1 line 143, §6.3.5 line 673, §6.0 line 304.
|
||||
|
||||
Whether one worker pool may span nodes is a load-bearing decision the doc never makes — Dynamo routes *between* workers; it does not own the NCCL mesh *inside* one. If pools are single-node by fiat, SP degree caps at ~8 GPUs, directly contradicting line 543's jumbo-step mitigation ("shrink jumbo step wall-time with SP"), capping MoT model scale — and `RayDistributedExecutor`, today's shipping multi-node path, is silently dropped: it appears in the §3.1 diagram and then never again in §6, §10, §11, or §12 (violating the plan's own "every phase deletes or freezes what it replaces" discipline). If pools may span nodes, the engine owns cross-node collective bring-up, NCCL-timeout fault domains, and a multi-node health/drain contract — none designed, and C1's recovery problem becomes a multi-node recovery problem. Either answer changes Phase 2/3 scope. "Node-group" appears once, undefined.
|
||||
|
||||
### M7. Policies carry per-request mutable state with no state-scoping contract — and the doc contradicts itself on when policies are resolved
|
||||
**Where:** §6.2.3 lines 412-417 vs §6.2.2 line 387 vs risk 2 line 1621; §6.4 lines 837-840.
|
||||
|
||||
The doc says policies are resolved at pipeline build (lines 412-413; risk 2: "resolved to bound methods at build time") *and* in `DenoiseLoop.init` (line 387) — a genuine contradiction on a load-bearing contract. It matters: AdaptiveGateCFG — a named CFGPolicy example and the Wan2.2 worked-example default — is per-request mutable state in shipped code (`denoising.py:338-343, 507-551`: `delta_cached`, `delta_cached_model_id`, gate counters). Build-time-resolved singletons mean request A's cached CFG delta gets applied to request B the moment Phase 2 interleaving lands — silent quality corruption no Phase-2 gate (load tests, latency budget) can catch. This is the *exact* failure mode §6.4 cites to justify interceptor state scoping ("silently corrupts under concurrent requests") — the contract was designed for the plugin tier and forgotten for the policy tier, which sits on a hotter path. Cheap fix (policy state into LoopState, same as plugins), but it must be in the spec.
|
||||
|
||||
### M8. The six-policy taxonomy does not factor the shipped step bodies — no step skeleton or cross-policy interaction contract is defined
|
||||
**Where:** §6.2.3 (policy table, line 424 claim); §6.2.2 lines 386-389; §6.4 lines 837-844.
|
||||
|
||||
The proposed step is three phases (forward → CFG combine → scheduler step); the shipped loops need ~six, with dependencies that cross policy boundaries. Verified examples:
|
||||
|
||||
- **Cosmos** conditioning-frame injection consumes the *sampler's* EDM coefficients, applies per-CFG-branch both pre-forward (input mix) and post-forward (x0 clamp), and the CFG combine runs in x0 space — ConditioningInjector × Sampler × CFGPolicy interleaved inside each branch, unownable by any one of them (`denoising.py:845-933`).
|
||||
- **TI2V** clamps latents *after* `scheduler.step` — a post-step constraint with no policy slot (`denoising.py:570-573`).
|
||||
- **Cosmos2.5** builds per-frame timestep vectors with a conditioned-frame override and re-clamps GT every step pre-forward.
|
||||
- **CausalDMD** renoises between steps choosing `add_noise` vs `add_noise_high` by expert boundary — Sampler × ExpertRouting (`causal_denoising.py:268-301`).
|
||||
- **AdaptiveGateCFG** must observe ExpertRouting's switch to invalidate its delta (today an inline `id(current_model)` check) — yet no channel for one policy to observe another is defined anywhere.
|
||||
- **LTX2** guidance is 1–4 runtime-decided passes whose branches alter the network via forward kwargs (`skip_cross_modal_attn`, `skip_video/audio_self_attn_blocks`) — colliding with BlockInterceptor's domain in a way the "two block-skippers conflict" pre-flight check cannot see, and breaking §6.4's per-CFG-branch state scoping, which assumes a fixed cond/uncond branch vocabulary (`ltx2_denoising.py:503-605, 620-631`).
|
||||
|
||||
None of the six policies covers prediction-space conversion, per-token timestep construction, post-step latent constraints, inter-step renoising, or chunk-boundary refresh. The fix is not abandoning policies — the Sampler registry is the natural home for some of this, and composition still strips the duplicated offload/attn-metadata/autocast/trajectory plumbing — but the design needs the fixed step skeleton with ordered, typed extension points and an explicit policy-interaction contract, worked through Cosmos2.5 and LTX2 *in the doc*. Until then, "a new model contributes policies + a graph spec; it does not edit shared loop code" (line 424) is asserted, not demonstrated.
|
||||
|
||||
### M9. OmniRequest cannot parameterize multi-loop graphs
|
||||
**Where:** §6.1 lines 318-334; §6.6 line 905; worked examples (c)(d) lines 943-950.
|
||||
|
||||
One flat `SamplingParams` + one flat `DiffusionParams` per request, while the design's own flagship examples are multi-loop graphs needing per-node knobs: LTX-2's refine loop has its own step count and guidance scale *today* as first-class fields (`fastvideo_args.py:204-205`, threaded through `compat.py` and `dynamo/examples/diffusers/worker.py:201-203`); a thinker and talker need different `max_tokens`/`temperature`/`stop`. No request→graph-node parameter binding is defined anywhere; the only escape hatch is line 905's per-model `ModelOptions` blocks — i.e., the `ltx2_*` field-leakage pattern the doc indicts at P3, with a type wrapper, regenerated into the OpenAI/CLI views that derive from the request schema (line 907). Needs a real decision — parameters keyed by graph-node id, or per-node override blocks validated against the PipelineSpec — made in Phase 0, because that schema ships first and external consumers build against it.
|
||||
|
||||
---
|
||||
|
||||
## Major — caches and weights
|
||||
|
||||
### M10. No feature-cache invalidation story under LoRA hot-swap — te-LoRAs make the embedding cache serve stale embeddings in the workflow cloud
|
||||
**Where:** §6.3.2 lines 570-574 vs §9.4 lines 1349-1380.
|
||||
|
||||
The only invalidation rule in the document is RL `update_weights` → `reset()`. But ComfyUI-grade LoRAs routinely patch the *text encoder* alongside the DiT (`comfy/lora.py` maintains `lora_te/lora_te1/lora_te2` key maps; `load_lora_for_models` takes a separate `strength_clip`), so a content-hash-keyed embedding cache returns embeddings computed under the wrong adapter state the moment two workflows share a prompt but differ in te-LoRA stacks — silent wrong output in the exact product (§9.4 "exact mode") whose trust claim is reproducibility. §11.8 even makes cross-request embedding reuse load-bearing as the radix-cache substitute. And once Punica-style batched multi-LoRA lands, requests with different adapter stacks coexist concurrently on one pool, so the cache must be key-*partitioned* by (encoder identity × adapter set × strengths), not flushed — a different design from the `EncoderCacheManager` reset() semantics being adopted, which come from a world where encoders are never patched per request. The key schema needs a weight-state epoch / adapter-set hash as a mandatory component, decided before Phase 3.
|
||||
|
||||
### M11. Checkpoint/LoRA patching mutates pool-shared weights — a pool-quiescing barrier the StepScheduler has no vocabulary for
|
||||
**Where:** §9.4 lines 1371-1380 vs §6.3.1 and §6.0 line 299.
|
||||
|
||||
Components are "one resident copy per worker pool"; patch/unpatch mutates that copy, which is global to every loop interleaved on the pool — yet step-interleaving is the engine's core Phase-2 value. Two interleaved loops requiring different patch states cannot coexist, so every cross-group transition is a drain barrier: finish in-flight steps, apply/undo `W += scale·BA` across 14–28 GB shard-consistently across TP/SP ranks (ComfyUI keeps weight backups for the undo — 2× weight memory or a CPU→GPU restore at PCIe seconds), re-admit. Under workflow-cloud traffic (long-tail checkpoints, per-request adapter stacks), transition frequency is the whole game — and the §6.3.1 cost model (lines 516-521) has no weight-state-transition term, no notion of weight state as schedulable state, and no quiesce-vs-queue policy, even though transition cost is exactly what A1 checkpoint-affinity routing must weigh. The §8.6 safe-point-swap pattern shows the doc knows the shape but never applies it here. §9.4 calls this "the one real new subsystem"; §12 carries no risk entry for it.
|
||||
|
||||
---
|
||||
|
||||
## Major — training/RL
|
||||
|
||||
### M12. "Step bodies are plain tensor programs, so autograd composes" is contradicted by the distillation code the substrate must absorb
|
||||
**Where:** §8.2 lines 1016-1019; §6.2.2; §6.3.2.
|
||||
|
||||
Self-forcing does not "drive `DenoiseLoop.step`": its rollout samples per-block exit indices broadcast across ranks, runs no-grad steps to the exit, runs exactly *one* grad-enabled forward, then a separate no-grad `store_kv=True` context-caching pass with context noise, gated by `start_gradient_frame`. None of this fits `init/step/finalize` + `StepResult(done, emit)` without grad-gating flags, per-step cache-write control, and per-block exit policies — training-only surface in substrate code, or the method keeps its own loop and the "3 copies → 1" dedup claim dies for the hardest case. The KV path needs grad/AC-aware semantics the engine pool lacks: today's causal model snapshots KV indices whenever `torch.is_grad_enabled()` so activation-checkpoint recompute doesn't double-advance the cache (`wan_causal.py:119-120,405-431`), and never recycles blocks mid-rollout — while §6.3.2 specs vLLM-style out-of-window block recycling, and §8.5's own profile taxonomy says "training forward … *no caches*," showing the grad+KV case was never designed. §8.3 explicitly stakes the architecture on ChunkKVPool serving self-forcing training.
|
||||
|
||||
(Note: the related forward-context-backward attack was refuted — the Phase-1 retirement of the global plus explicit metadata passing *helps* autograd composition. The surviving residue is the grad-window/cache-mode design above.)
|
||||
|
||||
### M13. Behavior Record cost is understated ~1.5 orders of magnitude for its own flagship case (MoE diffusion)
|
||||
**Where:** §8.5 lines 1156-1160; §5 miles row line 1088.
|
||||
|
||||
The miles ~60 MB/sample figure is per-token routing, one forward per generated token. Diffusion re-routes the *entire packed sequence at every denoise step, twice under CFG*: the record is steps × CFG × tokens × MoE-layers × top_k. For a Cosmos3-class request (Qwen3-VL-MoE config: 60 experts, top_k 4, ~24 sparse layers via `decoder_sparse_step=1`, ~50K packed tokens, 35-50 steps × 2 branches) that is ~1.3–1.9 GB/sample int32 — ~20–30 GB per 16-sample GRPO group, before latents. "Cheap because trajectory capture is already an OutputSpec feature" conflates plumbing cost with byte cost; at these sizes the Record forces a buffering/transport/storage design (GB-scale trajectories through connectors from disaggregated rollout fleets) that appears nowhere — not in §8.7's TrajectoryBuffer, not in §12, not in the known-gaps list. (The RNG-draws sub-claim was refuted: seeded generators in a single shared loop reproduce draws; uint8 expert IDs also cut 4×. The routing-record problem stands.)
|
||||
|
||||
### M14. The omni-RL pilot is a Phase-4 deliverable with no objective design
|
||||
**Where:** §8.7 lines 1236-1240; §10 Phase 4.
|
||||
|
||||
The section establishes *expressibility* (one trajectory, two segment types — true, and a real structural advantage over engine-per-stage stacks) and quietly upgrades it to a deliverable without posing the algorithm problem:
|
||||
|
||||
- **Scale mismatch:** token log-probs are O(1–10) nats over 10²–10³ tokens; per-step diffusion SDE log-probs are Gaussian densities over 10⁶–10⁷ latent dims — any joint clipped-ratio objective needs principled per-segment normalization that none of the cited recipes (FlowGRPO/DanceGRPO/NFT/AIPO/GSPO) provides; get it wrong and one modality silently dominates the shared trunk.
|
||||
- **Credit assignment:** the reasoner influences video reward only through *sampled discrete tokens* re-entering as conditioning — a non-differentiable boundary, so token segments get sparse trajectory-level REINFORCE signal while denoise segments get dense per-step ratios, both updating shared attention-trunk weights, with no interference analysis.
|
||||
- **Reasoning regression:** RL-updating the und pathway on video-reward-correlated signal risks degrading its reasoning; reference-model KL anchoring for hybrid episodes is never mentioned.
|
||||
|
||||
The entire treatment is the phrase "optimized with mixed objectives," and §12's 15 open questions contain nothing on it — for the capability marketed as "the capability nobody else has." Either it gets an algorithm sketch and an open-question entry with an owner, or the Phase-4 item should be demoted from "pilot" to "trajectory capture demonstrated."
|
||||
|
||||
---
|
||||
|
||||
## Major — the migration plan (the weakest section)
|
||||
|
||||
### M15. The "frozen legacy stack" premise is empirically false in this very repo
|
||||
**Where:** lines 5, 110, 1026; §11.4; risk 5.
|
||||
|
||||
The anti-third-stack defense is a declared freeze plus intent to delete — and this repo has already run that experiment and it failed within weeks. Verified from git: `fastvideo/train/` landed 2026-03-09 (#1159); since then **19 commits modified the "frozen" `fastvideo/training/`**, including a *brand-new* `cosmos2_5_training_pipeline.py` added to the legacy stack on 2026-05-11 (#1227) — **nine days after `training/AGENTS.md` explicitly forbade adding new models there**, and eleven days after the same model landed in `train/` (#1224). World-model training (#1179) and LongCat finetuning (#1244) also landed in the frozen stack in May; EMA bugfixes as recently as June 8-9; `AGENTS.md` still calls `training/` "authoritative for shipped models."
|
||||
|
||||
The doc invokes the training/-vs-train/ "lesson" but proposes nothing mechanically different from what was tried: no CI gate rejecting new files under legacy paths, no codeowner veto, no named owner per family, no calendar date for Phase 5. "Phase 5 is a scheduled deletion, not an aspiration" (risk 5) — but nothing in the document is scheduled. Under the same model-port pressure that broke the training/ freeze (measurably higher on the inference side), this freeze breaks the same way. Name the enforcement mechanism that did not exist last time, or the deprecation commitment is the prior failure restated with more confidence.
|
||||
|
||||
### M16. Phase dependency inversion: Phases 1–2 consume the substrate Phase 4 builds
|
||||
**Where:** §10 lines 1406-1446 vs §6.3.1 lines 487-489, §6.3.2; three lenses converged on this.
|
||||
|
||||
Phase 1 migrates causal Wan ("exercises chunk-KV"); Phase 2 ships "AR continuous batching" — which §6.3.1 *constitutively defines* as "(continuous batching; paged KV; chunked prefill)"; the CacheManager owning both lands in Phase 4, and risk 3 even defers the StepScheduler+KVPool prototype to "early in Phase 4," contradicting Phase 2. Compounding it: **no AR-pathway model exists on the new runtime before the Phase-4 Cosmos3 re-port** (Wan-causal is chunked denoise, not token AR; thinkers/talkers are Phase 4), so Phase 2's headline deliverable has neither a cache backing nor a workload — and none of Phase 2's gates (lines 1427-1430) tests AR batching.
|
||||
|
||||
The Phase-1 half is softenable: an interim per-request chunk-KV behind the unchanged `KVHandle` seam, with a Phase-4 allocator swap, is normal incremental staging — but the doc never states this, and its own "no third stack / every phase deletes what it replaces" principle cuts against unstated throwaway implementations. Fix structurally: pull a CacheManager v0 (chunk-KV slabs + minimal paged text-KV) into Phases 1–2, or move AR batching to Phase 4 and rewrite the Phase-2 gate to what it actually exercises.
|
||||
|
||||
### M17. Phase 4 re-ports a baseline that is not on main, and the plan schedules neither its merge nor its rebase
|
||||
**Where:** §10 Phase 0 line 1405, Phase 4 lines 1439-1446; §1 lines 42-49; Appendix.
|
||||
|
||||
`fastvideo/pipelines/basic/cosmos3/` on main contains only `__pycache__` — the design's forcing function exists solely as the unmerged 5-branch stacked chain (`feat/cosmos3-tier-a-port` → … → `feat/cosmos3-reasoning`). Phase 0's "Cosmos3 audio leaves `batch.extra`" cannot execute against main: it presupposes the chain is merged (a major-model review effort the plan never schedules) or means maintaining the migration on a side branch, continuously rebased across the most churn-heavy refactors in the repo's history (ForwardBatch→RequestState, loop inversion, executor→engine) — months of conflict-resolution work, unowned and unsized, on the artifact whose 150/150 bit-exactness is the design's proudest credential and whose parity suite the Phase-4 gate requires ("every phase ships green" cannot apply to a suite that is not in the tree). The plan sequences other in-flight work explicitly (`fastvideo/api/` in Phase 0, PR #1438 in Phase 1) but skips this. Needs an explicit merge milestone before Phase 0 touches the port.
|
||||
|
||||
### M18. G5's enforcement instrument has holes: ~6-7 shipped families have no SSIM test, and the CI-cost mitigation is incoherent for substrate PRs
|
||||
**Where:** G5 lines 128-129; Phase 0 gate line 1405; risk 6.
|
||||
|
||||
`fastvideo/tests/ssim/` covers ~14 of 20+ families. Cosmos(2/2.5), Hunyuan, Hunyuan15(+SR), HYWorld, MagiHuman, Waypoint, and MatrixGame-v1 have no SSIM test — "all SSIM suites unchanged" passes *vacuously* for roughly a third of shipped pipelines, exactly the ones sitting on the shared loop being refactored. And risk 6's "gated to touched families" mitigation is designed for model-local PRs; Phases 0–2 are by construction not model-local — the ForwardBatch adapter, loop inversion, and executor replacement sit under every family, so "touched families" = all of them on precisely the riskiest PRs. Either substrate PRs run the full GPU matrix (a cost the plan should budget — SSIM runs on Modal L40S today) or gating quietly degrades to sampling, which is how regressions slip through. Needs: a reference-seeding work item before Phase 1, or G5 restated as "zero regression for the SSIM-covered subset," plus a stated per-phase GPU-CI budget.
|
||||
|
||||
### M19. Phase 5's deletion milestone breaks the "frozen and untouched" legacy training/ stack
|
||||
**Where:** lines 5-6, 144, 1026-1027 vs Phase 5 line 1448.
|
||||
|
||||
The frozen stack is a live consumer of exactly the code Phase 5 deletes: `fastvideo/training/training_pipeline.py:39` imports `ComposedPipelineBase`/`ForwardBatch`/`LoRAPipeline`, holds `validation_pipeline: ComposedPipelineBase`, and its validation instantiates real legacy pipelines that run the legacy `DenoisingStage`; `distillation_pipeline.py:31` likewise. So Phase 5 cannot remove `ComposedPipelineBase` and `DenoisingStage` while leaving `training/` untouched — either the deletion milestone hollows to "delete except what legacy training/ needs" (the old path never dies — the very smell being fixed) or the scope statement is false and `training/` breaks on this plan's schedule. Relatedly, "loop inversion makes the step functions the single shared implementation" is arithmetically 3→2, not 3→1: the legacy inlined copies are out of scope forever. The doc needs an explicit answer: what happens to `fastvideo/training/` at Phase 5?
|
||||
|
||||
### M20. "Retire `fastvideo/forward_context.py` (Phase 1)" is infeasible as scheduled
|
||||
**Where:** §6.3.3 lines 618-621; Phase 1 lines 1412-1414; vs N2/N4; Appendix line 1791.
|
||||
|
||||
194 references across ~50 files. The global is read inside `fastvideo/attention/layer.py` — the shared Attention module on *every* family's hot path — and set in 27 places inside the frozen `training/` stack (8 module-level imports). Phase 1 migrates only Wan+Flux2; the other ~16 families run "unmodified" behind the legacy adapter (N4) and still set the global. So in Phase 1 the file cannot be deleted (touches the frozen stack, violating N2; breaks every unmigrated family), and `attention/layer.py` must serve both worlds simultaneously — a dual-sourcing branch in the hottest shared layer, undesigned. The honest description: Phase 1 *adds a second context mechanism beside the global*, and the global survives until Phase 5 at the earliest — where the deliverables list never mentions it. Appendix A states "retired Phase 1" as accomplished fact. Rewrite as "new-path-only StageContext; `forward_context` frozen for legacy consumers; deletion gated on Phase 5," and design the dual-mechanism cost.
|
||||
|
||||
### M21. §10 is a dependency ordering, not a plan — no timeline, no staffing, no sizing, and no policy for the ~1-2 new model ports per month that arrive during the migration
|
||||
**Where:** §10; N4 line 153; risk 1.
|
||||
|
||||
The scope — typed I/O, loop inversion + policies, extension system, async engine + StepScheduler + online-calibrated cost model, four-class CacheManager, PackedSeq/MoT layers, declarative parallelism compiler, workflow compiler, RL layer, Dynamo contract, config collapse — is plainly multi-engineer-years, with zero dates, headcount, per-phase sizing, or owners; "by Phase 2" decision deadlines (§11.1, risks 7/15) are unanchored because Phase 2 is not a date.
|
||||
|
||||
The sharper, unanswered problem is **inflow**: git shows ~1–2 new families landing per month (Flux2 Klein and Lucy Edit on 2026-06-09 alone; MatrixGame3 05-27; MagiHuman 05-12; Stable Audio 05-01; Gen3C 04-01…). Over multi-quarter Phases 0–4, another 10–15 models arrive, and the doc never says what they target: land them on legacy abstractions and the Phase-5 tail grows faster than phases retire it (negative net migration velocity); force them onto the new stack and every port blocks on machinery that doesn't exist until Phase 1/3/4. Either answer materially changes the plan; choosing neither means the terminal state recedes indefinitely. Minimum fix: per-phase engineer-month estimates, a named owner per phase, a calendar target for Phase 5, and an explicit "new ports target the new stack starting at Phase X" rule with its porting-velocity cost stated.
|
||||
|
||||
---
|
||||
|
||||
## Major — product/trust surfaces
|
||||
|
||||
### M22. Per-request plugin enablement is an unsandboxed third-party-code and noisy-neighbor surface; only workflow JSON is named untrusted
|
||||
**Where:** §6.4 lines 859-861 vs §12 input-hardening gap lines 1693-1695.
|
||||
|
||||
Entry-point plugins execute arbitrary code inside the serving engine, and the doc makes their selection part of the *request* (`diffusion.plugins=[{"name": "cache_dit", "Fn": 8, "Bn": 8}]`) in the same engine pitched as a multi-tenant cloud — and since the OpenAI protocol is *generated from the request schema* (lines 907-908), the field derives into the public API with no carve-out. Consequences forcing a design change: (a) **correctness** — a caller can attach a distribution-altering interceptor to a request the product has labeled "exact mode" (the §9.4 trust claim), or pass unvalidated kwargs into third-party code; (b) **isolation** — a `needs_eager` observer on one request drops compile/cudagraph capture for scopes shared with co-scheduled tenants (line 809), a noisy-neighbor vector with no cost attribution anywhere in the metrics design; (c) **supply chain** — entry-point resolution imports whatever package claims the name. The needed contract: enablement/allowlisting at DeployConfig scope only; requests merely parameterize pre-enabled plugins against per-plugin validated schemas; plugin overhead attributed per-request in the cost model. §12's input-hardening gap names only workflow JSON — a categorically different surface.
|
||||
|
||||
### M23. No versioning or stability contract for the serialized schemas shipped to external consumers mid-migration
|
||||
**Where:** §6.4 line 861; §6.6 lines 920-927; §10; open question 12.
|
||||
|
||||
By Phase 3 there are at least four externally consumed serialized surfaces: hub-published ModelSpec manifests (interchange with diffusers' `modular_model_index.json` — a format co-owned with an external party), compiled-workflow PipelineSpecs (content-hash-keyed in the weight-fleet cache — schema changes silently change hashes and invalidate fleet affinity), the OmniEvent streaming schema (Dreamverse's frontend; proposed as Dynamo ask A3's wire format), and per-model ModelOptions blocks. Phase 4 then lands PackedSeq, session-scoped inputs, and the Cosmos3 re-port — guaranteed churn after consumers exist. The migration plan gates *behavior* at every phase (SSIM, parity, load) and gates *interfaces* at none; the only versioning commitment in the document is hook-point names (open question 12 is scoped to hook points). Without per-surface decisions now — `schema_version` fields, frozen-vs-experimental tiers per phase, a deprecation window — Phase 4 either breaks published artifacts or gets paralyzed by accidental freezing. G5 protects only the Python `VideoGenerator` call.
|
||||
|
||||
---
|
||||
|
||||
## Minor (confirmed)
|
||||
|
||||
1. **ForwardBatch has 111 fields, not ~250** (AST-verified; stated twice, lines 33/188). P3 survives at 111, but the headline metric is inflated 2.3× in a doc that brands its pain points "evidence-backed" — it invites discounting of the numbers that *do* verify exactly (1381 lines and 35 probes both check out).
|
||||
2. **"Prediction is a table lookup" vs the design's own flagship features** (§6.3.1 vs §6.4): DBCache/FBCache/TaylorSeer decide per step from runtime residual similarity — a stochastic per-step cost multiplier unknowable at admission; VSA tile selection is content-dependent; and AR decode lengths are unbounded (the doc concedes vLLM "must guess decode lengths," then silently exempts its own AR group).
|
||||
3. **Worked example (g) is internally contradictory**: cache-dit + C1 + "identical trajectories" are pairwise incompatible under §8.5's own `distribution_altering` contract (§8.7 states the rule correctly: cache acceleration is C0). Matters because (g) is the template PR #1438 is told to target in Phase 1.
|
||||
4. **The Phase-2 Dreamverse gate is untestable as written**: at ~4.55 s GPU-saturating per 5 s clip (line 1263), "≥2 concurrent sessions per GPU at unchanged segment latency" is only passable under an unstated think-time/collision-rate assumption — the gate can be passed or failed at will by choosing the test's session behavior. More broadly, no quantitative multiplexing target (sessions/GPU under a stated load profile, GPU-utilization, cost/clip) exists anywhere, so there is no way to conclude after Phase 2 whether step-level scheduling earned its complexity over the §11.6-rejected simpler design.
|
||||
5. **The exec summary launders Dynamo contingencies into outcomes** (line 75: "each with a fallback — so Dynamo fronts both production serving and RL rollout fleets"): the body is honest (A1–A7 with fallbacks; §11.9; §12.15), but the asks are unfiled RFCs on an NVIDIA-governed roadmap; A5's own fallback "weakens fleet-scale async RL," and if A3 misses Phase 2, Dreamverse ships on the direct-WebSocket bypass and the production-hardened fallback becomes permanent — the exact "permanent workaround" dynamic §11.9 claims the direct relationship avoids. Ask-sequencing (§12.15) has no owner or decision dates.
|
||||
6. **diffusers as "convergent validation" cuts both ways** (see M2): its four-wrappers-per-family shape is the subclass forest again; the citation supports the rejected alternative as well as the chosen one.
|
||||
7. **Punica/ComfyUI LoRA semantics gap** — see M4.
|
||||
|
||||
---
|
||||
|
||||
## Fact-check corrections
|
||||
|
||||
70 concrete claims were checked; **none was fabricated**; 13 need correction. Everything else verified, including the claims most likely to be embellished: vLLM RFC #42770 (author/date/content/two-tier resolution), PR #42304 **merged** 2026-05-16 with `VLLM_USE_BREAKABLE_CUDAGRAPH`, vllm-omni RFC #4084, the Thinking Machines numbers (80/1000 unique outputs, divergence at token 103, 26s→42s, KL results), the Dynamo worker's `asyncio.Lock`, cache-dit, the cosmos-framework MoT details (PackedAttentionMoT, MoTDecoderLayer, ReasonerKVCache, MoE gen-MLP), miles/verl-omni/sglang-omni mechanics, sglang's cache-dit monkeypatch scars, and `enable_teacache` genuinely having no consumer.
|
||||
|
||||
| # | design.md says | Reality |
|
||||
|---|---|---|
|
||||
| 1 | "1381-line `DenoisingStage`" (lines 34, 201) | 1381 is the **file**; the class is ~670 lines (47–715) plus 6 subclasses in-file. The 35-probe count is exact for the file. |
|
||||
| 2 | "~250-field ForwardBatch" (33, 188) | **111 fields** (whole file incl. TrainingBatch/PreprocessBatch: ~153). |
|
||||
| 3 | "19 denoising-stage classes" (201) | **22** model/variant classes (+ base = 23); the list omits Magi-class and two other same-category stages predating the doc. |
|
||||
| 4 | "Cosmos2.5 clamping … hardcoded in the shared loop" (201) | Clamping lives in the `Cosmos25DenoisingStage` **subclass**; the Wan2.2 expert switch (`denoising.py:229-235, 352-376`) and TI2V inline VAE encode (`:239-268, 399-404, 570-572`) are in the shared loop as claimed. |
|
||||
| 5 | `SamplingParam` "~170 fields" (887) | **75**. The ~170 figure belongs to TrainingArgs (90 own + 81 inherited = 171). |
|
||||
| 6 | `FastVideoArgs` "~96 fields" (885) | **81** (TrainingArgs subclassing claim correct). |
|
||||
| 7 | "TP and SP (Ulysses/ring)" (192) | Main is **Ulysses-only** (`all_to_all_4D`); no ring-attention SP is wired into FastVideo. |
|
||||
| 8 | CFG "3 copies: `stages/conditioning.py` vs …" (993) | Right count, wrong citation: the inference-stack copy is in `denoising.py`, not `conditioning.py`. |
|
||||
| 9 | ComfyUI "~45 `comfy_extras` packs", "90+ blueprints" (1335-1339) | **117** packs (matching nodes.py's 117-entry registration list); **80** in-tree blueprints (the larger library ships via the registry). 64 core nodes, 39 API providers, GPL-3.0, FIFO-no-batching all verify. |
|
||||
| 10 | kv-router events "`{sequence_hash, block_hash, removed}`" (707, A1 733-739) | Paraphrase: actual shape is `KvCacheEventData::Stored{parent_hash, blocks[{block_hash, tokens_hash}]}` / `Removed` / `Cleared` (`protocols.rs:627-646`). Token-prefix-derived keying verifies. |
|
||||
| 11 | miles TIS clamp "to `[0.5, 2.0]`" (1086) | Configurable `[tis_clip_low, tis_clip]`, CLI defaults [0, 2.0]; the 0.5/2.0 pair comes from the MIS example config (`mis.yaml`). |
|
||||
| 12 | sglang-omni "`DllmScheduler` for a DiT talker" (269) | DllmScheduler serves the **LLaDA2-Uni thinker** (diffusion-LLM); the DiT talker is Ming-Omni's, on a different scheduler. |
|
||||
| 13 | `_iter_packed_batches` under `model/vfm/` (236); §11.3's claim that the port's "own status notes" list reasoning-KV/batching/streaming/prefix-reuse as "missing for production" | Lives at `cosmos_framework/inference/inference.py:66`. PORT_STATUS.md confirms 150/150 but contains no such missing-for-production list — that framing is the design doc's own and should not be attributed to the port's status notes. |
|
||||
|
||||
---
|
||||
|
||||
## Attacks that failed (the doc survives these)
|
||||
|
||||
The refute-by-default verifiers killed 36 findings, several of them attacks a hostile reviewer would lead with — worth knowing they don't land:
|
||||
|
||||
- **ChunkRollout/DenoiseLoop nesting is expressible** in the stated Stage/LoopStage/StepResult contracts ("one solver step / one token / one chunk" + composition).
|
||||
- **N1 vs engine-internal pools** is consistent on a careful read (N1 is about datacenter orchestration; §6.3.5 states the reconciliation).
|
||||
- **The trainer-scope line (N2 vs §8)** is drawn consistently — N2's own text enumerates exactly what §8 changes.
|
||||
- **G6 vs the ≤2% Phase-2 gate** is goal-vs-acceptance-gate, not contradiction (Phase 1 is gated bit-identical).
|
||||
- **The clean-room GPL posture holds**: sampler/scheduler math (DPM-Solver, Karras sigmas, flow-match shift) is published outside GPL sources.
|
||||
- **C2 for the video denoise path is fine**: batch-1 fixed shapes are trivially batch-invariant — the doc's own analysis at lines 1145-1147 is correct; the AR/image/sharding exposures are correctly identified there too.
|
||||
- **Self-forcing's cross-chunk gradients truncate by construction** (KV written under `no_grad` on detached context), so the engine KV pool is not blocked the way one might fear — the surviving residue is M12's grad-window cache mode.
|
||||
- **"Every phase deletes or freezes something" survives audit** at the phase-deliverable level (the failures are the specific items in M19/M20).
|
||||
- **The tier-1 ComfyUI vocabulary claim survives** blueprint-corpus measurement under the doc's actual claim (curated canonical workflows, not top-N node frequency).
|
||||
- **The sglang reconvergence deferral** is substantively defended in §11.1 with reasons valid under either outcome.
|
||||
- **WeightSyncPlan's "literal no-op"** is correctly scoped to colocated same-layout in the doc's own sentence; FSDP-vs-TP/SP is explicitly routed to in-place reshard.
|
||||
|
||||
---
|
||||
|
||||
## Ranked recommendations
|
||||
|
||||
1. **Design the abort/cancellation/OOM path with Phase 2** (C1) **and add memory as a budget axis with admission planning and preemption semantics** (M3). These two are the soundness conditions of the multiplexing bet; everything else in the execution plane sits on them.
|
||||
2. **Re-derive §6.3.2 from the real vLLM constraint** (M1). The two-pool→one-pool reversal was made on a false premise; either accept uniform page bytes (and redesign the slab story) or bring back two pools with an explicit fragmentation/deadlock argument.
|
||||
3. **Fix the migration plan's three structural defects**: CacheManager v0 into Phases 1–2 or AR batching out of Phase 2 (M16); a merge milestone for the cosmos3 chain before Phase 0 touches it (M17); a new-port inflow rule plus a freeze-enforcement mechanism that did not exist last time — CI path gate, codeowners, a date (M15, M21). Also reconcile Phase 5 with the frozen `training/` stack (M19) and restate the `forward_context` retirement honestly (M20).
|
||||
4. **Specify the step skeleton and the policy contracts** — ordered, typed extension points; a policy state-scoping rule (state in LoopState, like plugins); a policy-observation channel — and work the mapping through Cosmos2.5 and LTX2 in the doc (M7, M8). Decide per-node request parameter binding in Phase 0 (M9).
|
||||
5. **Give MoT a stated parallelism answer** (M5) and make the single-pool-spans-nodes decision explicit, including the fate of `RayDistributedExecutor` (M6).
|
||||
6. **Close the workflow-cloud trust/correctness holes before Phase 3**: adapter-aware feature-cache keys (M10), weight-state transitions as a scheduled, costed operation (M11), DeployConfig-scoped plugin allowlisting (M22), per-surface schema stability tiers (M23), and an honest assessment of Punica's fit (M4).
|
||||
7. **Right-size the RL claims**: design the grad+KV cache mode or scope self-forcing out of the shared loop (M12); budget the Behavior Record at real byte counts (M13); demote the omni-RL pilot or give it an objective sketch and an owner (M14); fix worked example (g).
|
||||
8. **Reclassify loop inversion as unprecedented at scheduler granularity** in risk 3 and drop the diffusers "validation" (M2). The bet may still be right — but it should be made with open eyes, and the parity-gate plan is then carrying more weight than the doc admits.
|
||||
9. **Correct the thirteen numbers above before circulating.** The doc's credibility rests on its "evidence-backed" brand; ~250-vs-111 is the kind of error that makes a reader re-check everything else — and most of everything else checks out.
|
||||
@@ -1,124 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Convert DreamX-World-5B autoregressive weights to FastVideo layout.
|
||||
|
||||
The HF repository stores one raw official ``model.safetensors`` whose keys match
|
||||
FastVideo's native ``DreamXWorldARTransformer3DModel``. The converter writes a
|
||||
Diffusers-like root with ``transformer/config.json`` and reusable Wan2.2
|
||||
components. Use ``--symlink-transformer`` locally to avoid duplicating the 21GB
|
||||
AR tensor file.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
TRANSFORMER_CONFIG: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldARTransformer3DModel",
|
||||
"model_type": "ti2v",
|
||||
"patch_size": [1, 2, 2],
|
||||
"text_len": 512,
|
||||
"num_attention_heads": 24,
|
||||
"attention_head_dim": 128,
|
||||
"in_channels": 48,
|
||||
"out_channels": 48,
|
||||
"text_dim": 4096,
|
||||
"freq_dim": 256,
|
||||
"ffn_dim": 14336,
|
||||
"num_layers": 30,
|
||||
"local_attn_size": 12,
|
||||
"sink_size": 3,
|
||||
"cross_attn_norm": True,
|
||||
"qk_norm": True,
|
||||
"eps": 1e-6,
|
||||
"add_control_adapter": True,
|
||||
"cam_method": "prope",
|
||||
"attn_compress": 4,
|
||||
"cam_self_attn_layers": list(range(30)),
|
||||
"num_frames_per_block": 3,
|
||||
}
|
||||
|
||||
MODEL_INDEX: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldARPipeline",
|
||||
"_diffusers_version": "0.31.0",
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"transformer": ["diffusers", "DreamXWorldARTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"],
|
||||
}
|
||||
|
||||
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
|
||||
|
||||
|
||||
def _source_safetensors(source: Path) -> Path:
|
||||
if source.is_file():
|
||||
return source
|
||||
path = source / "model.safetensors"
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Missing AR model.safetensors under {source}")
|
||||
return path
|
||||
|
||||
|
||||
def convert_transformer(source: Path, output: Path, symlink_transformer: bool) -> None:
|
||||
src = _source_safetensors(source)
|
||||
transformer_dir = output / "transformer"
|
||||
transformer_dir.mkdir(parents=True, exist_ok=True)
|
||||
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
|
||||
dst = transformer_dir / "model.safetensors"
|
||||
if dst.is_symlink() and not dst.exists():
|
||||
dst.unlink() # dangling symlink: replace so re-conversion self-heals
|
||||
if dst.exists():
|
||||
return
|
||||
if symlink_transformer:
|
||||
dst.symlink_to(src.resolve())
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
|
||||
|
||||
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
|
||||
src = component_source / component
|
||||
dst = output / component
|
||||
if not src.exists():
|
||||
raise FileNotFoundError(f"Missing reused component source: {src}")
|
||||
if dst.is_symlink() and not dst.exists():
|
||||
dst.unlink() # dangling symlink: replace so re-conversion self-heals
|
||||
if dst.exists():
|
||||
return
|
||||
if symlink:
|
||||
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
|
||||
elif src.is_dir():
|
||||
shutil.copytree(src, dst)
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
|
||||
|
||||
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
|
||||
if component_source is not None:
|
||||
for component in REUSED_COMPONENTS:
|
||||
_copy_or_link_component(component, component_source, output, symlink_components)
|
||||
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
|
||||
if missing:
|
||||
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
|
||||
return
|
||||
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--source", type=Path, required=True)
|
||||
parser.add_argument("--output", type=Path, required=True)
|
||||
parser.add_argument("--component-source", type=Path)
|
||||
parser.add_argument("--symlink-components", action="store_true")
|
||||
parser.add_argument("--symlink-transformer", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
convert_transformer(args.source, args.output, args.symlink_transformer)
|
||||
write_model_index(args.output, args.component_source, args.symlink_components)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,221 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Convert DreamX-World-5B-Cam raw transformer weights to FastVideo-loadable format.
|
||||
|
||||
The GD-ML/DreamX-World-5B-Cam repository stores the transformer as raw
|
||||
DreamX/Wan official shards. FastVideo's TransformerLoader expects a Diffusers-like
|
||||
transformer folder with a config.json and safetensors whose keys can be mapped by
|
||||
WanVideoConfig.param_names_mapping. This script performs the raw official ->
|
||||
Diffusers-like key rename and writes the DreamX 5B-Cam transformer config.
|
||||
|
||||
Example:
|
||||
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
|
||||
--source official_weights/dreamx_world \
|
||||
--output converted_weights/dreamx_world
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from huggingface_hub import save_torch_state_dict
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file
|
||||
|
||||
|
||||
OFFICIAL_TO_DIFFUSERS_MAPPING: dict[str, str] = {
|
||||
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.linear_1.\1",
|
||||
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.linear_2.\1",
|
||||
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.linear_1.\1",
|
||||
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.linear_2.\1",
|
||||
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_proj.\1",
|
||||
r"^img_emb\.proj\.0\.(.*)$": r"condition_embedder.image_embedder.norm1.\1",
|
||||
r"^img_emb\.proj\.1\.(.*)$": r"condition_embedder.image_embedder.ff.net.0.proj.\1",
|
||||
r"^img_emb\.proj\.3\.(.*)$": r"condition_embedder.image_embedder.ff.net.2.\1",
|
||||
r"^img_emb\.proj\.4\.(.*)$": r"condition_embedder.image_embedder.norm2.\1",
|
||||
r"^head\.modulation": r"scale_shift_table",
|
||||
r"^head\.head\.(.*)$": r"proj_out.\1",
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.attn1.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.attn1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.attn1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k_img\.(.*)$": r"blocks.\1.attn2.add_k_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v_img\.(.*)$": r"blocks.\1.attn2.add_v_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k_img\.(.*)$": r"blocks.\1.attn2.norm_added_k.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.net.0.proj.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.net.2.\2",
|
||||
r"^blocks\.(\d+)\.modulation": r"blocks.\1.scale_shift_table",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm2.\2",
|
||||
}
|
||||
|
||||
TRANSFORMER_CONFIG: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldTransformer3DModel",
|
||||
"patch_size": [1, 2, 2],
|
||||
"text_len": 512,
|
||||
"num_attention_heads": 24,
|
||||
"attention_head_dim": 128,
|
||||
"in_channels": 48,
|
||||
"out_channels": 48,
|
||||
"text_dim": 4096,
|
||||
"freq_dim": 256,
|
||||
"ffn_dim": 14336,
|
||||
"num_layers": 30,
|
||||
"cross_attn_norm": True,
|
||||
"qk_norm": "rms_norm_across_heads",
|
||||
"eps": 1e-6,
|
||||
"image_dim": None,
|
||||
"added_kv_proj_dim": None,
|
||||
"rope_max_seq_len": 1024,
|
||||
"add_control_adapter": True,
|
||||
"cam_method": "prope",
|
||||
"attn_compress": 1,
|
||||
"cam_self_attn_layers": None,
|
||||
}
|
||||
|
||||
MODEL_INDEX: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldPipeline",
|
||||
"_diffusers_version": "0.31.0",
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"],
|
||||
}
|
||||
|
||||
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
|
||||
|
||||
|
||||
def map_transformer_key(key: str) -> str:
|
||||
for pattern, replacement in OFFICIAL_TO_DIFFUSERS_MAPPING.items():
|
||||
if re.match(pattern, key):
|
||||
return re.sub(pattern, replacement, key)
|
||||
return key
|
||||
|
||||
|
||||
def _safetensor_files(source: Path) -> list[Path]:
|
||||
if source.is_file():
|
||||
if source.suffix != ".safetensors":
|
||||
raise ValueError(f"Only .safetensors files are supported, got {source}")
|
||||
return [source]
|
||||
|
||||
index_path = source / "diffusion_pytorch_model.safetensors.index.json"
|
||||
if index_path.exists():
|
||||
index = json.loads(index_path.read_text())
|
||||
return sorted({source / shard for shard in index["weight_map"].values()})
|
||||
|
||||
files = sorted(source.glob("*.safetensors"))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"No safetensors files found under {source}")
|
||||
return files
|
||||
|
||||
|
||||
def convert_transformer(source: Path, output: Path, max_shard_size: str) -> None:
|
||||
transformer_dir = output / "transformer"
|
||||
transformer_dir.mkdir(parents=True, exist_ok=True)
|
||||
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
|
||||
|
||||
converted: OrderedDict[str, torch.Tensor] = OrderedDict()
|
||||
for shard in _safetensor_files(source):
|
||||
print(f"loading {shard}")
|
||||
for key, tensor in load_file(shard, device="cpu").items():
|
||||
new_key = map_transformer_key(key)
|
||||
if new_key in converted:
|
||||
raise ValueError(f"Duplicate converted key: {new_key}")
|
||||
converted[new_key] = tensor
|
||||
|
||||
print(f"saving {len(converted)} tensors to {transformer_dir}")
|
||||
save_torch_state_dict(converted, transformer_dir, max_shard_size=max_shard_size)
|
||||
|
||||
|
||||
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
|
||||
src = component_source / component
|
||||
dst = output / component
|
||||
if not src.exists():
|
||||
raise FileNotFoundError(f"Missing reused component source: {src}")
|
||||
if dst.is_symlink() and not dst.exists():
|
||||
dst.unlink() # dangling symlink: replace so re-conversion self-heals
|
||||
if dst.exists():
|
||||
print(f"keeping existing {dst}")
|
||||
return
|
||||
if symlink:
|
||||
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
|
||||
print(f"linked {dst} -> {src}")
|
||||
elif src.is_dir():
|
||||
shutil.copytree(src, dst)
|
||||
print(f"copied {src} -> {dst}")
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
print(f"copied {src} -> {dst}")
|
||||
|
||||
|
||||
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
|
||||
if component_source is not None:
|
||||
for component in REUSED_COMPONENTS:
|
||||
_copy_or_link_component(component, component_source, output, symlink_components)
|
||||
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
|
||||
if missing:
|
||||
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
|
||||
print("pass --component-source <Wan2.2 Diffusers root> to copy or link reused components")
|
||||
return
|
||||
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
|
||||
print(f"wrote {output / 'model_index.json'}")
|
||||
|
||||
|
||||
def analyze(source: Path) -> None:
|
||||
total = 0
|
||||
unchanged = 0
|
||||
examples: list[tuple[str, str]] = []
|
||||
for shard in _safetensor_files(source):
|
||||
with safe_open(shard, framework="pt", device="cpu") as tensors:
|
||||
for key in tensors:
|
||||
total += 1
|
||||
new_key = map_transformer_key(key)
|
||||
unchanged += int(new_key == key)
|
||||
if len(examples) < 20 and new_key != key:
|
||||
examples.append((key, new_key))
|
||||
print(f"total_keys={total} unchanged_keys={unchanged}")
|
||||
for old, new in examples:
|
||||
print(f"{old} -> {new}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--source", type=Path, required=True, help="DreamX raw transformer directory or safetensors file")
|
||||
parser.add_argument("--output", type=Path, required=True, help="Output model root; transformer/ is created inside it")
|
||||
parser.add_argument("--max-shard-size", default="10GB")
|
||||
parser.add_argument(
|
||||
"--component-source",
|
||||
type=Path,
|
||||
help="Optional Wan2.2 Diffusers root whose scheduler/text_encoder/tokenizer/vae components are reused.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--symlink-components",
|
||||
action="store_true",
|
||||
help="Symlink reused components from --component-source instead of copying them.",
|
||||
)
|
||||
parser.add_argument("--analyze", action="store_true", help="Only print key mapping summary")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.analyze:
|
||||
analyze(args.source)
|
||||
else:
|
||||
convert_transformer(args.source, args.output, args.max_shard_size)
|
||||
write_model_index(args.output, args.component_source, args.symlink_components)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,189 +0,0 @@
|
||||
# DreamX World Port Status
|
||||
|
||||
## Summary
|
||||
|
||||
- model_family: `dreamx_world`
|
||||
- workload_types: `I2V camera-control compatibility shim`; `I2V autoregressive camera-control forcing`
|
||||
- official_ref: `https://github.com/AMAP-ML/DreamX-World`
|
||||
- official_ref_dir: `DreamX-World/`
|
||||
- hf_weights_path: `GD-ML/DreamX-World-5B-Cam`
|
||||
- local_weights_dir: `official_weights/dreamx_world`
|
||||
- source_layout: `raw_official`
|
||||
- local_tests_readme: `tests/local_tests/dreamx_world/README.md`
|
||||
|
||||
## Current Phase
|
||||
|
||||
- phase: `phase_11_post_parity_handoff`
|
||||
- status: `complete`
|
||||
- owner: `orchestrator`
|
||||
- last_updated: `2026-07-02`
|
||||
|
||||
## Component Matrix
|
||||
|
||||
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|
||||
|---|---|---|---|---|---|---|---|---|---|
|
||||
| transformer | dit | ported_dedicated | `DreamX-World/models/wan_transformer3d.py`; PRoPE helpers in `DreamX-World/models/prope_utils.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `Wan2_2Transformer3DModel.from_pretrained(... cam_method=prope, add_control_adapter=True)` | `fastvideo/models/dits/dreamx_world.py`; `fastvideo/configs/models/dits/dreamx_world.py`; DreamX pipeline config helper | native_prope_pass | real_conversion_pass | strict_load_and_forward_parity_pass | none |
|
||||
| vae | vae | reuse_pending | `DreamX-World/models/wan_vae3_8.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `AutoencoderKLWan3_8.from_pretrained(Wan2.2_VAE.pth)` | `fastvideo/models/vaes/wanvae.py`; DreamX VAE config helper | config_smoke_pass | raw_key_mapping_pass | encode_parity_pass | none |
|
||||
| text_encoder/tokenizer | encoder | reuse_pending | `DreamX-World/models/wan_text_encoder.py`; tokenizer via Wan2.2 base model | `DreamX-World/inference_dreamx5b.py::setup_models`, `WanT5EncoderModel` + tokenizer subpaths | `fastvideo/models/encoders/t5.py::UMT5EncoderModel`; DreamX UMT5 config helper | config_smoke_pass | staged_weight_load_pass | hidden_state_parity_pass | none |
|
||||
| scheduler | generic | reuse_proven | Diffusers `FlowMatchEulerDiscreteScheduler` | `DreamX-World/inference_dreamx5b.py::setup_models`, default `sampler_name=Flow` | `fastvideo/models/schedulers/scheduling_flow_match_euler_discrete.py` | pass | not_required | non_skip_pass | Q003 |
|
||||
| camera_conditioning | generic | port_pending | `DreamX-World/utils/inference_utils.py`, `DreamX-World/models/prope_utils.py`, `DreamX-World/wan/modules/camera_prope.py` | `DreamX-World/inference_dreamx5b.py::get_camera_sequence`, `pipeline(... control_camera_video=...)` | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` | pass | not_required | non_skip_pass | none |
|
||||
| pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_dreamxworld.py` | `DreamX-World/inference_dreamx5b.py::process_inference_from_json` | `fastvideo/pipelines/basic/dreamx_world/` plus config/preset/registry | pipeline_load_generate_smoke_pass | model_index_and_config_consistency_smoke_pass | pipeline_api_vs_worker_forward_parity_pass | none |
|
||||
| ar_transformer | dit | port_complete | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py::CausalWanModel` | `DreamX-World/inference_ar_forcing.py::load_pipeline` | `fastvideo/models/dits/dreamx_world_ar.py`; `fastvideo/configs/models/dits/dreamx_world.py::DreamXWorldARConfig` | tiny_official_forward_parity_pass | identity_conversion_pass | real_5b_strict_load_pass | none |
|
||||
| ar_pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_causal_camera.py` | `DreamX-World/inference_ar_forcing.py::main` | `fastvideo/pipelines/basic/dreamx_world/dreamx_world_ar_pipeline.py`; `fastvideo/pipelines/basic/dreamx_world/ar_denoising.py`; registry/preset/config | config_registry_pass | symlink_model_index_pass | short_full_generation_pass | none |
|
||||
|
||||
## Conversion State
|
||||
|
||||
- conversion_script: `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py`
|
||||
- converted_weights_dir: `converted_weights/dreamx_world`
|
||||
- source_layout: `raw_official`
|
||||
- strict_load_status: `pass`
|
||||
- conversion_script_status: `transformer_model_index_and_config_consistency_smoke_pass`
|
||||
- model_index_status: `smoke_pass`
|
||||
- passthrough_components: `Wan2.2 Diffusers scheduler, tokenizer, and text encoder are symlinked from official_weights/Wan2.2-TI2V-5B-Diffusers; VAE parity uses raw Wan2.2_VAE.pth with an explicit DreamX raw-to-FastVideo key mapper because official encode returns normalized latents.`
|
||||
- retry_history: `none`
|
||||
|
||||
## Parity Commands
|
||||
|
||||
| Scope | Command | Last Result | Notes |
|
||||
|---|---|---|---|
|
||||
| transformer | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s` | strict_load_and_forward_parity_pass | 2026-07-01: converted real 5B-Cam transformer shards strict-load into dedicated `DreamXWorldTransformer3DModel` with 0 shape mismatches; official-vs-FastVideo small-input fp32 forward parity passes on CUDA (`diff_max=0.072533`, `diff_mean=0.008014`). |
|
||||
| vae | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s` | encode_parity_pass | 2026-06-30: official DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE and encode parity passes after applying the same official latent normalization (`(mu - mean) / std`). |
|
||||
| text_encoder | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s` | hidden_state_parity_pass | 2026-06-30: official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA using staged Wan2.2 text encoder/tokenizer weights and reference-only `xfuser` stubs. |
|
||||
| scheduler | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s` | non_skip_pass | 2026-06-30: FastVideo FlowMatch scheduler matches official Diffusers timesteps and step output for DreamX default Flow sampler; `DreamXWorldPipeline` initializes FlowMatch with official `shift=3.0`. |
|
||||
| camera_conditioning | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s` | non_skip_pass | 2026-06-29: 3 parameterized cases passed against official reference on CPU. |
|
||||
| pipeline_config | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s` | pipeline_entry_preset_scheduler_modelinfo_and_camera_stage_smoke_pass | DreamX 5B-Cam PipelineConfig wires DiT/VAE/UMT5/Flow/TI2V settings and official `shift=3.0`; default preset is registered for `GD-ML/DreamX-World-5B-Cam`; local converted-style `model_index.json` resolves to `DreamXWorldPipeline`; the pipeline initializes FlowMatch, camera conditioning writes `batch.extra["dreamx_y_camera"]`, and generic denoising can pass it as `y_camera`. |
|
||||
| ar_transformer | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_ar_conversion.py tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py -q -rs` | 6_passed_0_skipped | 2026-07-02: AR converter writes symlinked transformer layout/model_index; tiny official `CausalWanModel` vs FastVideo `DreamXWorldARTransformer3DModel` forward parity passes; real 5B `model.safetensors` strict-loads with zero missing/unexpected keys from `/tmp/converted_dreamx_world_ar`. |
|
||||
| ar_pipeline_config | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -q -rs` | 10_passed_0_skipped | 2026-07-02: `DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, registry config selection for `GD-ML/DreamX-World-5B`, and `dreamx_world_5b_ar` preset pass. |
|
||||
| ar_full_generation | `PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py` | generated_video_pass | 2026-07-02: A40 short full-generation smoke passed from `/tmp/converted_dreamx_world_ar` with 64x64, 9 frames, 4 denoise steps, `output_type=pil`, `save_video=True`; MP4 saved at `outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4` and decoded as 9 frames of `(64, 64, 3)` uint8. |
|
||||
| ar_long_horizon | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py` | generated_video_pass | 2026-07-02: A40 long-horizon AR generation passed from `/tmp/converted_dreamx_world_ar` with 64x64, 1005 frames, 4 denoise steps, seed 4096; MP4 saved at `outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4` and decoded as 1005 frames of `(64, 64, 3)`; end-to-end generation latency was 231.68s after load. |
|
||||
| ar_ssim_default | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download` | 1_passed_0_skipped | 2026-07-02: A40 default AR SSIM reference seeded locally at `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/`; helper removes stale generated base MP4 before generation; default params are 192x192, 81 frames, 4 steps, seed 2048, min SSIM 0.98. |
|
||||
| ar_ssim_modal_l40s | `modal run /tmp/modal_dreamx_ar_ssim_git.py` | 1_passed_0_skipped | 2026-07-02: Modal L40S run checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded default `L40S_reference_videos` via `reference_videos_cli.py download`, used cached converted AR weights under `/root/data/dreamx_world_ar_converted`, seeded the missing AR L40S reference from generated output, reran a fresh generated-vs-reference compare successfully (`mean_ssim=1.0`), and exported the reference to Modal volume `hf-model-weights:dreamx_ar_ssim_l40s`. Downloaded local reference decodes as 81 frames of `(192, 192, 3)`. A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944. |
|
||||
| pipeline_smoke | `python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs` | 4_passed_0_skipped | 2026-06-30: combined smoke/parity passed. 2026-07-01: smoke alone passed with real `image_path` TI2V coverage (`3 passed`), validating image load, TI2V preprocessing, VAE first-frame encode under CPU offload, camera conditioning, and 1-step latent generation from `converted_weights/dreamx_world`. Tests force `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA` to avoid the local FlashAttention-4 cute ABI mismatch. |
|
||||
| basic_example | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA DREAMX_WORLD_MODEL_DIR=converted_weights/dreamx_world DREAMX_WORLD_IMAGE_PATH= DREAMX_WORLD_HEIGHT=64 DREAMX_WORLD_WIDTH=64 DREAMX_WORLD_NUM_FRAMES=9 DREAMX_WORLD_STEPS=1 DREAMX_WORLD_GUIDANCE=1.0 DREAMX_WORLD_OUTPUT_PATH=outputs_video/dreamx_world_example_smoke python examples/inference/basic/basic_dreamx_world.py` | generated_video_pass | 2026-06-30: example saved an MP4 under `outputs_video/dreamx_world_example_smoke`; imageio/ffmpeg decoded frame 0 as `(64, 64, 3)` uint8, fps 16, duration 0.56s. |
|
||||
|
||||
## Open Questions
|
||||
|
||||
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|
||||
|---|---|---|---|---|---|
|
||||
| Q001 | Should first PR expose only the `DreamX-World-5B-Cam` 5s camera-control mode and exclude AR long-horizon forcing? | user | prep | resolved | User approved starting with `DreamX-World-5B-Cam`; AR long-horizon is out of first-PR scope. |
|
||||
| Q002 | Does FastVideo's existing Wan2.2 TI2V transformer support DreamX PRoPE/control adapter with a small extension, or is a DreamX-specific DiT required? | component:transformer | Phase 3 | resolved | Project guidance prefers a separate DreamX DiT for maintainability. DreamX PRoPE/control adapter now lives in `fastvideo/models/dits/dreamx_world.py`; Wan DiT/config have no DreamX-specific fields or `y_camera` signature. |
|
||||
| Q003 | Which sampler is in first-PR scope: official default `Flow` only, or also `Flow_Unipc` and `Flow_DPM++`? | orchestrator | Phase 3 | resolved | First PR should support official default `Flow` only. FastVideo FlowMatch scheduler parity is non-skip PASS; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. |
|
||||
| Q004 | Which HF token env var should be used if rate limits or gated Wan2.2 base weights require auth? | user | Phase 5 | resolved | No auth was required for the completed local downloads; keep using env var names only if future gated repos require auth. |
|
||||
| Q005 | Should native FastVideo production code depend on DreamX reference-only packages such as `xfuser` or OpenCV? | user | Phase 3 | resolved | No. These packages may be used only for official reference/local parity setup; native FastVideo integration must remove that runtime requirement. |
|
||||
| Q006 | Should AR handoff require a full generated long-horizon video in this no-HF-token/no-GPU-budget pass? | user/runtime | pipeline | resolved | A40 long-horizon generation passes with 1005 frames at 64x64/4 steps. AR default SSIM passes locally on A40 and on Modal L40S with a 192x192/81-frame reference. HF upload/publication remains a separate operation if the reference dataset should be updated upstream. |
|
||||
| Q007 | Can A40 references stand in for L40S CI references? | quality | release | resolved | No. After deterministic AR noise fix, A40-vs-L40S reference mean SSIM is 0.7770, below the 0.98 same-device threshold. Publish the L40S-specific reference for CI. |
|
||||
|
||||
## Issues And Blockers
|
||||
|
||||
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|---|---|
|
||||
| I001 | prep | official_env | medium | Official import initially failed because `xfuser` was missing. | `ModuleNotFoundError: No module named 'xfuser'` from `python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b"` | prep | resolved | Installed `xfuser==0.4.1`; import progressed. |
|
||||
| I002 | prep | official_env | medium | Official import then failed because GUI OpenCV required missing system `libxcb.so.1`. | `ImportError: libxcb.so.1: cannot open shared object file` through `cv2` import in Diffusers ConsisID path. | prep | resolved | Installed `opencv-python-headless`; `import inference_dreamx5b` passed. |
|
||||
| I003 | prep | weights | medium | HF repo has raw official transformer shards and no Diffusers `model_index.json`. | `inspect_hf_layout.py GD-ML/DreamX-World-5B-Cam --json` returned `source_layout=raw_official`, `needs_conversion=yes`, `model_index_class=null`. | conversion | resolved | Downloaded raw DreamX shards to `official_weights/dreamx_world`; converted transformer to `converted_weights/dreamx_world/transformer`; symlinked reusable Wan2.2 Diffusers components; real 5B transformer strict-load passes. |
|
||||
| I004 | prep | dependencies | high | Official reference import required extra packages in the local environment, but FastVideo native runtime should not inherit those dependencies. | `xfuser==0.4.1` and `opencv-python-headless` were installed only to make `DreamX-World/inference_dreamx5b.py` import for reference/parity. | pipeline | resolved | Production DreamX FastVideo code uses native camera/image/video utilities and has no runtime `xfuser` or OpenCV import requirement; those packages remain reference-only local parity dependencies. |
|
||||
| I005 | parity | transformer | medium | Transformer full forward parity initially failed in bf16 official harness. | Official CUDA bf16 LayerNorm path was unstable; fp32 small-input harness avoids that dtype issue and compares against FastVideo with single-process SP identity patches. | component:transformer | resolved | Official-vs-FastVideo forward parity now passes on CUDA with `diff_max=0.072533`, `diff_mean=0.008014`. |
|
||||
| I006 | parity | vae/text_encoder | medium | VAE/text parity initially remained skipped after weights were staged. | Text official import needed a reference-only `xfuser` stub; VAE comparison initially used raw official normalized latents against FastVideo raw mu. | component:vae,component:text_encoder | resolved | Text hidden-state parity passes. VAE encode parity passes after raw key mapping and applying the official latent normalization to FastVideo output. |
|
||||
| I007 | quality | pipeline_ti2v | medium | Real image-path TI2V smoke initially failed when `vae_cpu_offload=True` because DenoisingStage encoded the first frame while the VAE weights remained on CPU. | DreamX SSIM first run failed with `RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same` at `fastvideo/pipelines/stages/denoising.py` VAE encode. | pipeline | resolved | DenoisingStage now moves the VAE to `local_device` before TI2V first-frame encode; image-path pipeline smoke and DreamX SSIM both pass. |
|
||||
| I008 | quality | ssim_helper | high | SSIM helper could compare against a stale generated base MP4 when a rerun saved the new video as `_1.mp4`. | Existing generated outputs made AR reference seeding appear to pass before a fresh generated-vs-reference compare. | quality | resolved | `run_text_to_video_similarity_test` and `run_image_to_video_similarity_test` now remove the stale generated base MP4 before generation. A40 and Modal L40S AR SSIM were rerun after the fix. |
|
||||
| I009 | quality | ar_denoising | high | AR denoising added CUDA noise without using the request seed when the original generator was CPU-backed. | Fresh reruns against old AR references produced mean SSIM near 0.05. | pipeline | resolved | `DreamXWorldARCausalDenoisingStage` now derives a device-local generator from the request seed for AR noise. Fresh same-device A40 and L40S reruns pass with mean SSIM 1.0 after reseeding references. |
|
||||
|
||||
## Escape Hatches
|
||||
|
||||
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|
|
||||
|
||||
## Decisions
|
||||
|
||||
| Date | Decision | Rationale | Impact |
|
||||
|---|---|---|---|
|
||||
| 2026-06-29 | First PR scope is `DreamX-World-5B-Cam` only. | Cam mode is closest to existing Wan2.2 TI2V support; AR forcing needs separate causal/KV pipeline work. | Component inventory and parity focus on `inference_dreamx5b.py` and `pipeline_dreamxworld.py`. |
|
||||
| 2026-06-29 | Do not install full DreamX requirements during prep. | Full requirements pin core FastVideo stack packages. | Installed only `xfuser==0.4.1` and `opencv-python-headless` to make official imports work. |
|
||||
| 2026-06-29 | Treat HF DreamX-World-5B-Cam weights as raw official transformer layout requiring conversion. | HF inspection found no `model_index.json`. | Phase 5 must create `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py` after component prototype/key dumps. |
|
||||
| 2026-06-29 | Do not add DreamX reference-only dependencies to FastVideo production requirements. | The current environment should remain the FastVideo environment; extra packages are only for official reference parity. | Native DreamX integration must avoid runtime `xfuser` and OpenCV requirements unless explicitly approved later. |
|
||||
| 2026-06-29 | Implement DreamX camera conditioning as native FastVideo utility. | It is weightless and removes the need to import DreamX reference utilities at production runtime. | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` now has non-skip parity against official action-to-PRoPE tensors. |
|
||||
| 2026-06-29 | First PR supports DreamX default `Flow` sampler only. | FastVideo FlowMatch Euler scheduler matches the official Diffusers scheduler for DreamX defaults. | Pipeline work can use FastVideo native FlowMatch scheduler; UniPC and DPM++ are out of first-PR scope. |
|
||||
| 2026-07-01 | Keep DreamX PRoPE/control adapter in a dedicated DreamX DiT class. | Project guidance is that putting too much DreamX behavior into Wan makes the model hard to manage. | `fastvideo/models/dits/dreamx_world.py` defines `DreamXWorldTransformer3DModel`, `DreamXWorldTransformerBlock`, and `DreamXPropeSelfAttention`; `fastvideo/configs/models/dits/dreamx_world.py` owns DreamX adapter config fields; Wan DiT/config are unchanged from DreamX. |
|
||||
| 2026-06-30 | Camera parity test loads official camera functions by file instead of importing the official `utils` package. | Official package initialization pulls unrelated dependencies that can require GUI OpenCV system libraries. | Camera parity remains non-skip without adding DreamX reference-only dependencies to FastVideo production requirements. |
|
||||
| 2026-06-30 | Add DreamX-World-5B-Cam model and pipeline config helpers plus a conversion script. | Official HF DreamX 5B-Cam transformer config is 30 layers, hidden size 3072, 24 heads, 48 latent channels, plus Wan2.2 48-channel VAE and UMT5-XXL text encoder. | DreamX helpers wire DiT/VAE/UMT5/Flow/TI2V settings; `dreamx_world_to_diffusers.py` writes a FastVideo-loadable transformer config plus renamed safetensors; strict-load smoke passes on a tiny official DreamX transformer and the real 5B converted shards. |
|
||||
| 2026-06-30 | Pass DreamX camera PRoPE condition through the FastVideo batch/denoising path. | DreamX transformer expects `y_camera={"viewmats", "K"}` at denoising time. | `DreamXWorldPipeline` is registered as a basic pipeline entry and initializes the official default FlowMatch scheduler; `dreamx_world_5b_cam` preset mirrors official 5B-Cam defaults; `DreamXWorldCameraConditioningStage` writes `batch.extra["dreamx_y_camera"]`; generic denoising filters and forwards it as `y_camera` only for compatible transformers. |
|
||||
|
||||
## Handoff Notes
|
||||
|
||||
- Prep, component parity, pipeline smoke/parity, and the basic example validation are complete for `DreamX-World-5B-Cam`.
|
||||
- Official reference clone is staged at `DreamX-World/` and ignored by git.
|
||||
- Workspace-local weights are staged: DreamX raw transformer shards under `official_weights/dreamx_world`, Wan2.2 raw base artifacts under `official_weights/Wan2.2-TI2V-5B`, and Wan2.2 Diffusers reusable components under `official_weights/Wan2.2-TI2V-5B-Diffusers`.
|
||||
- Camera conditioning parity is active and passing without weights.
|
||||
- Default Flow scheduler parity is active and passing without weights.
|
||||
- Transformer has corrected official 5B-Cam architecture in dedicated DreamX DiT/config files, native PRoPE/control-adapter, conversion mapping, real converted 5B strict-load, and official-vs-FastVideo forward parity passing on CUDA. VAE encode parity and text hidden-state parity pass on CUDA. Pipeline entry/registry, local model_info resolution, preset, config, FlowMatch scheduler init, camera stage, denoising `y_camera` kwarg smokes, independent CUDA pipeline smoke/parity, and a small saved-video basic example pass.
|
||||
- Full local DreamX component suite is non-skip PASS: `python -m pytest tests/local_tests/dreamx_world/ -q -rs` returned `26 passed` on 2026-07-01. Pipeline smoke/parity and SSIM quality regression are also non-skip PASS locally.
|
||||
- Keep `xfuser` and OpenCV as reference-only parity dependencies. Do not add them
|
||||
to FastVideo requirements or production imports.
|
||||
|
||||
|
||||
## Quality Regression
|
||||
|
||||
- status: `added`
|
||||
- test: `fastvideo/tests/ssim/test_dreamx_world_similarity.py`
|
||||
- command: `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs`
|
||||
- result: `1 passed, 0 skipped` on 2026-07-01
|
||||
- reference: Local A40/TORCH_SDPA reference seeded from the generated candidate under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`. The test uses a deterministic generated input image, 64x64 request dimensions, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. Full-quality params are present for 480x832/161 frames/30 steps.
|
||||
- note: Modal L40S seeding passed using the configured Modal profile and unauthenticated HF public downloads. HF upload/publication still requires `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` with write access; no token values were used or recorded.
|
||||
|
||||
## Final Handoff
|
||||
|
||||
```text
|
||||
final_handoff:
|
||||
prep_handoff_complete: yes
|
||||
conversion_status: pass
|
||||
components:
|
||||
- name: transformer
|
||||
reuse_or_port: ported_dedicated_dit
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: vae
|
||||
reuse_or_port: reused
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: text_encoder_tokenizer
|
||||
reuse_or_port: reused
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: scheduler
|
||||
reuse_or_port: reused
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: camera_conditioning
|
||||
reuse_or_port: ported
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: ar_transformer
|
||||
reuse_or_port: ported
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: ar_pipeline
|
||||
reuse_or_port: ported
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py plus AR smoke/SSIM commands listed above
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
pipeline_smoke: pass
|
||||
pipeline_parity: pass
|
||||
example_status: pass
|
||||
quality_regression: added
|
||||
local_tests_readme: tests/local_tests/dreamx_world/README.md
|
||||
port_state_file: tests/local_tests/dreamx_world/PORT_STATUS.md
|
||||
token_values_committed: no
|
||||
runtime_third_party_model_imports: none
|
||||
blockers: none
|
||||
escape_hatch: none
|
||||
```
|
||||
|
||||
| 2026-07-02 | Add DreamX-World-5B autoregressive support. | Official AR repo is raw single-safetensors layout and needs a dedicated causal/KV stage. | Added native AR DiT/config, identity converter, AR pipeline config/preset/registry, and targeted non-skip tests. Raw AR weights are staged at `/tmp/dreamx_world_ar_weights`; converted symlink layout at `/tmp/converted_dreamx_world_ar`. |
|
||||
| 2026-07-02 | Validate DreamX-World-5B AR short full generation on A40. | Targeted parity/config tests prove components, but end-to-end runtime can still fail at scheduler, RoPE cache, device, or decode boundaries. | `DreamXWorldARPipeline` generated and saved a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`; the saved MP4 decodes to 9 frames. |
|
||||
| 2026-07-02 | Validate DreamX-World-5B AR long-horizon and default SSIM on A40. | AR needs coverage beyond the 9-frame smoke to exercise longer KV/cache progression and a quality regression path. | 1005-frame 64x64/4-step generation passes and decodes; default AR SSIM test uses 192x192/81 frames because MS-SSIM requires short side >160. |
|
||||
| 2026-07-02 | Validate DreamX-World-5B AR default SSIM on Modal L40S. | CI references are device-specific; A40 alone is not enough for L40S reference coverage. | Modal checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded existing default L40S references, seeded the missing AR reference, reran SSIM successfully (`mean_ssim=1.0`), and exported the L40S reference back to the workspace. |
|
||||
@@ -1,254 +0,0 @@
|
||||
# DreamX World Local Tests
|
||||
|
||||
Local-only parity and smoke tests for the `dreamx_world` FastVideo port. These
|
||||
tests compare FastVideo against the official DreamX-World reference
|
||||
implementation and are not expected to run in CI unless explicitly promoted
|
||||
later.
|
||||
|
||||
Port progress, open questions, issues, and handoff notes live in
|
||||
`tests/local_tests/dreamx_world/PORT_STATUS.md`.
|
||||
|
||||
## Reference Assets
|
||||
|
||||
| Field | Value |
|
||||
|---|---|
|
||||
| Model family | `dreamx_world` |
|
||||
| First-PR scope | `DreamX-World-5B-Cam`; follow-up scope now includes `DreamX-World-5B` autoregressive forcing |
|
||||
| Out-of-scope variants | none for the DreamX-World 5B/Cam paths currently ported |
|
||||
| Workload types | I2V camera-control compatibility shim: image + prompt + action sequence to video |
|
||||
| Official reference | `https://github.com/AMAP-ML/DreamX-World` |
|
||||
| Local reference dir | `DreamX-World/` |
|
||||
| Official commit/version | `221875811ba31f7eac6c3025b215c09ad2cefd1d` |
|
||||
| HF weights | `GD-ML/DreamX-World-5B-Cam` |
|
||||
| HF revision | default |
|
||||
| Local weights dir | `official_weights/dreamx_world` |
|
||||
| Source layout | `raw_official` |
|
||||
| Needs conversion | `yes` |
|
||||
|
||||
Do not write token values in this file. Current token env var detected during
|
||||
prep: `none`.
|
||||
|
||||
## Shared Environment Setup
|
||||
|
||||
Run from the FastVideo repo root in the same conda/env used for FastVideo. Do
|
||||
not create a separate upstream environment for parity tests.
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/clone_reference_repo.py" \
|
||||
"https://github.com/AMAP-ML/DreamX-World.git" \
|
||||
"DreamX-World" \
|
||||
--commit "221875811ba31f7eac6c3025b215c09ad2cefd1d" \
|
||||
--update-gitignore
|
||||
```
|
||||
|
||||
DreamX-World does not expose a packaging file for editable install. During prep
|
||||
the official import check used `sys.path.insert(0, "DreamX-World")`.
|
||||
|
||||
Additional official deps installed into the current environment for imports:
|
||||
|
||||
```bash
|
||||
uv pip install xfuser==0.4.1
|
||||
uv pip install opencv-python-headless
|
||||
```
|
||||
|
||||
These packages are for running the official DreamX reference during local
|
||||
parity only. They must not become FastVideo production/runtime dependencies for
|
||||
the native `dreamx_world` pipeline.
|
||||
|
||||
Do not install the full `DreamX-World/requirements.txt` without explicit
|
||||
approval. It pins core FastVideo stack packages including `torch`, `torchvision`,
|
||||
`triton`, `flash_attn`, and `diffusers`.
|
||||
|
||||
## Official Environment Status
|
||||
|
||||
```text
|
||||
dependency_changes: installed official deps in current env
|
||||
official_env_status: imports_ok
|
||||
private_dep_stubs: none
|
||||
blocked_on: none
|
||||
```
|
||||
|
||||
Import check used during prep:
|
||||
|
||||
```bash
|
||||
python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b; print('imports_ok')"
|
||||
```
|
||||
|
||||
## Weight Setup
|
||||
|
||||
HF layout inspection found no root `model_index.json`; the repo contains
|
||||
`config.json`, a safetensors index, and three transformer safetensors shards.
|
||||
This is a raw official transformer layout and requires conversion before
|
||||
FastVideo can load it through `VideoGenerator.from_pretrained`.
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/inspect_hf_layout.py" \
|
||||
"GD-ML/DreamX-World-5B-Cam" \
|
||||
--json
|
||||
```
|
||||
|
||||
Weights have been staged workspace-locally. The raw DreamX transformer repo lives at `official_weights/dreamx_world`; Wan2.2 raw base artifacts live at `official_weights/Wan2.2-TI2V-5B`; Wan2.2 Diffusers reusable components live at `official_weights/Wan2.2-TI2V-5B-Diffusers`. To reproduce the DreamX download:
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
|
||||
"GD-ML/DreamX-World-5B-Cam" \
|
||||
"official_weights/dreamx_world"
|
||||
```
|
||||
|
||||
|
||||
|
||||
### DreamX-World-5B Autoregressive Setup
|
||||
|
||||
The AR repository `GD-ML/DreamX-World-5B` is also raw official layout: no
|
||||
`model_index.json`, root `config.json`, and a single `model.safetensors`. The
|
||||
current environment has the raw AR checkpoint staged outside the workspace at
|
||||
`/tmp/dreamx_world_ar_weights` to avoid workspace quota pressure. The converted
|
||||
FastVideo layout is staged at `/tmp/converted_dreamx_world_ar` with the 21GB
|
||||
transformer safetensors symlinked instead of copied.
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
|
||||
"GD-ML/DreamX-World-5B" \
|
||||
"/tmp/dreamx_world_ar_weights"
|
||||
|
||||
python scripts/checkpoint_conversion/dreamx_world_ar_to_diffusers.py \
|
||||
--source /tmp/dreamx_world_ar_weights \
|
||||
--output /tmp/converted_dreamx_world_ar \
|
||||
--component-source official_weights/Wan2.2-TI2V-5B-Diffusers \
|
||||
--symlink-components \
|
||||
--symlink-transformer
|
||||
```
|
||||
|
||||
AR production code uses `DreamXWorldARTransformer3DModel`,
|
||||
`DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, and
|
||||
`DreamXWorldARCausalDenoisingStage`. The AR DiT is a native FastVideo port of
|
||||
the official Apache-2.0 `CausalWanModel`; it has no production DreamX, Diffusers
|
||||
model-class, Transformers model-class, `xfuser`, or OpenCV import.
|
||||
|
||||
## Prototype And Conversion Artifacts
|
||||
|
||||
State-dict key/shape dumps are generated after FastVideo native prototypes exist
|
||||
and are used to build the conversion mapping.
|
||||
|
||||
```text
|
||||
official_key_dumps:
|
||||
transformer: converted_weights/dreamx_world/_mapping/transformer_official_keys.json
|
||||
fastvideo_key_dumps:
|
||||
transformer: converted_weights/dreamx_world/_mapping/transformer_fastvideo_keys.json
|
||||
conversion_script: scripts/checkpoint_conversion/dreamx_world_to_diffusers.py
|
||||
conversion_script_status: transformer_model_index_and_config_consistency_smoke_pass
|
||||
conversion_source_layout: raw_official
|
||||
converted_weights_dir: converted_weights/dreamx_world
|
||||
model_index_status: smoke_pass
|
||||
strict_load_status: pass
|
||||
```
|
||||
|
||||
The converter writes `transformer/` from raw DreamX shards. To create a full
|
||||
FastVideo-loadable diffusers-style root, pass a Wan2.2 Diffusers directory as the
|
||||
component source so reusable components are copied or symlinked before
|
||||
`model_index.json` is emitted:
|
||||
|
||||
```bash
|
||||
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
|
||||
--source official_weights/dreamx_world \
|
||||
--output converted_weights/dreamx_world \
|
||||
--component-source /path/to/Wan2.2-TI2V-5B-Diffusers \
|
||||
--symlink-components
|
||||
```
|
||||
|
||||
## Expected Parity Tests
|
||||
|
||||
Planned local tests for this family:
|
||||
|
||||
| Component | Official files / args | Test | Concerns | Status |
|
||||
|---|---|---|---|---|
|
||||
| transformer | `DreamX-World/models/wan_transformer3d.py`; instantiated in `DreamX-World/inference_dreamx5b.py` with `cam_method=prope`, `add_control_adapter=True` | `tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py` | FastVideo uses dedicated `DreamXWorldTransformer3DModel`/`DreamXWorldConfig` files for DreamX 5B-Cam config, PRoPE, conversion mapping, real converted 5B strict-load PASS, and official-vs-FastVideo small-input forward parity PASS on CUDA. Wan DiT/config have no DreamX-specific adapter fields. | strict_load_and_forward_parity_pass |
|
||||
| vae | `DreamX-World/models/wan_vae3_8.py`; `vae_type=AutoencoderKLWan3_8`, `vae_subpath=Wan2.2_VAE.pth` | `tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py` | DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE; encode parity passes after applying official latent normalization. | encode_parity_pass |
|
||||
| text_encoder/tokenizer | `DreamX-World/models/wan_text_encoder.py`; T5 path from Wan2.2 base model | `tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py` | Official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA with staged Wan2.2 weights/tokenizer. | hidden_state_parity_pass |
|
||||
| scheduler | Diffusers `FlowMatchEulerDiscreteScheduler`; selected by default `sampler_name=Flow` | `tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py` | First PR can support official default `Flow`; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. | non_skip_pass |
|
||||
| camera_conditioning | `DreamX-World/utils/inference_utils.py`, `models/prope_utils.py`, `wan/modules/camera_prope.py` | `tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py` | Action sequence to PRoPE/control input must match official tensor shapes and values. | non_skip_pass |
|
||||
| pipeline | `DreamX-World/pipeline/pipeline_dreamxworld.py`; call path in `DreamX-World/inference_dreamx5b.py` | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py` | DreamX PipelineConfig wires first-scope DiT/VAE/UMT5/Flow/TI2V settings, official `shift=3.0`, default preset values, FlowMatch scheduler initialization, and local `model_index.json` resolution; independent pipeline smoke covers real CUDA local load + latent generation, and parity compares public API output to worker-side explicit ForwardBatch execution. | pipeline_smoke_and_parity_pass |
|
||||
| ar_transformer | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py`; instantiated by `DreamX-World/inference_ar_forcing.py` as `CausalWanModel` with `local_attn_size=12`, `sink_size=3`, `attn_compress=4` | `tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py` | Native `DreamXWorldARTransformer3DModel` keeps official identity key layout; tiny official-vs-FastVideo forward parity passes; real 5B AR safetensors strict-load passes from `/tmp/converted_dreamx_world_ar`. | tiny_forward_parity_and_real_strict_load_pass |
|
||||
| ar_pipeline | `DreamX-World/pipeline/pipeline_causal_camera.py`; AR block/KV/context-noise loop | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; A40 short full-generation smoke | Dedicated `DreamXWorldARCausalDenoisingStage` implements blockwise KV forcing; raw HF repo must be converted before `VideoGenerator.from_pretrained` because it has no `model_index.json`; converted AR layout generated a 64x64/9-frame/4-step MP4 on A40. | config_registry_and_short_full_generation_pass |
|
||||
|
||||
Include reused components in parity. Reuse is accepted only after the FastVideo
|
||||
component definition and official instantiation arguments have both been checked
|
||||
and the component parity test passes non-skip.
|
||||
|
||||
Run the relevant tests with:
|
||||
|
||||
```bash
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_conversion.py -v -s
|
||||
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
|
||||
```
|
||||
|
||||
## Current Local Results
|
||||
|
||||
```bash
|
||||
python -m pytest tests/local_tests/dreamx_world/ -v -s
|
||||
# 2026-07-01: 26 passed, 0 skipped
|
||||
# 2026-07-02: AR targeted suite passed: 16 passed, 0 skipped
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py
|
||||
# 2026-07-02: DreamX-World-5B AR short full-generation smoke passed on A40
|
||||
# output: outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4
|
||||
# decoded: 9 frames, (64, 64, 3), uint8
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py
|
||||
# 2026-07-02: DreamX-World-5B AR long-horizon generation passed on A40
|
||||
# output: outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4
|
||||
# decoded: 1005 frames, (64, 64, 3)
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download
|
||||
# 2026-07-02: DreamX-World-5B AR default SSIM passed: 1 passed, 0 skipped
|
||||
# local A40 reference: fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/
|
||||
|
||||
modal run /tmp/modal_dreamx_ar_ssim_git.py
|
||||
# 2026-07-02: Modal L40S default SSIM passed: 1 passed, 0 skipped
|
||||
# checked out post-fix dreamx-world-5b-cam branch commit
|
||||
# first downloaded default L40S references with reference_videos_cli.py download
|
||||
# seeded missing AR L40S reference, then reran a fresh generated-vs-reference compare
|
||||
# Modal JSON: mean_ssim=1.0, min_ssim=1.0, max_ssim=1.0
|
||||
# local L40S reference: fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/DreamX-World-5B/TORCH_SDPA/
|
||||
# A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944
|
||||
|
||||
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
|
||||
# 2026-06-30: 4 passed, 0 skipped
|
||||
# 2026-07-01: smoke image-path TI2V coverage passed separately with 3 passed, 0 skipped
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs
|
||||
# 2026-07-01: 1 passed, 0 skipped
|
||||
```
|
||||
|
||||
`camera_conditioning`, default `Flow` scheduler, DreamX component/pipeline configs, default preset, pipeline entry/registry, FlowMatch scheduler initialization, DreamX camera stage, denoising `y_camera` pass-through, conversion model-index/config-consistency checks, real converted 5B dedicated DreamX transformer strict-load and forward parity on CUDA, VAE encode parity on CUDA, text hidden-state parity on CUDA, native DreamX PRoPE branch structure smoke, AR transformer parity/strict-load, AR config/registry, and AR short full-generation smoke are non-skip PASS results. The full local DreamX component suite, independent pipeline smoke/parity suite, image-path TI2V smoke, and SSIM quality regression currently have zero skips. The basic 5B-Cam example was run against `converted_weights/dreamx_world` with a 64x64/9-frame/1-step saved-video smoke, and imageio decoded the generated MP4 successfully; the AR pipeline separately saved and decoded a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`.
|
||||
|
||||
## Review Notes
|
||||
|
||||
- Required before handoff: non-skip PASS for each required component parity
|
||||
test, including reused components that own weights or numerical behavior.
|
||||
- First PR scope originally targeted `DreamX-World-5B-Cam`; scope was later
|
||||
expanded to include `DreamX-World-5B` AR support with a separate causal/KV
|
||||
pipeline.
|
||||
- AR support has targeted parity/config coverage, short full-generation smoke,
|
||||
1005-frame A40 long-horizon generation, local A40 default SSIM coverage,
|
||||
and Modal L40S default SSIM coverage. L40S validation was rerun after fixing
|
||||
stale generated-output comparison in the SSIM helper and deterministic AR
|
||||
noise seeding. HF reference publication remains a separate token-gated
|
||||
release operation.
|
||||
- FastVideo production code must not require `xfuser` or OpenCV just because the
|
||||
official reference import needed them. Port camera/action preprocessing and
|
||||
sequence-parallel behavior into existing FastVideo-native utilities or keep
|
||||
reference-only imports inside local parity tests.
|
||||
- Review agents should verify setup commands still match the PR, then run the
|
||||
listed parity tests or report the exact blocker.
|
||||
|
||||
|
||||
## Quality Regression
|
||||
|
||||
Quality regression is added in `fastvideo/tests/ssim/test_dreamx_world_similarity.py`. The 5B-Cam default test uses the workspace converted model root, `TORCH_SDPA`, a deterministic generated conditioning image, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. A local A40 reference was seeded under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`, and the test passed non-skip locally on 2026-07-01. The AR default test uses `/tmp/converted_dreamx_world_ar` or `/root/data/dreamx_world_ar_converted`, `TORCH_SDPA`, 192x192, 81 frames, 4 steps, seed 2048, and min SSIM 0.98; local A40 and Modal L40S references were seeded under `fastvideo/tests/ssim/reference_videos/default/{A40,L40S}_reference_videos/DreamX-World-5B/TORCH_SDPA/`, and both tests passed non-skip on 2026-07-02 after fresh generated-output cleanup was added to the SSIM helper. Modal L40S validation checked out `post-fix dreamx-world-5b-cam branch commit`, seeded the missing AR L40S reference, reran the test, and wrote `mean_ssim=1.0`. A cross-device A40-vs-L40S reference spot check produced mean SSIM 0.7770, so CI should use the L40S-specific reference rather than the A40 artifact. Full-quality params are present for 5B-Cam 480x832/161 frames/30 steps and AR 704x1280/1005 frames/4 steps; publishing references to the HF dataset remains a release operation requiring a write-capable HF token env var, never a raw token value.
|
||||
@@ -1 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -1,58 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B autoregressive conversion smoke tests."""
|
||||
|
||||
import json
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
|
||||
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
|
||||
from scripts.checkpoint_conversion.dreamx_world_ar_to_diffusers import (
|
||||
MODEL_INDEX,
|
||||
REUSED_COMPONENTS,
|
||||
TRANSFORMER_CONFIG,
|
||||
convert_transformer,
|
||||
write_model_index,
|
||||
)
|
||||
|
||||
|
||||
def test_dreamx_world_ar_converter_writes_symlinked_transformer_and_model_index(tmp_path):
|
||||
source = tmp_path / "raw"
|
||||
source.mkdir()
|
||||
raw_tensor = source / "model.safetensors"
|
||||
raw_tensor.write_bytes(b"placeholder")
|
||||
component_source = tmp_path / "wan22"
|
||||
output = tmp_path / "dreamx_ar"
|
||||
|
||||
for component in REUSED_COMPONENTS:
|
||||
component_dir = component_source / component
|
||||
component_dir.mkdir(parents=True)
|
||||
(component_dir / "config.json").write_text("{}\n")
|
||||
|
||||
convert_transformer(source, output, symlink_transformer=True)
|
||||
write_model_index(output, component_source, symlink_components=True)
|
||||
|
||||
assert (output / "transformer" / "model.safetensors").is_symlink()
|
||||
model_index = json.loads((output / "model_index.json").read_text())
|
||||
assert model_index == MODEL_INDEX
|
||||
assert model_index["_class_name"] == "DreamXWorldARPipeline"
|
||||
assert model_index["transformer"] == ["diffusers", "DreamXWorldARTransformer3DModel"]
|
||||
|
||||
|
||||
def test_dreamx_world_ar_transformer_config_matches_pipeline_dit_config():
|
||||
dit_config = make_dreamx_world_5b_ar_dit_config()
|
||||
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldARTransformer3DModel"
|
||||
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
|
||||
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
|
||||
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
|
||||
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
|
||||
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
|
||||
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
|
||||
assert TRANSFORMER_CONFIG["local_attn_size"] == dit_config.local_attn_size
|
||||
assert TRANSFORMER_CONFIG["sink_size"] == dit_config.sink_size
|
||||
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
|
||||
assert tuple(TRANSFORMER_CONFIG["cam_self_attn_layers"]) == dit_config.cam_self_attn_layers
|
||||
|
||||
|
||||
def test_dreamx_world_ar_model_index_component_classes_are_registered():
|
||||
for component in ("scheduler", "text_encoder", "transformer", "vae"):
|
||||
class_name = MODEL_INDEX[component][1]
|
||||
assert class_name in _LEGACY_FAST_VIDEO_MODELS
|
||||
@@ -1,178 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B autoregressive transformer parity.
|
||||
|
||||
Coverage scope: both. The tiny forward parity compares FastVideo's native AR DiT
|
||||
against the official DreamX ``CausalWanModel`` implementation with identical
|
||||
weights. The real 5B checkpoint gate strict-loads the downloaded safetensors
|
||||
through the FastVideo model class.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARArchConfig, DreamXWorldARConfig
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
|
||||
from fastvideo.models.dits.dreamx_world_ar import DreamXWorldARTransformer3DModel
|
||||
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
|
||||
from fastvideo.models.loader.utils import get_param_names_mapping
|
||||
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
CONVERTED_AR_DIR = Path(os.getenv("DREAMX_WORLD_AR_CONVERTED_DIR", "/tmp/converted_dreamx_world_ar"))
|
||||
PARITY_SCOPE = "both"
|
||||
|
||||
|
||||
def _tiny_config() -> DreamXWorldARConfig:
|
||||
return DreamXWorldARConfig(
|
||||
arch_config=DreamXWorldARArchConfig(
|
||||
num_attention_heads=1,
|
||||
attention_head_dim=8,
|
||||
in_channels=4,
|
||||
out_channels=4,
|
||||
ffn_dim=16,
|
||||
num_layers=1,
|
||||
text_dim=8,
|
||||
freq_dim=8,
|
||||
text_len=4,
|
||||
local_attn_size=2,
|
||||
sink_size=1,
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=(0,),
|
||||
))
|
||||
|
||||
|
||||
def _load_official_tiny():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official DreamX reference missing: {OFFICIAL_REF_DIR}")
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
try:
|
||||
from wan.modules import attention as official_attention
|
||||
from wan.modules import causal_camera_model_2_2_prope_infinity as causal_module
|
||||
from wan.modules import model_2_2 as official_model_2_2
|
||||
from wan.modules.causal_camera_model_2_2_prope_infinity import CausalWanModel
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official AR transformer: {exc}")
|
||||
official_attention.FLASH_ATTN_2_AVAILABLE = False
|
||||
official_attention.FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
def _sdpa_same_dtype(q, k, v, **kwargs):
|
||||
del kwargs
|
||||
out = torch.nn.functional.scaled_dot_product_attention(
|
||||
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=0.0)
|
||||
return out.transpose(1, 2).contiguous()
|
||||
|
||||
official_attention.attention = _sdpa_same_dtype
|
||||
official_attention.flash_attention = _sdpa_same_dtype
|
||||
official_model_2_2.flash_attention = _sdpa_same_dtype
|
||||
causal_module.attention = _sdpa_same_dtype
|
||||
return CausalWanModel(
|
||||
model_type="ti2v",
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=4,
|
||||
in_dim=4,
|
||||
dim=8,
|
||||
ffn_dim=16,
|
||||
freq_dim=8,
|
||||
text_dim=8,
|
||||
out_dim=4,
|
||||
num_heads=1,
|
||||
num_layers=1,
|
||||
local_attn_size=2,
|
||||
sink_size=1,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=(0,),
|
||||
).eval()
|
||||
|
||||
|
||||
def _make_inputs():
|
||||
torch.manual_seed(123)
|
||||
x = [torch.randn(4, 1, 4, 4)]
|
||||
t = torch.zeros(1, 4, dtype=torch.long)
|
||||
context = [torch.randn(2, 8)]
|
||||
camera = {
|
||||
"viewmats": torch.eye(4).reshape(1, 1, 4, 4).repeat(1, 4, 1, 1),
|
||||
"K": torch.eye(3).reshape(1, 1, 3, 3).repeat(1, 4, 1, 1),
|
||||
}
|
||||
kv_cache = [{
|
||||
"k": torch.zeros(1, 8, 1, 8),
|
||||
"v": torch.zeros(1, 8, 1, 8),
|
||||
"global_end_index": torch.tensor([0]),
|
||||
"local_end_index": torch.tensor([0]),
|
||||
"prope_k": torch.zeros(1, 8, 1, 8),
|
||||
"prope_v": torch.zeros(1, 8, 1, 8),
|
||||
"prope_global_end_index": torch.tensor([0]),
|
||||
"prope_local_end_index": torch.tensor([0]),
|
||||
}]
|
||||
cross_cache = [{
|
||||
"k": torch.zeros(1, 4, 1, 8),
|
||||
"v": torch.zeros(1, 4, 1, 8),
|
||||
"is_init": False,
|
||||
}]
|
||||
return x, t, context, camera, kv_cache, cross_cache
|
||||
|
||||
|
||||
def test_dreamx_world_ar_tiny_forward_matches_official():
|
||||
official = _load_official_tiny()
|
||||
# The official init_weights zero-inits the output head (head.head.weight and
|
||||
# biases), so both models would output exactly zero and the comparison would
|
||||
# pass vacuously. Randomize the head (deterministically) before copying the
|
||||
# state dict so the outputs reflect the internal computation.
|
||||
generator = torch.Generator().manual_seed(7)
|
||||
with torch.no_grad():
|
||||
official.head.head.weight.normal_(std=0.5, generator=generator)
|
||||
official.head.head.bias.normal_(std=0.5, generator=generator)
|
||||
fastvideo = DreamXWorldARTransformer3DModel(_tiny_config(), {}).eval()
|
||||
fastvideo.load_state_dict(official.state_dict(), strict=True)
|
||||
|
||||
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
|
||||
official_out = official(x=x, t=t, context=context, seq_len=4, y_camera=camera,
|
||||
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
|
||||
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
|
||||
fastvideo_out = fastvideo(x=x, t=t, context=context, seq_len=4, y_camera=camera,
|
||||
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
|
||||
assert official_out.abs().max() > 0, "official output is all-zero; parity comparison is vacuous"
|
||||
assert_close(fastvideo_out, official_out, atol=1e-5, rtol=1e-5)
|
||||
|
||||
|
||||
def test_dreamx_world_ar_5b_config_matches_official_shape():
|
||||
config = make_dreamx_world_5b_ar_dit_config()
|
||||
assert config.num_layers == 30
|
||||
assert config.num_attention_heads == 24
|
||||
assert config.attention_head_dim == 128
|
||||
assert config.hidden_size == 3072
|
||||
assert config.ffn_dim == 14336
|
||||
assert config.local_attn_size == 12
|
||||
assert config.sink_size == 3
|
||||
assert config.attn_compress == 4
|
||||
assert config.cam_self_attn_layers == tuple(range(30))
|
||||
|
||||
|
||||
def test_dreamx_world_ar_converted_5b_transformer_strict_loads():
|
||||
transformer_dir = CONVERTED_AR_DIR / "transformer"
|
||||
if not transformer_dir.exists():
|
||||
pytest.skip(f"Converted AR transformer missing: {transformer_dir}")
|
||||
with torch.device("meta"):
|
||||
model = DreamXWorldARTransformer3DModel(make_dreamx_world_5b_ar_dit_config(), {})
|
||||
incompatible = load_model_from_full_model_state_dict(
|
||||
model,
|
||||
safetensors_weights_iterator(resolve_safetensors_files(str(transformer_dir)), to_cpu=True),
|
||||
device=torch.device("cpu"),
|
||||
param_dtype=torch.bfloat16,
|
||||
strict=True,
|
||||
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
|
||||
training_mode=False,
|
||||
)
|
||||
assert incompatible.missing_keys == []
|
||||
assert incompatible.unexpected_keys == []
|
||||
assert not any(param.is_meta for param in model.parameters())
|
||||
@@ -1,95 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World camera-conditioning parity against the official reference.
|
||||
|
||||
Coverage scope: implementation_subcomponent. This verifies the weightless
|
||||
action-sequence to PRoPE camera tensor path used by DreamX-World-5B-Cam.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _load_official_functions():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
try:
|
||||
import importlib.util
|
||||
|
||||
pose_path = OFFICIAL_REF_DIR / "utils" / "pose_utils.py"
|
||||
pose_spec = importlib.util.spec_from_file_location(
|
||||
"dreamx_world_pose_utils", pose_path)
|
||||
if pose_spec is None or pose_spec.loader is None:
|
||||
raise RuntimeError(f"Cannot load DreamX pose_utils: {pose_path}")
|
||||
pose_module = importlib.util.module_from_spec(pose_spec)
|
||||
pose_spec.loader.exec_module(pose_module)
|
||||
|
||||
source = (OFFICIAL_REF_DIR / "utils" / "inference_utils.py").read_text()
|
||||
source = source.replace(
|
||||
"from .pose_utils import interpolate_camera_poses\n", "")
|
||||
namespace = {"interpolate_camera_poses": pose_module.interpolate_camera_poses}
|
||||
exec(compile(source, str(OFFICIAL_REF_DIR / "utils" / "inference_utils.py"), "exec"), namespace)
|
||||
except Exception as exc: # noqa: BLE001 - local parity should skip missing reference deps.
|
||||
pytest.skip(f"Cannot load DreamX camera reference: {exc}")
|
||||
return namespace["ActionToPoseFromID"], namespace["GetPoseEmbedsFromPosesPrope"]
|
||||
|
||||
|
||||
def _official_camera_condition(
|
||||
action_seq: list[str],
|
||||
action_speed_list: list[float],
|
||||
*,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
action_to_pose, get_pose_embeds = _load_official_functions()
|
||||
duration = -(-num_frames // len(action_seq))
|
||||
poses = action_to_pose(action_seq, action_speed_list, duration=duration)[:num_frames]
|
||||
condition, _ = get_pose_embeds(poses, height, width, len(poses), False, 0, dtype=dtype, device="cpu")
|
||||
return condition
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("action_seq", "action_speed_list", "num_frames"),
|
||||
[
|
||||
(["w"], [4], 81),
|
||||
(["wj", "d"], [4, 6], 121),
|
||||
(["i", "k", "l"], [3, 5, 2], 85),
|
||||
],
|
||||
)
|
||||
def test_dreamx_world_camera_conditioning_matches_official(action_seq, action_speed_list, num_frames):
|
||||
dtype = torch.float32
|
||||
official = _official_camera_condition(
|
||||
action_seq,
|
||||
action_speed_list,
|
||||
num_frames=num_frames,
|
||||
height=704,
|
||||
width=1280,
|
||||
dtype=dtype,
|
||||
)
|
||||
fastvideo = build_dreamx_camera_condition(
|
||||
action_seq,
|
||||
action_speed_list,
|
||||
num_frames=num_frames,
|
||||
height=704,
|
||||
width=1280,
|
||||
dtype=dtype,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
assert official.keys() == fastvideo.keys() == {"viewmats", "K"}
|
||||
for key in ("viewmats", "K"):
|
||||
assert official[key].shape == fastvideo[key].shape
|
||||
diff = (official[key] - fastvideo[key]).abs()
|
||||
print(f"{key}: shape={tuple(fastvideo[key].shape)} diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
|
||||
assert_close(fastvideo[key], official[key], atol=1e-5, rtol=1e-5)
|
||||
@@ -1,79 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World conversion script smoke tests."""
|
||||
|
||||
import json
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
|
||||
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
|
||||
|
||||
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import (
|
||||
MODEL_INDEX,
|
||||
REUSED_COMPONENTS,
|
||||
TRANSFORMER_CONFIG,
|
||||
_copy_or_link_component,
|
||||
write_model_index,
|
||||
)
|
||||
|
||||
|
||||
def test_dreamx_world_converter_writes_full_model_index_with_reused_components(tmp_path):
|
||||
component_source = tmp_path / "wan22"
|
||||
output = tmp_path / "dreamx"
|
||||
output.mkdir()
|
||||
|
||||
for component in REUSED_COMPONENTS:
|
||||
component_dir = component_source / component
|
||||
component_dir.mkdir(parents=True)
|
||||
(component_dir / "config.json").write_text("{}\n")
|
||||
|
||||
write_model_index(output, component_source, symlink_components=True)
|
||||
|
||||
model_index = json.loads((output / "model_index.json").read_text())
|
||||
assert model_index == MODEL_INDEX
|
||||
assert model_index["_class_name"] == "DreamXWorldPipeline"
|
||||
assert model_index["transformer"] == ["diffusers", "DreamXWorldTransformer3DModel"]
|
||||
for component in REUSED_COMPONENTS:
|
||||
assert (output / component).is_symlink()
|
||||
|
||||
|
||||
def test_dreamx_world_transformer_config_has_camera_adapter_enabled():
|
||||
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldTransformer3DModel"
|
||||
assert TRANSFORMER_CONFIG["add_control_adapter"] is True
|
||||
assert TRANSFORMER_CONFIG["cam_method"] == "prope"
|
||||
assert TRANSFORMER_CONFIG["num_layers"] == 30
|
||||
|
||||
|
||||
def test_dreamx_world_converter_transformer_config_matches_pipeline_dit_config():
|
||||
dit_config = make_dreamx_world_5b_cam_dit_config()
|
||||
|
||||
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
|
||||
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
|
||||
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
|
||||
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
|
||||
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
|
||||
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
|
||||
assert TRANSFORMER_CONFIG["cross_attn_norm"] == dit_config.cross_attn_norm
|
||||
assert TRANSFORMER_CONFIG["qk_norm"] == dit_config.qk_norm
|
||||
assert TRANSFORMER_CONFIG["add_control_adapter"] == dit_config.add_control_adapter
|
||||
assert TRANSFORMER_CONFIG["cam_method"] == dit_config.cam_method
|
||||
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
|
||||
assert TRANSFORMER_CONFIG["cam_self_attn_layers"] == dit_config.cam_self_attn_layers
|
||||
|
||||
|
||||
def test_dreamx_world_model_index_component_classes_are_registered():
|
||||
for component in ("scheduler", "text_encoder", "transformer", "vae"):
|
||||
class_name = MODEL_INDEX[component][1]
|
||||
assert class_name in _LEGACY_FAST_VIDEO_MODELS
|
||||
|
||||
|
||||
def test_dreamx_world_copy_or_link_component_keeps_broken_symlink(tmp_path):
|
||||
component_source = tmp_path / "wan22"
|
||||
src = component_source / "scheduler"
|
||||
src.mkdir(parents=True)
|
||||
output = tmp_path / "dreamx"
|
||||
output.mkdir()
|
||||
dst = output / "scheduler"
|
||||
dst.symlink_to(tmp_path / "missing_scheduler", target_is_directory=True)
|
||||
|
||||
_copy_or_link_component("scheduler", component_source, output, symlink=True)
|
||||
|
||||
assert dst.is_symlink()
|
||||
@@ -1,306 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World pipeline config and conditioning smoke tests."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.api.presets import get_preset
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.pipeline_registry import PipelineType, import_pipeline_classes
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
|
||||
DreamXCamera,
|
||||
_interpolate_camera_poses,
|
||||
build_dreamx_camera_condition,
|
||||
)
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import (
|
||||
DREAMX_Y_CAMERA_KEY,
|
||||
DreamXWorldCameraConditioningStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.registry import get_default_preset, get_model_info, get_pipeline_config_cls_from_name
|
||||
|
||||
|
||||
def test_dreamx_world_5b_cam_pipeline_config_wires_first_scope_components():
|
||||
config = DreamXWorld5BCamPipelineConfig()
|
||||
|
||||
assert config.flow_shift == 3.0
|
||||
assert config.ti2v_task is True
|
||||
assert config.expand_timesteps is True
|
||||
assert config.dit_config.expand_timesteps is True
|
||||
assert config.dit_config.num_layers == 30
|
||||
assert config.dit_config.add_control_adapter is True
|
||||
assert config.dit_config.cam_method == "prope"
|
||||
|
||||
assert config.vae_config.load_encoder is True
|
||||
assert config.vae_config.load_decoder is True
|
||||
assert config.vae_config.z_dim == 48
|
||||
assert config.vae_config.scale_factor_temporal == 4
|
||||
assert config.vae_config.scale_factor_spatial == 16
|
||||
|
||||
assert len(config.text_encoder_configs) == 1
|
||||
text_config = config.text_encoder_configs[0]
|
||||
assert text_config.prefix == "umt5"
|
||||
assert text_config.vocab_size == 256384
|
||||
assert text_config.d_model == 4096
|
||||
assert config.text_encoder_precisions == ("bf16",)
|
||||
|
||||
|
||||
def test_dreamx_world_pipeline_registry_discovers_entrypoint():
|
||||
pipelines = import_pipeline_classes(PipelineType.BASIC)
|
||||
|
||||
assert pipelines["basic"]["DreamXWorldPipeline"] is DreamXWorldPipeline
|
||||
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_local_model_index_resolves_model_info(tmp_path):
|
||||
model_dir = tmp_path / "DreamX-World-5B-Cam-converted"
|
||||
model_dir.mkdir()
|
||||
for component in ("scheduler", "text_encoder", "tokenizer", "transformer", "vae"):
|
||||
(model_dir / component).mkdir()
|
||||
(model_dir / "model_index.json").write_text(json.dumps({
|
||||
"_class_name": "DreamXWorldPipeline",
|
||||
"_diffusers_version": "0.31.0",
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"],
|
||||
}) + "\n")
|
||||
|
||||
info = get_model_info(str(model_dir), pipeline_type=PipelineType.BASIC, workload_type=WorkloadType.I2V)
|
||||
|
||||
assert info.pipeline_cls is DreamXWorldPipeline
|
||||
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
|
||||
|
||||
def test_dreamx_world_model_path_resolves_pipeline_config():
|
||||
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B-Cam") is DreamXWorld5BCamPipelineConfig
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_default_preset_is_registered():
|
||||
preset_name = get_default_preset("GD-ML/DreamX-World-5B-Cam")
|
||||
preset = get_preset(preset_name, "dreamx_world")
|
||||
|
||||
assert preset.name == "dreamx_world_5b_cam"
|
||||
assert preset.workload_type == "i2v"
|
||||
assert preset.defaults["height"] == 480
|
||||
assert preset.defaults["width"] == 832
|
||||
assert preset.defaults["num_frames"] == 161
|
||||
assert preset.defaults["num_inference_steps"] == 30
|
||||
assert preset.defaults["guidance_scale"] == 5.0
|
||||
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_pipeline_initializes_official_flow_scheduler():
|
||||
pipeline = DreamXWorldPipeline.__new__(DreamXWorldPipeline)
|
||||
pipeline.modules = {}
|
||||
fastvideo_args = SimpleNamespace(pipeline_config=DreamXWorld5BCamPipelineConfig())
|
||||
|
||||
pipeline.initialize_pipeline(fastvideo_args)
|
||||
|
||||
scheduler = pipeline.modules["scheduler"]
|
||||
assert isinstance(scheduler, FlowMatchEulerDiscreteScheduler)
|
||||
assert scheduler.config.shift == 3.0
|
||||
|
||||
def test_dreamx_world_camera_conditioning_stage_sets_y_camera_extra():
|
||||
batch = ForwardBatch(
|
||||
data_type="t2v",
|
||||
action_list=["wj", "d"],
|
||||
action_speed_list=[4, 6],
|
||||
num_frames=17,
|
||||
height=704,
|
||||
width=1280,
|
||||
latents=torch.zeros(1, 16, 5, 44, 80),
|
||||
)
|
||||
stage = DreamXWorldCameraConditioningStage()
|
||||
|
||||
out = stage.forward(batch, fastvideo_args=object())
|
||||
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
|
||||
expected = build_dreamx_camera_condition(
|
||||
["wj", "d"],
|
||||
[4, 6],
|
||||
num_frames=17,
|
||||
height=704,
|
||||
width=1280,
|
||||
dtype=torch.float32,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
assert set(y_camera) == {"viewmats", "K"}
|
||||
for key, expected_value in expected.items():
|
||||
assert y_camera[key].shape == (1, *expected_value.shape)
|
||||
torch.testing.assert_close(y_camera[key][0], expected_value)
|
||||
assert stage.verify_output(out, fastvideo_args=object()).is_valid()
|
||||
|
||||
|
||||
def test_dreamx_world_denoising_kwargs_filter_for_y_camera():
|
||||
y_camera = {"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)}
|
||||
stage = DenoisingStage.__new__(DenoisingStage)
|
||||
|
||||
def accepts_y_camera(hidden_states, encoder_hidden_states, timestep, y_camera=None):
|
||||
return y_camera
|
||||
|
||||
def no_y_camera(hidden_states, encoder_hidden_states, timestep):
|
||||
return hidden_states
|
||||
|
||||
assert stage.prepare_extra_func_kwargs(accepts_y_camera, {"y_camera": y_camera}) == {"y_camera": y_camera}
|
||||
assert stage.prepare_extra_func_kwargs(no_y_camera, {"y_camera": y_camera}) == {}
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_ar_pipeline_config_wires_components():
|
||||
config = DreamXWorld5BARPipelineConfig()
|
||||
assert config.is_causal is True
|
||||
assert config.flow_shift == 5.0
|
||||
assert config.dmd_denoising_steps == (1000, 750, 500, 250)
|
||||
assert config.warp_denoising_step is True
|
||||
assert config.context_noise == 0.1
|
||||
assert config.dit_config.arch_config.local_attn_size == 12
|
||||
assert config.dit_config.arch_config.sink_size == 3
|
||||
assert config.dit_config.arch_config.attn_compress == 4
|
||||
|
||||
|
||||
def test_dreamx_world_ar_pipeline_registry_and_preset():
|
||||
from fastvideo.api.presets import get_preset
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name, get_preset_selection
|
||||
|
||||
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B") is DreamXWorld5BARPipelineConfig
|
||||
preset_name, family = get_preset_selection("GD-ML/DreamX-World-5B")
|
||||
assert (preset_name, family) == ("dreamx_world_5b_ar", "dreamx_world")
|
||||
preset = get_preset("dreamx_world_5b_ar", "dreamx_world")
|
||||
assert preset.defaults["num_inference_steps"] == 4
|
||||
assert DreamXWorldARPipeline.pipeline_config_cls is DreamXWorld5BARPipelineConfig
|
||||
|
||||
|
||||
def test_dreamx_world_camera_conditioning_stage_expands_scalar_speed():
|
||||
batch = ForwardBatch(
|
||||
data_type="t2v",
|
||||
action_list=["w", "d"],
|
||||
action_speed_list=2.0,
|
||||
num_frames=17,
|
||||
height=704,
|
||||
width=1280,
|
||||
latents=torch.zeros(1, 16, 5, 44, 80),
|
||||
)
|
||||
stage = DreamXWorldCameraConditioningStage()
|
||||
|
||||
out = stage.forward(batch, fastvideo_args=object())
|
||||
|
||||
assert set(out.extra[DREAMX_Y_CAMERA_KEY]) == {"viewmats", "K"}
|
||||
|
||||
|
||||
def test_dreamx_world_camera_interpolation_handles_single_camera():
|
||||
camera = DreamXCamera(
|
||||
fx=0.8,
|
||||
fy=0.8,
|
||||
cx=0.5,
|
||||
cy=0.5,
|
||||
w2c_mat=np.eye(4, dtype=np.float64),
|
||||
)
|
||||
|
||||
out = _interpolate_camera_poses(
|
||||
[camera],
|
||||
src_indices=np.array([0.0]),
|
||||
tgt_indices=np.array([0.0, 1.0, 2.0]),
|
||||
)
|
||||
|
||||
assert out == [camera, camera, camera]
|
||||
|
||||
|
||||
def test_dreamx_world_ar_cache_initializes_camera_self_attention_entries():
|
||||
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
|
||||
transformer = SimpleNamespace(
|
||||
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn), SimpleNamespace(cam_self_attn=cam_self_attn)],
|
||||
num_attention_heads=2,
|
||||
attention_head_dim=4,
|
||||
)
|
||||
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
|
||||
stage.transformer = transformer
|
||||
stage.num_transformer_blocks = 2
|
||||
stage.local_attn_size = 6
|
||||
|
||||
caches = stage._initialize_kv_cache(
|
||||
batch_size=1,
|
||||
dtype=torch.float32,
|
||||
device=torch.device("cpu"),
|
||||
frame_seq_length=7,
|
||||
)
|
||||
|
||||
assert len(caches) == 2
|
||||
assert caches[0]["k"].shape == (1, 42, 2, 4)
|
||||
assert caches[0]["prope_k"].shape == (1, 42, 3, 5)
|
||||
assert caches[0]["prope_v"].shape == (1, 42, 3, 5)
|
||||
assert int(caches[0]["prope_global_end_index"].item()) == 0
|
||||
assert int(caches[0]["prope_local_end_index"].item()) == 0
|
||||
|
||||
|
||||
def test_dreamx_world_ar_context_noise_fraction_maps_to_scheduler_timestep():
|
||||
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(0.1) == 100
|
||||
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(100) == 100
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_ar_context_update_advances_camera_cache_indices():
|
||||
class DummyTransformer:
|
||||
def __call__(self, *, hidden_states, encoder_hidden_states, timestep, y_camera, kv_cache, crossattn_cache,
|
||||
current_start):
|
||||
del encoder_hidden_states, y_camera, crossattn_cache
|
||||
assert current_start == 0
|
||||
assert timestep.unique().tolist() == [100]
|
||||
new_tokens = timestep.shape[1]
|
||||
for cache in kv_cache:
|
||||
cache["local_end_index"] += new_tokens
|
||||
cache["global_end_index"] += new_tokens
|
||||
cache["prope_local_end_index"] += new_tokens
|
||||
cache["prope_global_end_index"] += new_tokens
|
||||
cache["k"][:, :new_tokens] = 1
|
||||
cache["prope_k"][:, :new_tokens] = 1
|
||||
return hidden_states
|
||||
|
||||
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
|
||||
cache_transformer = SimpleNamespace(
|
||||
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn)],
|
||||
num_attention_heads=2,
|
||||
attention_head_dim=4,
|
||||
)
|
||||
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
|
||||
stage.transformer = cache_transformer
|
||||
stage.num_transformer_blocks = 1
|
||||
stage.local_attn_size = 6
|
||||
caches = stage._initialize_kv_cache(
|
||||
batch_size=1,
|
||||
dtype=torch.float32,
|
||||
device=torch.device("cpu"),
|
||||
frame_seq_length=2,
|
||||
)
|
||||
# Keep the cache allocation source separate from the callable transformer used by _update_context_cache.
|
||||
stage.transformer = DummyTransformer()
|
||||
|
||||
stage._update_context_cache(
|
||||
block_latents=torch.zeros(1, 4, 3, 2, 2),
|
||||
context=[torch.zeros(2, 4)],
|
||||
camera_block={"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)},
|
||||
kv_cache=caches,
|
||||
crossattn_cache=[{}],
|
||||
start=0,
|
||||
frame_seq_length=2,
|
||||
target_dtype=torch.float32,
|
||||
autocast_enabled=False,
|
||||
context_noise=0.1,
|
||||
)
|
||||
|
||||
assert int(caches[0]["local_end_index"].item()) == 6
|
||||
assert int(caches[0]["prope_local_end_index"].item()) == 6
|
||||
assert caches[0]["k"][:, :6].sum().item() == 48
|
||||
assert caches[0]["prope_k"][:, :6].sum().item() == 90
|
||||
@@ -1,53 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World default Flow scheduler parity.
|
||||
|
||||
Coverage scope: implementation_subcomponent. DreamX-World-5B-Cam defaults to
|
||||
Diffusers FlowMatchEulerDiscreteScheduler for sampler_name=Flow. This test
|
||||
checks that FastVideo's native FlowMatchEulerDiscreteScheduler matches the
|
||||
timestep schedule and Euler step used by the official default path.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
import inspect
|
||||
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler as OfficialFlowMatchEulerDiscreteScheduler
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler as FastVideoFlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _scheduler_kwargs(cls):
|
||||
config_path = REPO_ROOT / "DreamX-World" / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
|
||||
config = OmegaConf.load(config_path)
|
||||
raw_kwargs = OmegaConf.to_container(config["scheduler_kwargs"])
|
||||
signature = inspect.signature(cls)
|
||||
return {key: value for key, value in raw_kwargs.items() if key in signature.parameters}
|
||||
|
||||
|
||||
def test_dreamx_world_flow_scheduler_timesteps_and_step_match():
|
||||
official = OfficialFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(OfficialFlowMatchEulerDiscreteScheduler))
|
||||
fastvideo = FastVideoFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(FastVideoFlowMatchEulerDiscreteScheduler))
|
||||
|
||||
official.set_timesteps(50, device="cpu", mu=1)
|
||||
fastvideo.set_timesteps(50, device="cpu", mu=1)
|
||||
assert_close(fastvideo.timesteps, official.timesteps, atol=0, rtol=0)
|
||||
assert_close(fastvideo.sigmas, official.sigmas, atol=0, rtol=0)
|
||||
|
||||
torch.manual_seed(7)
|
||||
sample = torch.randn(1, 4, 2, 8, 8)
|
||||
model_output = torch.randn_like(sample)
|
||||
timestep = official.timesteps[3]
|
||||
|
||||
official_prev = official.step(model_output, timestep, sample, return_dict=False)[0]
|
||||
fastvideo_prev = fastvideo.step(model_output, fastvideo.timesteps[3], sample, return_dict=False)[0]
|
||||
diff = (official_prev - fastvideo_prev).abs()
|
||||
print(f"scheduler diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
|
||||
assert_close(fastvideo_prev, official_prev, atol=0, rtol=0)
|
||||
@@ -1,159 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World Wan T5 encoder reuse parity scaffold.
|
||||
|
||||
Coverage scope: implementation_subcomponent. It records the official
|
||||
WanT5EncoderModel loading path and FastVideo T5 target for later activation
|
||||
with staged Wan2.2 base text encoder/tokenizer weights.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.configs.pipelines.dreamx_world import (
|
||||
DreamXWorld5BCamPipelineConfig,
|
||||
make_dreamx_world_5b_cam_text_encoder_config,
|
||||
)
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
|
||||
WAN_DIFFUSERS_DIR = Path(os.getenv("DREAMX_WORLD_WAN_DIFFUSERS_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B-Diffusers"))
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _add_official_to_path():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
if str(OFFICIAL_REF_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
|
||||
|
||||
def _install_xfuser_stub() -> None:
|
||||
if "xfuser" in sys.modules:
|
||||
return
|
||||
xfuser = types.ModuleType("xfuser")
|
||||
core = types.ModuleType("xfuser.core")
|
||||
distributed = types.ModuleType("xfuser.core.distributed")
|
||||
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
|
||||
distributed.get_sequence_parallel_rank = lambda: 0
|
||||
distributed.get_sequence_parallel_world_size = lambda: 1
|
||||
distributed.get_sp_group = lambda: None
|
||||
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
|
||||
distributed.init_distributed_environment = lambda *args, **kwargs: None
|
||||
distributed.initialize_model_parallel = lambda *args, **kwargs: None
|
||||
distributed.model_parallel_is_initialized = lambda: False
|
||||
|
||||
class XFuserLongContextAttention:
|
||||
def __call__(self, *args, **kwargs):
|
||||
raise RuntimeError("xfuser stub cannot execute attention")
|
||||
|
||||
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
|
||||
sys.modules.update({
|
||||
"xfuser": xfuser,
|
||||
"xfuser.core": core,
|
||||
"xfuser.core.distributed": distributed,
|
||||
"xfuser.core.long_ctx_attention": long_ctx,
|
||||
})
|
||||
|
||||
|
||||
def _text_kwargs():
|
||||
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
|
||||
return OmegaConf.to_container(config["text_encoder_kwargs"])
|
||||
|
||||
|
||||
def _patch_single_process_text_parallel(monkeypatch):
|
||||
import fastvideo.layers.linear as fastvideo_linear
|
||||
import fastvideo.layers.vocab_parallel_embedding as fastvideo_embedding
|
||||
import fastvideo.models.encoders.t5 as fastvideo_t5
|
||||
|
||||
for module in (fastvideo_t5, fastvideo_embedding, fastvideo_linear):
|
||||
if hasattr(module, "get_tp_rank"):
|
||||
monkeypatch.setattr(module, "get_tp_rank", lambda: 0)
|
||||
if hasattr(module, "get_tp_world_size"):
|
||||
monkeypatch.setattr(module, "get_tp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(fastvideo_embedding, "tensor_model_parallel_all_reduce", lambda x: x)
|
||||
|
||||
|
||||
def _load_official_text_encoder(device, dtype):
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
text_path = WAN_BASE_DIR / "models_t5_umt5-xxl-enc-bf16.pth"
|
||||
if not text_path.exists():
|
||||
pytest.skip(f"Wan2.2 text encoder weights missing: {text_path}")
|
||||
try:
|
||||
from models import WanT5EncoderModel
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official DreamX text encoder: {exc}")
|
||||
model = WanT5EncoderModel.from_pretrained(
|
||||
str(text_path), additional_kwargs=_text_kwargs(), low_cpu_mem_usage=True, torch_dtype=dtype
|
||||
)
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_text_encoder(device, dtype, monkeypatch):
|
||||
text_encoder_path = WAN_DIFFUSERS_DIR / "text_encoder"
|
||||
if not text_encoder_path.exists():
|
||||
pytest.skip(f"Wan2.2 Diffusers text encoder missing: {text_encoder_path}")
|
||||
_patch_single_process_text_parallel(monkeypatch)
|
||||
pipeline_config = DreamXWorld5BCamPipelineConfig()
|
||||
pipeline_config.text_encoder_configs[0]._fsdp_shard_conditions = []
|
||||
args = FastVideoArgs(
|
||||
model_path=str(text_encoder_path),
|
||||
pipeline_config=pipeline_config,
|
||||
text_encoder_cpu_offload=(device.type == "cpu"),
|
||||
)
|
||||
args.model_paths = {}
|
||||
return TextEncoderLoader().load(str(text_encoder_path), args).to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def test_dreamx_world_text_encoder_config_matches_umt5_xxl_shape():
|
||||
config = make_dreamx_world_5b_cam_text_encoder_config()
|
||||
assert config.vocab_size == 256384
|
||||
assert config.d_model == 4096
|
||||
assert config.d_kv == 64
|
||||
assert config.d_ff == 10240
|
||||
assert config.num_heads == 64
|
||||
assert config.num_layers == 24
|
||||
assert config.relative_attention_num_buckets == 32
|
||||
assert config.dropout_rate == 0.0
|
||||
assert config.text_len == 512
|
||||
assert config.prefix == "umt5"
|
||||
|
||||
|
||||
def test_dreamx_world_fastvideo_text_encoder_loads_staged_weights(monkeypatch):
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
model = _load_fastvideo_text_encoder(device, torch.bfloat16, monkeypatch)
|
||||
assert model.__class__.__name__ == "UMT5EncoderModel"
|
||||
assert next(model.parameters()).device.type == device.type
|
||||
assert next(model.parameters()).dtype == torch.bfloat16
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for text encoder parity.")
|
||||
def test_dreamx_world_text_encoder_parity_scaffold(monkeypatch):
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
official = _load_official_text_encoder(device, dtype)
|
||||
fastvideo = _load_fastvideo_text_encoder(device, dtype, monkeypatch)
|
||||
tokenizer_path = WAN_BASE_DIR / "google" / "umt5-xxl"
|
||||
if not tokenizer_path.exists():
|
||||
pytest.skip(f"Wan2.2 tokenizer missing: {tokenizer_path}")
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(tokenizer_path))
|
||||
batch = tokenizer(["A quiet forest trail at sunrise."], padding="max_length", max_length=512, return_tensors="pt")
|
||||
input_ids = batch.input_ids.to(device)
|
||||
attention_mask = batch.attention_mask.to(device)
|
||||
with torch.inference_mode():
|
||||
official_hidden = official(input_ids, attention_mask=attention_mask)[0].float().cpu()
|
||||
fastvideo_hidden = fastvideo(input_ids, attention_mask=attention_mask).last_hidden_state.float().cpu()
|
||||
assert official_hidden.shape == fastvideo_hidden.shape
|
||||
assert_close(fastvideo_hidden, official_hidden, atol=1e-3, rtol=1e-3)
|
||||
@@ -1,331 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World transformer parity scaffold.
|
||||
|
||||
Coverage scope: both. The official side loads DreamX-World-5B-Cam through
|
||||
Wan2_2Transformer3DModel.from_pretrained with PRoPE camera control enabled.
|
||||
The FastVideo side strict-loads the converted DreamX transformer weights into
|
||||
the native DreamX-World DiT implementation.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.configs.models.dits.dreamx_world import (
|
||||
DreamXWorldArchConfig, DreamXWorldConfig)
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.models.dits.dreamx_world import (
|
||||
DreamXPropeSelfAttention, DreamXWorldTransformer3DModel,
|
||||
DreamXWorldTransformerBlock)
|
||||
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
|
||||
from fastvideo.models.loader.utils import get_param_names_mapping
|
||||
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
|
||||
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import map_transformer_key
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
LOCAL_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "dreamx_world"))
|
||||
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
|
||||
CONVERTED_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / "dreamx_world"))
|
||||
PARITY_SCOPE = "both"
|
||||
|
||||
|
||||
def _install_xfuser_stub() -> None:
|
||||
if "xfuser" in sys.modules:
|
||||
return
|
||||
xfuser = types.ModuleType("xfuser")
|
||||
core = types.ModuleType("xfuser.core")
|
||||
distributed = types.ModuleType("xfuser.core.distributed")
|
||||
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
|
||||
distributed.get_sequence_parallel_rank = lambda: 0
|
||||
distributed.get_sequence_parallel_world_size = lambda: 1
|
||||
distributed.get_sp_group = lambda: None
|
||||
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
|
||||
distributed.init_distributed_environment = lambda *args, **kwargs: None
|
||||
distributed.initialize_model_parallel = lambda *args, **kwargs: None
|
||||
distributed.model_parallel_is_initialized = lambda: False
|
||||
|
||||
class XFuserLongContextAttention:
|
||||
def __call__(self, *args, **kwargs):
|
||||
raise RuntimeError("xfuser stub cannot execute attention")
|
||||
|
||||
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
|
||||
sys.modules.update({
|
||||
"xfuser": xfuser,
|
||||
"xfuser.core": core,
|
||||
"xfuser.core.distributed": distributed,
|
||||
"xfuser.core.long_ctx_attention": long_ctx,
|
||||
})
|
||||
|
||||
|
||||
def _make_tiny_dreamx_config() -> DreamXWorldConfig:
|
||||
return DreamXWorldConfig(
|
||||
arch_config=DreamXWorldArchConfig(
|
||||
num_attention_heads=1,
|
||||
attention_head_dim=8,
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
ffn_dim=32,
|
||||
num_layers=1,
|
||||
cross_attn_norm=True,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=None,
|
||||
))
|
||||
|
||||
|
||||
def _add_official_to_path() -> None:
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
if str(OFFICIAL_REF_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
|
||||
|
||||
def _official_transformer_kwargs() -> dict:
|
||||
config_path = OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
|
||||
if not config_path.exists():
|
||||
pytest.skip(f"DreamX Wan config missing: {config_path}")
|
||||
config = OmegaConf.load(config_path)
|
||||
kwargs = OmegaConf.to_container(config["transformer_additional_kwargs"])
|
||||
kwargs["cam_method"] = "prope"
|
||||
kwargs["add_control_adapter"] = True
|
||||
return kwargs
|
||||
|
||||
|
||||
def _load_official_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
if not LOCAL_WEIGHTS_DIR.exists():
|
||||
pytest.skip(f"DreamX transformer weights missing: {LOCAL_WEIGHTS_DIR}")
|
||||
try:
|
||||
from models import Wan2_2Transformer3DModel
|
||||
except Exception as exc: # noqa: BLE001 - local parity should skip missing refs.
|
||||
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
|
||||
model = Wan2_2Transformer3DModel.from_pretrained(
|
||||
str(LOCAL_WEIGHTS_DIR),
|
||||
transformer_additional_kwargs=_official_transformer_kwargs(),
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
model = _load_fastvideo_transformer_strict(torch.device("cpu"), dtype)
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_transformer_strict(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
transformer_dir = CONVERTED_WEIGHTS_DIR / "transformer"
|
||||
if not transformer_dir.exists():
|
||||
pytest.skip(f"Converted DreamX transformer missing: {transformer_dir}")
|
||||
safetensors_files = resolve_safetensors_files(str(transformer_dir))
|
||||
config = make_dreamx_world_5b_cam_dit_config()
|
||||
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
|
||||
original_get_sp_world_size = fastvideo_dreamx.get_sp_world_size
|
||||
fastvideo_dreamx.get_sp_world_size = lambda: 1
|
||||
try:
|
||||
with torch.device("meta"):
|
||||
model = DreamXWorldTransformer3DModel(config=config, hf_config={})
|
||||
finally:
|
||||
fastvideo_dreamx.get_sp_world_size = original_get_sp_world_size
|
||||
incompatible = load_model_from_full_model_state_dict(
|
||||
model,
|
||||
safetensors_weights_iterator(safetensors_files, to_cpu=True),
|
||||
device=device,
|
||||
param_dtype=dtype,
|
||||
strict=True,
|
||||
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
|
||||
training_mode=False,
|
||||
)
|
||||
assert incompatible.missing_keys == []
|
||||
assert incompatible.unexpected_keys == []
|
||||
assert not any(param.is_meta for param in model.parameters())
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _make_inputs(device: torch.device, dtype: torch.dtype):
|
||||
torch.manual_seed(1234)
|
||||
num_frames = 5
|
||||
height = 64
|
||||
width = 64
|
||||
latent_frames = (num_frames - 1) // 4 + 1
|
||||
latent_h = height // 16
|
||||
latent_w = width // 16
|
||||
x = torch.randn(1, 48, latent_frames, latent_h, latent_w, device=device, dtype=dtype)
|
||||
context = [torch.randn(16, 4096, device=device, dtype=dtype)]
|
||||
seq_len = math.ceil((latent_h * latent_w) / 4 * latent_frames)
|
||||
timestep = torch.full((1, seq_len), 250, device=device, dtype=torch.long)
|
||||
camera = build_dreamx_camera_condition(
|
||||
["w"], [4], num_frames=num_frames, height=height, width=width, dtype=dtype, device=device
|
||||
)
|
||||
camera = {key: value.unsqueeze(0) for key, value in camera.items()}
|
||||
return {"x": [x[0]], "context": context, "t": timestep, "seq_len": seq_len, "y_camera": camera}
|
||||
|
||||
|
||||
def _run_official(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
|
||||
with torch.inference_mode():
|
||||
output = model(**inputs)
|
||||
if isinstance(output, list):
|
||||
output = torch.stack(output, dim=0)
|
||||
assert torch.is_tensor(output), f"official output is not a tensor: {type(output)}"
|
||||
return output.detach().float().cpu()
|
||||
|
||||
|
||||
def _run_fastvideo(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
|
||||
hidden_states = torch.stack(inputs["x"], dim=0)
|
||||
encoder_hidden_states = torch.stack([
|
||||
torch.cat([inputs["context"][0], inputs["context"][0].new_zeros(512 - inputs["context"][0].shape[0], 4096)])
|
||||
])
|
||||
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
output = model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=inputs["t"],
|
||||
y_camera=inputs["y_camera"],
|
||||
)
|
||||
assert torch.is_tensor(output), f"FastVideo output is not a tensor: {type(output)}"
|
||||
return output.detach().float().cpu()
|
||||
|
||||
|
||||
def test_dreamx_world_conversion_mapping_strict_load_smoke(monkeypatch):
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
try:
|
||||
from models import Wan2_2Transformer3DModel
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
|
||||
|
||||
official = Wan2_2Transformer3DModel(
|
||||
dim=8,
|
||||
ffn_dim=32,
|
||||
num_heads=1,
|
||||
num_layers=1,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
)
|
||||
official_state = official.state_dict()
|
||||
diffusers_like_state = {
|
||||
map_transformer_key(key): value.detach().clone()
|
||||
for key, value in official_state.items()
|
||||
}
|
||||
|
||||
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
|
||||
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
|
||||
with torch.device("meta"):
|
||||
fastvideo = DreamXWorldTransformer3DModel(config=_make_tiny_dreamx_config(), hf_config={})
|
||||
|
||||
incompatible = load_model_from_full_model_state_dict(
|
||||
fastvideo,
|
||||
iter(diffusers_like_state.items()),
|
||||
device=torch.device("cpu"),
|
||||
param_dtype=torch.float32,
|
||||
strict=True,
|
||||
param_names_mapping=get_param_names_mapping(fastvideo.param_names_mapping),
|
||||
training_mode=False,
|
||||
)
|
||||
assert incompatible.missing_keys == []
|
||||
assert incompatible.unexpected_keys == []
|
||||
assert not any(param.is_meta for param in fastvideo.parameters())
|
||||
|
||||
|
||||
def test_dreamx_world_5b_cam_dit_config_matches_official_shape():
|
||||
config = make_dreamx_world_5b_cam_dit_config()
|
||||
assert config.num_layers == 30
|
||||
assert config.num_attention_heads == 24
|
||||
assert config.attention_head_dim == 128
|
||||
assert config.hidden_size == 3072
|
||||
assert config.ffn_dim == 14336
|
||||
assert config.add_control_adapter is True
|
||||
assert config.cam_method == "prope"
|
||||
assert config.attn_compress == 1
|
||||
|
||||
|
||||
def test_dreamx_world_converted_5b_transformer_strict_loads():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
_load_fastvideo_transformer_strict(device, torch.bfloat16)
|
||||
|
||||
|
||||
def test_dreamx_world_fastvideo_prope_branch_smoke():
|
||||
block = DreamXWorldTransformerBlock(
|
||||
8,
|
||||
32,
|
||||
1,
|
||||
cross_attn_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
layer_idx=0,
|
||||
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA,),
|
||||
)
|
||||
assert block.cam_self_attn is not None
|
||||
assert [
|
||||
name for name, _ in block.named_parameters()
|
||||
if name.startswith("cam_self_attn.")
|
||||
][:8] == [
|
||||
"cam_self_attn.q_proj.weight",
|
||||
"cam_self_attn.q_proj.bias",
|
||||
"cam_self_attn.k_proj.weight",
|
||||
"cam_self_attn.k_proj.bias",
|
||||
"cam_self_attn.v_proj.weight",
|
||||
"cam_self_attn.v_proj.bias",
|
||||
"cam_self_attn.out_proj.weight",
|
||||
"cam_self_attn.out_proj.bias",
|
||||
]
|
||||
|
||||
module = DreamXPropeSelfAttention(
|
||||
dim=8,
|
||||
attn_dim=8,
|
||||
num_heads=1,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
).eval()
|
||||
assert module.num_heads == 1
|
||||
assert module.head_dim == 8
|
||||
assert tuple(module.out_proj.weight.shape) == (8, 8)
|
||||
assert torch.count_nonzero(module.out_proj.weight) == 0
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for transformer parity.")
|
||||
def test_dreamx_world_transformer_parity_scaffold(monkeypatch):
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.float32
|
||||
inputs = _make_inputs(device, dtype)
|
||||
|
||||
official = _load_official_transformer(device, dtype)
|
||||
official_out = _run_official(official, inputs)
|
||||
del official
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
import fastvideo.attention.layer as attention_layer
|
||||
import fastvideo.distributed.communication_op as communication_op
|
||||
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
|
||||
monkeypatch.setattr(attention_layer, "get_sp_parallel_rank", lambda: 0)
|
||||
monkeypatch.setattr(attention_layer, "get_sp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_to_all_4D", lambda tensor, scatter_dim=2, gather_dim=1: tensor)
|
||||
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_gather", lambda tensor, dim=-1: tensor)
|
||||
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(communication_op, "get_sp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(fastvideo_dreamx, "sequence_model_parallel_shard", lambda tensor, dim=1: (tensor, tensor.shape[dim]))
|
||||
monkeypatch.setattr(
|
||||
fastvideo_dreamx,
|
||||
"sequence_model_parallel_all_gather_with_unpad",
|
||||
lambda tensor, original_seq_len, dim=1: tensor.narrow(dim, 0, original_seq_len),
|
||||
)
|
||||
fastvideo = _load_fastvideo_transformer(device, dtype)
|
||||
fastvideo_out = _run_fastvideo(fastvideo, inputs)
|
||||
assert official_out.shape == fastvideo_out.shape
|
||||
diff = (official_out - fastvideo_out).abs()
|
||||
print(f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
|
||||
assert_close(fastvideo_out, official_out, atol=1e-1, rtol=1e-1)
|
||||
@@ -1,220 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World Wan2.2 VAE reuse parity scaffold.
|
||||
|
||||
Coverage scope: implementation_subcomponent. The official side uses
|
||||
AutoencoderKLWan3_8 from DreamX, while the FastVideo side targets the native
|
||||
Wan VAE. This remains a scaffold until Wan2.2 base VAE weights are staged.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import re
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_vae_config
|
||||
from fastvideo.models.vaes.wanvae import AutoencoderKLWan
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _add_official_to_path():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
if str(OFFICIAL_REF_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
|
||||
|
||||
def _install_xfuser_stub() -> None:
|
||||
if "xfuser" in sys.modules:
|
||||
return
|
||||
xfuser = types.ModuleType("xfuser")
|
||||
core = types.ModuleType("xfuser.core")
|
||||
distributed = types.ModuleType("xfuser.core.distributed")
|
||||
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
|
||||
distributed.get_sequence_parallel_rank = lambda: 0
|
||||
distributed.get_sequence_parallel_world_size = lambda: 1
|
||||
distributed.get_sp_group = lambda: None
|
||||
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
|
||||
distributed.init_distributed_environment = lambda *args, **kwargs: None
|
||||
distributed.initialize_model_parallel = lambda *args, **kwargs: None
|
||||
distributed.model_parallel_is_initialized = lambda: False
|
||||
|
||||
class XFuserLongContextAttention:
|
||||
def __call__(self, *args, **kwargs):
|
||||
raise RuntimeError("xfuser stub cannot execute attention")
|
||||
|
||||
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
|
||||
sys.modules.update({
|
||||
"xfuser": xfuser,
|
||||
"xfuser.core": core,
|
||||
"xfuser.core.distributed": distributed,
|
||||
"xfuser.core.long_ctx_attention": long_ctx,
|
||||
})
|
||||
|
||||
|
||||
def _map_residual_subkey(prefix: str, sub: str) -> str | None:
|
||||
if sub == "residual.0.gamma":
|
||||
return f"{prefix}.norm1.gamma"
|
||||
match = re.match(r"^residual\.2\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.conv1.{match.group(1)}"
|
||||
if sub == "residual.3.gamma":
|
||||
return f"{prefix}.norm2.gamma"
|
||||
match = re.match(r"^residual\.6\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.conv2.{match.group(1)}"
|
||||
match = re.match(r"^shortcut\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.conv_shortcut.{match.group(1)}"
|
||||
return None
|
||||
|
||||
|
||||
def _map_attention_subkey(prefix: str, sub: str) -> str | None:
|
||||
if sub == "norm.gamma":
|
||||
return f"{prefix}.norm.gamma"
|
||||
match = re.match(r"^(to_qkv|proj)\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.{match.group(1)}.{match.group(2)}"
|
||||
return None
|
||||
|
||||
|
||||
def _map_resample_subkey(prefix: str, sub: str) -> str | None:
|
||||
match = re.match(r"^resample\.1\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.resample.1.{match.group(1)}"
|
||||
match = re.match(r"^time_conv\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.time_conv.{match.group(1)}"
|
||||
return None
|
||||
|
||||
|
||||
def _map_dreamx_raw_vae_key(key: str) -> str | None:
|
||||
match = re.match(r"^(conv1|conv2)\.(weight|bias)$", key)
|
||||
if match:
|
||||
prefix = "quant_conv" if match.group(1) == "conv1" else "post_quant_conv"
|
||||
return f"{prefix}.{match.group(2)}"
|
||||
match = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
|
||||
if match:
|
||||
return f"{match.group(1)}.conv_in.{match.group(2)}"
|
||||
match = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
|
||||
if match:
|
||||
return f"{match.group(1)}.norm_out.gamma"
|
||||
match = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
|
||||
if match:
|
||||
return f"{match.group(1)}.conv_out.{match.group(2)}"
|
||||
match = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
|
||||
if match:
|
||||
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.0", match.group(2))
|
||||
match = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
|
||||
if match:
|
||||
return _map_attention_subkey(f"{match.group(1)}.mid_block.attentions.0", match.group(2))
|
||||
match = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
|
||||
if match:
|
||||
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.1", match.group(2))
|
||||
match = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
|
||||
if match:
|
||||
stage = int(match.group(1))
|
||||
block = int(match.group(2))
|
||||
sub = match.group(3)
|
||||
if block in (0, 1):
|
||||
return _map_residual_subkey(f"encoder.down_blocks.{stage}.resnets.{block}", sub)
|
||||
if block == 2:
|
||||
return _map_resample_subkey(f"encoder.down_blocks.{stage}.downsampler", sub)
|
||||
return None
|
||||
match = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
|
||||
if match:
|
||||
stage = int(match.group(1))
|
||||
block = int(match.group(2))
|
||||
sub = match.group(3)
|
||||
if block in (0, 1, 2):
|
||||
return _map_residual_subkey(f"decoder.up_blocks.{stage}.resnets.{block}", sub)
|
||||
if block == 3:
|
||||
return _map_resample_subkey(f"decoder.up_blocks.{stage}.upsampler", sub)
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _vae_kwargs():
|
||||
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
|
||||
return OmegaConf.to_container(config["vae_kwargs"])
|
||||
|
||||
|
||||
def _load_official_vae(device, dtype):
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
|
||||
if not vae_path.exists():
|
||||
pytest.skip(f"Wan2.2 base VAE weights missing: {vae_path}")
|
||||
try:
|
||||
from models import AutoencoderKLWan3_8
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official DreamX VAE: {exc}")
|
||||
model = AutoencoderKLWan3_8.from_pretrained(str(vae_path), additional_kwargs=_vae_kwargs())
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_vae(device, dtype):
|
||||
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
|
||||
if not vae_path.exists():
|
||||
pytest.skip(f"Wan2.2 raw VAE weights missing: {vae_path}")
|
||||
config = make_dreamx_world_5b_cam_vae_config()
|
||||
config.load_encoder = True
|
||||
config.load_decoder = True
|
||||
model = AutoencoderKLWan(config).to(device=device, dtype=dtype)
|
||||
raw_state = torch.load(str(vae_path), map_location="cpu", weights_only=True)
|
||||
mapped_state = {}
|
||||
for key, value in raw_state.items():
|
||||
mapped_key = _map_dreamx_raw_vae_key(key)
|
||||
if mapped_key is None:
|
||||
raise AssertionError(f"Unmapped DreamX raw VAE key: {key}")
|
||||
mapped_state[mapped_key] = value
|
||||
model.load_state_dict(mapped_state, strict=True)
|
||||
return model.eval()
|
||||
|
||||
|
||||
def _normalize_fastvideo_vae_latent(latent: torch.Tensor) -> torch.Tensor:
|
||||
config = make_dreamx_world_5b_cam_vae_config()
|
||||
mean = torch.tensor(config.latents_mean, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
|
||||
std = torch.tensor(config.latents_std, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
|
||||
return (latent - mean) / std
|
||||
|
||||
|
||||
def test_dreamx_world_vae_config_matches_wan22_shape():
|
||||
config = make_dreamx_world_5b_cam_vae_config()
|
||||
assert config.z_dim == 48
|
||||
assert config.in_channels == 12
|
||||
assert config.out_channels == 12
|
||||
assert config.base_dim == 160
|
||||
assert config.decoder_base_dim == 256
|
||||
assert config.scale_factor_temporal == 4
|
||||
assert config.scale_factor_spatial == 16
|
||||
assert config.patch_size == 2
|
||||
assert config.is_residual is True
|
||||
assert config.clip_output is False
|
||||
assert len(config.latents_mean) == 48
|
||||
assert len(config.latents_std) == 48
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for VAE parity.")
|
||||
def test_dreamx_world_vae_encode_parity_scaffold():
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
official = _load_official_vae(device, dtype)
|
||||
fastvideo = _load_fastvideo_vae(device, dtype)
|
||||
torch.manual_seed(123)
|
||||
video = torch.randn(1, 3, 5, 64, 64, device=device, dtype=dtype).clamp(-1, 1)
|
||||
with torch.inference_mode():
|
||||
official_latent = official.encode(video).latent_dist.mean.float().cpu()
|
||||
fastvideo_latent = _normalize_fastvideo_vae_latent(fastvideo.encode(video).mean).float().cpu()
|
||||
assert official_latent.shape == fastvideo_latent.shape
|
||||
assert_close(fastvideo_latent, official_latent, atol=5e-2, rtol=5e-2)
|
||||
@@ -1,112 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World pipeline parity checks.
|
||||
|
||||
This test compares the FastVideo pipeline's DreamX-specific conditioning and
|
||||
single-step scheduler path against an explicit hand-rolled pass using the same
|
||||
loaded modules. It is intentionally local and deterministic: component parity
|
||||
against the official DreamX repository lives in ``tests/local_tests/dreamx_world``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
MODEL_DIR = Path(os.getenv("DREAMX_WORLD_MODEL_DIR", "converted_weights/dreamx_world"))
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="DreamX-World pipeline parity requires CUDA",
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def _run_worker_forward_batch(worker_wrapper: Any, request_kwargs: dict[str, Any]) -> torch.Tensor:
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.utils import shallow_asdict
|
||||
|
||||
fastvideo_args = worker_wrapper.worker.fastvideo_args
|
||||
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
|
||||
sampling_param.update({
|
||||
key: value
|
||||
for key, value in request_kwargs.items()
|
||||
if key not in {"prompt", "output_path"}
|
||||
})
|
||||
sampling_param.prompt = request_kwargs["prompt"]
|
||||
|
||||
latents_size = [
|
||||
(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8,
|
||||
sampling_param.width // 8,
|
||||
]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
)
|
||||
output_batch = worker_wrapper.worker.pipeline.forward(batch, fastvideo_args)
|
||||
assert output_batch.output is not None
|
||||
return output_batch.output.detach().cpu()
|
||||
|
||||
def _close_generator(generator: Any) -> None:
|
||||
generator.shutdown()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def test_dreamx_world_one_step_pipeline_latent_matches_manual_pass() -> None:
|
||||
if not MODEL_DIR.exists():
|
||||
pytest.fail(f"DreamX-World converted model directory is missing: {MODEL_DIR}")
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
common_kwargs = dict(
|
||||
prompt="a quiet road through a futuristic city at sunrise",
|
||||
output_path="outputs_video/dreamx_world_parity",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=64,
|
||||
width=64,
|
||||
num_frames=9,
|
||||
num_inference_steps=1,
|
||||
guidance_scale=1.0,
|
||||
action_list=["w"],
|
||||
action_speed_list=[2.0],
|
||||
seed=123,
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
str(MODEL_DIR),
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
output_type="latent",
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
try:
|
||||
result = cast(dict[str, Any], generator.generate_video(**common_kwargs))
|
||||
pipeline_latents = cast(torch.Tensor, result["samples"]).detach().cpu()
|
||||
|
||||
manual_latents = generator.executor.collective_rpc(
|
||||
_run_worker_forward_batch,
|
||||
kwargs={"request_kwargs": common_kwargs},
|
||||
)[0]
|
||||
finally:
|
||||
_close_generator(generator)
|
||||
|
||||
assert_close(pipeline_latents, manual_latents, atol=0.0, rtol=0.0)
|
||||
@@ -1,159 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Smoke tests for the DreamX-World-5B-Cam pipeline."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
MODEL_DIR = Path(os.getenv("DREAMX_WORLD_MODEL_DIR", "converted_weights/dreamx_world"))
|
||||
|
||||
|
||||
|
||||
|
||||
def _write_smoke_image(path: Path) -> None:
|
||||
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
|
||||
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
|
||||
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
|
||||
image.save(path)
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="DreamX-World pipeline smoke requires CUDA",
|
||||
)
|
||||
|
||||
|
||||
def test_dreamx_world_typed_surface_preflight() -> None:
|
||||
import fastvideo.registry as registry
|
||||
from fastvideo.api.presets import get_preset, get_presets_for_family
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import (
|
||||
DreamXWorldPipeline,
|
||||
EntryClass,
|
||||
)
|
||||
|
||||
assert DreamXWorldPipeline.__name__ == "DreamXWorldPipeline"
|
||||
assert EntryClass is DreamXWorldPipeline
|
||||
assert DreamXWorldPipeline._required_config_modules == [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
default_preset, model_family = registry.get_preset_selection(
|
||||
"GD-ML/DreamX-World-5B-Cam"
|
||||
)
|
||||
assert model_family == "dreamx_world"
|
||||
assert default_preset == "dreamx_world_5b_cam"
|
||||
|
||||
info = registry.get_model_info(
|
||||
"GD-ML/DreamX-World-5B-Cam",
|
||||
workload_type=WorkloadType.I2V,
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
assert info.pipeline_cls is DreamXWorldPipeline
|
||||
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
|
||||
|
||||
names = {p.name for p in get_presets_for_family("dreamx_world")}
|
||||
assert "dreamx_world_5b_cam" in names
|
||||
preset = get_preset("dreamx_world_5b_cam", "dreamx_world")
|
||||
assert preset.defaults["num_inference_steps"] == 30
|
||||
assert preset.defaults["height"] == 480
|
||||
assert preset.defaults["width"] == 832
|
||||
assert preset.defaults["num_frames"] == 161
|
||||
assert preset.defaults["guidance_scale"] == 5.0
|
||||
|
||||
cfg = DreamXWorld5BCamPipelineConfig()
|
||||
assert cfg.flow_shift == 3.0
|
||||
assert cfg.ti2v_task is True
|
||||
assert cfg.expand_timesteps is True
|
||||
assert cfg.dit_config.arch_config.add_control_adapter is True
|
||||
assert cfg.dit_config.arch_config.cam_method == "prope"
|
||||
|
||||
|
||||
def test_dreamx_world_camera_stage_writes_y_camera() -> None:
|
||||
from types import SimpleNamespace
|
||||
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import (
|
||||
DREAMX_Y_CAMERA_KEY,
|
||||
DreamXWorldCameraConditioningStage,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt="camera smoke",
|
||||
latents=torch.zeros(1, 48, 3, 8, 8, dtype=torch.bfloat16, device="cuda"),
|
||||
num_frames=9,
|
||||
height=64,
|
||||
width=64,
|
||||
action_list=["w", "d"],
|
||||
action_speed_list=[2.0, 1.0],
|
||||
)
|
||||
out = DreamXWorldCameraConditioningStage().forward(batch, cast(Any, SimpleNamespace()))
|
||||
|
||||
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
|
||||
assert set(y_camera) == {"viewmats", "K"}
|
||||
assert y_camera["viewmats"].shape == (1, 3, 4, 4)
|
||||
assert y_camera["K"].shape == (1, 3, 3, 3)
|
||||
assert y_camera["viewmats"].device.type == "cuda"
|
||||
assert y_camera["viewmats"].dtype == torch.bfloat16
|
||||
|
||||
|
||||
def test_dreamx_world_pipeline_load_generate_latent_smoke(tmp_path: Path) -> None:
|
||||
if not MODEL_DIR.exists():
|
||||
pytest.fail(f"DreamX-World converted model directory is missing: {MODEL_DIR}")
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
image_path = tmp_path / "dreamx_world_smoke_input.png"
|
||||
_write_smoke_image(image_path)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
str(MODEL_DIR),
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
output_type="latent",
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
try:
|
||||
result = generator.generate_video(
|
||||
prompt="a quiet road through a futuristic city at sunrise",
|
||||
output_path="outputs_video/dreamx_world_smoke",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=64,
|
||||
width=64,
|
||||
num_frames=9,
|
||||
num_inference_steps=1,
|
||||
guidance_scale=1.0,
|
||||
image_path=str(image_path),
|
||||
action_list=["w"],
|
||||
action_speed_list=[2.0],
|
||||
seed=0,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
assert isinstance(result, dict)
|
||||
samples = cast(dict[str, Any], result)["samples"]
|
||||
assert torch.is_tensor(samples)
|
||||
assert samples.ndim == 5
|
||||
assert samples.shape[1] == 48
|
||||
assert torch.isfinite(samples).all()
|
||||
@@ -0,0 +1,3 @@
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.pyo
|
||||
+105
@@ -0,0 +1,105 @@
|
||||
# Handoff — GPU bring-up of the v2 torch backend
|
||||
|
||||
**For: an agent on a GPU box, branched from `will/mini-fastvideo`.**
|
||||
**Your job:** take the *written-not-run* `cuda` backend to *runs-and-generates*, then commit + push.
|
||||
|
||||
Everything below is committed on `will/mini-fastvideo` and CPU-tested (**204 tests pass**). The torch
|
||||
path was authored on a machine with **no GPU and no torch**, so it is grounded in the real
|
||||
`fastvideo` APIs and cross-checked against the source, but **never executed**. That's what you finish.
|
||||
|
||||
---
|
||||
|
||||
## 0. Orientation (read these first, in order)
|
||||
|
||||
1. **`v2/README.md`** — what the whole v2 mini is (the `(recipe, runtime)` runtime; "architecture is
|
||||
real, kernels are toys"). The "Honest scope" paragraph says exactly what's wired.
|
||||
2. **`v2/platform/backends/GPU_BRINGUP.md`** — *your checklist*: the ordered 10-step bring-up + the
|
||||
risk table (A–G), each tied to a `# BRINGUP` marker in the source. **This handoff is orientation +
|
||||
process; GPU_BRINGUP.md is the work.**
|
||||
3. This file — the meta-instructions (verify bar, commit/push, gotchas).
|
||||
|
||||
## 1. What's already done (commits on this branch)
|
||||
|
||||
```
|
||||
d6d0580a [fix] correct GPU adapters against real fastvideo API (cross-check findings)
|
||||
b8d78f40 [feat] real torch/CUDA backend (written-not-run) behind the cuda cells
|
||||
27791b51 [feat] static-buffer capture form for the cudagraph step body (Path A)
|
||||
ae6a170d [feat] piecewise CUDA-graph capture/replay at the step boundary (Path A)
|
||||
9308d87e [feat] route diffusion loops through the kernel table
|
||||
7490c590 [feat] multi-backend dispatch substrate (device/arch/kernel registries)
|
||||
```
|
||||
|
||||
The dispatch substrate (two tuple-keyed registries `COMPONENTS(kind,device,variant)` +
|
||||
`KERNELS(op,device,arch,variant)`, a detected `Platform`, numpy terminal + parity oracle), the
|
||||
universal kernel seam (diffusion loops go through `model.platform.kernels`), the
|
||||
piecewise cudagraph lifecycle, and the torch backend cells are all in place. On a GPU box,
|
||||
`Platform.detect()` returns a `cuda` platform and resolves the torch cells instead of the numpy toys —
|
||||
**the inference loops/policies/scheduler are unchanged**; only the resolved implementations differ.
|
||||
|
||||
## 2. The files you'll touch
|
||||
|
||||
| File | What it is |
|
||||
|---|---|
|
||||
| `v2/platform/backends/torch_adapters.py` | `TorchWanDiT` / `TorchWanVAE` / `TorchT5Encoder` — wrap the real `fastvideo.models.*` (named by each card's `load_id`) to the mini's duck-typed surface. Built via the real FastVideo loaders. |
|
||||
| `v2/platform/backends/torch_kernels.py` | torch `flow_match_step` / `flow_sde_step` (plain elementwise — there is **no** fused solver kernel in fastvideo-kernel; don't look for one). |
|
||||
| `v2/platform/backends/torch_cuda.py` | registers the `cuda` cells as lazy trampolines (torch imported only inside builder bodies). |
|
||||
| `v2/card/specs.py` | `ComponentSpec.checkpoint` — the per-component weights source (empty on toys; **you fill it in**). |
|
||||
|
||||
The surface the adapters must honor (what the loops call):
|
||||
`dit(latent, text_embed, sigma) -> velocity` · `vae.decode(latent)` / `vae.encode(video)` ·
|
||||
`text_encoder.encode(text)`. The CPU toys in `v2/models/backend.py` are the reference behavior.
|
||||
|
||||
## 3. Your task (the gating items — full detail in GPU_BRINGUP.md)
|
||||
|
||||
1. **Env:** install `torch` + the parent `fastvideo` package + weights. (`fastvideo` source lives at
|
||||
`/Users/willlin/src/FastVideo`.)
|
||||
2. **Risk A — the one blocking gap:** the builders call `_load_via_fastvideo(...)` → the real loaders
|
||||
need a **`FastVideoArgs`**, which `_fastvideo_args(spec)` builds minimally from `spec.checkpoint`.
|
||||
Confirm/extend its fields (model config, precision, parallelism). And stamp `ComponentSpec.checkpoint`
|
||||
onto the wan21 card — a tiny helper that maps a model root onto the three components is the cleanest
|
||||
way (the toy cards leave it `""`).
|
||||
3. **Work the risk list (A–G in GPU_BRINGUP.md).** The *interface* contracts were cross-checked as
|
||||
matching (DiT returns bare velocity; `timestep=sigma*1000`; `encode().mode()`; `.last_hidden_state`;
|
||||
no fused solver kernel) — confirm them numerically. The *construction* layer was fixed (real loaders,
|
||||
`set_forward_context`, latent normalization, UMT5-from-config). What's left is box-dependent:
|
||||
`FastVideoArgs` fields, `shift_factor` placement/sign, exact tokenizer kwargs, FSDP sharding.
|
||||
4. **Bring up in order:** build each component in isolation → one DiT step → one solver step → VAE
|
||||
decode → full t2v → SDE stochastic sampling → cudagraph capture (last).
|
||||
|
||||
## 4. The verification bar (how you know it's right)
|
||||
|
||||
- **CPU suite must stay green:** `python3 -m pytest v2/ -q` → still **204 passed**. The torch path is
|
||||
gated `available=False` off-GPU; importing the backends must never import torch. If you break either,
|
||||
you broke the substrate. (`v2/tests/test_torch_backend.py` pins these.)
|
||||
- **Parity oracle is the spec:** the substrate's whole point is that a real backend matches the numpy
|
||||
reference on the consistency ladder. On GPU, compare a full generation against a known-good fastvideo
|
||||
output — use the parent repo's SSIM regression harness (`fastvideo/tests/ssim/`). Target C4 (SSIM /
|
||||
artifact quality); component/trajectory parity (C0/C1) is bit-level vs the reference pipeline.
|
||||
- **Don't trust "it ran" — trust "it matched."** A wrong `timestep` scale or `shift_factor` produces
|
||||
plausible-but-wrong video, not a crash (risks B/D). Diff against a reference, don't eyeball.
|
||||
|
||||
## 5. Commit + push
|
||||
|
||||
- **You are on a GPU branch** (branched from `will/mini-fastvideo`). Commit your bring-up fixes there,
|
||||
focused by concern (e.g. one commit per confirmed risk), in the existing style (`[fix]`/`[feat] …`).
|
||||
- **NEVER add Claude as a co-author** (repo policy, `/Users/willlin/src/.claude/CLAUDE.md`).
|
||||
- **Do not rewrite or force-push** the six commits above — build on top.
|
||||
- When the CPU suite is green **and** a GPU generation matches the reference, **push your branch.**
|
||||
- If you launch inference with wandb logging enabled, log in with the token in the project
|
||||
`CLAUDE.md` (`/Users/willlin/src/.claude/CLAUDE.md`) — **do not paste it into any committed file.**
|
||||
|
||||
## 6. Gotchas (don't relearn these the hard way)
|
||||
|
||||
- **No fused solver kernel exists.** `fastvideo-kernel` ships only attention/norm/quant primitives;
|
||||
the cuda `flow_match_step`/`flow_sde_step` are plain torch by design. Don't hunt for a `.cu` solver.
|
||||
- **`from_pretrained` is not the loader.** `WanTransformer3DModel`/`AutoencoderKLWan` have none — the
|
||||
real path is the `*Loader().load(model_path, fastvideo_args)` classes in
|
||||
`fastvideo/models/loader/component_loader.py`. The loader resolves the class from the checkpoint
|
||||
config (this is what makes UMT5-vs-T5 correct without hardcoding).
|
||||
- **T5 needs `set_forward_context`.** A bare encoder forward reads stale/None global context.
|
||||
- **The loop surface stays numpy** for bring-up; adapters marshal numpy↔torch at the boundary. A
|
||||
torch-native surface (latent on-device through forward→combine→solver) is the **perf follow-up**
|
||||
(Risk G), not bring-up — don't rewrite `cfg.combine`/`precision.cast`/the samplers yet.
|
||||
- **cudagraph capture is last.** The wan21 loop declares `breakable_cudagraph`; the v2 capturer models
|
||||
the lifecycle with a numpy `StaticWorkspace`. Capturing a real `torch.cuda.CUDAGraph` is GPU-only
|
||||
work and the riskiest step — leave it until inference is verified.
|
||||
+148
@@ -0,0 +1,148 @@
|
||||
# FastVideo v2 - Inference Runtime Scope
|
||||
|
||||
**Status:** source of truth for `v2/`.
|
||||
|
||||
`v2` is the model-native inference runtime for FastVideo. It owns model cards,
|
||||
programs, loops, runtime execution, serving, cache/memory policy, backend dispatch,
|
||||
compile/cudagraph integration, and inference parity checks.
|
||||
|
||||
`v2` does **not** own training, finetuning, distillation, RL, optimizer steps, or
|
||||
checkpoint production. Training remains in the existing FastVideo stacks:
|
||||
|
||||
- `fastvideo/train/` - the new modular trainer.
|
||||
- `fastvideo/training/` - the legacy shipped training pipelines.
|
||||
|
||||
`v2` may record how a checkpoint was produced through `RecipeSpec` metadata
|
||||
(`method`, `parents`, `assumes_loop`, `assumes_precision`), because inference must
|
||||
know which runtime loop and precision policy a post-training checkpoint expects.
|
||||
That metadata is provenance, not a v2 training API.
|
||||
|
||||
## Design Center
|
||||
|
||||
FastVideo v2 is video-generation first. Wan/LTX-style diffusion video inference
|
||||
is the baseline path, and unified models such as BAGEL/Cosmos3 are first-class:
|
||||
one resident model may run AR, diffusion, VAE, and codec loops in one request.
|
||||
Audio and TTS are supported as additional modalities on the same stage/loop
|
||||
model, not as the reason to build a separate universal serving framework.
|
||||
|
||||
The core abstraction should stay small:
|
||||
|
||||
- `ModelCard` declares the resident components, loops, capabilities, precision,
|
||||
caches, and sampling defaults a checkpoint needs.
|
||||
- `Program` is an ordered list of typed nodes passing values through named
|
||||
slots. It is not a general DAG, Walk graph, or declarative control-flow IR.
|
||||
- `Loop` owns model semantics. The runtime drives the loop, handles admission,
|
||||
cancellation, streaming, cache access, and backend dispatch.
|
||||
|
||||
Do not add a public contract field until a runtime path consumes it. Future
|
||||
optimizations such as richer stage placement, multi-GPU transport, or paged KV
|
||||
should start behind a concrete Wan/BAGEL/Cosmos/Qwen use case and graduate only
|
||||
after they simplify at least two model recipes.
|
||||
|
||||
## Scope
|
||||
|
||||
In scope:
|
||||
|
||||
- Python inference entrypoint through `v2.VideoGenerator`.
|
||||
- Typed `ModelCard` declarations for components, loops, capabilities, parity,
|
||||
sampling defaults, precision, and checkpoint layout.
|
||||
- Driven inference loops such as diffusion denoise, AR decode, causal/world
|
||||
continuation, VAE/audio decode, and multi-stage programs.
|
||||
- Runtime execution through `Engine` and `AsyncEngine`.
|
||||
- Serving through OpenAI-compatible HTTP/SSE surfaces and deployment cards.
|
||||
- Backend dispatch through the CPU toy backend, accelerator stand-ins, and the
|
||||
real torch/CUDA backend.
|
||||
- Inference acceleration features such as FP8/NVFP4 loading, Sage/Flash/SDPA
|
||||
attention backend selection, `torch.compile`, cudagraph capture, cache policy,
|
||||
and component placement.
|
||||
- Inference parity and regression tests.
|
||||
|
||||
Out of scope:
|
||||
|
||||
- Training methods, optimizers, loss functions, RL rewards, rollout trainers,
|
||||
weight-sync training loops, and behavior records for policy updates.
|
||||
- Training examples under `v2_examples/`.
|
||||
- Any CLI/API that advertises v2 as a trainer.
|
||||
- A universal graph runtime, Walk/state-machine authoring layer, or parallelism
|
||||
vocabulary that is not consumed by the current inference runtime.
|
||||
|
||||
## Core Model
|
||||
|
||||
The atomic inference artifact is a `(recipe, runtime)` pair:
|
||||
|
||||
- `RecipeSpec` records what the weights assume: parent checkpoints, post-training
|
||||
method name, required loop, and required precision.
|
||||
- `ModelCard` declares the runtime surface: components, loops, capabilities,
|
||||
caches, precision, parallelism, sampling defaults, and checkpoint manifest.
|
||||
- `Program` composes component nodes and loop nodes into a user-facing task as
|
||||
an ordered named-slot stage list.
|
||||
- `ModelInstance` is the resident loaded card with shared components, caches,
|
||||
weight versions, and optional captured graphs.
|
||||
|
||||
This keeps post-training artifacts serveable without making `v2` responsible for
|
||||
creating them.
|
||||
|
||||
## Execution Model
|
||||
|
||||
Loops are model-owned state machines:
|
||||
|
||||
```python
|
||||
state = loop.init(req, model, ctx)
|
||||
while True:
|
||||
plan = loop.next(state)
|
||||
if isinstance(plan, Done):
|
||||
break
|
||||
result = ctx.execute(plan)
|
||||
state = loop.advance(state, result)
|
||||
return loop.finalize(state)
|
||||
```
|
||||
|
||||
The loop owns semantics. The runtime owns execution, admission, cancellation,
|
||||
streaming, cache access, graph capture, and backend dispatch.
|
||||
|
||||
Serving is pooled run-to-completion. `AsyncEngine` bounds concurrency by pool
|
||||
slots; each request runs its program to completion. The synchronous `Engine` is
|
||||
the offline path used by tests and `VideoGenerator`.
|
||||
|
||||
## Package Layout
|
||||
|
||||
```text
|
||||
v2/
|
||||
video_generator.py public inference facade
|
||||
registry.py model id -> card builder registry
|
||||
core/
|
||||
card/ ModelCard, specs, ModelInstance
|
||||
loop/ loop contracts, driver, sampler, policies
|
||||
program/ task programs and workflows
|
||||
request/ request params, tasks, outputs, sessions
|
||||
parity/ inference parity helpers
|
||||
parallel/ named parallel plans
|
||||
recipes/ model-specific cards, loops, and programs
|
||||
runtime/ Engine, AsyncEngine, cache, memory, cudagraph, transport
|
||||
serving/ HTTP/SSE server and deployment adapters
|
||||
platform/ backend/device/kernel dispatch
|
||||
_vendor/ vendored FastVideo model/loader/config pieces for inference
|
||||
tests/ v2 inference/runtime/serving/parity tests
|
||||
```
|
||||
|
||||
There is intentionally no `v2/training/` package.
|
||||
|
||||
## Current Inference Path
|
||||
|
||||
The torch backend builds real components from stamped checkpoint paths, keeps
|
||||
components in eval mode, and dispatches inference through the same cards and loops
|
||||
used by the CPU tests. Wan2.1 T2V inference is the primary real path today.
|
||||
Wan/FastWan inference supports:
|
||||
|
||||
- real Wan component loading through vendored component loaders,
|
||||
- FP8 post-load quantization for `FastVideo/FastWan-QAD-FP8-1.3B`,
|
||||
- attention backend selection, including SageAttention when installed,
|
||||
- `torch.compile` for inference DiT modules,
|
||||
- on-device latent residency for cards that set `device_io=True`.
|
||||
|
||||
## Boundary Rule
|
||||
|
||||
If a change adds training behavior, it belongs in `fastvideo/train/` or
|
||||
`fastvideo/training/`, not in `v2/`. If inference needs to consume the result of
|
||||
that training, add or update a v2 card, loop, registry entry, checkpoint loader,
|
||||
sampling defaults, and inference tests.
|
||||
@@ -0,0 +1,89 @@
|
||||
"""v2 - the FastVideo inference runtime (see v2/README.md).
|
||||
|
||||
> 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.
|
||||
|
||||
v2 is inference-only. Training, finetuning, distillation, RL, and optimizer
|
||||
loops belong to ``fastvideo/train`` or ``fastvideo/training``. v2 only records
|
||||
checkpoint provenance in recipe metadata so inference can bind weights to the
|
||||
right loop and precision policy.
|
||||
|
||||
The core is numpy-only and CPU-testable; heavy Wan/LTX neural forwards become lazy torch
|
||||
adapters (see ``v2/platform/backends/``) that are off the test path.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.core.enums import (
|
||||
Capability,
|
||||
ConsistencyLevel,
|
||||
ExecutionProfile,
|
||||
LoopKind,
|
||||
WorkUnitKind,
|
||||
)
|
||||
from v2.core.card import (
|
||||
CapabilityMatrix,
|
||||
ComponentSpec,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ModelInstance,
|
||||
ParitySpec,
|
||||
RecipeSpec,
|
||||
load_card,
|
||||
)
|
||||
from v2.core.program import ComponentNode, ModelLoopNode, Program, ProgramKind, when_opt, when_task
|
||||
from v2.core.request import (
|
||||
DiffusionParams,
|
||||
Output,
|
||||
Request,
|
||||
SamplingParams,
|
||||
Session,
|
||||
TaskType,
|
||||
make_request,
|
||||
)
|
||||
from v2.runtime import AsyncEngine, Engine
|
||||
|
||||
__version__ = "0.2.0"
|
||||
|
||||
__all__ = [
|
||||
"ModelCard",
|
||||
"ComponentSpec",
|
||||
"LoopSpec",
|
||||
"RecipeSpec",
|
||||
"ParitySpec",
|
||||
"CapabilityMatrix",
|
||||
"ModelInstance",
|
||||
"load_card",
|
||||
"Engine",
|
||||
"AsyncEngine",
|
||||
"Program",
|
||||
"ProgramKind",
|
||||
"ComponentNode",
|
||||
"ModelLoopNode",
|
||||
"when_task",
|
||||
"when_opt",
|
||||
"Request",
|
||||
"Session",
|
||||
"Output",
|
||||
"make_request",
|
||||
"TaskType",
|
||||
"SamplingParams",
|
||||
"DiffusionParams",
|
||||
"LoopKind",
|
||||
"WorkUnitKind",
|
||||
"ConsistencyLevel",
|
||||
"ExecutionProfile",
|
||||
"Capability",
|
||||
"VideoGenerator",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
# Lazy: the GPU entrypoint imports torch / fastvideo, so resolve it only on access — plain
|
||||
# ``import v2`` (and the CPU-only mini) stay torch-free.
|
||||
if name == "VideoGenerator":
|
||||
from v2.video_generator import VideoGenerator
|
||||
return VideoGenerator
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
@@ -0,0 +1 @@
|
||||
"""Vendored fastvideo code (copied, standalone). Internal layout mirrors upstream for diffing; v2-native code must not edit these ad hoc."""
|
||||
@@ -0,0 +1,28 @@
|
||||
"""Slim vendored API config surface for the v2 VideoGenerator.
|
||||
|
||||
Only the inference-config dataclasses (schema) + result types are vendored. The fastvideo
|
||||
parser / presets / overrides modules are intentionally NOT vendored — they pull the fastvideo
|
||||
pipeline runtime, which v2 replaces. See v2/README.md (vendoring)."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._vendor.api.results import GenerationResult
|
||||
from v2._vendor.api.schema import (
|
||||
CompileConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
"EngineConfig",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"SamplingConfig",
|
||||
"GenerationResult",
|
||||
]
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class ConfigValidationError(ValueError):
|
||||
"""Validation error that keeps track of the nested config path."""
|
||||
|
||||
def __init__(self, path: str, message: str):
|
||||
self.path = path
|
||||
self.message = message
|
||||
super().__init__(str(self))
|
||||
|
||||
def __str__(self) -> str:
|
||||
if self.path:
|
||||
return f"{self.path}: {self.message}"
|
||||
return self.message
|
||||
@@ -0,0 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from v2._vendor.api.sampling_param import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class MatrixGame2SamplingParam(SamplingParam):
|
||||
height: int = 352
|
||||
width: int = 640
|
||||
num_frames: int = 57
|
||||
fps: int = 25
|
||||
guidance_scale: float = 1.0
|
||||
num_inference_steps: int = 3
|
||||
negative_prompt: str | None = None
|
||||
@@ -0,0 +1,233 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Track which GenerationRequest fields the user explicitly provided.
|
||||
|
||||
When translating a GenerationRequest into a legacy SamplingParam we must
|
||||
distinguish user-provided values (which should override model defaults)
|
||||
from schema defaults (which should NOT override model defaults).
|
||||
|
||||
The mechanism: a single ``_fastvideo_explicit_paths`` set stored on the
|
||||
root ``GenerationRequest``. It holds dotted leaf paths (e.g.
|
||||
``"sampling.guidance_scale"``) the user has touched, either via raw
|
||||
config at bind time or via attribute assignment at runtime. A patched
|
||||
``__setattr__`` on the request dataclass types records assignments into
|
||||
this set.
|
||||
|
||||
The set holds leaf paths only. Nested dataclass or mapping assignments
|
||||
are flattened to their leaves at record time.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Mapping
|
||||
import dataclasses
|
||||
from typing import Any, cast
|
||||
|
||||
from v2._vendor.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationPlan,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
PlannedStage,
|
||||
RequestRuntimeConfig,
|
||||
RunConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
)
|
||||
|
||||
EXPLICIT_PATHS_ATTR = "_fastvideo_explicit_paths"
|
||||
|
||||
_TRACKING_ROOT_ATTR = "_fastvideo_request_tracking_root"
|
||||
_TRACKING_PATH_ATTR = "_fastvideo_request_tracking_path"
|
||||
_TRACKING_PATCHED_ATTR = "_fastvideo_request_tracking_patched"
|
||||
|
||||
_TRACKED_REQUEST_TYPES = (
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
SamplingConfig,
|
||||
RequestRuntimeConfig,
|
||||
OutputConfig,
|
||||
ContinuationState,
|
||||
PlannedStage,
|
||||
GenerationPlan,
|
||||
)
|
||||
|
||||
|
||||
def bind_generation_request_raw(
|
||||
request: GenerationRequest,
|
||||
raw: Mapping[str, Any] | None,
|
||||
) -> GenerationRequest:
|
||||
"""Install explicit-path tracking on *request*.
|
||||
|
||||
*raw* is the parsed config dict (YAML/JSON/kwargs); every leaf key
|
||||
in it becomes an explicit path. Subsequent attribute assignments on
|
||||
*request* or its nested dataclasses are recorded automatically via a
|
||||
patched ``__setattr__``.
|
||||
"""
|
||||
_ensure_request_tracking()
|
||||
# Disable recording while we walk the tree to install roots.
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, None)
|
||||
_set_tracking_roots(request, request, "")
|
||||
paths: set[str] = set()
|
||||
_record_value_paths(raw or {}, "", paths)
|
||||
object.__setattr__(request, EXPLICIT_PATHS_ATTR, paths)
|
||||
return request
|
||||
|
||||
|
||||
def bind_run_config_raw(
|
||||
config: RunConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> RunConfig:
|
||||
request_raw = raw.get("request")
|
||||
if isinstance(request_raw, Mapping):
|
||||
bind_generation_request_raw(config.request, request_raw)
|
||||
else:
|
||||
bind_generation_request_raw(config.request, {})
|
||||
return config
|
||||
|
||||
|
||||
def bind_serve_config_raw(
|
||||
config: ServeConfig,
|
||||
raw: Mapping[str, Any],
|
||||
) -> ServeConfig:
|
||||
default_request_raw = raw.get("default_request")
|
||||
if isinstance(default_request_raw, Mapping):
|
||||
bind_generation_request_raw(config.default_request, default_request_raw)
|
||||
else:
|
||||
bind_generation_request_raw(config.default_request, {})
|
||||
return config
|
||||
|
||||
|
||||
def get_explicit_paths(request: GenerationRequest) -> frozenset[str]:
|
||||
"""Return a snapshot of the explicit paths set on *request*."""
|
||||
paths = getattr(request, EXPLICIT_PATHS_ATTR, None)
|
||||
if isinstance(paths, set | frozenset):
|
||||
return frozenset(paths)
|
||||
return frozenset()
|
||||
|
||||
|
||||
def reset_tracking_roots(request: GenerationRequest) -> None:
|
||||
"""Re-install tracking roots after a deepcopy or manual clone.
|
||||
|
||||
The paths set itself deepcopies correctly; we only need to repoint
|
||||
the tracking root on nested dataclasses at the new root.
|
||||
"""
|
||||
_ensure_request_tracking()
|
||||
_set_tracking_roots(request, request, "")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Path recording
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _record_value_paths(
|
||||
value: Any,
|
||||
prefix: str,
|
||||
out: set[str],
|
||||
) -> None:
|
||||
"""Add every leaf path under *value* to *out*.
|
||||
|
||||
A leaf is any terminal value (non-dataclass, non-mapping, or empty
|
||||
mapping/dataclass). ``prefix`` is the dotted path at which *value*
|
||||
sits. When called with an empty ``prefix`` (the root), leaves are
|
||||
recorded at their own key.
|
||||
"""
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
dc_fields = dataclasses.fields(value)
|
||||
if not dc_fields:
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
return
|
||||
for field in dc_fields:
|
||||
child = getattr(value, field.name)
|
||||
path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
_record_value_paths(child, path, out)
|
||||
return
|
||||
if isinstance(value, Mapping):
|
||||
if not value:
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
return
|
||||
for key, child in value.items():
|
||||
path = f"{prefix}.{key}" if prefix else key
|
||||
_record_value_paths(child, path, out)
|
||||
return
|
||||
if prefix:
|
||||
out.add(prefix)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# __setattr__ patching
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _ensure_request_tracking() -> None:
|
||||
for config_type in _TRACKED_REQUEST_TYPES:
|
||||
_patch_tracking_setattr(config_type)
|
||||
|
||||
|
||||
def _patch_tracking_setattr(config_type: type[Any]) -> None:
|
||||
if getattr(config_type, _TRACKING_PATCHED_ATTR, False):
|
||||
return
|
||||
|
||||
original_setattr = cast(
|
||||
Callable[[Any, str, Any], None],
|
||||
config_type.__setattr__,
|
||||
)
|
||||
field_names = {field.name for field in dataclasses.fields(config_type)}
|
||||
|
||||
def _tracking_setattr(self: Any, name: str, value: Any) -> None:
|
||||
if name.startswith("_fastvideo_") or name not in field_names:
|
||||
original_setattr(self, name, value)
|
||||
return
|
||||
|
||||
original_setattr(self, name, value)
|
||||
|
||||
root = getattr(self, _TRACKING_ROOT_ATTR, None)
|
||||
if root is None:
|
||||
return
|
||||
paths = getattr(root, EXPLICIT_PATHS_ATTR, None)
|
||||
if not isinstance(paths, set):
|
||||
return
|
||||
|
||||
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
|
||||
path = f"{prefix}.{name}" if prefix else name
|
||||
# Wholesale dataclass replacement: install roots on the new
|
||||
# instance so its future mutations are tracked too.
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
_set_tracking_roots(root, value, path)
|
||||
_record_value_paths(value, path, paths)
|
||||
|
||||
type.__setattr__(config_type, "__setattr__", _tracking_setattr)
|
||||
setattr(config_type, _TRACKING_PATCHED_ATTR, True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tree walk to set tracking root/path on nested dataclasses
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _set_tracking_roots(
|
||||
root: GenerationRequest,
|
||||
obj: Any,
|
||||
prefix: str,
|
||||
) -> None:
|
||||
if not dataclasses.is_dataclass(obj) or isinstance(obj, type):
|
||||
return
|
||||
object.__setattr__(obj, _TRACKING_ROOT_ATTR, root)
|
||||
object.__setattr__(obj, _TRACKING_PATH_ATTR, prefix)
|
||||
for field in dataclasses.fields(obj):
|
||||
child = getattr(obj, field.name)
|
||||
child_path = f"{prefix}.{field.name}" if prefix else field.name
|
||||
if dataclasses.is_dataclass(child) and not isinstance(child, type):
|
||||
_set_tracking_roots(root, child, child_path)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"EXPLICIT_PATHS_ATTR",
|
||||
"bind_generation_request_raw",
|
||||
"bind_run_config_raw",
|
||||
"bind_serve_config_raw",
|
||||
"get_explicit_paths",
|
||||
"reset_tracking_roots",
|
||||
]
|
||||
@@ -0,0 +1,173 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from collections.abc import Mapping
|
||||
|
||||
from v2._vendor.api.schema import ContinuationState
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationResult:
|
||||
prompt: str | None = None
|
||||
prompt_index: int | None = None
|
||||
samples: Any | None = None
|
||||
frames: Any | None = None
|
||||
audio: Any | None = None
|
||||
audio_sample_rate: int | None = None
|
||||
size: tuple[int, int, int] | None = None
|
||||
generation_time: float | None = None
|
||||
logging_info: Any | None = None
|
||||
trajectory: Any | None = None
|
||||
trajectory_timesteps: Any | None = None
|
||||
trajectory_decoded: Any | None = None
|
||||
video_path: str | None = None
|
||||
peak_memory_mb: float | None = None
|
||||
state: ContinuationState | None = None
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
@classmethod
|
||||
def from_legacy_result(
|
||||
cls,
|
||||
result: Mapping[str, Any],
|
||||
) -> GenerationResult:
|
||||
prompt = result.get("prompt")
|
||||
if prompt is None:
|
||||
prompt = result.get("prompts")
|
||||
|
||||
extra = {
|
||||
key: value
|
||||
for key, value in result.items() if key not in {
|
||||
"prompt",
|
||||
"prompt_index",
|
||||
"prompts",
|
||||
"samples",
|
||||
"frames",
|
||||
"audio",
|
||||
"audio_sample_rate",
|
||||
"size",
|
||||
"generation_time",
|
||||
"logging_info",
|
||||
"trajectory",
|
||||
"trajectory_timesteps",
|
||||
"trajectory_decoded",
|
||||
"video_path",
|
||||
"peak_memory_mb",
|
||||
"state",
|
||||
}
|
||||
}
|
||||
|
||||
return cls(
|
||||
prompt=prompt,
|
||||
prompt_index=result.get("prompt_index"),
|
||||
samples=result.get("samples"),
|
||||
frames=result.get("frames"),
|
||||
audio=result.get("audio"),
|
||||
audio_sample_rate=result.get("audio_sample_rate"),
|
||||
size=result.get("size"),
|
||||
generation_time=result.get("generation_time"),
|
||||
logging_info=result.get("logging_info"),
|
||||
trajectory=result.get("trajectory"),
|
||||
trajectory_timesteps=result.get("trajectory_timesteps"),
|
||||
trajectory_decoded=result.get("trajectory_decoded"),
|
||||
video_path=result.get("video_path"),
|
||||
peak_memory_mb=result.get("peak_memory_mb"),
|
||||
state=result.get("state"),
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
def to_legacy_dict(self) -> dict[str, Any]:
|
||||
result = {
|
||||
"prompts": self.prompt,
|
||||
"samples": self.samples,
|
||||
"frames": self.frames,
|
||||
"audio": self.audio,
|
||||
"audio_sample_rate": self.audio_sample_rate,
|
||||
"size": self.size,
|
||||
"generation_time": self.generation_time,
|
||||
"logging_info": self.logging_info,
|
||||
"trajectory": self.trajectory,
|
||||
"trajectory_timesteps": self.trajectory_timesteps,
|
||||
"trajectory_decoded": self.trajectory_decoded,
|
||||
"video_path": self.video_path,
|
||||
"peak_memory_mb": self.peak_memory_mb,
|
||||
}
|
||||
if self.prompt_index is not None:
|
||||
result["prompt_index"] = self.prompt_index
|
||||
result["prompt"] = self.prompt
|
||||
if self.state is not None:
|
||||
result["state"] = self.state
|
||||
result.update(self.extra)
|
||||
return result
|
||||
|
||||
|
||||
# Alias the canonical result type; matches the public docs.
|
||||
VideoResult = GenerationResult
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoProgressEvent:
|
||||
"""Per-step progress event emitted by :meth:`VideoGenerator.generate_async`.
|
||||
|
||||
Consumers treat these as best-effort telemetry; ``total_steps`` is
|
||||
the count the pipeline reported at the start of the run, not a
|
||||
rolling estimate.
|
||||
"""
|
||||
|
||||
step: int
|
||||
total_steps: int
|
||||
stage: str = "denoise"
|
||||
"""Logical stage name (``denoise`` | ``refine`` | ``decode`` | …)."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoPartialEvent:
|
||||
"""Chunk of decoded frames ready for streaming.
|
||||
|
||||
Emitted only on the streaming path; the aggregated code path never
|
||||
yields partials. ``frames`` is a numpy ``(N, H, W, 3)`` uint8
|
||||
ndarray; ``index`` is a monotonic chunk index starting at 0.
|
||||
"""
|
||||
|
||||
frames: Any
|
||||
index: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoFinalEvent:
|
||||
"""Terminal event carrying the generated video and metadata.
|
||||
|
||||
Exactly one ``VideoFinalEvent`` is emitted per request. When
|
||||
``request.output.return_state`` is True the event also carries the
|
||||
:class:`ContinuationState` the caller needs to resume.
|
||||
"""
|
||||
|
||||
video_bytes: bytes | None = None
|
||||
tensor: Any | None = None
|
||||
frames: Any | None = None
|
||||
metadata: dict[str, Any] = field(default_factory=dict)
|
||||
continuation_state: ContinuationState | None = None
|
||||
result: VideoResult | None = None
|
||||
"""The full :class:`VideoResult` for callers that want everything.
|
||||
|
||||
Streaming consumers typically only care about ``frames`` /
|
||||
``continuation_state``; keeping the full result here avoids a
|
||||
second code path."""
|
||||
|
||||
|
||||
VideoEvent = VideoProgressEvent | VideoPartialEvent | VideoFinalEvent
|
||||
"""Union of every event :meth:`VideoGenerator.generate_async` yields.
|
||||
|
||||
Consumers match by ``isinstance`` rather than ``type`` so subclasses
|
||||
(e.g. a future ``VideoAudioSegmentEvent``) slot in without breaking
|
||||
existing code."""
|
||||
|
||||
__all__ = [
|
||||
"GenerationResult",
|
||||
"VideoEvent",
|
||||
"VideoFinalEvent",
|
||||
"VideoPartialEvent",
|
||||
"VideoProgressEvent",
|
||||
"VideoResult",
|
||||
]
|
||||
@@ -0,0 +1,411 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from v2._vendor.logger import init_logger
|
||||
from v2._vendor.utils import StoreBoolean
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from v2._vendor.api.schema import ContinuationState
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingParam:
|
||||
"""
|
||||
Sampling parameters for video generation.
|
||||
"""
|
||||
# All fields below are copied from ForwardBatch
|
||||
data_type: str = "video"
|
||||
|
||||
# Image inputs
|
||||
image_path: str | None = None
|
||||
pil_image: Any | None = None
|
||||
|
||||
# Video inputs
|
||||
video_path: str | None = None
|
||||
|
||||
# Optional pre-generated diffusion latents. Used by parity/debug harnesses
|
||||
# and advanced callers that need deterministic latent reuse.
|
||||
latents: Any | None = None
|
||||
|
||||
# Action control inputs (Matrix-Game)
|
||||
mouse_cond: Any | None = None # Shape: (B, T, 2)
|
||||
keyboard_cond: Any | None = None # Shape: (B, T, K)
|
||||
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
|
||||
|
||||
# Camera control inputs (HYWorld)
|
||||
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
|
||||
prompt_attention_mask: list = field(default_factory=list)
|
||||
negative_attention_mask: list = field(default_factory=list)
|
||||
|
||||
# Camera/action control inputs (GameCraft)
|
||||
camera_states: Any | None = None # Plücker coordinates [B, T_video, 6, H, W]
|
||||
camera_trajectory: str | None = None
|
||||
action_list: list[str] | None = None
|
||||
action_speed_list: list[float] | None = None
|
||||
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
|
||||
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
|
||||
|
||||
# Camera control inputs (LingBotWorld)
|
||||
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
|
||||
# Refine inputs (LongCat 480p->720p upscaling)
|
||||
# Path-based refine (load stage1 video from disk, e.g. MP4)
|
||||
refine_from: str | None = None # Path to stage1 video (480p output from distill)
|
||||
t_thresh: float = 0.5 # Threshold for timestep scheduling in refinement
|
||||
spatial_refine_only: bool = False # If True, only spatial (no temporal doubling)
|
||||
num_cond_frames: int = 0 # Number of conditioning frames
|
||||
# In-memory refine input (for two-stage pipeline where stage1 frames are already in memory)
|
||||
# This mirrors LongCat's demo where a list of frames (e.g. np.ndarray or PIL.Image)
|
||||
# is passed directly to the refinement pipeline instead of reloading from disk.
|
||||
stage1_video: Any | None = None
|
||||
|
||||
# Text inputs
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
max_sequence_length: int | None = None
|
||||
prompt_path: str | None = None
|
||||
output_path: str = "outputs/"
|
||||
output_video_name: str | None = None
|
||||
|
||||
# Batch info
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 1024
|
||||
|
||||
# Original dimensions (before VAE scaling)
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
fps: int = 24
|
||||
|
||||
# Denoising parameters
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
|
||||
# GEN3C camera control
|
||||
trajectory_type: str | None = None
|
||||
movement_distance: float | None = None
|
||||
camera_rotation: str | None = None
|
||||
|
||||
# LTX-2 multi-modal CFG and STG.
|
||||
# Class-level defaults match the *distilled* LTX-2 schedule
|
||||
# (mirrors ``FastVideo-internal/.../LTX2DistilledSamplingParam``):
|
||||
# the distilled model expects neutral guidance scales — modality 1,
|
||||
# rescale 0, STG 0 — and explicit-CFG callers (full LTX-2) opt back
|
||||
# in by selecting the ``LTX2_BASE`` preset, which overrides these
|
||||
# to mod=3.0 / rescale=0.7 / stg=1.0 in its ``defaults`` dict.
|
||||
# cfg_scale defaults stay at 1.0 (CFG off) so
|
||||
# ``ForwardBatch.__post_init__`` doesn't force CFG on non-LTX-2
|
||||
# models that never override these fields.
|
||||
ltx2_cfg_scale_video: float = 1.0
|
||||
ltx2_cfg_scale_audio: float = 1.0
|
||||
ltx2_modality_scale_video: float = 1.0
|
||||
ltx2_modality_scale_audio: float = 1.0
|
||||
ltx2_rescale_scale: float = 0.0
|
||||
ltx2_stg_scale_video: float = 0.0
|
||||
ltx2_stg_scale_audio: float = 0.0
|
||||
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
|
||||
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
|
||||
|
||||
# LTX-2 image / video / continuation conditioning. These flow from
|
||||
# generate_video(...) kwargs through ``sampling_param.update(kwargs)``
|
||||
# onto the ForwardBatch fields of the same name. ``ltx2_image_crf``
|
||||
# gates the conditioning-image H.264 re-encode; the streaming
|
||||
# session controller passes ``ltx2_image_crf=0.0`` because it
|
||||
# conditions on already-decoded VAE-quality frames.
|
||||
ltx2_images: list[tuple[str, int, float]] | None = None
|
||||
ltx2_image_crf: float = 33.0
|
||||
ltx2_conditioning_latent_stage1: Any | None = None
|
||||
ltx2_conditioning_latent_stage2: Any | None = None
|
||||
ltx2_video_conditions: list[tuple[list[str], int, float]] | None = None
|
||||
|
||||
# Stable Audio (T2A): clip start/end in seconds. Honored by
|
||||
# `StableAudioConditioningStage` + `StableAudioDecodingStage`. Other
|
||||
# families ignore them.
|
||||
audio_start_in_s: float | None = None
|
||||
audio_end_in_s: float | None = None
|
||||
|
||||
# Stable Audio audio-to-audio (variation):
|
||||
# `init_audio` -- a path or `[B, C, samples]` waveform at the model
|
||||
# sample rate; the pipeline encodes it via the VAE
|
||||
# and uses it as the starting latent.
|
||||
# `init_audio_strength` -- 0..1, higher = closer to the reference
|
||||
# (matches the convention of Stability's
|
||||
# commercial Stable Audio 2.0 UI). 1.0 ~=
|
||||
# VAE round-trip, 0.0 ~= plain T2A.
|
||||
# `init_noise_level` -- legacy raw `sigma_max` override (0.3..500,
|
||||
# higher = more freedom). Kept for callers
|
||||
# that already use it; prefer `init_audio_strength`.
|
||||
init_audio: Any = None
|
||||
init_audio_strength: float | None = None
|
||||
init_noise_level: float | None = None
|
||||
|
||||
# Stable Audio inpainting (RePaint-style): `inpaint_audio` is the
|
||||
# reference clip, `inpaint_mask` is a [samples] tensor in {0, 1} where
|
||||
# 1 means *keep the reference* and 0 means *regenerate*.
|
||||
inpaint_audio: Any = None
|
||||
inpaint_mask: Any = None
|
||||
|
||||
# Continuation state carried across streaming/multi-segment calls.
|
||||
continuation_state: ContinuationState | None = None
|
||||
# When True, the pipeline returns a ContinuationState on the result so
|
||||
# the caller can resume from the generated segment.
|
||||
return_continuation_state: bool = False
|
||||
|
||||
# Misc
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
return_trajectory_latents: bool = False # returns all latents for each timestep
|
||||
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.data_type = "video" if self.num_frames > 1 else "image"
|
||||
|
||||
def check_sampling_param(self):
|
||||
if self.prompt_path and not self.prompt_path.endswith(".txt"):
|
||||
raise ValueError("prompt_path must be a txt file")
|
||||
|
||||
def update(self, source_dict: dict[str, Any]) -> None:
|
||||
valid_fields = {f.name for f in fields(self)}
|
||||
unknown = [key for key in source_dict if key not in valid_fields]
|
||||
if unknown:
|
||||
raise ValueError(f"{type(self).__name__}.update() received unknown field(s): "
|
||||
f"{sorted(unknown)}. All kwargs must correspond to declared "
|
||||
f"SamplingParam fields. If a kwarg is meant to flow into "
|
||||
f"ForwardBatch.extra (e.g. LTX2 audio conditioning), route it "
|
||||
f"via VideoGenerator._BATCH_EXTRA_PASSTHROUGH_KEYS instead.")
|
||||
for key, value in source_dict.items():
|
||||
setattr(self, key, value)
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> SamplingParam:
|
||||
sampling_param = cls._from_preset(model_path)
|
||||
if sampling_param is not None:
|
||||
return sampling_param
|
||||
|
||||
logger.warning(
|
||||
"Couldn't find a preset for %s."
|
||||
" Using the default sampling param.",
|
||||
model_path,
|
||||
)
|
||||
return cls()
|
||||
|
||||
@classmethod
|
||||
def _from_preset(
|
||||
cls,
|
||||
model_path: str,
|
||||
) -> SamplingParam | None:
|
||||
"""Build a SamplingParam from preset defaults.
|
||||
|
||||
Returns ``None`` when no preset is configured for
|
||||
*model_path*, letting the caller fall back to the legacy
|
||||
subclass lookup.
|
||||
"""
|
||||
from v2.registry import get_preset_selection
|
||||
|
||||
try:
|
||||
preset_name, model_family = get_preset_selection(model_path)
|
||||
except (ValueError, RuntimeError):
|
||||
return None
|
||||
if preset_name is None or model_family is None:
|
||||
return None
|
||||
|
||||
from v2._vendor.api.presets import get_preset
|
||||
|
||||
preset = get_preset(preset_name, model_family)
|
||||
sp = cls()
|
||||
valid_fields = {f.name for f in fields(cls)}
|
||||
for key, value in preset.defaults.items():
|
||||
if key in valid_fields:
|
||||
setattr(sp, key, copy.deepcopy(value))
|
||||
sp.__post_init__()
|
||||
return sp
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: Any) -> Any:
|
||||
"""Add CLI arguments for SamplingParam fields"""
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
default=SamplingParam.prompt,
|
||||
help="Text prompt for video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--negative-prompt",
|
||||
type=str,
|
||||
default=SamplingParam.negative_prompt,
|
||||
help="Negative text prompt for video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--prompt-path",
|
||||
type=str,
|
||||
default=SamplingParam.prompt_path,
|
||||
help="Path to a text file containing the prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
default=SamplingParam.output_path,
|
||||
help="Path to save the generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-video-name",
|
||||
type=str,
|
||||
default=SamplingParam.output_video_name,
|
||||
help="Name of the output video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-videos-per-prompt",
|
||||
type=int,
|
||||
default=SamplingParam.num_videos_per_prompt,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=SamplingParam.seed,
|
||||
help="Random seed for generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=SamplingParam.num_frames,
|
||||
help="Number of frames to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=SamplingParam.height,
|
||||
help="Height of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=SamplingParam.width,
|
||||
help="Width of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=SamplingParam.fps,
|
||||
help="Frames per second for saved video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=SamplingParam.num_inference_steps,
|
||||
help="Number of denoising steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=SamplingParam.guidance_scale,
|
||||
help="Classifier-free guidance scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
default=SamplingParam.guidance_rescale,
|
||||
help="Guidance rescale factor",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--boundary-ratio",
|
||||
type=float,
|
||||
default=SamplingParam.boundary_ratio,
|
||||
help="Boundary timestep ratio",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save-video",
|
||||
action="store_true",
|
||||
default=SamplingParam.save_video,
|
||||
help="Whether to save the video to disk",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--no-save-video",
|
||||
action="store_false",
|
||||
dest="save_video",
|
||||
help="Don't save the video to disk",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-frames",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Whether to return the raw frames",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--image-path",
|
||||
type=str,
|
||||
default=SamplingParam.image_path,
|
||||
help="Path to input image for image-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video-path",
|
||||
type=str,
|
||||
default=SamplingParam.video_path,
|
||||
help="Path to input video for video-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--refine-from",
|
||||
type=str,
|
||||
default=SamplingParam.refine_from,
|
||||
help="Path to stage1 video for refinement (LongCat 480p->720p)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--t-thresh",
|
||||
type=float,
|
||||
default=SamplingParam.t_thresh,
|
||||
help="Threshold for timestep scheduling in refinement (default: 0.5)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--spatial-refine-only",
|
||||
action=StoreBoolean,
|
||||
default=SamplingParam.spatial_refine_only,
|
||||
help="Only perform spatial super-resolution (no temporal doubling)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-cond-frames",
|
||||
type=int,
|
||||
default=SamplingParam.num_cond_frames,
|
||||
help="Number of conditioning frames for refinement",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--moba-config-path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Path to a JSON file containing V-MoBA specific configurations.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-latents",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_latents,
|
||||
help="Whether to return the trajectory",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--return-trajectory-decoded",
|
||||
action="store_true",
|
||||
default=SamplingParam.return_trajectory_decoded,
|
||||
help="Whether to return the decoded trajectory",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheParams:
|
||||
cache_type: str = "none"
|
||||
@@ -0,0 +1,307 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Literal
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServerConfig:
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 8000
|
||||
output_dir: str = "outputs/"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParallelismConfig:
|
||||
tp_size: int = -1
|
||||
sp_size: int = -1
|
||||
hsdp_replicate_dim: int = 1
|
||||
hsdp_shard_dim: int = -1
|
||||
dist_timeout: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class OffloadConfig:
|
||||
dit: bool = True
|
||||
dit_layerwise: bool = True
|
||||
text_encoder: bool = True
|
||||
image_encoder: bool = True
|
||||
vae: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompileConfig:
|
||||
"""Typed ``torch.compile`` configuration.
|
||||
|
||||
``backend``/``fullgraph``/``mode``/``dynamic`` are the four most
|
||||
common ``torch.compile`` knobs. ``extras`` holds any remaining
|
||||
``torch.compile`` kwargs (e.g. ``options``, ``disable``).
|
||||
|
||||
The ``enabled`` switch covers the DiT transformer path (including
|
||||
``transformer_2`` and the LTX-2 stage-2 ``transformer_refine``).
|
||||
Per-component flags below are independent overlays — set to ``True``
|
||||
to compile that component, ``None`` to leave it eager. Each
|
||||
``*_kwargs`` dict overrides the master ``backend``/``fullgraph``/
|
||||
``mode``/``dynamic``/``extras`` for that component when non-empty;
|
||||
leaving it empty inherits the master kwargs.
|
||||
"""
|
||||
|
||||
enabled: bool = False
|
||||
backend: str | None = None
|
||||
fullgraph: bool | None = None
|
||||
mode: str | None = None
|
||||
dynamic: bool | None = None
|
||||
extras: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
text_encoder_enabled: bool | None = None
|
||||
vae_enabled: bool | None = None
|
||||
audio_vae_enabled: bool | None = None
|
||||
|
||||
dit_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
text_encoder_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
vae_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
audio_vae_kwargs: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class QuantizationConfig:
|
||||
text_encoder_quant: str | None = None
|
||||
transformer_quant: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class EngineConfig:
|
||||
num_gpus: int = 1
|
||||
execution_backend: Literal["mp", "ray"] = "mp"
|
||||
parallelism: ParallelismConfig = field(default_factory=ParallelismConfig)
|
||||
offload: OffloadConfig = field(default_factory=OffloadConfig)
|
||||
compile: CompileConfig = field(default_factory=CompileConfig)
|
||||
enable_stage_verification: bool = True
|
||||
use_fsdp_inference: bool = False
|
||||
disable_autocast: bool = False
|
||||
quantization: QuantizationConfig | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ComponentConfig:
|
||||
config_root: str | None = None
|
||||
pipeline_config_path: str | None = None
|
||||
text_encoder_weights: str | None = None
|
||||
transformer_weights: str | None = None
|
||||
transformer_2_weights: str | None = None
|
||||
vae_weights: str | None = None
|
||||
upsampler_weights: str | None = None
|
||||
lora_path: str | None = None
|
||||
override_pipeline_cls_name: str | None = None
|
||||
override_transformer_cls_name: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineSelection:
|
||||
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
|
||||
preset: str | None = None
|
||||
preset_version: int | None = None
|
||||
components: ComponentConfig = field(default_factory=ComponentConfig)
|
||||
vae_tiling: bool | None = None
|
||||
"""Tile-based VAE decode. ``None`` keeps the model's default."""
|
||||
preset_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
experimental: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GeneratorConfig:
|
||||
model_path: str
|
||||
revision: str | None = None
|
||||
trust_remote_code: bool = False
|
||||
engine: EngineConfig = field(default_factory=EngineConfig)
|
||||
pipeline: PipelineSelection = field(default_factory=PipelineSelection)
|
||||
|
||||
|
||||
@dataclass
|
||||
class InputConfig:
|
||||
prompt_path: str | None = None
|
||||
image_path: str | list[str] | None = None
|
||||
video_path: str | list[str] | None = None
|
||||
pil_image: Any | None = None
|
||||
pose: str | None = None
|
||||
mouse_cond: Any | None = None
|
||||
keyboard_cond: Any | None = None
|
||||
grid_sizes: Any | None = None
|
||||
c2ws_plucker_emb: Any | None = None
|
||||
refine_from: str | None = None
|
||||
stage1_video: Any | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingConfig:
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 1024
|
||||
num_frames: int = 125
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
height_sr: int = 1072
|
||||
width_sr: int = 1920
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 50
|
||||
num_inference_steps_sr: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_scale_2: float | None = None
|
||||
guidance_rescale: float = 0.0
|
||||
true_cfg_scale: float | None = None
|
||||
boundary_ratio: float | None = None
|
||||
sigmas: list[float] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestRuntimeConfig:
|
||||
enable_teacache: bool = False
|
||||
return_trajectory_latents: bool = False
|
||||
return_trajectory_decoded: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class OutputConfig:
|
||||
output_path: str = "outputs/"
|
||||
output_video_name: str | None = None
|
||||
save_video: bool = True
|
||||
return_frames: bool = True
|
||||
return_state: bool = False
|
||||
|
||||
|
||||
@dataclass
|
||||
class ContinuationState:
|
||||
kind: str
|
||||
payload: dict[str, Any]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PlannedStage:
|
||||
name: str
|
||||
kind: str
|
||||
source: str | None = None
|
||||
overrides: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationPlan:
|
||||
stages: list[PlannedStage]
|
||||
final_stage: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class GenerationRequest:
|
||||
prompt: str | list[str] | None = None
|
||||
negative_prompt: str | None = None
|
||||
inputs: InputConfig = field(default_factory=InputConfig)
|
||||
sampling: SamplingConfig = field(default_factory=SamplingConfig)
|
||||
runtime: RequestRuntimeConfig = field(default_factory=RequestRuntimeConfig)
|
||||
output: OutputConfig = field(default_factory=OutputConfig)
|
||||
stage_overrides: dict[str, Any] = field(default_factory=dict)
|
||||
state: ContinuationState | None = None
|
||||
plan: GenerationPlan | None = None
|
||||
extensions: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RunConfig:
|
||||
generator: GeneratorConfig
|
||||
request: GenerationRequest
|
||||
|
||||
|
||||
@dataclass
|
||||
class WarmupConfig:
|
||||
enabled: bool = True
|
||||
prompt: str = ("A cinematic drone shot over coastal cliffs at sunrise, "
|
||||
"golden light, gentle ocean waves, ultra detailed")
|
||||
timeout_seconds: int = 2400
|
||||
|
||||
|
||||
@dataclass
|
||||
class GpuPoolConfig:
|
||||
num_workers: int | None = None
|
||||
enable_audio_reencode: bool = True
|
||||
conditioning_num_frames: int = 9
|
||||
conditioning_end_offset: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptEnhancerConfig:
|
||||
enabled: bool = False
|
||||
provider: Literal["cerebras", "groq"] = "cerebras"
|
||||
model: str = "gpt-oss-120b"
|
||||
timeout_ms: int = 20000
|
||||
system_prompt_dir: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptSafetyConfig:
|
||||
enabled: bool = False
|
||||
classifier_path: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamingConfig:
|
||||
session_timeout_seconds: int = 300
|
||||
generation_segment_cap: int = 6
|
||||
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
|
||||
warmup: WarmupConfig = field(default_factory=WarmupConfig)
|
||||
pool: GpuPoolConfig = field(default_factory=GpuPoolConfig)
|
||||
prompt: PromptEnhancerConfig = field(default_factory=PromptEnhancerConfig)
|
||||
safety: PromptSafetyConfig = field(default_factory=PromptSafetyConfig)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServeConfig:
|
||||
"""Typed serve config loaded from ``fastvideo serve --config``.
|
||||
|
||||
``default_request`` is a full :class:`GenerationRequest` — the same type
|
||||
clients POST to ``/v1/videos``. At request time the server merges it into
|
||||
the incoming body as the operator-pinned baseline.
|
||||
|
||||
Important nuance: only fields the operator **explicitly wrote** in the
|
||||
serve YAML/JSON count as defaults. Although the in-memory object is
|
||||
fully populated (schema defaults fill every unset field), the merge
|
||||
walks ``_fastvideo_explicit_paths`` — populated during parse — so
|
||||
unset fields are *not* forced onto requests. Per-request precedence:
|
||||
|
||||
body (client-explicit) > default_request (operator-explicit)
|
||||
> hardcoded fallback (e.g. ``fps=24``)
|
||||
|
||||
See :func:`v2._vendor.api.compat.explicit_request_updates` for the
|
||||
projection and ``entrypoints/openai/video_api.py::_build_generation_kwargs``
|
||||
for the merge.
|
||||
"""
|
||||
generator: GeneratorConfig
|
||||
server: ServerConfig = field(default_factory=ServerConfig)
|
||||
default_request: GenerationRequest = field(default_factory=GenerationRequest)
|
||||
streaming: StreamingConfig | None = None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"CompileConfig",
|
||||
"ComponentConfig",
|
||||
"ContinuationState",
|
||||
"EngineConfig",
|
||||
"GenerationPlan",
|
||||
"GenerationRequest",
|
||||
"GeneratorConfig",
|
||||
"GpuPoolConfig",
|
||||
"InputConfig",
|
||||
"OffloadConfig",
|
||||
"OutputConfig",
|
||||
"ParallelismConfig",
|
||||
"PipelineSelection",
|
||||
"PlannedStage",
|
||||
"PromptEnhancerConfig",
|
||||
"PromptSafetyConfig",
|
||||
"QuantizationConfig",
|
||||
"RequestRuntimeConfig",
|
||||
"RunConfig",
|
||||
"SamplingConfig",
|
||||
"ServeConfig",
|
||||
"ServerConfig",
|
||||
"StreamingConfig",
|
||||
"WarmupConfig",
|
||||
]
|
||||
@@ -0,0 +1,58 @@
|
||||
# `fastvideo/attention/` — Attention Backends
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Backend registry + selector wrapping FlashAttn / SageAttn / SageAttn3 / SDPA / VSA / VMoBA / SLA / BSA.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
attention/
|
||||
├── __init__.py # Exports DistributedAttention, LocalAttention, get_attn_backend
|
||||
├── layer.py # DistributedAttention, DistributedAttention_VSA, LocalAttention
|
||||
├── selector.py # get_attn_backend (cached) + env-var override
|
||||
├── backends/
|
||||
│ ├── abstract.py # AttentionBackend / AttentionMetadata / AttentionMetadataBuilder
|
||||
│ ├── flash_attn.py # FA2/FA3
|
||||
│ ├── sage_attn.py # SageAttention v1
|
||||
│ ├── sage_attn3.py # SageAttention v3
|
||||
│ ├── sdpa.py # torch SDPA fallback
|
||||
│ ├── video_sparse_attn.py # VSA (paper: Video Sparse Attention)
|
||||
│ ├── vmoba.py # Video-MoBA
|
||||
│ ├── sla.py # Sliding-window (STA)
|
||||
│ └── bsa_attn.py # Block-sparse
|
||||
└── utils/
|
||||
├── flash_attn_cute.py
|
||||
└── flash_attn_no_pad.py
|
||||
```
|
||||
|
||||
## Selection Order
|
||||
|
||||
`get_attn_backend()` resolves via:
|
||||
|
||||
1. Env-var override `FASTVIDEO_ATTENTION_BACKEND` (see `STR_BACKEND_ENV_VAR` in `fastvideo/utils.py`).
|
||||
2. Per-platform default from `fastvideo/platforms/`.
|
||||
3. Heuristic fallback to SDPA.
|
||||
|
||||
The result is `@lru_cache`d. Tests that need a specific backend must use the
|
||||
`global_force_attn_backend(...)` context manager from `selector.py`, never set
|
||||
the env var mid-process.
|
||||
|
||||
## Adding a Backend
|
||||
|
||||
1. Subclass `AttentionBackend` in `backends/<name>.py`.
|
||||
2. Implement `AttentionMetadata` + `AttentionMetadataBuilder` for the new path.
|
||||
3. Register the enum value in `fastvideo/platforms/interface.py` (`AttentionBackendEnum`).
|
||||
4. Wire string → class resolution in `selector.py`.
|
||||
5. Verify the new backend works with `DistributedAttention` (sequence parallel)
|
||||
and `LocalAttention` (single-rank). If it cannot support SP, document the
|
||||
gap in the backend file's module docstring.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Calling `torch.nn.functional.scaled_dot_product_attention` directly inside a
|
||||
model's forward — go through `DistributedAttention` / `LocalAttention`.
|
||||
- Reading `os.environ[STR_BACKEND_ENV_VAR]` from arbitrary call sites. Use
|
||||
`get_env_variable_attn_backend()`.
|
||||
- Caching backend instances per-module. The selector cache is process-wide; do
|
||||
not duplicate it.
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from v2._vendor.attention.backends.abstract import (AttentionBackend, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from v2._vendor.attention.layer import (DistributedAttention, DistributedAttention_VSA, LocalAttention)
|
||||
from v2._vendor.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
# "AttentionState",
|
||||
"get_attn_backend",
|
||||
]
|
||||
@@ -0,0 +1,177 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/backends/abstract.py
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
|
||||
|
||||
if TYPE_CHECKING:
|
||||
pass
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class AttentionBackend(ABC):
|
||||
"""Abstract class for attention backends."""
|
||||
# For some attention backends, we allocate an output tensor before
|
||||
# calling the custom op. When piecewise cudagraph is enabled, this
|
||||
# makes sure the output tensor is allocated inside the cudagraph.
|
||||
accept_output_buffer: bool = False
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_name() -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_impl_cls() -> type["AttentionImpl"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
# @staticmethod
|
||||
# @abstractmethod
|
||||
# def get_state_cls() -> Type["AttentionState"]:
|
||||
# raise NotImplementedError
|
||||
|
||||
# @classmethod
|
||||
# def make_metadata(cls, *args, **kwargs) -> "AttentionMetadata":
|
||||
# return cls.get_metadata_cls()(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass
|
||||
class AttentionMetadata:
|
||||
"""Attention metadata for prefill and decode batched together."""
|
||||
# Current step of diffusion process
|
||||
current_timestep: int
|
||||
VSA_sparsity: float = field(default=0.0, kw_only=True)
|
||||
|
||||
def __getattr__(self, name: str) -> Any:
|
||||
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
|
||||
|
||||
def asdict_zerocopy(self, skip_fields: set[str] | None = None) -> dict[str, Any]:
|
||||
"""Similar to dataclasses.asdict, but avoids deepcopying."""
|
||||
if skip_fields is None:
|
||||
skip_fields = set()
|
||||
# Note that if we add dataclasses as fields, they will need
|
||||
# similar handling.
|
||||
return {field.name: getattr(self, field.name) for field in fields(self) if field.name not in skip_fields}
|
||||
|
||||
|
||||
T = TypeVar("T", bound=AttentionMetadata)
|
||||
|
||||
|
||||
class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
"""Abstract class for attention metadata builders."""
|
||||
|
||||
@abstractmethod
|
||||
def __init__(self) -> None:
|
||||
"""Create the builder, remember some configuration and parameters."""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def prepare(self) -> None:
|
||||
"""Prepare for one batch."""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def build(
|
||||
self,
|
||||
**kwargs: Any,
|
||||
) -> AttentionMetadata:
|
||||
"""Build attention metadata with on-device tensors."""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class AttentionLayer(Protocol):
|
||||
|
||||
_k_scale: torch.Tensor
|
||||
_v_scale: torch.Tensor
|
||||
_k_scale_float: float
|
||||
_v_scale_float: float
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
...
|
||||
|
||||
|
||||
class AttentionImpl(ABC, Generic[T]):
|
||||
|
||||
@abstractmethod
|
||||
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:
|
||||
raise NotImplementedError
|
||||
|
||||
def preprocess_qkv(self, qkv: torch.Tensor, attn_metadata: T) -> torch.Tensor:
|
||||
"""Preprocess QKV tensor before performing attention operation.
|
||||
|
||||
Default implementation returns the tensor unchanged.
|
||||
Subclasses can override this to implement custom preprocessing
|
||||
like reshaping, tiling, scaling, or other transformations.
|
||||
|
||||
Called AFTER all_to_all for distributed attention
|
||||
|
||||
Args:
|
||||
qkv: The query-key-value tensor
|
||||
attn_metadata: Metadata for the attention operation
|
||||
|
||||
Returns:
|
||||
Processed QKV tensor
|
||||
"""
|
||||
return qkv
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: T,
|
||||
) -> torch.Tensor:
|
||||
"""Postprocess the output tensor after the attention operation.
|
||||
|
||||
Default implementation returns the tensor unchanged.
|
||||
Subclasses can override this to implement custom postprocessing
|
||||
like untiling, scaling, or other transformations.
|
||||
|
||||
Called BEFORE all_to_all for distributed attention
|
||||
|
||||
Args:
|
||||
output: The output tensor from the attention operation
|
||||
attn_metadata: Metadata for the attention operation
|
||||
|
||||
Returns:
|
||||
Postprocessed output tensor
|
||||
"""
|
||||
|
||||
return output
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: T,
|
||||
) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,125 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from v2._vendor.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from v2._vendor.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_project_root = Path(__file__).resolve().parent.parent.parent.parent
|
||||
_kernel_root = _project_root / "fastvideo-kernel"
|
||||
_kernel_python_root = _kernel_root / "python"
|
||||
_attn_qat_infer: Callable[..., torch.Tensor] | None = None
|
||||
_attn_qat_infer_import_attempted = False
|
||||
|
||||
|
||||
def _ensure_kernel_paths() -> None:
|
||||
for path in (_project_root, _kernel_root, _kernel_python_root):
|
||||
path_str = str(path)
|
||||
if path_str not in sys.path:
|
||||
sys.path.insert(0, path_str)
|
||||
|
||||
|
||||
def _get_attn_qat_infer() -> Callable[..., torch.Tensor] | None:
|
||||
global _attn_qat_infer
|
||||
global _attn_qat_infer_import_attempted
|
||||
|
||||
if _attn_qat_infer_import_attempted:
|
||||
return _attn_qat_infer
|
||||
|
||||
_attn_qat_infer_import_attempted = True
|
||||
_ensure_kernel_paths()
|
||||
|
||||
try:
|
||||
# Prefer the in-repo kernel implementation during local development.
|
||||
_attn_qat_infer = importlib.import_module("attn_qat_infer").sageattn_blackwell
|
||||
except ImportError:
|
||||
_attn_qat_infer = None
|
||||
|
||||
return _attn_qat_infer
|
||||
|
||||
|
||||
def is_attn_qat_infer_available() -> bool:
|
||||
return _get_attn_qat_infer() is not None
|
||||
|
||||
|
||||
class AttnQatInferBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "ATTN_QAT_INFER"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["AttnQatInferImpl"]:
|
||||
return AttnQatInferImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder[AttentionMetadata]"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class AttnQatInferImpl(AttentionImpl[AttentionMetadata]):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
dropout_p = extra_impl_args.get("dropout_p", 0.0)
|
||||
if dropout_p > 0:
|
||||
raise NotImplementedError(f"attn_qat_infer does not support dropout (got dropout_p={dropout_p}). "
|
||||
"The QAT inference kernel applies no stochastic dropout.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
attn_qat_infer = _get_attn_qat_infer()
|
||||
if attn_qat_infer is None:
|
||||
raise ImportError("attn_qat_infer is not available. Please ensure the "
|
||||
"attn_qat_infer kernel package is installed.")
|
||||
|
||||
query = query.transpose(1, 2).contiguous()
|
||||
key = key.transpose(1, 2).contiguous()
|
||||
value = value.transpose(1, 2).contiguous()
|
||||
|
||||
output = attn_qat_infer(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
attn_mask=None,
|
||||
is_causal=self.causal,
|
||||
sm_scale=self.softmax_scale,
|
||||
)
|
||||
return output.transpose(1, 2).contiguous()
|
||||
@@ -0,0 +1,154 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import importlib
|
||||
import sys
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from v2._vendor.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from v2._vendor.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_project_root = Path(__file__).resolve().parent.parent.parent.parent
|
||||
_kernel_root = _project_root / "fastvideo-kernel"
|
||||
_kernel_python_root = _kernel_root / "python"
|
||||
_attn_qat_train_attention: Callable[..., torch.Tensor] | None = None
|
||||
_attn_qat_train_import_attempted = False
|
||||
|
||||
|
||||
def _ensure_kernel_paths() -> None:
|
||||
for path in (_project_root, _kernel_root, _kernel_python_root):
|
||||
path_str = str(path)
|
||||
if path_str not in sys.path:
|
||||
sys.path.insert(0, path_str)
|
||||
|
||||
|
||||
def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
|
||||
global _attn_qat_train_attention
|
||||
global _attn_qat_train_import_attempted
|
||||
|
||||
if _attn_qat_train_import_attempted:
|
||||
return _attn_qat_train_attention
|
||||
|
||||
_attn_qat_train_import_attempted = True
|
||||
_ensure_kernel_paths()
|
||||
|
||||
try:
|
||||
_attn_qat_train_attention = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train").attention
|
||||
except ImportError:
|
||||
_attn_qat_train_attention = None
|
||||
|
||||
return _attn_qat_train_attention
|
||||
|
||||
|
||||
def is_attn_qat_train_available() -> bool:
|
||||
return _get_attn_qat_train_attention() is not None
|
||||
|
||||
|
||||
def attn_qat_train(q_BLHD: torch.Tensor,
|
||||
k_BLHD: torch.Tensor,
|
||||
v_BLHD: torch.Tensor,
|
||||
is_causal: bool = False,
|
||||
sm_scale: float | None = None) -> torch.Tensor:
|
||||
attention = _get_attn_qat_train_attention()
|
||||
if attention is None:
|
||||
raise ImportError("fastvideo_kernel.triton_kernels.attn_qat_train is not available. "
|
||||
"Please ensure the FastVideo kernel package is installed.")
|
||||
|
||||
q_BHLD = q_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
k_BHLD = k_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
v_BHLD = v_BLHD.permute(0, 2, 1, 3).contiguous()
|
||||
|
||||
use_qat_qkv_backward = True
|
||||
smooth_k = False
|
||||
warp_specialize = True
|
||||
is_qat = True
|
||||
two_level_quant_p_sage3 = False
|
||||
fake_quant_p_bwd = True
|
||||
use_high_prec_o = True
|
||||
smooth_q = False
|
||||
if sm_scale is None:
|
||||
sm_scale = 1.0 / (q_BHLD.shape[-1]**0.5)
|
||||
use_global_sf_qkv = False
|
||||
use_global_sf_p = False
|
||||
|
||||
o_BHLD = attention(
|
||||
q_BHLD,
|
||||
k_BHLD,
|
||||
v_BHLD,
|
||||
is_causal,
|
||||
sm_scale,
|
||||
use_qat_qkv_backward,
|
||||
smooth_k,
|
||||
warp_specialize,
|
||||
is_qat,
|
||||
two_level_quant_p_sage3,
|
||||
fake_quant_p_bwd,
|
||||
use_high_prec_o,
|
||||
smooth_q,
|
||||
use_global_sf_p,
|
||||
use_global_sf_qkv,
|
||||
)
|
||||
return o_BHLD.permute(0, 2, 1, 3).contiguous()
|
||||
|
||||
|
||||
class AttnQatTrainBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "ATTN_QAT_TRAIN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["AttnQatTrainImpl"]:
|
||||
return AttnQatTrainImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder[AttentionMetadata]"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class AttnQatTrainImpl(AttentionImpl[AttentionMetadata]):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
dropout_p = extra_impl_args.get("dropout_p", 0.0)
|
||||
if dropout_p > 0:
|
||||
raise NotImplementedError(f"attn_qat_train does not support dropout (got dropout_p={dropout_p}). "
|
||||
"The QAT training kernel applies no stochastic dropout.")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
return attn_qat_train(query, key, value, is_causal=self.causal, sm_scale=self.softmax_scale)
|
||||
@@ -0,0 +1,740 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Bidirectional Sparse Attention (BSA) backend for FastVideo.
|
||||
|
||||
Pure-PyTorch reference implementation from:
|
||||
"Bidirectional Sparse Attention for Faster Video Diffusion Training"
|
||||
(arXiv:2509.01085)
|
||||
|
||||
BSA sparsifies both queries (pruning redundant tokens per block) and
|
||||
key-value pairs (keeping only relevant KV blocks per query block).
|
||||
|
||||
This is a training-free inference backend: it works with any model
|
||||
trained with full attention by applying BSA sparsity at inference time.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from v2._vendor.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from v2._vendor.distributed import get_sp_group
|
||||
from v2._vendor.logger import init_logger
|
||||
|
||||
try:
|
||||
from v2._vendor.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_varlen_func_impl, )
|
||||
|
||||
FLASH_ATTN_AVAILABLE = True
|
||||
except ImportError:
|
||||
flash_attn_varlen_func_impl = None
|
||||
FLASH_ATTN_AVAILABLE = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
BSA_TILE_SIZE = (4, 4, 4)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cached index helpers (same pattern as VSA)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_tile_partition_indices(
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Map raster-order tokens to tile-contiguous order."""
|
||||
T, H, W = dit_seq_shape
|
||||
ts, hs, ws = tile_size
|
||||
indices = torch.arange(T * H * W, device=device, dtype=torch.long).reshape(T, H, W)
|
||||
ls = []
|
||||
for t in range(math.ceil(T / ts)):
|
||||
for h in range(math.ceil(H / hs)):
|
||||
for w in range(math.ceil(W / ws)):
|
||||
ls.append(indices[
|
||||
t * ts:min(t * ts + ts, T),
|
||||
h * hs:min(h * hs + hs, H),
|
||||
w * ws:min(w * ws + ws, W),
|
||||
].flatten())
|
||||
return torch.cat(ls, dim=0)
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def get_reverse_tile_partition_indices(
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
tile_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
) -> torch.LongTensor:
|
||||
"""Inverse mapping: tile-contiguous order back to raster order."""
|
||||
return torch.argsort(get_tile_partition_indices(dit_seq_shape, tile_size, device))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BSA core operations
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _prune_queries(
|
||||
q_blocks: torch.Tensor,
|
||||
keep_ratio: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
||||
"""
|
||||
Prune redundant query tokens within each block.
|
||||
|
||||
Scores tokens by cosine similarity to the block center.
|
||||
Keeps the LEAST similar (most informative) tokens.
|
||||
|
||||
Args:
|
||||
q_blocks: [B, N_heads, N_blocks, block_size, D]
|
||||
keep_ratio: fraction of tokens to keep
|
||||
|
||||
Returns:
|
||||
sparse_q: [B, N_heads, N_blocks, keep_size, D]
|
||||
keep_indices: [B, N_heads, N_blocks, keep_size]
|
||||
keep_size: int
|
||||
"""
|
||||
B, H, N, S, D = q_blocks.shape
|
||||
keep_size = max(1, int(S * keep_ratio))
|
||||
|
||||
if keep_size >= S:
|
||||
idx = torch.arange(S, device=q_blocks.device)
|
||||
idx = idx.view(1, 1, 1, S).expand(B, H, N, S)
|
||||
return q_blocks, idx, S
|
||||
|
||||
center_idx = S // 2
|
||||
center = q_blocks[:, :, :, center_idx:center_idx + 1, :]
|
||||
|
||||
q_norm = F.normalize(q_blocks, dim=-1)
|
||||
c_norm = F.normalize(center, dim=-1)
|
||||
similarity = (q_norm * c_norm).sum(dim=-1) # [B, H, N, S]
|
||||
|
||||
# lowest similarity = most distinctive = keep
|
||||
_, indices = similarity.topk(keep_size, dim=-1, largest=False)
|
||||
indices, _ = indices.sort(dim=-1)
|
||||
|
||||
idx_expand = indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
|
||||
sparse_q = torch.gather(q_blocks, 3, idx_expand)
|
||||
|
||||
return sparse_q, indices, keep_size
|
||||
|
||||
|
||||
def _select_kv_blocks(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
cumulative_threshold: float,
|
||||
min_kv_blocks: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Dynamically select KV blocks for each query block.
|
||||
|
||||
Mean-pools to block level, computes block attention scores,
|
||||
admits blocks in descending order until cumulative mass
|
||||
exceeds threshold.
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
cumulative_threshold: e.g. 0.9
|
||||
min_kv_blocks: minimum blocks to keep
|
||||
|
||||
Returns:
|
||||
kv_mask: [B, H, N, N] boolean
|
||||
"""
|
||||
B, H, N, _, D = sparse_q.shape
|
||||
|
||||
q_repr = sparse_q.mean(dim=3)
|
||||
k_repr = k_blocks.mean(dim=3)
|
||||
|
||||
scores = torch.matmul(q_repr, k_repr.transpose(-1, -2)) / (D**0.5)
|
||||
block_attn = F.softmax(scores, dim=-1)
|
||||
|
||||
sorted_attn, sorted_idx = block_attn.sort(dim=-1, descending=True)
|
||||
cumsum = sorted_attn.cumsum(dim=-1)
|
||||
|
||||
keep_sorted = torch.ones_like(cumsum, dtype=torch.bool)
|
||||
keep_sorted[..., 1:] = cumsum[..., :-1] < cumulative_threshold
|
||||
|
||||
min_mask = torch.zeros_like(keep_sorted)
|
||||
min_mask[..., :min(min_kv_blocks, N)] = True
|
||||
keep_sorted = keep_sorted | min_mask
|
||||
|
||||
kv_mask = torch.zeros_like(block_attn, dtype=torch.bool)
|
||||
kv_mask.scatter_(-1, sorted_idx, keep_sorted)
|
||||
|
||||
return kv_mask
|
||||
|
||||
|
||||
def _compute_sparse_attention(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute attention for each query block against selected KV blocks.
|
||||
|
||||
Handles per-batch and per-head KV masks correctly.
|
||||
Uses flash_attn_varlen_func when available on GPU.
|
||||
Falls back to pure-PyTorch reference on CPU.
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
v_blocks: [B, H, N, Sk, D]
|
||||
kv_mask: [B, H, N, N] boolean (per-batch, per-head)
|
||||
|
||||
Returns:
|
||||
output: [B, H, N, Sq, D]
|
||||
"""
|
||||
if FLASH_ATTN_AVAILABLE and sparse_q.is_cuda:
|
||||
return _compute_sparse_attention_flash(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
else:
|
||||
return _compute_sparse_attention_reference(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
|
||||
|
||||
def _compute_sparse_attention_reference(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Pure-PyTorch fallback with per-batch, per-head mask support."""
|
||||
B, H, N, Sq, D = sparse_q.shape
|
||||
output = torch.zeros_like(sparse_q)
|
||||
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for qb in range(N):
|
||||
selected = kv_mask[b, h, qb] # [N] boolean
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
# [num_sel * Sk, D]
|
||||
sel_k = k_blocks[b, h, sel_idx].reshape(-1, D)
|
||||
sel_v = v_blocks[b, h, sel_idx].reshape(-1, D)
|
||||
|
||||
q = sparse_q[b, h, qb] # [Sq, D]
|
||||
scores = torch.matmul(q, sel_k.transpose(-1, -2)) / (D**0.5)
|
||||
weights = F.softmax(scores, dim=-1)
|
||||
output[b, h, qb] = torch.matmul(weights, sel_v)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _compute_sparse_attention_flash(
|
||||
sparse_q: torch.Tensor,
|
||||
k_blocks: torch.Tensor,
|
||||
v_blocks: torch.Tensor,
|
||||
kv_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
FlashAttention implementation with per-batch, per-head mask support.
|
||||
|
||||
Strategy: check if all heads share the same mask. If so, use a single
|
||||
FlashAttention call per batch (fast path). If not, process each head
|
||||
separately (correct path).
|
||||
|
||||
Args:
|
||||
sparse_q: [B, H, N, Sq, D]
|
||||
k_blocks: [B, H, N, Sk, D]
|
||||
v_blocks: [B, H, N, Sk, D]
|
||||
kv_mask: [B, H, N, N] boolean
|
||||
|
||||
Returns:
|
||||
output: [B, H, N, Sq, D]
|
||||
"""
|
||||
B, H, N, Sq, D = sparse_q.shape
|
||||
Sk = k_blocks.shape[3]
|
||||
device = sparse_q.device
|
||||
output = torch.zeros_like(sparse_q)
|
||||
|
||||
for b in range(B):
|
||||
# Check if all heads share the same mask for this batch element
|
||||
# Compare each head's mask to head 0's mask
|
||||
head0_mask = kv_mask[b, 0] # [N, N]
|
||||
all_heads_same = all(torch.equal(kv_mask[b, h], head0_mask) for h in range(1, H))
|
||||
|
||||
if all_heads_same:
|
||||
# Fast path: all heads share the same mask, single FA call
|
||||
_flash_attn_single_mask(
|
||||
sparse_q[b],
|
||||
k_blocks[b],
|
||||
v_blocks[b],
|
||||
head0_mask,
|
||||
output[b],
|
||||
H,
|
||||
N,
|
||||
Sq,
|
||||
Sk,
|
||||
D,
|
||||
device,
|
||||
)
|
||||
else:
|
||||
# Per-head path: process each head individually
|
||||
for h in range(H):
|
||||
head_mask = kv_mask[b, h] # [N, N]
|
||||
# Process single head: squeeze head dim, run FA, put back
|
||||
_flash_attn_single_head(
|
||||
sparse_q[b, h],
|
||||
k_blocks[b, h],
|
||||
v_blocks[b, h],
|
||||
head_mask,
|
||||
output,
|
||||
b,
|
||||
h,
|
||||
N,
|
||||
Sq,
|
||||
Sk,
|
||||
D,
|
||||
device,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _flash_attn_single_mask(
|
||||
sparse_q_b: torch.Tensor, # [H, N, Sq, D]
|
||||
k_blocks_b: torch.Tensor, # [H, N, Sk, D]
|
||||
v_blocks_b: torch.Tensor, # [H, N, Sk, D]
|
||||
mask: torch.Tensor, # [N, N] boolean
|
||||
output_b: torch.Tensor, # [H, N, Sq, D] (modified in-place)
|
||||
H: int,
|
||||
N: int,
|
||||
Sq: int,
|
||||
Sk: int,
|
||||
D: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Run FlashAttention for all heads sharing the same KV mask."""
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
cu_seqlens_q = [0]
|
||||
cu_seqlens_k = [0]
|
||||
active_blocks = []
|
||||
|
||||
for qb in range(N):
|
||||
selected = mask[qb] # [N] boolean
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
active_blocks.append(qb)
|
||||
num_kv_tokens = sel_idx.shape[0] * Sk
|
||||
|
||||
# [H, Sq, D] -> [Sq, H, D]
|
||||
q_block = sparse_q_b[:, qb].permute(1, 0, 2)
|
||||
q_list.append(q_block)
|
||||
|
||||
# [H, num_sel, Sk, D] -> [num_kv_tokens, H, D]
|
||||
sel_k = k_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
|
||||
sel_v = v_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
|
||||
k_list.append(sel_k)
|
||||
v_list.append(sel_v)
|
||||
|
||||
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
|
||||
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
|
||||
|
||||
if not q_list:
|
||||
return
|
||||
|
||||
flat_q = torch.cat(q_list, dim=0)
|
||||
flat_k = torch.cat(k_list, dim=0)
|
||||
flat_v = torch.cat(v_list, dim=0)
|
||||
|
||||
# Compute max_seqlen_k from the Python list before moving to GPU to
|
||||
# avoid a `.item()` round-trip that would force a host/device sync.
|
||||
max_seqlen_q = Sq
|
||||
max_seqlen_k = max(b - a for a, b in zip(cu_seqlens_k[:-1], cu_seqlens_k[1:], strict=False))
|
||||
|
||||
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
|
||||
|
||||
orig_dtype = flat_q.dtype
|
||||
compute_dtype = orig_dtype
|
||||
if compute_dtype not in (torch.float16, torch.bfloat16):
|
||||
compute_dtype = torch.bfloat16
|
||||
flat_q = flat_q.to(compute_dtype)
|
||||
flat_k = flat_k.to(compute_dtype)
|
||||
flat_v = flat_v.to(compute_dtype)
|
||||
|
||||
flat_out = flash_attn_varlen_func_impl(
|
||||
flat_q,
|
||||
flat_k,
|
||||
flat_v,
|
||||
cu_seqlens_q_t,
|
||||
cu_seqlens_k_t,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
if compute_dtype != orig_dtype:
|
||||
flat_out = flat_out.to(orig_dtype)
|
||||
|
||||
idx = 0
|
||||
for qb in active_blocks:
|
||||
block_out = flat_out[idx:idx + Sq] # [Sq, H, D]
|
||||
output_b[:, qb] = block_out.permute(1, 0, 2) # [H, Sq, D]
|
||||
idx += Sq
|
||||
|
||||
|
||||
def _flash_attn_single_head(
|
||||
sparse_q_bh: torch.Tensor, # [N, Sq, D]
|
||||
k_blocks_bh: torch.Tensor, # [N, Sk, D]
|
||||
v_blocks_bh: torch.Tensor, # [N, Sk, D]
|
||||
mask: torch.Tensor, # [N, N] boolean
|
||||
output: torch.Tensor, # [B, H, N, Sq, D] (modified in-place)
|
||||
b: int,
|
||||
h: int,
|
||||
N: int,
|
||||
Sq: int,
|
||||
Sk: int,
|
||||
D: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
"""Run FlashAttention for a single head with its own KV mask."""
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
cu_seqlens_q = [0]
|
||||
cu_seqlens_k = [0]
|
||||
active_blocks = []
|
||||
|
||||
for qb in range(N):
|
||||
selected = mask[qb]
|
||||
sel_idx = selected.nonzero(as_tuple=True)[0]
|
||||
|
||||
if sel_idx.shape[0] == 0:
|
||||
continue
|
||||
|
||||
active_blocks.append(qb)
|
||||
num_kv_tokens = sel_idx.shape[0] * Sk
|
||||
|
||||
# [Sq, D] -> [Sq, 1, D] (single head)
|
||||
q_block = sparse_q_bh[qb].unsqueeze(1)
|
||||
q_list.append(q_block)
|
||||
|
||||
# [num_sel, Sk, D] -> [num_kv_tokens, 1, D]
|
||||
sel_k = k_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
|
||||
sel_v = v_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
|
||||
k_list.append(sel_k)
|
||||
v_list.append(sel_v)
|
||||
|
||||
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
|
||||
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
|
||||
|
||||
if not q_list:
|
||||
return
|
||||
|
||||
flat_q = torch.cat(q_list, dim=0)
|
||||
flat_k = torch.cat(k_list, dim=0)
|
||||
flat_v = torch.cat(v_list, dim=0)
|
||||
|
||||
# Compute max_seqlen_k from the Python list before moving to GPU to
|
||||
# avoid a `.item()` round-trip that would force a host/device sync.
|
||||
max_seqlen_q = Sq
|
||||
max_seqlen_k = max(b - a for a, b in zip(cu_seqlens_k[:-1], cu_seqlens_k[1:], strict=False))
|
||||
|
||||
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
|
||||
|
||||
orig_dtype = flat_q.dtype
|
||||
compute_dtype = orig_dtype
|
||||
if compute_dtype not in (torch.float16, torch.bfloat16):
|
||||
compute_dtype = torch.bfloat16
|
||||
flat_q = flat_q.to(compute_dtype)
|
||||
flat_k = flat_k.to(compute_dtype)
|
||||
flat_v = flat_v.to(compute_dtype)
|
||||
|
||||
flat_out = flash_attn_varlen_func_impl(
|
||||
flat_q,
|
||||
flat_k,
|
||||
flat_v,
|
||||
cu_seqlens_q_t,
|
||||
cu_seqlens_k_t,
|
||||
max_seqlen_q,
|
||||
max_seqlen_k,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
if compute_dtype != orig_dtype:
|
||||
flat_out = flat_out.to(orig_dtype)
|
||||
|
||||
idx = 0
|
||||
for qb in active_blocks:
|
||||
block_out = flat_out[idx:idx + Sq] # [Sq, 1, D]
|
||||
output[b, h, qb] = block_out.squeeze(1) # [Sq, D]
|
||||
idx += Sq
|
||||
|
||||
|
||||
def _reconstruct_pruned(
|
||||
sparse_output: torch.Tensor,
|
||||
keep_indices: torch.Tensor,
|
||||
block_size: int,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Scatter sparse output back to full block size.
|
||||
Pruned positions get nearest kept token's output.
|
||||
|
||||
Handles per-batch, per-head indices correctly.
|
||||
|
||||
Args:
|
||||
sparse_output: [B, H, N, keep_size, D]
|
||||
keep_indices: [B, H, N, keep_size]
|
||||
block_size: original tokens per block
|
||||
|
||||
Returns:
|
||||
full_output: [B, H, N, block_size, D]
|
||||
"""
|
||||
B, H, N, keep_size, D = sparse_output.shape
|
||||
device = sparse_output.device
|
||||
|
||||
if keep_size >= block_size:
|
||||
return sparse_output
|
||||
|
||||
full_output = torch.zeros(B, H, N, block_size, D, device=device, dtype=sparse_output.dtype)
|
||||
|
||||
# Scatter kept tokens
|
||||
idx_expand = keep_indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
|
||||
full_output.scatter_(3, idx_expand, sparse_output)
|
||||
|
||||
# Fill pruned positions with nearest kept token (vectorized)
|
||||
all_pos = torch.arange(block_size, device=device)
|
||||
|
||||
for b in range(B):
|
||||
for h in range(H):
|
||||
for n in range(N):
|
||||
kept = keep_indices[b, h, n] # [keep_size]
|
||||
|
||||
# Distance from every position to every kept position
|
||||
dists = (all_pos.view(-1, 1) - kept.view(1, -1)).abs()
|
||||
nearest_local_idx = dists.argmin(dim=1) # [block_size]
|
||||
|
||||
# Identify pruned positions
|
||||
is_pruned = torch.ones(block_size, dtype=torch.bool, device=device)
|
||||
is_pruned[kept] = False
|
||||
pruned_indices = is_pruned.nonzero(as_tuple=True)[0]
|
||||
|
||||
if pruned_indices.numel() > 0:
|
||||
src_indices = nearest_local_idx[pruned_indices]
|
||||
full_output[b, h, n, pruned_indices] = sparse_output[b, h, n, src_indices]
|
||||
|
||||
return full_output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# FastVideo backend classes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BSAAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = False
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [64, 128]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "BSA_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["BSAAttentionImpl"]:
|
||||
return BSAAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["BSAAttentionMetadata"]:
|
||||
return BSAAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["BSAAttentionMetadataBuilder"]:
|
||||
return BSAAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class BSAAttentionMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
dit_seq_shape: tuple[int, int, int]
|
||||
total_seq_length: int
|
||||
num_blocks: int
|
||||
block_size: int
|
||||
tile_partition_indices: torch.LongTensor
|
||||
reverse_tile_partition_indices: torch.LongTensor
|
||||
# BSA-specific config
|
||||
query_keep_ratio: float
|
||||
kv_cumulative_threshold: float
|
||||
min_kv_blocks: int
|
||||
|
||||
|
||||
class BSAAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
bsa_query_keep_ratio: float = 0.5,
|
||||
bsa_kv_cumulative_threshold: float = 0.9,
|
||||
bsa_min_kv_blocks: int = 4,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> "BSAAttentionMetadata":
|
||||
# Ensure patching does not drop tokens silently.
|
||||
assert all(r % p == 0 for r, p in zip(raw_latent_shape, patch_size, strict=False)), (
|
||||
"raw_latent_shape must be divisible by patch_size for BSA", )
|
||||
|
||||
dit_seq_shape = (
|
||||
raw_latent_shape[0] // patch_size[0],
|
||||
raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2],
|
||||
)
|
||||
|
||||
total_seq_length = math.prod(dit_seq_shape)
|
||||
block_size = math.prod(BSA_TILE_SIZE)
|
||||
# Require exact tiling to avoid reshape failures later.
|
||||
assert all(d % t == 0 for d, t in zip(dit_seq_shape, BSA_TILE_SIZE, strict=False)), (
|
||||
"dit_seq_shape must be divisible by BSA_TILE_SIZE", )
|
||||
num_blocks = total_seq_length // block_size
|
||||
|
||||
tile_partition_indices = get_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
|
||||
reverse_tile_partition_indices = get_reverse_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
|
||||
|
||||
return BSAAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
dit_seq_shape=dit_seq_shape,
|
||||
total_seq_length=total_seq_length,
|
||||
num_blocks=num_blocks,
|
||||
block_size=block_size,
|
||||
tile_partition_indices=tile_partition_indices,
|
||||
reverse_tile_partition_indices=reverse_tile_partition_indices,
|
||||
query_keep_ratio=bsa_query_keep_ratio,
|
||||
kv_cumulative_threshold=bsa_kv_cumulative_threshold,
|
||||
min_kv_blocks=bsa_min_kv_blocks,
|
||||
)
|
||||
|
||||
|
||||
class BSAAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.prefix = prefix
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
if num_kv_heads is not None and num_kv_heads != num_heads:
|
||||
raise ValueError("BSA backend does not support grouped-query attention")
|
||||
if causal:
|
||||
raise ValueError("BSA backend is bidirectional; causal=True is unsupported")
|
||||
if softmax_scale is not None:
|
||||
expected_scale = 1.0 / math.sqrt(self.head_size)
|
||||
if not math.isclose(softmax_scale, expected_scale, rel_tol=1e-4, abs_tol=1e-5):
|
||||
raise ValueError("softmax_scale must be default (1/sqrt(d)) for BSA")
|
||||
try:
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
except (AssertionError, RuntimeError):
|
||||
self.sp_size = 1
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Reorder tokens from raster order to tile-contiguous order."""
|
||||
# qkv: [B, L, num_heads, D]
|
||||
return qkv[:, attn_metadata.tile_partition_indices]
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""Reorder tokens from tile-contiguous order back to raster order."""
|
||||
return output[:, attn_metadata.reverse_tile_partition_indices]
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: BSAAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
BSA attention forward pass.
|
||||
|
||||
Input tensors are already in tile-contiguous order from preprocess_qkv.
|
||||
|
||||
Args:
|
||||
query: [B, L, num_heads, D] (tile-ordered)
|
||||
key: [B, L, num_heads, D] (tile-ordered)
|
||||
value: [B, L, num_heads, D] (tile-ordered)
|
||||
attn_metadata: BSA metadata
|
||||
|
||||
Returns:
|
||||
output: [B, L, num_heads, D] (tile-ordered)
|
||||
"""
|
||||
B, L, H, D = query.shape
|
||||
block_size = attn_metadata.block_size
|
||||
num_blocks = attn_metadata.num_blocks
|
||||
assert num_blocks * block_size == L, "Sequence length must match tiling"
|
||||
|
||||
# Reshape to [B, H, L, D] for attention computation
|
||||
q = query.transpose(1, 2).contiguous() # [B, H, L, D]
|
||||
k = key.transpose(1, 2).contiguous()
|
||||
v = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Reshape into blocks: [B, H, num_blocks, block_size, D]
|
||||
q_blocks = q.view(B, H, num_blocks, block_size, D)
|
||||
k_blocks = k.view(B, H, num_blocks, block_size, D)
|
||||
v_blocks = v.view(B, H, num_blocks, block_size, D)
|
||||
|
||||
# --- Query sparsification ---
|
||||
sparse_q, keep_indices, keep_size = _prune_queries(q_blocks, attn_metadata.query_keep_ratio)
|
||||
|
||||
# --- KV block selection ---
|
||||
kv_mask = _select_kv_blocks(
|
||||
sparse_q,
|
||||
k_blocks,
|
||||
attn_metadata.kv_cumulative_threshold,
|
||||
attn_metadata.min_kv_blocks,
|
||||
)
|
||||
|
||||
# --- Sparse attention ---
|
||||
sparse_output = _compute_sparse_attention(sparse_q, k_blocks, v_blocks, kv_mask)
|
||||
|
||||
# --- Reconstruct pruned positions ---
|
||||
full_output = _reconstruct_pruned(sparse_output, keep_indices, block_size)
|
||||
|
||||
# Reshape back: [B, H, num_blocks, block_size, D] -> [B, H, L, D] -> [B, L, H, D]
|
||||
hidden_states = full_output.view(B, H, L, D).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,341 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from v2._vendor.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.v2._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 `v2._vendor.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 v2._vendor.attention.backends.abstract import (
|
||||
AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder,
|
||||
)
|
||||
from v2._vendor.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
|
||||
# The FP4 path uses a dedicated custom_op wrapper (flash_attn_fp4_func) so that
|
||||
# torch.compile treats the CuTeDSL kernel as an opaque boundary.
|
||||
try:
|
||||
from v2._vendor.attention.utils.flash_attn_cute import flash_attn_fp4_func
|
||||
_FA4_FP4_AVAILABLE = True
|
||||
except ImportError:
|
||||
flash_attn_fp4_func = None
|
||||
_FA4_FP4_AVAILABLE = False
|
||||
|
||||
|
||||
def _nvfp4_quantize_for_fa4(tensor_4d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Quantize a (batch, seqlen, nheads, headdim) BF16 tensor to FP4.
|
||||
|
||||
Returns:
|
||||
fp4_tensor: torch.float4_e2m1fn_x2, shape (batch, seqlen_padded, nheads, headdim//2)
|
||||
where seqlen_padded is seqlen rounded up to multiple of 128.
|
||||
Caller should slice [:, :orig_seqlen] before passing to FA4.
|
||||
sf_tensor: torch.uint8, shape (32, 4, rest_m, 4, rest_k, nheads, batch) with stride[3]=1
|
||||
"""
|
||||
from flashinfer.quantization import nvfp4_quantize, SfLayout
|
||||
|
||||
batch, seqlen, nheads, headdim = tensor_4d.shape
|
||||
sf_vec_size = 16
|
||||
|
||||
# Pad seqlen to multiple of 128 (required by nvfp4_quantize layout_128x4)
|
||||
tile_m = 128
|
||||
seqlen_padded = (seqlen + tile_m - 1) // tile_m * tile_m
|
||||
if seqlen_padded != seqlen:
|
||||
tensor_4d = F.pad(tensor_4d, (0, 0, 0, 0, 0, seqlen_padded - seqlen))
|
||||
|
||||
# Quantize with nheads squashed into K dimension so M=batch*seqlen (divisible by 128)
|
||||
# and K=nheads*headdim. This ensures 128-row SF tiles align with seqlen boundaries.
|
||||
t2d = tensor_4d.reshape(batch * seqlen_padded, nheads * headdim)
|
||||
one = torch.ones(1, device=t2d.device, dtype=torch.float32)
|
||||
fp4_data, sf_data = nvfp4_quantize(t2d, one, sfLayout=SfLayout.layout_128x4, do_shuffle=False)
|
||||
|
||||
# FP4 data: (batch*seqlen, nheads*headdim/2) → (batch, seqlen, nheads, headdim/2)
|
||||
fp4_tensor = (fp4_data.reshape(batch, seqlen_padded, nheads,
|
||||
headdim // 2).view(torch.int8).view(torch.float4_e2m1fn_x2))
|
||||
|
||||
# SF layout conversion: nvfp4_quantize layout_128x4 → FA4 MMA layout
|
||||
# layout_128x4 buffer: [mTile, kTile, 32, 4, 4]
|
||||
# FA4 expects: (32, 4, rest_m, 4, rest_k, nheads, batch) with stride[3]=1
|
||||
atom_m0, atom_m1, atom_k = 32, 4, 4
|
||||
rest_m = seqlen_padded // tile_m
|
||||
sf_k_per_head = headdim // sf_vec_size # 8 for headdim=128
|
||||
rest_k = sf_k_per_head // atom_k # 2
|
||||
|
||||
total_m_tiles = batch * rest_m
|
||||
total_k_tiles = (nheads * sf_k_per_head) // atom_k
|
||||
|
||||
sf_swizzled = sf_data.reshape(total_m_tiles, total_k_tiles, atom_m0, atom_m1, atom_k)
|
||||
sf_decomposed = sf_swizzled.reshape(batch, rest_m, nheads, rest_k, atom_m0, atom_m1, atom_k)
|
||||
sf_canonical = sf_decomposed.permute(0, 2, 1, 3, 4, 5, 6).contiguous()
|
||||
sf_mma = sf_canonical.permute(4, 5, 2, 6, 3, 1, 0)
|
||||
|
||||
return fp4_tensor, sf_mma
|
||||
|
||||
|
||||
class FlashAttentionBackend(AttentionBackend):
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "FLASH_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["FlashAttentionImpl"]:
|
||||
return FlashAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor, key_len: int) -> torch.Tensor:
|
||||
# Normalize attn_mask to [B, key_len] where True means valid token.
|
||||
if attn_mask.dim() == 4:
|
||||
attn_mask = attn_mask[:, 0, 0, :]
|
||||
elif attn_mask.dim() == 3:
|
||||
attn_mask = attn_mask[:, 0, :]
|
||||
elif attn_mask.dim() != 2:
|
||||
raise ValueError(f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}")
|
||||
|
||||
# SDPA additive mask convention: valid=0, masked=-inf/large negative.
|
||||
key_padding_mask = attn_mask if attn_mask.dtype == torch.bool else attn_mask >= 0
|
||||
|
||||
if key_padding_mask.shape[-1] != key_len:
|
||||
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
|
||||
f"expected {key_len}, got {key_padding_mask.shape[-1]}")
|
||||
return key_padding_mask
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlashAttnMetadata(AttentionMetadata):
|
||||
current_timestep: int
|
||||
attn_mask: torch.Tensor | None = None
|
||||
|
||||
|
||||
class FlashAttnMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
attn_mask: torch.Tensor,
|
||||
) -> FlashAttnMetadata:
|
||||
return FlashAttnMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
|
||||
|
||||
|
||||
class FlashAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
self.nvfp4_fa4 = extra_impl_args.get("nvfp4_fa4", False) or os.environ.get("FASTVIDEO_NVFP4_FA4", "0") == "1"
|
||||
if self.nvfp4_fa4:
|
||||
cap = torch.cuda.get_device_capability()
|
||||
assert cap in [(10, 0), (10, 3)], (f"NVFP4 FA4 requires Blackwell (sm100a/sm103a), got sm{cap[0]}{cap[1]}")
|
||||
assert _FA4_FP4_AVAILABLE, ("NVFP4 FA4 requires flash-attention-fp4 (flash_attn.cute). "
|
||||
"Install via instructions in docs/inference/optimizations.md")
|
||||
logger.info("NVFP4 FA4 enabled for FlashAttentionImpl (quant_qk only)")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: FlashAttnMetadata,
|
||||
):
|
||||
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
|
||||
from v2._vendor.attention.utils.flash_attn_no_pad import (
|
||||
flash_attn_no_pad,
|
||||
flash_attn_varlen_qk_no_pad,
|
||||
)
|
||||
|
||||
attn_mask = attn_metadata.attn_mask
|
||||
|
||||
# flash_attn_no_pad packs q/k/v as one tensor and assumes equal q/k
|
||||
# sequence lengths. Cross-attention can violate this.
|
||||
if query.shape[1] != key.shape[1]:
|
||||
query_padding_mask = torch.ones(
|
||||
(query.shape[0], query.shape[1]),
|
||||
dtype=torch.bool,
|
||||
device=query.device,
|
||||
)
|
||||
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, key.shape[1]).to(device=key.device)
|
||||
|
||||
return flash_attn_varlen_qk_no_pad(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_padding_mask=query_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
causal=self.causal,
|
||||
dropout_p=0.0,
|
||||
softmax_scale=self.softmax_scale,
|
||||
)
|
||||
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
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)
|
||||
|
||||
else:
|
||||
# Route through the compilable wrapper so dynamo sees a
|
||||
# registered op (no graph break) for FA2/FA3; identical
|
||||
# kernel + numerics, op runs eager internally.
|
||||
output = flash_attn_func_compilable(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal,
|
||||
)
|
||||
return output
|
||||
|
||||
def _forward_nvfp4(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor) -> torch.Tensor:
|
||||
"""FP4 flash attention with quantized Q and K, BF16 V."""
|
||||
orig_seqlen_q = query.shape[1]
|
||||
orig_seqlen_k = key.shape[1]
|
||||
|
||||
# Quantize Q/K to FP4 (internally pads to multiple of 128 for SF layout)
|
||||
q_fp4, q_sf = _nvfp4_quantize_for_fa4(query)
|
||||
k_fp4, k_sf = _nvfp4_quantize_for_fa4(key)
|
||||
|
||||
# Pass original seqlen to FA4 — the kernel handles non-multiple-of-128
|
||||
# via boundary masking. FP4/SF data is padded to 128-multiple but FA4
|
||||
# only attends to orig_seqlen positions, avoiding softmax bias on padding.
|
||||
q_fp4 = q_fp4[:, :orig_seqlen_q]
|
||||
k_fp4 = k_fp4[:, :orig_seqlen_k]
|
||||
|
||||
output = flash_attn_fp4_func(
|
||||
q_fp4,
|
||||
k_fp4,
|
||||
value,
|
||||
q_sf,
|
||||
k_sf,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal,
|
||||
)
|
||||
if isinstance(output, tuple):
|
||||
output = output[0]
|
||||
return output
|
||||
@@ -0,0 +1,64 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
from sageattention import sageattn
|
||||
|
||||
from v2._vendor.attention.backends.abstract import ( # FlashAttentionMetadata,
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata)
|
||||
from v2._vendor.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SageAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> list[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SAGE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type["SageAttentionImpl"]:
|
||||
return SageAttentionImpl
|
||||
|
||||
# @staticmethod
|
||||
# def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
class SageAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: int | None = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
output = sageattn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
# since input is (batch_size, seq_len, head_num, head_dim)
|
||||
tensor_layout="NHD",
|
||||
is_causal=self.causal)
|
||||
return output
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user