Compare commits
75
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4a33540c08 | ||
|
|
ba83e9bef7 | ||
|
|
21e59268d8 | ||
|
|
139577639f | ||
|
|
db866dfb04 | ||
|
|
53d63afa76 | ||
|
|
685702e330 | ||
|
|
9e8d515842 | ||
|
|
1302a27470 | ||
|
|
e1a2c9ede8 | ||
|
|
8ee7270368 | ||
|
|
26c46b901a | ||
|
|
fcc5a26b46 | ||
|
|
6450f9fa26 | ||
|
|
4d079c65ed | ||
|
|
55fcf27759 | ||
|
|
3149d0ec89 | ||
|
|
1f4b729c4a | ||
|
|
ac29750b55 | ||
|
|
50fd379f3a | ||
|
|
fa9d8040e1 | ||
|
|
9794be7202 | ||
|
|
e4b14e124d | ||
|
|
60ca94796c | ||
|
|
060bdd85b6 | ||
|
|
0a18b6a031 | ||
|
|
2e5bad1d6c | ||
|
|
705f580542 | ||
|
|
cbe67337c3 | ||
|
|
113478fd74 | ||
|
|
9b2007c058 | ||
|
|
1542b7a03d | ||
|
|
e4b7261b66 | ||
|
|
b39b68c6f5 | ||
|
|
5f698f03f5 | ||
|
|
799a9b69a2 | ||
|
|
27b320afd8 | ||
|
|
c74e9cb4fe | ||
|
|
0a2c985b56 | ||
|
|
8b343db604 | ||
|
|
24b76d372c | ||
|
|
f7256ab197 | ||
|
|
9de0e2bcd1 | ||
|
|
5f33f80d6c | ||
|
|
e7d4aaeef4 | ||
|
|
2a085a5c39 | ||
|
|
261d3007a9 | ||
|
|
b2da3b2b4d | ||
|
|
7ec4d01afc | ||
|
|
8b0e2bbb4b | ||
|
|
7fd58f26ea | ||
|
|
15b6278e9f | ||
|
|
cd0fbbd3a0 | ||
|
|
02151584cd | ||
|
|
51bb5e349a | ||
|
|
1bdbf1d5b3 | ||
|
|
de295403c7 | ||
|
|
34834eb4f4 | ||
|
|
bd36d69680 | ||
|
|
94b84b03b9 | ||
|
|
a16ce1a07f | ||
|
|
37544c68a5 | ||
|
|
88e4a4355c | ||
|
|
f1dead1361 | ||
|
|
cfc2ad7487 | ||
|
|
97fedaf432 | ||
|
|
3ec511419f | ||
|
|
145082cf74 | ||
|
|
71caa58ba9 | ||
|
|
8240970ea6 | ||
|
|
44d10b5c79 | ||
|
|
9d874aa609 | ||
|
|
8e828680c6 | ||
|
|
735ee6aa56 | ||
|
|
068a532b31 |
@@ -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).
|
||||
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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()
|
||||
@@ -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,ui/package-lock.json,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.
|
||||
@@ -0,0 +1,3 @@
|
||||
__pycache__/
|
||||
*.pyc
|
||||
*.pyo
|
||||
+107
@@ -0,0 +1,107 @@
|
||||
# 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 all diffusion loops + RL recompute 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 (all diffusion loops + RL recompute 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 loops/policies/scheduler/training 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 rollout → 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 training/inference with wandb, 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.
|
||||
- **`mse_grad_step` raises on the cuda DiT** (Risk F). RL/distill *training* on GPU is a separate
|
||||
workstream; inference + SDE *rollout* bring-up don't need it.
|
||||
- **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.
|
||||
+607
@@ -0,0 +1,607 @@
|
||||
# FastVideo v2 — A Model-Native Runtime for the (Recipe, Runtime) Era
|
||||
|
||||
**Status:** source of truth. This README is the single design document for v2 — it supersedes and absorbs
|
||||
the four prior design docs (`design.md` v19 strategic proposal, `designv2.md` model-plane thesis,
|
||||
`design_v3.md` unconstrained north star, `designv4.md` as-built + joint-RL validation), which have been
|
||||
removed. It carries their load-bearing ideas, reflects **what is actually built and running in `v2/`**, and
|
||||
states the **forward roadmap** (including the M\* paper comparison — see §18 and
|
||||
[`.agents/exploration/mstar-v2-roadmap.md`](../.agents/exploration/mstar-v2-roadmap.md)).
|
||||
|
||||
**What v2 is.** A model-native serving **and** training substrate for composite multimodal models — video,
|
||||
image, audio, omni/MoT, world models, VLAs — where the atomic unit is a typed **(recipe, runtime) pair**
|
||||
owned by a `ModelCard`, every iterative computation is a **driven loop**, and one scheduler runs the *steps*
|
||||
of all loops in one currency. The control flow, contracts, scheduling, caching, parity gates, and training
|
||||
math are real and CPU-testable (216 tests, 34 files, two runners); the heavy neural forwards run either on a
|
||||
numpy **toy backend** (laptop, no GPU) or the real **torch backend** (`platform/backends/torch_backend.py`),
|
||||
selected by the `platform/` dispatch substrate with **no change to loops/scheduler/caches/parity/training**.
|
||||
As of this writing **20+ real model families are GPU-verified on H100**, and the omni trio
|
||||
(BAGEL, Qwen2.5-Omni, Cosmos3) runs real via **vllm-omni** (§17).
|
||||
|
||||
---
|
||||
|
||||
## Table of contents
|
||||
1. [The thesis](#1-the-thesis) · 2. [Three signature ideas](#2-three-signature-ideas) ·
|
||||
3. [Planes & dependency order](#3-planes--dependency-order) · 4. [Model Plane](#4-model-plane--the-center) ·
|
||||
5. [The driven-loop contract](#5-the-driven-loop-contract) · 6. [Runtime & scheduler](#6-runtime--scheduler) ·
|
||||
7. [Memory, cache, transport, compile](#7-memory-cache-transport-compile) · 8. [Parallelism](#8-parallelism-as-a-model-contract) ·
|
||||
9. [Correctness — parity as a typed gate](#9-correctness--parity-as-a-typed-gate) · 10. [Training & RL](#10-training--rl-on-the-same-loops) ·
|
||||
11. [Weight-sharing topologies & stress tests](#11-weight-sharing-topologies--the-stress-test-catalog) ·
|
||||
12. [Request/session/artifact + programs/workflows](#12-request-session-artifact--programsworkflows) ·
|
||||
13. [Serving & fleet](#13-serving--fleet) · 14. [Extensions](#14-extensions) ·
|
||||
15. [Model taxonomy](#15-the-model-taxonomy-what-the-runtime-must-serve) ·
|
||||
16. [Package layout (actual)](#16-package-layout-actual) · 17. [Current status & GPU bring-up](#17-current-status--gpu-bring-up) ·
|
||||
18. [Roadmap](#18-roadmap) · 19. [Reference synthesis](#19-reference-synthesis) ·
|
||||
20. [Honest unknowns & falsifiers](#20-honest-unknowns--falsifiers)
|
||||
|
||||
---
|
||||
|
||||
## 1. The thesis
|
||||
|
||||
Three facts about video/omni generation dictate the architecture:
|
||||
|
||||
- **A deployable model is a post-training artifact.** Unlike an LLM (where inference optimizes frozen weights
|
||||
post-hoc), a *usable* video model is *created* by training: step distillation is mandatory for latency, low
|
||||
precision needs QAT, causal/world models are made by distillation + self-forcing. Every inference capability
|
||||
is therefore a **(recipe, runtime) pair** — the weights and the loop that produced-and-assumes them are one
|
||||
versioned object, a `ModelCard`.
|
||||
- **Video/omni systems are loop systems.** The work is iteration — denoise timesteps, AR decode, chunked
|
||||
rollout, VAE tiles, encoder chunks, audio tokens, reward batches, optimizer steps, media chunks — not a
|
||||
single `forward()`. A runtime that reduces everything to `forward()` cannot schedule, batch, cancel, stream,
|
||||
reserve memory for, or capture the behavior of what actually runs.
|
||||
- **Omni models share weights across loop types within one request.** Cosmos3's text reasoner and multimodal
|
||||
denoiser are the *same resident weights*, driven by an AR loop then a diffusion loop in one request. This
|
||||
cannot be a DAG of separate engines (that doubles 30B+ of weights and severs the shared KV/denoise state); it
|
||||
must be one resident instance running many loops. The hard part — the differentiation — is making those loops
|
||||
**runtime-visible, step-scheduled, batchable, and cost-priced** (vllm-omni's `bagel_single_stage`/`lance`
|
||||
prove the sharing is *expressible*, but bury it in one opaque `DIFFUSION` stage the scheduler never sees
|
||||
inside).
|
||||
|
||||
**The one invariant, stated once:**
|
||||
|
||||
```
|
||||
Model cards own components, loops, recipes, and parity.
|
||||
Programs compose loops into tasks. Workflows compose models into pipelines.
|
||||
The scheduler executes the steps of loops as WorkUnits under one budget.
|
||||
Caches are correct by key, not by hope.
|
||||
Training records behavior on the same loops it serves.
|
||||
Deployment places and routes; products stream artifacts; neither defines the model.
|
||||
```
|
||||
|
||||
Everything below is the elaboration of that invariant.
|
||||
|
||||
---
|
||||
|
||||
## 2. Three signature ideas
|
||||
|
||||
### 2.1 The (recipe, runtime) pair is a first-class, versioned, typed object
|
||||
A model is not a checkpoint. It is a `ModelCard` owning, as one versioned unit: the **components** (weights,
|
||||
loaders, layouts), the **loops** it can run, the **recipe** that produced the weights (distillation/QAT/RL
|
||||
method, parents, `assumes_loop`, `assumes_precision`), and the **parity contract** binding the train-forward
|
||||
to the serve-forward at a declared consistency level. You cannot ship the weights without the loop they assume
|
||||
(`assumes_loop`), or change the loop without re-proving parity. This turns "we do training and inference in one
|
||||
repo" from an org chart into a typed guarantee.
|
||||
|
||||
### 2.2 Driven loops — the model owns control flow, the runtime owns execution
|
||||
A loop is a serializable state machine `init → (next → advance)* → finalize`: the model describes *the next
|
||||
step it needs* (`next` is kernel-free — it returns a typed `WorkPlan` thunk), the runtime decides *when and
|
||||
with whom that step runs* (`await ctx.execute(plan)` — the inversion point: admission, batching, placement,
|
||||
streaming, behavior capture), and the model folds the result back (`advance`) and decides what to do next.
|
||||
Content-adaptive decisions (cache-dit skips, EOS, VSA tile selection) are ordinary control flow in the model;
|
||||
per-request state lives in a typed `LoopState`, **never in module globals**, so interleaving requests through
|
||||
one instance cannot smear state — the failure mode that makes naive loop-inversion dangerous is *structurally*
|
||||
excluded (and proven by the interleave gate, §9).
|
||||
|
||||
### 2.3 One vocabulary spans every weight-sharing topology
|
||||
The Card/Loop/Program split is not specialized to one sharing pattern. The **same primitives** express the
|
||||
whole spread — and every one is just `shared_weight_components` bindings on `LoopSpec`s plus hand-off nodes in
|
||||
a `Program`, **no new primitive**:
|
||||
|
||||
| Topology | Components | Loops | Sharing | Recipe |
|
||||
|---|---|---|---|---|
|
||||
| Single diffusion | `transformer` | `diffusion_denoise` | — | `wan21` |
|
||||
| **MoT omni** | **one** `transformer` | `ar_decode` + `diffusion_denoise` | both loops → **same** component | `cosmos3`, `bagel` |
|
||||
| **Joint LM+gen RL** | `llm` + `transformer` | `ar_decode`→`llm`, `diffusion_denoise`→`transformer` | **disjoint** experts, one request, jointly RL'd | `unified` |
|
||||
| **Cascade omni-speech** | `thinker`+`talker`+`vocoder` | `ar_decode`→`ar_decode`→`audio_decode` | **disjoint**, **chained** | `qwen_omni` |
|
||||
| **N-way joint RL** | N `refiner_i` + `transformer` | N×`ar_decode` + `diffusion_denoise` | **disjoint**, N-way jointly RL'd | `multi_expert` |
|
||||
|
||||
The signature claim, validated by the §11 stress tests: *the weight-sharing graph is data on the card, not
|
||||
structure in the runtime.* (BAGEL's real MoT is *partial* sharing — co-resident experts sharing attention,
|
||||
separate FFNs — expressed via the expert-routing policy over one resident instance; a fifth point on the same
|
||||
axis, still no new primitive.)
|
||||
|
||||
---
|
||||
|
||||
## 3. Planes & dependency order
|
||||
|
||||
```
|
||||
Products: Python · CLI · OpenAI server · ComfyUI · Dreamverse · RTC · Trainer (thin: validate, request, subscribe)
|
||||
Request / Session / Artifact / Stream v2/request/, v2/runtime/session.py
|
||||
Program Plane (typed loop programs; cross-model Workflows) v2/program/
|
||||
┌──────────────── Model Plane (CENTER) ────────────────┐ v2/card/, v2/recipes/*/card.py
|
||||
│ ModelCard: components · loops · recipe · parity · │
|
||||
│ capabilities · caches · parallelism · precision │
|
||||
└───────────────────────┬───────────────────────────────┘
|
||||
┌────────────────────────┼────────────────────────┐ same loops, different capture
|
||||
│ Runtime / Scheduler │ Training / RL │ v2/runtime/ · v2/training/
|
||||
│ WorkUnits · GPU-time budget │ rollout·reward·sync │
|
||||
└────────────────────────┼────────────────────────┘
|
||||
Memory · Cache · Transport · Compile v2/memory, v2/cache, v2/transport, v2/runtime/cudagraph.py
|
||||
Parallelism (named axes → DeviceMesh, validated, part of the cache key) v2/parallel/
|
||||
Platform dispatch (COMPONENTS/KERNELS registries → toy | torch backend) v2/platform/
|
||||
Deployment / Fleet (DeploymentCard → LocalFleet | Dynamo; never the core) v2/deploy/, v2/serving/
|
||||
```
|
||||
|
||||
**Enforced boundaries:** `card/` imports no product/runtime; `runtime/` executes `card/` loops but defines no
|
||||
semantics; `training/` requires behavior records but forks no loop and **the engine never imports `training`**
|
||||
(verified — the only `training` mention in `runtime/` is the comment documenting this rule); cross-model
|
||||
`Workflow` orchestration sits *above* the engine (no change to the single-instance hot path). `parity/` is a
|
||||
first-class package, not a test folder.
|
||||
|
||||
---
|
||||
|
||||
## 4. Model Plane — the center
|
||||
|
||||
```python
|
||||
class ModelCard:
|
||||
model_id: str # "fastwan-1.3b-nvfp4-4step"
|
||||
family: str
|
||||
components: dict[str, ComponentSpec]
|
||||
loops: dict[str, LoopSpec]
|
||||
capabilities: CapabilityMatrix # text_to_video, image_to_video, reasoning_text, vae_decode, ...
|
||||
recipe: RecipeSpec # what produced these weights (§2.1)
|
||||
parity: ParitySpec # train-forward ≡ serve-forward, to a declared level (§9)
|
||||
caches: dict[str, CacheContract]
|
||||
parallelism: ParallelismContract
|
||||
precision: PrecisionContract
|
||||
```
|
||||
|
||||
The card is both a **declarative contract** (validatable before any GPU touches it) and a **runtime factory**
|
||||
(it instantiates components, binds loops, resolves caches). `card.validate()` checks every loop's components +
|
||||
cache policies exist and that `recipe.assumes_loop` is a declared loop.
|
||||
|
||||
- **`RecipeSpec`** — `method` (`dmd2`/`self_forcing`/`diffusion_nft`/`unified_rl`/`base`/…), `parents`
|
||||
(teacher/base ids), `assumes_loop`, `assumes_precision`, `consistency_required`. `assumes_loop` is the teeth:
|
||||
a 4-step distilled model cannot be served under a 50-step sampler without a typed mismatch error.
|
||||
- **`ComponentSpec`** — `component_id`, `kind` (`dit`/`vae`/`text_encoder`/`reasoner_tower`/…), `load_id` (the
|
||||
real module to load on the torch backend), `factory` (the toy stand-in), `checkpoint` (weights path/HF id),
|
||||
`required_for`/`optional_for`/`resident_for` task sets, precision/placement policies. (Note: `required_for`
|
||||
is currently declared on every card but not yet consumed by the executor — see §18 roadmap P0.)
|
||||
- **`LoopSpec`** — `loop_id`, `kind` (`LoopKind.DIFFUSION_DENOISE`/`AR_DECODE`/`CHUNK_ROLLOUT`/`AUDIO_DECODE`/…),
|
||||
`work_unit_kind`, `step_cost_model` (predicted GPU-time per step — §6), `shared_weight_components`,
|
||||
`cache_policy`, `loop_factory`.
|
||||
|
||||
A `ModelInstance` is a resident, loaded card: component instances, model state, caches, compiled graphs, a
|
||||
parallel plan. **A request may run several of the card's loops against one `ModelInstance`** — that single
|
||||
sentence is what makes omni native, and two loops binding the same `shared_weight_components` get the *same
|
||||
live object* via `instance.component()` (no weight duplication, no DAG split).
|
||||
|
||||
---
|
||||
|
||||
## 5. The driven-loop contract
|
||||
|
||||
```python
|
||||
class Loop(Protocol): # v2/loop/contracts.py — protocols, not ABCs
|
||||
def init(self, req, model, ctx) -> LoopState # per-request state (seeded rng, latents…)
|
||||
def next(self, st) -> WorkPlan | Done # describe the next step (KERNEL-FREE: a run() thunk)
|
||||
def advance(self, st, result) -> LoopState # fold result; capture behavior under ROLLOUT
|
||||
def finalize(self, st) -> LoopResult # outputs + metrics + behavior
|
||||
```
|
||||
|
||||
The runtime's `LoopRunner` is the only place iteration lives: `next` → `await ctx.execute(plan)` → `advance`,
|
||||
emitting `plan.emits` as stream chunks. Why this contract is right, against the failure modes:
|
||||
|
||||
- **Content-adaptive steps are natural** — `next()` reads `state` (and `advance` already folded in the last
|
||||
`StepResult`), so cache-dit's skip, AR's EOS, and VSA's tile selection are ordinary control flow; `next()`
|
||||
is still kernel-free (it *describes* work), which is all the scheduler needs.
|
||||
- **Cross-request state safety is structural** — all per-request state is in `LoopState`; there are no
|
||||
module-level residual/KV globals, so interleaving cannot smear state (the §9.3 interleave gate proves it).
|
||||
- **Serializable ⇒ resumable/migratable** — a `LoopState` is a resume point (preempt, migrate, crash-recover).
|
||||
|
||||
**Policies decompose the step body** (CFG/flow-shift/precision/expert-routing) — composed *into* a loop, not
|
||||
branched *inside* it. **CFG is a policy over one shared denoise body** (proven by vllm-omni's `CFGParallelMixin`
|
||||
unifying 2-forward / batched / cfg-parallel under one predict/combine pair), expressed as **three layers**:
|
||||
(1) `CFGPolicy` is *in-loop* — branch vocabulary, combine formula, per-request mutable state (the adaptive-gate
|
||||
cached delta is the canonical state case; batched-vs-2-forward is a dispatch detail inside one policy); (2)
|
||||
`cfgp` is a *parallelism axis* that shards branches across ranks and runs the same rank-invariant combine;
|
||||
(3) companions are an *orchestrator pattern* upstream of diffusion. Two caveats: `combine` runs in the step
|
||||
body's numeric space (Cosmos combines in x0-space post-EDM, not noise-space), and embedded-guidance (Flux) is a
|
||||
degenerate single-branch identity-combine policy, not "no CFG". A family whose math is genuinely braided ships
|
||||
a **custom `next`/`advance`** using samplers/CFG as a *library* — the runtime requires only the four methods.
|
||||
|
||||
---
|
||||
|
||||
## 6. Runtime & scheduler
|
||||
|
||||
**One WorkUnit, one currency.** Every `ctx.execute(plan)` is a `WorkUnit` — the smallest schedulable action
|
||||
with a resource reservation and a loop boundary. Kinds: `AR_TOKEN`, `AR_PREFILL`, `DIFFUSION_STEP`,
|
||||
`DIFFUSION_WINDOW`, `CHUNK_STEP`, `ENCODER_CHUNK`, `VAE_TILE`, `AUDIO_CHUNK`, `REWARD_BATCH`, `LOGPROB_BATCH`,
|
||||
`TRANSFER`, `CACHE_IO`, `GRAPH_CAPTURE` (a 13-kind taxonomy). Tokens are *one kind*, not the scheduler — the
|
||||
generalization of vLLM's token scheduler that diffusion forces.
|
||||
|
||||
**The budget currency is predicted GPU-time, not counts.** A bidirectional denoise step re-attends the full
|
||||
latent at O(L²) with zero KV amortization; an AR decode step is ~O(context) against a cache. Counting "steps"
|
||||
or "tokens" puts items three orders of magnitude apart in one bucket. Each WorkUnit converts to GPU-seconds via
|
||||
`LoopSpec.step_cost_model`, online-calibrated by the Profiler. **The same cost model is the object published to
|
||||
the fleet** (§13) — internal budget and Dynamo routing input are one thing.
|
||||
|
||||
**Admission rule (the soundness condition of multiplexing):** do not admit a waiting WorkUnit unless *every*
|
||||
resource it requests can be reserved — compute budget AND memory (resident + worst-case peak) AND cache blocks
|
||||
AND transfer bandwidth AND graph-capture shape AND output sinks. Infeasible requests fail fast
|
||||
(`AdmissionInfeasible`); budget is refunded on completion. Honesty caveats kept from contact with reality:
|
||||
admission uses the *conservative baseline* (cache-dit skips, VSA tiles, AR length are unknowable in advance —
|
||||
budgeted at the cap, refunded on early EOS); a denoise step is *indivisible* (mitigated by cost-class pools +
|
||||
SP-within-a-node + SLO classes).
|
||||
|
||||
**Scheduler in layers** (each testable on a fake pool, no GPU): `RequestScheduler` → `LoopScheduler` →
|
||||
`BatchScheduler` (groups compatible WorkPlans by `(instance, loop_kind, shape_sig, precision, parallel_plan,
|
||||
graph_key)`) → `PlacementScheduler` → `TransferScheduler` → `AdmissionController`. **SPMD consistency**: rank-0
|
||||
decides and broadcasts (the same channel as the abort broadcast — scheduling and failure isolation share one
|
||||
mechanism). **Cancellation is common-path** (vibe-directing makes abandoning in-flight work normal): it takes
|
||||
effect at the next step boundary, drops queued WorkUnits, releases `LoopState` + cache handles.
|
||||
|
||||
---
|
||||
|
||||
## 7. Memory, cache, transport, compile
|
||||
|
||||
**Cache correctness is a contract.** `CacheKey` carries every output-semantic field — `model_id`,
|
||||
`component_id`, `loop_id`, per-component `weights_version`, `adapter_versions`, `precision`,
|
||||
`parallel_plan_hash`, `shape_sig`, `layout_sig`, `scheduler_sig`, `guidance_sig`, `seed`, `input_hashes`,
|
||||
`step_index`, `contract_version`. **If a field can change output semantics, it is in the key.** Incorrect reuse
|
||||
is worse than no reuse: the key is *partitioned* by `adapter_versions` (a te-LoRA-differing request doesn't
|
||||
serve stale embeddings), and a weight-sync bumps only the affected component's `weights_version` — so a
|
||||
transformer sync **does not flush the frozen text-encoder's feature cache** (a K-sample RL group encodes its
|
||||
shared prompt once).
|
||||
|
||||
**Per-class pools (the granularity reality).** No single unified block pool — cache classes differ by 150–500×
|
||||
in natural granularity (text-KV page ≈ 64 KB/layer; causal-video latent-chunk slab ≈ 9.6–32 MB/layer). Each
|
||||
class gets a statically budgeted pool: paged text-KV (`ar_decode`), slab chunk-KV (`chunk_rollout`, with a
|
||||
training mode that disables mid-rollout recycling), feature caches (content-hash keyed, ref-counted),
|
||||
residual caches (cache-dit, scoped per `LoopState`), weight/adapter cache. **KV is the minority case** — a pure
|
||||
bidirectional deployment allocates none of it.
|
||||
|
||||
**Memory / transport / compile.** Tagged pools with sleep/wake by tag (CuMem-style, component-granular for RL).
|
||||
Transport is manifest-based and pluggable: in-proc reference → SHM → CUDA IPC → NCCL/NIXL → object-store;
|
||||
KV-bearing edges speak a `KVConnector`-shaped protocol (`chunk_ready` readiness + credit-based flow control,
|
||||
sglang-omni's model). Compile: CUDA graphs + `torch.compile` keyed on `(model, component, loop, work_kind,
|
||||
shape_sig, precision, parallel_plan, backend)` — **never full-graph across the engine**; per-step piecewise
|
||||
capture is wired (`runtime/cudagraph.py`, declared on 19 cards) with a static-buffer discipline and
|
||||
version-eviction on weight sync.
|
||||
|
||||
---
|
||||
|
||||
## 8. Parallelism as a model contract
|
||||
|
||||
Parallelism is not a launch flag — it affects cache keys, scheduling, transport, capture, and parity, so it
|
||||
lives on the card. `ParallelPlan` axes (`v2/parallel/plan.py`):
|
||||
`("dp","tp","sp","cp","cfgp","pp_patch","vae","ep","fsdp","role","replica")`. Declarative, validated
|
||||
(`validation.py`: `cfgp ≤ 2`; `pp_patch` is **invalid for causal/AR** because stale KV breaks causality;
|
||||
ownership conflicts like a `BatchedCFG` policy *and* a `cfgp` group are build errors), compiled to a PyTorch
|
||||
`DeviceMesh` via a `ParallelDims`-style builder. **Pre-flight or it fails at load, never halfway.** Degree-one
|
||||
axes exist as trivial groups so component code needs no special cases. **Pools are single-node**; multi-node
|
||||
scale is *multiple pools* fronted by the fleet (§13). Note: Wan/LTX shipped weights parallelize via **sequence
|
||||
parallelism (`sp`)**, not TP (they use `ReplicatedLinear`); real `ColumnParallelLinear` lives in `flux2`.
|
||||
|
||||
---
|
||||
|
||||
## 9. Correctness — parity as a typed gate
|
||||
|
||||
Every card carries a `ParitySpec`. Parity is **measured, never assumed**, by a `ParityAligner` observer:
|
||||
record named taps per step/block from a reference, replay with fixed seeds, report the first divergence beyond
|
||||
per-tap tolerance.
|
||||
|
||||
**The consistency ladder:**
|
||||
```
|
||||
C0 component parity — VAE/encoder/transformer-block/scheduler-step in isolation
|
||||
C1 loop parity — full denoise trajectory / AR logits, fixed seed
|
||||
C2 behavioral identity — the train-forward and serve-forward agree on the quantity the RL objective uses:
|
||||
· likelihood-based (GRPO/UniRL): per-step log-prob identity ⇒ PPO ratio == 1
|
||||
· likelihood-free (DiffusionNFT): seeded final-sample + prediction-space identity — NO log-probs to match
|
||||
C3 distribution parity — rollout distribution under allowed nondeterminism (defined; not yet consumed — §18)
|
||||
C4 artifact quality — SSIM/reward/human-preference (gates product claims; needs the eval system)
|
||||
```
|
||||
The **C2 split** is load-bearing and the lesson of the landed RL stack: the shipped DiffusionNFT is
|
||||
likelihood-free (no log-probs; "log-prob identity" is undefined for it), while UniRL is likelihood-based — both
|
||||
are demonstrated, on opposite halves of the rung.
|
||||
|
||||
**The interleave gate (§9.3) — the bet loop-inversion lives or dies on.** Loop inversion's real hazard is
|
||||
cross-request state smearing under interleaving, and a batch-of-1 gate is *structurally blind* to it. So a
|
||||
**batch-of-N interleave parity test** is a *required* gate: N concurrent requests interleaved at step
|
||||
granularity must be **bit-identical** to the same requests run serially — and it holds across every model, the
|
||||
MoT omni cards, the two-loop unified program, the three-loop cascade, and heterogeneous WorkUnit kinds. (A
|
||||
buggy module-global interceptor *breaks* the gate; the per-request one passes.) **Three execution profiles, one
|
||||
loop definition:** serve (no-grad, graphed, cached), rollout (serve + behavior capture), train (grad,
|
||||
checkpointed) — they differ only in grad mode and capture; the ladder measures the gap.
|
||||
|
||||
---
|
||||
|
||||
## 10. Training & RL on the same loops
|
||||
|
||||
```
|
||||
serve : request → program → loop → WorkUnits → artifacts
|
||||
rollout : prompt batch → program → loop → WorkUnits → BehaviorRecords → rewards → update
|
||||
```
|
||||
|
||||
The loop kernel is shared; the only difference is capture and training policy. **This is the moat — the one
|
||||
place a serving-only runtime structurally cannot follow.** The rollout forward *is* the serve forward plus
|
||||
capture, so every serving optimization (distilled samplers, cache-dit skips, CFG-parallel, paged/feature
|
||||
caches, step batching) is automatically a rollout optimization, and there is one numerics surface (the ladder
|
||||
*measures* the gap rather than a correction layer *papering over* it). The industry's two-runtime tax
|
||||
(verl-omni re-implements Wan inside vLLM-Omni + a correction layer; miles' TIS/MIS/bitwise-logprobs/R3 are
|
||||
mismatch patches) is exactly what collocation deletes — viable at FastVideo's 1–30B FSDP2 scale.
|
||||
|
||||
- **`BehaviorRecord`** — captured at generation time: seeds, scheduler trajectory, timesteps, latents-or-refs,
|
||||
log-probs *where applicable*, sampled/action tokens, guidance, reward in/out, precision, parallel plan,
|
||||
`weights_version`. Sized honestly — an opt-in instrument for goldens, not always-on.
|
||||
- **`WeightSyncPlan`** ships a **role**, not "the weights" (student / EMA / decay-blended old-policy /
|
||||
reference / teacher / critic), with a **per-component scope** so a sync versions and cache-invalidates one
|
||||
expert in isolation. Lifecycle (the RL flywheel's hardest correctness): freeze admission → drain/boundary-stop
|
||||
in-flight loops → transfer → bump version + invalidate that component's caches → resume.
|
||||
|
||||
Methods (a faithful CPU port — NFT is line-for-line vs the source — carrying none of the GPU/FSDP/checkpoint
|
||||
infra):
|
||||
|
||||
| Method | Consistency | Roles | Notes |
|
||||
|---|---|---|---|
|
||||
| `finetune` | C1 | student | plain flow-match regression |
|
||||
| `dmd2` | C2 (free) | student + fake-score critic + teacher | distribution-matching distillation |
|
||||
| `diffusion_nft` | C2 (free) | student + **old** (decay-blended) + reference | samples from *old*, not student |
|
||||
| `self_forcing` | C2 (free) | student + teacher | causal/chunked student |
|
||||
| `unified_rl` | C2 (based) | student (llm+transformer) + reference | §11 — joint LM+gen RL |
|
||||
| `joint_multi_rl` | C2 (based) | N refiners + generator | N-way joint RL |
|
||||
| `workflow_rl` | C2 (based) | two instances | end-to-end RL across a cross-model workflow |
|
||||
|
||||
---
|
||||
|
||||
## 11. Weight-sharing topologies & the stress-test catalog
|
||||
|
||||
The central question — *does the Card/Loop/Program split generalize beyond serving + MoT, or will joint
|
||||
multi-expert RL / cross-model pipelines / interactive sessions / joint A/V / content-adaptive compute / hot
|
||||
weight-sync force a redesign?* — was answered by a battery of stress tests. **Every frontier capability landed
|
||||
as a new card / method / loop / workflow / session-driver / controller, with NO new runtime primitive** (the
|
||||
only real bug any test surfaced — a no-op generator gradient — was a fix in the sampler *library*). Condensed:
|
||||
|
||||
- **Joint LM+generator RL** (UniRL/PromptRL) — one reward → token policy-gradient on the LM *and* FlowGRPO PPO
|
||||
on the DiT; dual log-prob capture (categorical + Gaussian SDE); likelihood-based C2 (per-step identity ⇒
|
||||
ratio == 1); two independently-versioned weight-sync plans; SDE rollout sampler gated behind `sde_rollout` so
|
||||
the serve path is byte-for-byte unchanged.
|
||||
- **Qwen-Omni cascade** — three disjoint experts, three loop types (`ar_decode→ar_decode→audio_decode`),
|
||||
chained cross-stage conditioning, streaming codec→waveform.
|
||||
- **Cross-model Workflow** (T2I→I2V) — composition across *distinct* model instances; each model keeps its own
|
||||
interleave-parity guarantee; a `workflow_id` is a first-class servable in the same namespace as a `model_id`,
|
||||
registered via `WorkflowRegistry`. Plus **nested workflows** (recursive) and **non-linear shapes** (fan-out,
|
||||
best-of-N feedback).
|
||||
- **N-way joint RL** — generalizes joint RL to arbitrary N (per-component sync, dict grad-targets); surfaced a
|
||||
*credit-assignment* finding (per-expert reward clean; shared reward noisy) — a reward-shaping choice, not a
|
||||
substrate change.
|
||||
- **Interactive world-model session** — persistent cross-request state, transactional step-boundary
|
||||
cancellation, no cross-session smearing.
|
||||
- **End-to-end RL over a workflow** — one *final-video* reward trains an *earlier* model; proven causal by a
|
||||
control (constant reward ⇒ nothing moves).
|
||||
- **Heterogeneous WorkUnit co-scheduling** — `VAE_TILE` interleaves bit-identically with `DIFFUSION_STEP`
|
||||
through one budget (the §20 falsifier's *mechanism* half).
|
||||
- **Joint A/V** (LTX-2 T2VS, per-modality guidance), **content-adaptive compute** (cache-dit skip + early-exit,
|
||||
ragged step counts that still pass the interleave gate), **hot weight-sync under in-flight serving**
|
||||
(drain-correct), **served reward model** (`REWARD_BATCH`), **speculative decoding** (exact, lower-latency
|
||||
AR), the **RL→distill flywheel** (RL-improve → distill from the RL'd teacher → faster card with provenance),
|
||||
and the **adapter plane** (per-request LoRA/ControlNet over one base).
|
||||
|
||||
---
|
||||
|
||||
## 12. Request / session / artifact + programs/workflows
|
||||
|
||||
Typed runtime objects, not IDs in a batch: **`Request`** (task is *declared*, never inferred; `inputs:
|
||||
list[ModalPart]`; AR `sampling` vs `diffusion` params; `OutputSpec`), **`Session`** (long-lived interactive
|
||||
context: prompt memory, media streams, persistent cross-request chunk-KV), **`Artifact`** (named, typed, with
|
||||
provenance — `VideoArtifact`, `AudioArtifact(sample_rate)`, `TextArtifact`, … — killing the `extra["audio"]`
|
||||
pattern), **`Stream`** (one ordered event channel), **`CancelScope`**. A **`Program`** composes one card's
|
||||
loops into a task DAG (`ComponentNode` for kernel-free seams, `ModelLoopNode` to drive a loop to completion);
|
||||
a **`Workflow`** composes *across* model instances (each stage a full `engine.run`, artifacts threaded
|
||||
stage→stage) — the crossing is a Workflow boundary, not a program loop step, so each model keeps its parity
|
||||
guarantee. Workflows **compile, they are not the runtime** (a ComfyUI graph maps onto a `Program`; unknown
|
||||
nodes become `ExternalNode`s or a coverage rejection — never silent wrongness).
|
||||
|
||||
---
|
||||
|
||||
## 13. Serving & fleet
|
||||
|
||||
Per the standing instruction "don't completely rely on Dynamo; we still need our own version," v2 ships a
|
||||
complete stack and treats Dynamo as one optional backend: **`AsyncEngine`** (queue, lifecycle, live SSE
|
||||
streaming, step-boundary cancellation) + an OpenAI-compatible server on stdlib asyncio (`serving/http.py`:
|
||||
`/v1/chat` SSE, `/v1/images`, `/v1/videos` job+poll, `/v1/models`, `/health`, `/metrics`); a
|
||||
**`DisaggregatedRunner`** proven **bit-identical to inline** + role/stage pools; connectors with
|
||||
`chunk_ready` + credit-based flow control; **our own `LocalFleet`** (cost/affinity/least-loaded routing,
|
||||
health/drain) and a **`DynamoWorkerAdapter`** exporting the *same* `DeploymentCard` + cost model (one object,
|
||||
two consumers) so Dynamo *can* front us but never *defines* the core. The clean line: the fleet owns global
|
||||
routing / cold start / role-pool scaling / SLO placement / failover; the engine owns model load / loop
|
||||
execution / local scheduling / local memory+cache / parity / WorkUnit batching.
|
||||
|
||||
---
|
||||
|
||||
## 14. Extensions
|
||||
|
||||
Versioned hook points assembled at loop build (an unused hook is *literally absent* from the hot path), wrapping
|
||||
`ctx.execute(plan)`. **Observers (read-only):** `ParityAligner`, `Profiler` (calibrates the cost model),
|
||||
`NaNWatch`, `ActivationTrace`. **Interceptors (compute-altering):** `StepInterceptor` (step-skip / cached
|
||||
prediction) and `BlockInterceptor` (cache-dit DBCache/FBCache/TaylorSeer). State lives in
|
||||
`LoopState.plugin_state[id]`, keyed **per request and per CFG branch** — the structural fix for module-global
|
||||
residual state that silently corrupts cache-dit/TeaCache forks under concurrency. **Capability negotiation:** a
|
||||
4-step distilled card *rejects* a residual-skip interceptor rather than producing garbage. **Trust boundary:**
|
||||
plugins are enabled at deploy scope only; requests only *parameterize* pre-enabled plugins through validated
|
||||
schemas. This is the seam M\* (§18) calls "extensible — integrate FastVideo-STA / xDiT / Inferix / FlashDrive";
|
||||
v2 already has the mechanism.
|
||||
|
||||
---
|
||||
|
||||
## 15. The model taxonomy (what the runtime must serve)
|
||||
|
||||
| Paradigm | Examples | Loop shape | State | Output |
|
||||
|---|---|---|---|---|
|
||||
| Bidirectional video diffusion | Wan2.1/2.2, Hunyuan(15), LongCat, Cosmos2/2.5, LTX-2 | N denoise steps over full clip | latents, CFG branches, block caches | video |
|
||||
| Few-step distilled | DMD/FastWan, TurboWan, rCM | 1–4 denoise steps | latents | video |
|
||||
| Causal/AR video | Wan-Causal(-DMD), MatrixGame2/3, LongCat-VC, SF-Wan | outer chunk loop × inner denoise | DiT KV (slab chunk-KV) | video (streamable) |
|
||||
| Interactive world models | MatrixGame, GameCraft, HYWorld, Gen3C, LingBotWorld | chunk loop driven by live actions | KV + action/camera cond | video stream |
|
||||
| Image | Flux2, SD3.5, Qwen-Image, Kandinsky5 | N denoise steps | latents | image (batchable) |
|
||||
| Audio / Joint A/V | Stable Audio; LTX-2 (video+audio), Cosmos3 t2vs | denoise over (joint) latents | per-modality CFG | audio / video+audio |
|
||||
| Multi-stage refinement | LTX-2 (base→upsample→refine), Hunyuan15-SR | pipeline of loops | inter-stage latents | video |
|
||||
| AR token decode | Cosmos3 reasoner; omni thinkers/talkers | token loop until EOS | paged text-KV | text / codec tokens |
|
||||
| Vocoder / one-shot | LTX-2 vocoder, audio codec → wav | single forward / chunked | none | audio |
|
||||
| Hybrid MoT (omni) | Cosmos3, BAGEL | AR loop then/while denoise, shared weights | text-KV + denoise state + packed seq | text+video+audio+action |
|
||||
|
||||
A runtime that serves the MoT row serves everything above it; a DAG-of-engines cannot (the reasoner and
|
||||
denoiser share weights).
|
||||
|
||||
---
|
||||
|
||||
## 16. Package layout (actual)
|
||||
|
||||
```
|
||||
v2/
|
||||
card/ ModelCard, ComponentSpec, LoopSpec, RecipeSpec, ParitySpec, instance, load_card
|
||||
loop/ contracts (LoopState/WorkPlan/StepResult/Done), driver (LoopRunner), policies (cfg/flowshift/
|
||||
precision/routing), sampler (flow-match + FlowGRPO SDE)
|
||||
program/ ComponentNode/ModelLoopNode/Program; Workflow + WorkflowRegistry (cross-model; §12)
|
||||
request/ requests, params (DiffusionParams + sde_rollout), tasks (TaskType), streams, artifacts, cancel
|
||||
runtime/ engine (run/run_serial/run_interleaved, workflow-aware), async_engine, scheduler, admission,
|
||||
cudagraph (piecewise capture), disaggregated (DisaggregatedRunner), pools, session (WorldModelSession), context
|
||||
cache/ keys (CacheKey, content_hash), classes (feature/residual/slab_kv/paged_kv), manager
|
||||
memory/ allocator, reservations, refundable budget
|
||||
transport/ manifests + connectors (in-proc; chunk_ready + credit flow)
|
||||
parallel/ plan (axis vocab), mesh, validation
|
||||
parity/ aligner, ladder, interleave_gate, compare_outputs
|
||||
extend/ observers (Profiler/NaNWatch), interceptors (cache-dit), registry, base
|
||||
platform/ Platform.detect(); COMPONENTS(kind,device,variant) + KERNELS(op,device,arch,variant) registries;
|
||||
backends/{toy.py (numpy reference + parity oracle), torch_backend.py (real GPU)}
|
||||
loader/ v2-owned component-loader seam (currently delegates to fastvideo; vendored cutover later)
|
||||
models/ v2-namespaced model code (re-export stubs over fastvideo until vendored; see vendoring memory)
|
||||
recipes/ the concrete cards/programs/loops — 35+ families: wan21, ltx2, wan_causal, sfwan22, fastwan,
|
||||
turbowan, flux2, sd35, kandinsky5, hunyuan_video(15), longcat, cosmos2/25/3, gen3c, hyworld,
|
||||
hunyuangamecraft, matrixgame2/3, lingbotworld, lucy_edit, stable_audio, wan_fun_control,
|
||||
bagel, qwen_omni, omni, unified, multi_expert, image_video, tiled, adaptive, adapters,
|
||||
speculative, reward
|
||||
training/ rollout, behavior, rewards (+ServedRewardScorer), weight_sync (+WeightSyncController),
|
||||
flywheel, methods/{finetune,dmd2,diffusion_nft,self_forcing,unified_rl,joint_multi_rl,workflow_rl}
|
||||
serving/ AsyncEngine glue, OpenAI server (http.py)
|
||||
deploy/ DeploymentCard (card.py), LocalFleet (fleet.py), DynamoWorkerAdapter (dynamo.py)
|
||||
distributed/ single-process dist-init seam (1×1 device mesh)
|
||||
tests/ 34 files, 216 tests — `pytest v2/tests/` OR `python3 v2/run_tests.py` (zero deps)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 17. Current status & GPU bring-up
|
||||
|
||||
**The kernels are no longer toys-only.** The `platform/` substrate selects backends through two tuple-keyed
|
||||
registries (`COMPONENTS(kind, device, variant)` + `KERNELS(op, device, arch, variant)`) that a detected
|
||||
`Platform` resolves, with the numpy **toy backend** as the terminal fallback rung *and* parity oracle. On a GPU
|
||||
box, `platform/backends/torch_backend.py` provides real `TorchComponent` adapters wrapping the real model code
|
||||
(resolved from each card's `load_id`, weights from `ComponentSpec.checkpoint`) — and the loops, scheduler,
|
||||
caches, parity, training, and workflows are **unchanged**, exactly as the (recipe, runtime) separation promised.
|
||||
|
||||
**Verified on H100 this session:**
|
||||
- **20+ real model families GPU-verified** end-to-end via the torch backend (`VideoGenerator.from_pretrained` +
|
||||
`generate_video`), including SF-Wan (self-forcing causal, CFG-free 4-step DMD) and LTX-2 two-stage SR
|
||||
(frame-verified) — see [`../v2_debug_videos/vlm.md`](../v2_debug_videos/vlm.md).
|
||||
- **The omni trio runs real via vllm-omni** (in an isolated venv; FastVideo's env untouched): **BAGEL-7B-MoT**
|
||||
(two-stage MoT, on-prompt 1024² image), **Qwen2.5-Omni-7B** (thinker→talker→code2wav, coherent text + 24kHz
|
||||
speech, 2-GPU), **Cosmos3-Nano** (`Cosmos3OmniDiffusersPipeline` T2V, 720p, frame-verified). Recipe +
|
||||
box-specific flags (`VLLM_USE_FLASHINFER_SAMPLER=0`, `VLLM_USE_DEEP_GEMM=0`, guardrails-off) are saved in the
|
||||
`vllm-omni-bringup` memory; outputs in [`../v2_debug_videos/omni/`](../v2_debug_videos/omni/).
|
||||
|
||||
**What is genuinely not built yet** (named, not hidden): real *distributed* parallelism inside one pool
|
||||
(collectives are stubbed; multi-node is *multiple* pools fronted by the fleet); the ComfyUI workflow compiler;
|
||||
WebRTC realtime *wire* (the interactive *session* logic is built); full LLM-grade AR serving (radix trees,
|
||||
chunked-prefill sophistication); C3 (batch-invariance) and C4 (quality/preference + the eval system); a
|
||||
torch-native loop surface (loops still marshal numpy↔torch at the boundary — a perf follow-up).
|
||||
|
||||
```bash
|
||||
cd /home/scratch.willlin_ent/FastVideo
|
||||
python3 -m pytest v2/tests/ -q # 216 tests
|
||||
python3 v2/run_tests.py # same suite, zero deps
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 18. Roadmap
|
||||
|
||||
The forward plan has two tracks: (A) **adopt the valuable ideas from the M\* paper** (the closest external work
|
||||
to v2's thesis), and (B) **finish the GPU port**. The full M\* gap-analysis (28-agent workflow, adversarially
|
||||
verified against the code) is in [`.agents/exploration/mstar-v2-roadmap.md`](../.agents/exploration/mstar-v2-roadmap.md).
|
||||
|
||||
**M\* in one line.** *M\*: A Modular, Extensible, Serving System for Multimodal Models* (arXiv 2606.12688) is
|
||||
essentially v2's thesis one step more mature: "every composite model is a dataflow graph; every request is a
|
||||
*Walk* over it." It beats vLLM-Omni/SGLang-Omni on exactly the models v2 now runs (BAGEL, Qwen3-Omni) and
|
||||
explicitly names **FastVideo** as an integratable technique. **v2 already implements the harder half** (its
|
||||
`Program` *is* M\*'s graph; `shared_weight_components` *is* cross-Walk node sharing) and **exceeds** M\* on a
|
||||
validated cost model, the interleave **bit-parity** gate, integrated training, and the `extend/` plugin seam.
|
||||
The insight: v2's substrate is ~80% built but parts are **inert** (authored metadata never wired to an
|
||||
executor).
|
||||
|
||||
**Prioritized (P0 = highest leverage / lowest risk; parity-safe, CPU-toy-testable):**
|
||||
|
||||
| Pri | Item | Action |
|
||||
|---|---|---|
|
||||
| **P0** | Min-components per request | Consume `required_for` in `Program.active_nodes`; **fix the real bug**: `runtime/engine.py:88` uses `self.program.nodes` while `runtime/disaggregated.py:96` uses `active_nodes(request)` — the two runners disagree. Deliver via the registry/card builder so all cards inherit it. → M\*'s "execute the minimum components per request," cutting wasted reasoner/AR steps on single-modality BAGEL/Cosmos3 requests. |
|
||||
| **P0** | Real EOS + declarative `DynamicLoop` | `recipes/omni/ar_loop.py` docstring claims EOS-stop but `next()` only checks `max_tokens`. Honor `eos_id` + `req.sampling.stop`; add `LoopSpec.dynamic_stop` + `register_loop_stop`. Also **training-enabling** (world-model rollout horizon). |
|
||||
| **P1** | CFG/branch as a label over one paged KV pool | Make `PagedKVCache` a real `(namespace,label)` store over one budget; reuse `CacheKey.guidance_sig` for the hash (NOT `partition_field`). → the measured BAGEL win (AR path only; diffusion has no KV). |
|
||||
| **P1** | `extend/` plugin: FastVideo-STA / Inferix | Expose this repo's sparse/sliding-tile attention + block-diffusion as `Interceptor`/`EngineKind` plugins — the paper's named integration, highest paper-alignment, low risk (seam exists). |
|
||||
| **P1** | `ParitySpec.output_determinism` | Close the dormant C3 rung so SDE/stochastic RL can declare its parity contract honestly. |
|
||||
| **P2** | Pluggable data plane; per-(node,Walk) placement; declarative TP/SP degrees; loop-spanning CUDA graphs | Gated on the multi-GPU runtime; the live 2-GPU Qwen-Omni bring-up is the natural first test. Keep the cheap declarative halves now (`EngineKind` tag, placement key, populate `parallel_plan_hash` on the serving cache path). |
|
||||
|
||||
**Do NOT regress** (where v2 already meets or beats M\*): the validated per-loop cost model, the interleave
|
||||
bit-parity gate, the C0–C4 consistency ladder, the integrated training plane (RL→distill flywheel,
|
||||
drain-correct hot weight-sync), CPU-toy parity for the whole stack, the `extend/` plugin seam, Dynamo
|
||||
citizenship.
|
||||
|
||||
**GPU-port track:** wire real distributed collectives (one-pool TP/SP); torch-native loop surface (drop the
|
||||
numpy↔torch boundary marshalling); admission budgeting of capture cost; per-stream workspace pool for a
|
||||
concurrent executor; the AR/vocoder GPU op families; the GPU training surface (`mse_grad_step`).
|
||||
|
||||
---
|
||||
|
||||
## 19. Reference synthesis
|
||||
|
||||
What v2 takes and what it constrains, per surveyed system (the design is a synthesis, never a copy):
|
||||
|
||||
| Source | Take | Constrain / reject |
|
||||
|---|---|---|
|
||||
| Cosmos3 (official + port) | Shared instance across reason/diffusion/action/sound; packed multimodal sequences; component+scheduler parity matrices | A strong `ModelCard`, not the framework; no Cosmos branching in the global runtime |
|
||||
| vLLM core | Running-first scheduling, reservation-before-admission, model-owned state, KV/encoder cache managers, CuMem sleep/wake, CUDA-graph dispatch, KV-connector split | Token scheduling is one WorkUnit kind; **never** full-graph compile |
|
||||
| sglang `multimodal_gen` | Role pools, request lifecycle, capacity dispatch, transfer manifests, disagg state machine, cache-dit integration | No giant mutable `Req`/`ForwardBatch` as the API; not single-item diffusion scheduling |
|
||||
| vLLM-Omni | Frozen pipeline-spec ⟂ deploy-YAML split; `OmniConnectorBase` + `chunk_ready`; `SupportsStepExecution` as loop-inversion prior art (we generalize to always-on); `CFGParallelMixin` proves CFG-as-policy; 3 cache subsystems confirm per-class pools | Expresses MoT only as one **opaque** request-scheduled stage (no step visibility, no cross-request batching); cross-stage KV is a *copy*; **no cost model** |
|
||||
| sglang-omni | The `next/wait_for/merge_fn/stream_to` edge vocabulary; Relay + **credit-based flow control** | Stages own disjoint weights; hybrid only as AR-stage→DiT-stage |
|
||||
| Dynamo | Fleet routing, disagg role pools, KV-aware routing, KVBM, SLA planner, cold-start weight streaming | Orchestrates engines; never the core. Export a `DeploymentCard` + cost model to it |
|
||||
| diffusers Modular | `ComponentSpec`/`modular_model_index.json` interchange; Guiders ≈ CFG policies | A Python pipeline interpreter is not the perf boundary; loop blocks own their iteration (not inversion) |
|
||||
| xDiT / PipeFusion / USP | DiT parallelism catalog (USP, ring/ulysses, PipeFusion, CFG-parallel, DistVAE) + world-size validation | Parallelism lives in the runtime + card, not a wrapper-per-model; `pp_patch` invalid for causal |
|
||||
| TorchTitan | Named mesh axes, `ParallelDims` validation, ModelSpec discipline, batch-invariance utils | Adopt the discipline, not the stack; `WeightSyncPlan` owns layout (DCP/TorchStore don't reshard) |
|
||||
| verl-omni / miles / cosmos-rl | Rollout adapters, per-step capture, async rewards, group-relative advantage, TIS/MIS, deterministic modes, per-payload weight-version, AIPO off-policy masking | The **two-runtime tax is the thing to delete**; capture behavior *in* the serving loop |
|
||||
| UniRL-Zero / PromptRL | Joint LM-refiner + flow-generator RL under one reward; FlowGRPO SDE/ODE with per-step log-probs; group advantage → token-PG + PPO; prompt-only vs joint ablation | A card + a method, not a bespoke trainer; SDE sampler gated behind `sde_rollout`; PPO ratio rests on the likelihood-based C2 gate |
|
||||
| ComfyUI | Workflow graph, node-signature cache, model memory management | Compile to `Program`; dynamic node execution is not the core; GPL hygiene |
|
||||
| Dreamverse / LiveKit | Sessions, prompt memory, typed media IPC, cancellation, duty-cycle capacity, preference-data flywheel; realtime frame/PTS streaming | Product/session behavior is first-class in the request plane, never merged into the model core; RTC only when triggers fire (<100ms interactive) |
|
||||
| **M\* (2606.12688)** | The Walk-Graph framing (named Walks + state machine), per-(node,Walk) placement, CFG-as-cache-label over one paged pool, the "extensible: integrate FastVideo/xDiT/Inferix" call-out | v2 already has the harder half + exceeds on cost/parity/training; adopt the declarative authoring layer where it's inert (§18) |
|
||||
|
||||
---
|
||||
|
||||
## 20. Honest unknowns & falsifiers
|
||||
|
||||
An ambitious design is not an unfalsifiable one. The bets, with the experiment that kills each:
|
||||
|
||||
- **Step-level scheduling must *pay* for video.** Runtime-owned diffusion iteration has narrow precedent
|
||||
(vllm-omni's opt-in `SupportsStepExecution`); v2 makes it the always-on universal contract. **Falsifier:** on
|
||||
a real duty-cycle trace, if step-level scheduling does not beat a request-level baseline (≥2 concurrent
|
||||
sessions/GPU, p95 within SLO), degrade to request-level dispatch and keep only the loop contract's
|
||||
streaming/cancellation/behavior seams (which still justify it). The contract is safe even if the scheduling
|
||||
bet loses — that is the insurance.
|
||||
- **The general WorkUnit scheduler may be over-general.** The *mechanism* half is validated (`VAE_TILE` and
|
||||
`REWARD_BATCH` interleave/schedule through one budget); the *economic* half (does it pay vs an in-loop call)
|
||||
is a GPU-port measurement.
|
||||
- **Cost-model admission is a modeling bet** — argued on its narrow window (many small concurrent jobs),
|
||||
measured on the port.
|
||||
- **Quality is unmeasured.** C4 (artifact quality / preference) and the eval system it needs do not exist yet,
|
||||
and they gate every product claim ("fast mode is equivalent", RL reward validity, distillation comparisons).
|
||||
|
||||
**Final position.** 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 by test, and the interleave gate is non-negotiable;
|
||||
training records behavior on the same loops it serves. The weight-sharing topology, the composition graph, the
|
||||
training recipe, the reward, and the session/sync lifecycle are all **data** over cards, loops, workflows, and
|
||||
controllers — so a new frontier capability is a card or a driver, not a rewrite.
|
||||
@@ -0,0 +1,92 @@
|
||||
"""v2 — a scoped, CPU-testable realization of the model-native 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; the interleave gate is non-negotiable.
|
||||
> Training records behavior on the same loops it serves.
|
||||
|
||||
Phase 1 supports Wan2.1-1.3B (T2V) and LTX2.3 (2-stage distilled), plus four training
|
||||
methods on Wan2.1-1.3B (finetuning, DMD2, DiffusionNFT, self-forcing). The spine is
|
||||
omni-ready (multi-loop ModelInstance, ar_decode/chunk_rollout loop kinds) for the phase-2
|
||||
Cosmos3 + vllm-omni omni ports.
|
||||
|
||||
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._enums import (
|
||||
Capability,
|
||||
ConsistencyLevel,
|
||||
ExecutionProfile,
|
||||
LoopKind,
|
||||
WorkUnitKind,
|
||||
)
|
||||
from v2.card import (
|
||||
CapabilityMatrix,
|
||||
ComponentSpec,
|
||||
CostModel,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ModelInstance,
|
||||
ParitySpec,
|
||||
RecipeSpec,
|
||||
load_card,
|
||||
)
|
||||
from v2.program import ComponentNode, ModelLoopNode, Program, ProgramKind, when_opt, when_task
|
||||
from v2.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",
|
||||
"CostModel",
|
||||
"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,79 @@
|
||||
"""Shared vocabulary enums.
|
||||
|
||||
These live in a leaf module so both ``card/`` (which references loop kinds and
|
||||
consistency levels in its specs) and ``loop/``/``runtime/`` can import them
|
||||
without a circular dependency. ``card/`` imports no runtime; this is pure vocabulary.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class LoopKind(str, Enum):
|
||||
"""The kind of iterative computation a LoopSpec describes."""
|
||||
DIFFUSION_DENOISE = "diffusion_denoise" # N solver steps over full clip (Wan, LTX-2)
|
||||
CHUNK_ROLLOUT = "chunk_rollout" # causal chunk loop × inner denoise (self-forcing, world models)
|
||||
AR_DECODE = "ar_decode" # token loop until EOS (reasoner/thinker/talker — phase 2)
|
||||
VAE_TILE = "vae_tile" # tiled VAE encode/decode
|
||||
ENCODER = "encoder" # one-shot encoder (text/vision) — degenerate single-step loop
|
||||
AUDIO_DECODE = "audio_decode" # vocoder / codec decode
|
||||
TRAIN_FORWARD = "train_forward" # grad-enabled forward for a training method
|
||||
|
||||
|
||||
class WorkUnitKind(str, Enum):
|
||||
"""The smallest schedulable action — the scheduler's currency unit.
|
||||
|
||||
Tokens are one kind among many, not the scheduler itself: this generalizes
|
||||
vLLM's token scheduler to the work units diffusion needs.
|
||||
"""
|
||||
AR_PREFILL = "ar_prefill"
|
||||
AR_TOKEN = "ar_token"
|
||||
DIFFUSION_STEP = "diffusion_step"
|
||||
DIFFUSION_WINDOW = "diffusion_window"
|
||||
CHUNK_STEP = "chunk_step"
|
||||
ENCODER_CHUNK = "encoder_chunk"
|
||||
VAE_TILE = "vae_tile"
|
||||
AUDIO_CHUNK = "audio_chunk"
|
||||
REWARD_BATCH = "reward_batch" # RL
|
||||
LOGPROB_BATCH = "logprob_batch" # RL (likelihood-based)
|
||||
TRANSFER = "transfer"
|
||||
CACHE_IO = "cache_io"
|
||||
GRAPH_CAPTURE = "graph_capture"
|
||||
|
||||
|
||||
class ConsistencyLevel(str, Enum):
|
||||
"""The consistency ladder. RL methods declare their required rung."""
|
||||
C0 = "C0" # component parity (VAE/encoder/block/scheduler-step in isolation)
|
||||
C1 = "C1" # loop parity (full denoise trajectory / AR logits, fixed seed)
|
||||
C2 = "C2" # behavioral identity (train-forward ≡ serve-forward on the RL objective's quantity)
|
||||
C3 = "C3" # distribution parity (rollout distribution under allowed nondeterminism)
|
||||
C4 = "C4" # artifact quality (SSIM / reward agreement / human preference)
|
||||
|
||||
@property
|
||||
def rank(self) -> int:
|
||||
return {"C0": 0, "C1": 1, "C2": 2, "C3": 3, "C4": 4}[self.value]
|
||||
|
||||
|
||||
class ExecutionProfile(str, Enum):
|
||||
"""Three forwards, one loop definition — differ only in grad mode + capture."""
|
||||
SERVE = "serve" # no-grad, graphed, cached, possibly quantized
|
||||
ROLLOUT = "rollout" # serve profile + behavior capture
|
||||
TRAIN = "train" # grad, checkpointed, FSDP-gathered
|
||||
|
||||
|
||||
class Capability(str, Enum):
|
||||
"""CapabilityMatrix entries (model plane)."""
|
||||
TEXT_TO_VIDEO = "text_to_video"
|
||||
IMAGE_TO_VIDEO = "image_to_video"
|
||||
VIDEO_TO_VIDEO = "video_to_video"
|
||||
TEXT_TO_IMAGE = "text_to_image"
|
||||
TEXT_TO_VIDEO_SOUND = "text_to_video_sound"
|
||||
AUDIO_TO_VIDEO = "audio_to_video"
|
||||
TEXT_TO_SPEECH = "text_to_speech" # thinker→talker→vocoder (Qwen-Omni): reason + speak
|
||||
REASONING_TEXT = "reasoning_text"
|
||||
ACTION_CONDITIONING = "action_conditioning"
|
||||
STREAMING_VIDEO_CONTINUATION = "streaming_video_continuation"
|
||||
VAE_ENCODE = "vae_encode"
|
||||
VAE_DECODE = "vae_decode"
|
||||
POLICY_ROLLOUT = "policy_rollout" # can serve as an RL rollout engine
|
||||
LOGPROB_RECOMPUTE = "logprob_recompute"
|
||||
@@ -0,0 +1,20 @@
|
||||
"""Shared lightweight type aliases for v2.
|
||||
|
||||
The core (card/loop/runtime/cache/parity/program/request/training) is numpy-only
|
||||
and CPU-testable, so tensors are typed structurally as ``TensorLike``: a
|
||||
``numpy.ndarray`` on the CPU test path, a ``torch.Tensor`` on GPU. The core never
|
||||
imports torch — only model-component adapters do, lazily (see
|
||||
``v2/card/components.py``).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
# A tensor-like object. numpy.ndarray (CPU tests) or torch.Tensor (GPU). The core
|
||||
# only relies on duck-typed ops provided by the active backend (see card/backend).
|
||||
TensorLike = Any
|
||||
|
||||
Shape = tuple[int, ...]
|
||||
|
||||
# A stable content hash (hex) used for cache keys and provenance.
|
||||
Hash = str
|
||||
@@ -0,0 +1,6 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""``api`` facade — the typed config dataclasses the VideoGenerator consumes. Re-exported so v2 code
|
||||
imports ``v2.api`` instead of ``fastvideo.api``; a vendored cutover will replace these with v2-native configs."""
|
||||
from fastvideo.api import ( # noqa: F401
|
||||
EngineConfig, GenerationRequest, GenerationResult, GeneratorConfig, OffloadConfig, OutputConfig, SamplingConfig,
|
||||
)
|
||||
Vendored
+11
@@ -0,0 +1,11 @@
|
||||
"""Cache plane — correct by key, per-class pools."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.cache.classes import FeatureCache, PagedKVCache, ResidualCache, Slab, SlabKVCache, make_pool
|
||||
from v2.cache.keys import CacheKey, CachePolicy, content_hash
|
||||
from v2.cache.manager import CacheManager
|
||||
|
||||
__all__ = [
|
||||
"CacheKey", "CachePolicy", "content_hash", "CacheManager", "FeatureCache", "ResidualCache", "SlabKVCache",
|
||||
"PagedKVCache", "Slab", "make_pool"
|
||||
]
|
||||
Vendored
+185
@@ -0,0 +1,185 @@
|
||||
"""Per-class cache pools.
|
||||
|
||||
No single unified block pool: cache classes differ by 150-500x in natural granularity, so
|
||||
each gets its own statically budgeted pool behind one ``CacheHandle``. Four classes:
|
||||
* ``FeatureCache`` — content-hash keyed, partitioned by adapter+weights (text/vision encoders)
|
||||
* ``ResidualCache`` — cache-dit residuals, scoped per request AND per CFG branch
|
||||
* ``SlabKVCache`` — chunk-KV slabs (self-forcing / world models); training mode disables recycle
|
||||
* ``PagedKVCache`` — paged text-KV stub for ar_decode (phase-2 omni); minority case, lazy
|
||||
|
||||
KV is the minority case — a pure bidirectional deployment (Wan/LTX T2V) allocates none of it.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from v2.cache.keys import CacheKey, CachePolicy
|
||||
|
||||
|
||||
def _nbytes(value: Any) -> int:
|
||||
if hasattr(value, "nbytes"):
|
||||
return int(value.nbytes)
|
||||
return 0
|
||||
|
||||
|
||||
class _Pool:
|
||||
|
||||
def __init__(self, policy: CachePolicy):
|
||||
self.policy = policy
|
||||
self.used_bytes = 0
|
||||
self.hits = 0
|
||||
self.misses = 0
|
||||
|
||||
|
||||
class FeatureCache(_Pool):
|
||||
"""Content-hash keyed, budget-aware FIFO/LRU.
|
||||
|
||||
Partitioned by ``adapter_versions``/``weights_version`` through the CacheKey, so two
|
||||
workflows sharing a prompt but differing in te-LoRA stack never serve stale embeddings.
|
||||
"""
|
||||
|
||||
def __init__(self, policy: CachePolicy):
|
||||
super().__init__(policy)
|
||||
self._store: OrderedDict[str, tuple[Any, int, CacheKey]] = OrderedDict()
|
||||
|
||||
def get(self, key: CacheKey) -> Any | None:
|
||||
if not self.policy.reuse_across_requests:
|
||||
return None
|
||||
h = key.hash
|
||||
if h in self._store:
|
||||
self.hits += 1
|
||||
self._store.move_to_end(h) # LRU
|
||||
return self._store[h][0]
|
||||
self.misses += 1
|
||||
return None
|
||||
|
||||
def put(self, key: CacheKey, value: Any) -> None:
|
||||
nb = _nbytes(value)
|
||||
h = key.hash
|
||||
if h in self._store:
|
||||
self.used_bytes -= self._store[h][1]
|
||||
self._store[h] = (value, nb, key)
|
||||
self._store.move_to_end(h)
|
||||
self.used_bytes += nb
|
||||
self._evict()
|
||||
|
||||
def _evict(self) -> None:
|
||||
while self.used_bytes > self.policy.max_bytes and self._store:
|
||||
_h, (_v, nb, _k) = self._store.popitem(last=False) # FIFO/LRU oldest
|
||||
self.used_bytes -= nb
|
||||
|
||||
def invalidate_weights(self, version: str) -> None:
|
||||
"""RL update_weights bumps weight epoch → drop entries from older epochs (wholesale)."""
|
||||
drop = [h for h, (_v, _nb, k) in self._store.items() if k.weights_version != version]
|
||||
for h in drop:
|
||||
_v, nb, _k = self._store.pop(h)
|
||||
self.used_bytes -= nb
|
||||
|
||||
def invalidate_components(self, components: set[str]) -> None:
|
||||
"""Drop only entries produced by the changed components (partition, not flush):
|
||||
a transformer-only weight sync must NOT evict text-encoder embeddings."""
|
||||
drop = [h for h, (_v, _nb, k) in self._store.items() if k.component_id in components]
|
||||
for h in drop:
|
||||
_v, nb, _k = self._store.pop(h)
|
||||
self.used_bytes -= nb
|
||||
|
||||
|
||||
class ResidualCache(_Pool):
|
||||
"""cache-dit residual store, scoped per ``LoopState`` AND per CFG branch.
|
||||
|
||||
Keyed by (namespace, branch, name) where namespace is the request/loop id. This is the
|
||||
structural fix for the module-global residual state that corrupts cache-dit forks under
|
||||
concurrency: two interleaved requests have disjoint namespaces.
|
||||
"""
|
||||
|
||||
def __init__(self, policy: CachePolicy):
|
||||
super().__init__(policy)
|
||||
self._store: dict[tuple[str, str, str], Any] = {}
|
||||
|
||||
def put(self, namespace: str, branch: str, name: str, value: Any) -> None:
|
||||
self._store[(namespace, branch, name)] = value
|
||||
|
||||
def get(self, namespace: str, branch: str, name: str) -> Any | None:
|
||||
v = self._store.get((namespace, branch, name))
|
||||
if v is not None:
|
||||
self.hits += 1
|
||||
else:
|
||||
self.misses += 1
|
||||
return v
|
||||
|
||||
def clear_namespace(self, namespace: str) -> None:
|
||||
for k in [k for k in self._store if k[0] == namespace]:
|
||||
del self._store[k]
|
||||
|
||||
|
||||
@dataclass
|
||||
class Slab:
|
||||
chunk_index: int
|
||||
k: Any
|
||||
v: Any
|
||||
|
||||
|
||||
class SlabKVCache(_Pool):
|
||||
"""Chunk-KV slabs for causal/world-model rollout.
|
||||
|
||||
``training_mode`` disables mid-rollout recycling so activation-checkpoint recompute
|
||||
doesn't double-advance the cache (self-forcing).
|
||||
"""
|
||||
|
||||
def __init__(self, policy: CachePolicy):
|
||||
super().__init__(policy)
|
||||
self._store: dict[str, list[Slab]] = {}
|
||||
self.window = max(1, policy.per_component.get("window", 1 << 30))
|
||||
self.training_mode = policy.training_mode_disables_recycle
|
||||
|
||||
def append(self, namespace: str, slab: Slab) -> None:
|
||||
slabs = self._store.setdefault(namespace, [])
|
||||
slabs.append(slab)
|
||||
self.used_bytes += _nbytes(slab.k) + _nbytes(slab.v)
|
||||
if not self.training_mode and len(slabs) > self.window:
|
||||
dropped = slabs.pop(0) # sliding-window recycle (inference only)
|
||||
self.used_bytes -= _nbytes(dropped.k) + _nbytes(dropped.v)
|
||||
|
||||
def get(self, namespace: str) -> list[Slab]:
|
||||
return self._store.get(namespace, [])
|
||||
|
||||
def clear_namespace(self, namespace: str) -> None:
|
||||
for slab in self._store.pop(namespace, []):
|
||||
self.used_bytes -= _nbytes(slab.k) + _nbytes(slab.v)
|
||||
|
||||
|
||||
class PagedKVCache(_Pool):
|
||||
"""Paged text-KV stub for ar_decode (phase-2 omni). Minority case, materialized lazily."""
|
||||
|
||||
def __init__(self, policy: CachePolicy):
|
||||
super().__init__(policy)
|
||||
self.total_blocks = max(1, policy.max_bytes // max(policy.block_bytes, 1))
|
||||
self.free = self.total_blocks
|
||||
self._alloc: dict[str, int] = {}
|
||||
|
||||
def allocate(self, namespace: str, n_blocks: int) -> bool:
|
||||
if n_blocks > self.free:
|
||||
return False
|
||||
self._alloc[namespace] = self._alloc.get(namespace, 0) + n_blocks
|
||||
self.free -= n_blocks
|
||||
return True
|
||||
|
||||
def free_namespace(self, namespace: str) -> None:
|
||||
self.free += self._alloc.pop(namespace, 0)
|
||||
|
||||
|
||||
_CLASS_REGISTRY = {
|
||||
"feature": FeatureCache,
|
||||
"residual": ResidualCache,
|
||||
"slab_kv": SlabKVCache,
|
||||
"paged_kv": PagedKVCache,
|
||||
}
|
||||
|
||||
|
||||
def make_pool(policy: CachePolicy) -> _Pool:
|
||||
cls = _CLASS_REGISTRY.get(policy.class_name)
|
||||
if cls is None:
|
||||
raise KeyError(f"unknown cache class {policy.class_name!r} (have {list(_CLASS_REGISTRY)})")
|
||||
return cls(policy)
|
||||
Vendored
+81
@@ -0,0 +1,81 @@
|
||||
"""CacheKey — cache correctness is a contract.
|
||||
|
||||
If a field can change output semantics, it is in the key (incorrect reuse is worse than no reuse).
|
||||
The serving hazard this kills: a request that shares a prompt but differs in te-LoRA stack must not
|
||||
serve stale embeddings — so the key is *partitioned* by ``adapter_versions``, not flushed. An RL
|
||||
``update_weights`` bumps ``weights_version`` and invalidates wholesale.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
|
||||
def content_hash(obj: Any) -> str:
|
||||
"""Stable content hash for feature-cache keys (text/vision embeddings).
|
||||
|
||||
Lets K rollout samples of one prompt reuse a single text encode.
|
||||
"""
|
||||
h = hashlib.sha256()
|
||||
if isinstance(obj, str):
|
||||
h.update(obj.encode("utf-8"))
|
||||
elif isinstance(obj, bytes | bytearray):
|
||||
h.update(obj)
|
||||
elif hasattr(obj, "tobytes") and hasattr(obj, "shape"): # numpy / torch tensor
|
||||
h.update(str(getattr(obj, "shape", "")).encode())
|
||||
h.update(str(getattr(obj, "dtype", "")).encode())
|
||||
try:
|
||||
h.update(obj.tobytes())
|
||||
except Exception:
|
||||
h.update(repr(obj).encode())
|
||||
else:
|
||||
h.update(repr(obj).encode())
|
||||
return h.hexdigest()[:32]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CacheKey:
|
||||
model_id: str
|
||||
component_id: str
|
||||
loop_id: str | None = None
|
||||
weights_version: str = "v0"
|
||||
adapter_versions: tuple[tuple[str, str], ...] = () # sorted (adapter_id, version) pairs
|
||||
precision: str = "float32"
|
||||
parallel_plan_hash: str = ""
|
||||
shape_sig: str = ""
|
||||
layout_sig: str = ""
|
||||
scheduler_sig: str | None = None
|
||||
guidance_sig: str | None = None
|
||||
seed: int | None = None
|
||||
input_hashes: tuple[tuple[str, str], ...] = ()
|
||||
step_index: int | None = None
|
||||
contract_version: str = "v0"
|
||||
|
||||
@property
|
||||
def hash(self) -> str:
|
||||
return hashlib.sha256(repr(self).encode()).hexdigest()[:24]
|
||||
|
||||
def partition_field(self) -> tuple:
|
||||
"""Fields that *partition* (not flush) a feature cache: adapters + weights."""
|
||||
return (self.weights_version, self.adapter_versions)
|
||||
|
||||
@staticmethod
|
||||
def adapters(d: dict[str, str] | None) -> tuple[tuple[str, str], ...]:
|
||||
return tuple(sorted((d or {}).items()))
|
||||
|
||||
@staticmethod
|
||||
def hashes(d: dict[str, str] | None) -> tuple[tuple[str, str], ...]:
|
||||
return tuple(sorted((d or {}).items()))
|
||||
|
||||
|
||||
@dataclass
|
||||
class CachePolicy:
|
||||
"""Runtime config for one cache class pool."""
|
||||
class_name: str # "feature" | "residual" | "slab_kv" | "paged_kv"
|
||||
max_bytes: int = 1 << 30
|
||||
block_bytes: int = 1 << 16
|
||||
eviction: str = "lru" # "lru" | "fifo" | "none"
|
||||
reuse_across_requests: bool = True
|
||||
per_component: dict[str, int] = field(default_factory=dict)
|
||||
training_mode_disables_recycle: bool = False
|
||||
Vendored
+75
@@ -0,0 +1,75 @@
|
||||
"""CacheManager — per-class pools with static budgets, behind one handle.
|
||||
|
||||
Static partitioning makes cross-class fragmentation impossible (jumbo slab traffic cannot
|
||||
strand text-KV pages and vice versa). Each class gets a budget carved at init from the card's
|
||||
CacheContracts. ``invalidate_weights`` implements the wholesale RL weight-epoch bump.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.cache.classes import _Pool, make_pool
|
||||
from v2.cache.keys import CachePolicy
|
||||
|
||||
|
||||
class CacheManager:
|
||||
|
||||
def __init__(self, policies: list[CachePolicy] | None = None):
|
||||
self._pools: dict[str, _Pool] = {}
|
||||
for p in (policies or []):
|
||||
self._pools[p.class_name] = make_pool(p)
|
||||
|
||||
@classmethod
|
||||
def from_card(cls, card) -> CacheManager:
|
||||
"""Build the per-class pools a card declares (KV pools materialize only if declared)."""
|
||||
policies = []
|
||||
for cc in card.caches.values():
|
||||
policies.append(
|
||||
CachePolicy(
|
||||
class_name=cc.cache_class,
|
||||
max_bytes=cc.max_bytes,
|
||||
block_bytes=cc.block_bytes,
|
||||
eviction=cc.eviction,
|
||||
reuse_across_requests=cc.reuse_across_requests,
|
||||
per_component=dict(cc.per_component),
|
||||
training_mode_disables_recycle=cc.training_mode_disables_recycle,
|
||||
))
|
||||
return cls(policies)
|
||||
|
||||
def pool(self, class_name: str) -> _Pool:
|
||||
if class_name not in self._pools:
|
||||
raise KeyError(f"no cache pool for class {class_name!r}; card did not declare it")
|
||||
return self._pools[class_name]
|
||||
|
||||
def has(self, class_name: str) -> bool:
|
||||
return class_name in self._pools
|
||||
|
||||
def invalidate_weights(self, version: str) -> None:
|
||||
for pool in self._pools.values():
|
||||
if hasattr(pool, "invalidate_weights"):
|
||||
pool.invalidate_weights(version)
|
||||
|
||||
def invalidate_components(self, components) -> None:
|
||||
"""Component-scoped invalidation: only drop caches for changed components."""
|
||||
comps = set(components)
|
||||
for pool in self._pools.values():
|
||||
if hasattr(pool, "invalidate_components"):
|
||||
pool.invalidate_components(comps)
|
||||
|
||||
def clear_namespace(self, namespace: str) -> None:
|
||||
"""Release a finished request's per-request caches (residual/slab)."""
|
||||
for pool in self._pools.values():
|
||||
if hasattr(pool, "clear_namespace"):
|
||||
pool.clear_namespace(namespace)
|
||||
if hasattr(pool, "free_namespace"):
|
||||
pool.free_namespace(namespace)
|
||||
|
||||
def stats(self) -> dict[str, Any]:
|
||||
return {
|
||||
name: {
|
||||
"used_bytes": p.used_bytes,
|
||||
"hits": p.hits,
|
||||
"misses": p.misses
|
||||
}
|
||||
for name, p in self._pools.items()
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Model Plane — the (recipe, runtime) pair as a typed card."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import Capability, ConsistencyLevel, ExecutionProfile, LoopKind, WorkUnitKind
|
||||
from v2.card.instance import ModelInstance, load_card
|
||||
from v2.card.specs import (
|
||||
CacheContract,
|
||||
CapabilityMatrix,
|
||||
CardValidationError,
|
||||
CheckpointManifest,
|
||||
ComponentSpec,
|
||||
CostModel,
|
||||
DataRef,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ParallelismContract,
|
||||
ParitySpec,
|
||||
ParityTestSpec,
|
||||
PrecisionContract,
|
||||
RecipeSpec,
|
||||
SamplingDefaults,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ModelCard",
|
||||
"ComponentSpec",
|
||||
"LoopSpec",
|
||||
"RecipeSpec",
|
||||
"ParitySpec",
|
||||
"ParityTestSpec",
|
||||
"CheckpointManifest",
|
||||
"CapabilityMatrix",
|
||||
"CostModel",
|
||||
"CacheContract",
|
||||
"SamplingDefaults",
|
||||
"ParallelismContract",
|
||||
"PrecisionContract",
|
||||
"DataRef",
|
||||
"CardValidationError",
|
||||
"ModelInstance",
|
||||
"load_card",
|
||||
"Capability",
|
||||
"ConsistencyLevel",
|
||||
"ExecutionProfile",
|
||||
"LoopKind",
|
||||
"WorkUnitKind",
|
||||
]
|
||||
@@ -0,0 +1,126 @@
|
||||
"""ModelInstance — a resident, loaded card: component instances, model state, caches,
|
||||
compiled graphs, and a parallel plan. One request may run several of the card's loops
|
||||
against one ``ModelInstance``.
|
||||
|
||||
This is what makes omni native: when two loops bind the same component
|
||||
(``shared_weight_components``), ``component()`` returns the *exact same* live object —
|
||||
no duplicated weights, shared live state. See v2/README.md.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.platform import Platform
|
||||
from v2.card.specs import ModelCard
|
||||
|
||||
|
||||
class ModelInstance:
|
||||
"""The live, resident form of a ModelCard. Also serves as the ``ModelState``
|
||||
passed to ``Loop.init`` (the loop reads components through it)."""
|
||||
|
||||
def __init__(self,
|
||||
card: ModelCard,
|
||||
parallel_plan: Any = None,
|
||||
cache_manager: Any = None,
|
||||
weights_version: str = "v0",
|
||||
platform: Any = None):
|
||||
self.card = card
|
||||
self.parallel_plan = parallel_plan
|
||||
self.caches = cache_manager
|
||||
self.weights_version = weights_version
|
||||
# The detected (device, arch). Resolves component/kernel implementations through the two
|
||||
# backend registries; defaults to CPU/numpy. Swapping this to a GPU platform changes only
|
||||
# the backend, not the loops/policies/training.
|
||||
self.platform = platform if platform is not None else Platform.cpu()
|
||||
# Piecewise CUDA-graph cache for loops declaring graph_capture="breakable_cudagraph". Lazily
|
||||
# created by the runtime (kept as a plain attr so card/ stays free of any runtime import);
|
||||
# a captured graph is tied to this instance's resident weights.
|
||||
self.graphs: Any = None
|
||||
self.adapter_versions: dict[str, str] = {}
|
||||
# Per-component weight versions: a component's version changes only when IT is synced,
|
||||
# so a transformer-only RL sync never invalidates the frozen text-encoder's feature cache.
|
||||
self.component_versions: dict[str, str] = {cid: weights_version for cid in card.components}
|
||||
self._components: dict[str, Any] = {}
|
||||
self._loops: dict[str, Any] = {}
|
||||
self._asleep: set[str] = set()
|
||||
|
||||
# --- components: shared by reference (the MoT requirement) ---------------- #
|
||||
def component(self, component_id: str) -> Any:
|
||||
if component_id in self._asleep:
|
||||
raise RuntimeError(f"component {component_id!r} is asleep; wake it first")
|
||||
if component_id not in self._components:
|
||||
spec = self.card.components.get(component_id)
|
||||
if spec is None:
|
||||
raise KeyError(f"component {component_id!r} not declared on card {self.card.model_id!r}")
|
||||
# The single materialization seam: the platform resolves (kind, device, variant) through
|
||||
# the COMPONENTS registry, falling back to spec.factory as the cpu/numpy terminal rung.
|
||||
self._components[component_id] = self.platform.build_component(spec, self)
|
||||
return self._components[component_id]
|
||||
|
||||
def has_component(self, component_id: str) -> bool:
|
||||
return component_id in self.card.components
|
||||
|
||||
# --- loops: one stateless Loop object per (instance, loop_id) ------------- #
|
||||
def loop(self, loop_id: str) -> Any:
|
||||
if loop_id not in self._loops:
|
||||
spec = self.card.loops.get(loop_id)
|
||||
if spec is None:
|
||||
raise KeyError(f"loop {loop_id!r} not declared on card {self.card.model_id!r}")
|
||||
if spec.loop_factory is None:
|
||||
raise RuntimeError(f"loop {loop_id!r} has no loop_factory")
|
||||
self._loops[loop_id] = spec.loop_factory()
|
||||
return self._loops[loop_id]
|
||||
|
||||
# --- sleep/wake by component (CuMem tags = component names) --------------- #
|
||||
def sleep(self, component_ids: list[str]) -> None:
|
||||
for cid in component_ids:
|
||||
self._components.pop(cid, None)
|
||||
self._asleep.add(cid)
|
||||
|
||||
def wake(self, component_ids: list[str]) -> None:
|
||||
for cid in component_ids:
|
||||
self._asleep.discard(cid)
|
||||
|
||||
# --- weight sync: bump version + invalidate caches ----------------------- #
|
||||
def version_of(self, component_id: str) -> str:
|
||||
"""The component's own weights version (defaults to the instance version)."""
|
||||
return self.component_versions.get(component_id, self.weights_version)
|
||||
|
||||
def set_weights_version(self, version: str, components: list[str] | None = None) -> None:
|
||||
"""Publish a new weights version. If ``components`` is given, only those components' versions
|
||||
bump and only their caches are invalidated (partition, not flush) — so a transformer-only RL
|
||||
weight sync leaves the frozen text-encoder's feature cache intact."""
|
||||
self.weights_version = version
|
||||
changed = components if components is not None else list(self.card.components.keys())
|
||||
for c in changed:
|
||||
self.component_versions[c] = version
|
||||
if self.caches is not None and hasattr(self.caches, "invalidate_components"):
|
||||
self.caches.invalidate_components(set(changed))
|
||||
# Evict captured CUDA graphs for the synced components too (else they leak on a real box;
|
||||
# version-in-key already makes them unreachable). Duck-typed → card/ imports no runtime.
|
||||
if self.graphs is not None and hasattr(self.graphs, "invalidate"):
|
||||
self.graphs.invalidate(set(changed))
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (f"ModelInstance(card={self.card.model_id!r}, weights={self.weights_version!r}, "
|
||||
f"resident={sorted(self._components)})")
|
||||
|
||||
|
||||
def load_card(card: ModelCard,
|
||||
parallel_plan: Any = None,
|
||||
cache_manager: Any = None,
|
||||
*,
|
||||
validate: bool = True,
|
||||
platform: Any = None) -> ModelInstance:
|
||||
"""The card-as-factory entrypoint: the card is a runtime factory.
|
||||
|
||||
``platform`` selects the backend (device, arch); when omitted it is detected (CPU/numpy here,
|
||||
CUDA on a torch+GPU box). The same card loads on any backend — only the resolved component/kernel
|
||||
implementations differ.
|
||||
"""
|
||||
if validate:
|
||||
card.validate()
|
||||
return ModelInstance(card,
|
||||
parallel_plan=parallel_plan,
|
||||
cache_manager=cache_manager,
|
||||
platform=platform if platform is not None else Platform.detect())
|
||||
@@ -0,0 +1,293 @@
|
||||
"""ModelCard and its sub-specs — the Model Plane.
|
||||
|
||||
The atomic unit is the (recipe, runtime) pair, owned by a typed ``ModelCard``. The card
|
||||
is both a declarative contract (strict enough to validate before any GPU touches it — see
|
||||
``ModelCard.validate``) and a runtime factory (it instantiates components, binds loops, and
|
||||
resolves caches — see ``card/instance.py``).
|
||||
|
||||
Boundary: ``card/`` imports no product/runtime. It depends only on the shared leaf modules
|
||||
(``_enums``, ``_types``) and references ``ParallelPlan`` under ``TYPE_CHECKING`` so there is
|
||||
no runtime coupling to ``parallel/``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from collections.abc import Callable
|
||||
|
||||
from v2._enums import Capability, ConsistencyLevel, LoopKind, WorkUnitKind
|
||||
|
||||
if TYPE_CHECKING: # avoid runtime card -> parallel coupling
|
||||
pass
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Cost model — the budget currency, REQUIRED on every LoopSpec #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class CostModel:
|
||||
"""Predicted GPU-time per WorkUnit at (shape, precision, policy).
|
||||
|
||||
The budget currency is predicted GPU-time, not counts. One object, two consumers:
|
||||
the scheduler's internal budget and the fleet's routing input. Calibrated online by
|
||||
the Profiler observer.
|
||||
"""
|
||||
kind: WorkUnitKind
|
||||
# synthetic seconds = base + per_unit * work_units(shape) * policy_factor
|
||||
base_seconds: float = 1.0e-3
|
||||
per_unit_seconds: float = 1.0e-6
|
||||
coefficients: dict[str, float] = field(default_factory=dict)
|
||||
|
||||
def predict(self, work_units: float, policy_factor: float = 1.0) -> float:
|
||||
"""Conservative baseline cost: admission is never optimistic."""
|
||||
return (self.base_seconds + self.per_unit_seconds * float(work_units)) * policy_factor
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Precision, parallelism, parity, cache, checkpoint contracts #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class PrecisionContract:
|
||||
default_dtype: str = "float32"
|
||||
component_overrides: dict[str, str] = field(default_factory=dict)
|
||||
quantization_scheme: str | None = None # "nvfp4" | "int8" | None
|
||||
training_precision: str = "float32" # distinct from serving precision
|
||||
|
||||
def dtype_for(self, component_id: str) -> str:
|
||||
return self.component_overrides.get(component_id, self.default_dtype)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParallelismContract:
|
||||
valid_plans: list = field(default_factory=list) # list[ParallelPlan]
|
||||
default_plan: Any = None # ParallelPlan | None
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheContract:
|
||||
"""A cache class the card declares it needs."""
|
||||
cache_class: str # "feature" | "residual" | "slab_kv" | "paged_kv"
|
||||
max_bytes: int = 1 << 30
|
||||
block_bytes: int = 1 << 16 # page bytes (paged) / slab granule (slab)
|
||||
eviction: str = "lru" # "lru" | "fifo" | "none"
|
||||
reuse_across_requests: bool = True # paged/feature=True; slab depends on mode
|
||||
per_component: dict[str, int] = field(default_factory=dict)
|
||||
training_mode_disables_recycle: bool = False # chunk-KV training mode
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParityTestSpec:
|
||||
"""A named tap + tolerance + ladder level."""
|
||||
name: str
|
||||
level: ConsistencyLevel
|
||||
tap: str # named activation tap, e.g. "block.0.out"
|
||||
rtol: float = 0.0
|
||||
atol: float = 0.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParitySpec:
|
||||
"""The (recipe, runtime) honesty contract. Measured, never assumed."""
|
||||
consistency_levels: list[ConsistencyLevel] = field(default_factory=lambda: [ConsistencyLevel.C1])
|
||||
tests: list[ParityTestSpec] = field(default_factory=list)
|
||||
tap_tolerances: dict[str, float] = field(default_factory=dict)
|
||||
interleave_required: bool = True # the batch-of-N gate is non-negotiable
|
||||
|
||||
@property
|
||||
def max_level(self) -> ConsistencyLevel:
|
||||
return max(self.consistency_levels, key=lambda c: c.rank) if self.consistency_levels else ConsistencyLevel.C0
|
||||
|
||||
|
||||
@dataclass
|
||||
class DataRef:
|
||||
"""What a recipe trained on, for governance/reproduction."""
|
||||
dataset_id: str = ""
|
||||
revision: str = ""
|
||||
description: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecipeSpec:
|
||||
"""The provenance half of the (recipe, runtime) pair.
|
||||
|
||||
``assumes_loop`` and ``assumes_precision`` are the teeth: a 4-step distilled
|
||||
model whose ``assumes_loop = "ddim_4step"`` cannot be served under a 50-step
|
||||
sampler without a typed mismatch error (enforced in ModelCard.validate).
|
||||
"""
|
||||
method: str = "base" # base | dmd2 | self_forcing | diffusion_nft | attn_qat_nvfp4
|
||||
parents: list[str] = field(default_factory=list) # teacher / base model_ids
|
||||
data_contract: DataRef = field(default_factory=DataRef)
|
||||
assumes_loop: str = "" # loop_id this recipe's weights require
|
||||
assumes_precision: str = "float32"
|
||||
consistency_required: ConsistencyLevel = ConsistencyLevel.C1
|
||||
|
||||
|
||||
@dataclass
|
||||
class ComponentSpec:
|
||||
"""A weight-bearing (or processing) component.
|
||||
|
||||
Omni-ready fields ``resident_for`` / ``optional_for`` / ``required_for`` turn the
|
||||
Cosmos3 lazy-sound-VAE problem into a declaration, not an ``if env_var`` inside
|
||||
``forward``.
|
||||
"""
|
||||
component_id: str
|
||||
kind: str # dit | vae | text_encoder | audio_vae | reasoner_tower | ...
|
||||
load_id: str = "" # "module:Class" for the real (torch) adapter
|
||||
config_schema: type | None = None
|
||||
io_schema: tuple[type | None, type | None] = (None, None)
|
||||
precision_policy: str | None = None
|
||||
placement_policy: str = "colocated"
|
||||
parallel_constraints: dict[str, Any] = field(default_factory=dict)
|
||||
parity_tests: list[ParityTestSpec] = field(default_factory=list)
|
||||
# omni-ready:
|
||||
resident_for: list[str] = field(default_factory=list) # loop_ids that keep this resident mid-request
|
||||
optional_for: set[str] = field(default_factory=set) # tasks that don't need it
|
||||
required_for: set[str] = field(default_factory=set) # tasks that require it
|
||||
# v2 wiring: a factory producing the live component (toy numpy or torch adapter)
|
||||
factory: Callable[..., Any] | None = None
|
||||
# GPU backend: weights source (HF id or local path) for the real torch adapter resolved from
|
||||
# ``load_id``. Empty for the CPU toy (its factory needs no weights); a GPU deployment fills it in.
|
||||
checkpoint: str = ""
|
||||
# GPU backend: optional explicit torch-adapter class "module:Class" (a TorchComponent subclass
|
||||
# constructed as cls(module, device=, dtype=)). Lets a NEW architecture declare its own adapter on
|
||||
# the card instead of editing the shared backend dispatch — so a port is a self-contained recipe
|
||||
# package. Empty -> the backend's built-in per-kind dispatch (Wan/LTX2) by module class name.
|
||||
adapter: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoopSpec:
|
||||
"""Describes one iterative computation the card can run.
|
||||
|
||||
The model owns loop *semantics*; the runtime owns loop *lifecycle*. A single
|
||||
``ModelInstance`` may run several of these against shared components — the MoT
|
||||
requirement (``shared_weight_components``).
|
||||
"""
|
||||
loop_id: str
|
||||
kind: LoopKind
|
||||
work_unit_kind: WorkUnitKind
|
||||
step_cost_model: CostModel # REQUIRED
|
||||
state_schema: type | None = None # the typed LoopState
|
||||
step_schema: type | None = None # the typed WorkPlan a step emits
|
||||
result_schema: type | None = None # the typed StepResult
|
||||
behavior_schema: type | None = None # what to capture for RL (None if not training-relevant)
|
||||
extension_schema: type | None = None # per-model LoopState extension (Cosmos3PackedSeq, etc.)
|
||||
cache_policy: list[str] = field(default_factory=list) # cache class names this loop draws from
|
||||
valid_parallel_plans: list = field(default_factory=list)
|
||||
graph_capture: str = "eager" # eager | breakable_cudagraph
|
||||
# omni-ready:
|
||||
shared_weight_components: list[str] = field(default_factory=list)
|
||||
allows_interleaving: bool = True
|
||||
# v2: the Loop implementation factory (built at bind time)
|
||||
loop_factory: Callable[..., Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class CheckpointManifest:
|
||||
"""Explicit declared components + key maps — no name-detector guessing."""
|
||||
upstream_source: str = ""
|
||||
revision: str = ""
|
||||
component_ownership: dict[str, list[str]] = field(default_factory=dict) # component_id -> file globs
|
||||
key_mappings: dict[str, str] = field(default_factory=dict) # ckpt_key -> component.param
|
||||
required_for: dict[str, set[str]] = field(default_factory=dict) # component_id -> tasks requiring it
|
||||
optional_for: dict[str, set[str]] = field(default_factory=dict)
|
||||
conversion_version: str = "v0"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CapabilityMatrix:
|
||||
capabilities: frozenset[Capability] = frozenset()
|
||||
|
||||
def has(self, cap: Capability) -> bool:
|
||||
return cap in self.capabilities
|
||||
|
||||
@classmethod
|
||||
def of(cls, *caps: Capability) -> CapabilityMatrix:
|
||||
return cls(frozenset(caps))
|
||||
|
||||
|
||||
@dataclass
|
||||
class SamplingDefaults:
|
||||
"""Per-model default generation params (the v2 mirror of fastvideo's per-model ``InferencePreset``
|
||||
defaults). Applied by the entrypoint when the caller didn't specify a value; ``None`` => fall back to
|
||||
the generic default. These are the user-facing knobs surfaced at request time."""
|
||||
num_steps: int | None = None
|
||||
guidance_scale: float | None = None
|
||||
guidance_per_modality: dict[str, float] = field(default_factory=dict) # joint A/V, e.g. {"video":3,"audio":7}
|
||||
height: int | None = None
|
||||
width: int | None = None
|
||||
num_frames: int | None = None
|
||||
fps: int | None = None
|
||||
negative_prompt: str | None = None
|
||||
shift: float | None = None
|
||||
sigmas: tuple[float, ...] | None = None
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# ModelCard — the (recipe, runtime) pair as one versioned, validatable object #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class CardValidationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelCard:
|
||||
model_id: str
|
||||
family: str
|
||||
components: dict[str, ComponentSpec] = field(default_factory=dict)
|
||||
loops: dict[str, LoopSpec] = field(default_factory=dict)
|
||||
capabilities: CapabilityMatrix = field(default_factory=CapabilityMatrix)
|
||||
recipe: RecipeSpec = field(default_factory=RecipeSpec)
|
||||
parity: ParitySpec = field(default_factory=ParitySpec)
|
||||
caches: dict[str, CacheContract] = field(default_factory=dict)
|
||||
parallelism: ParallelismContract = field(default_factory=ParallelismContract)
|
||||
precision: PrecisionContract = field(default_factory=PrecisionContract)
|
||||
checkpoint: CheckpointManifest = field(default_factory=CheckpointManifest)
|
||||
sampling_defaults: SamplingDefaults = field(default_factory=SamplingDefaults)
|
||||
# On a GPU box, keep this model's components' I/O on-device (torch tensors) instead of marshalling
|
||||
# numpy<->torch at every loop step — the latent stays resident for the whole denoise loop. Opt-in
|
||||
# per recipe: set True only when the model's loop+program are array-agnostic (see v2/platform/array_ns).
|
||||
device_io: bool = False
|
||||
|
||||
def validate(self) -> ModelCard:
|
||||
"""Strict enough to validate before any GPU touches it.
|
||||
|
||||
Returns self so it can be chained. Raises CardValidationError on any
|
||||
contract violation — the (recipe, runtime) binding is enforced here.
|
||||
"""
|
||||
errs: list[str] = []
|
||||
|
||||
# 1) recipe.assumes_loop must exist (the (recipe, runtime) binding)
|
||||
if self.recipe.assumes_loop and self.recipe.assumes_loop not in self.loops:
|
||||
errs.append(f"recipe.assumes_loop={self.recipe.assumes_loop!r} is not a declared loop "
|
||||
f"(have {sorted(self.loops)}) — a recipe cannot assume a loop the card does not run")
|
||||
|
||||
# 2) recipe.assumes_precision must be consistent with the precision contract
|
||||
if (self.recipe.assumes_precision and self.recipe.assumes_precision
|
||||
not in (self.precision.default_dtype, *self.precision.component_overrides.values())):
|
||||
errs.append(f"recipe.assumes_precision={self.recipe.assumes_precision!r} not present in precision contract")
|
||||
|
||||
# 3) every loop's shared_weight_components must be declared components
|
||||
for lid, loop in self.loops.items():
|
||||
for comp in loop.shared_weight_components:
|
||||
if comp not in self.components:
|
||||
errs.append(f"loop {lid!r} shares weight component {comp!r} which is not declared")
|
||||
for cc in loop.cache_policy:
|
||||
if cc not in self.caches:
|
||||
errs.append(f"loop {lid!r} references cache class {cc!r} not in card.caches")
|
||||
if loop.step_cost_model is None:
|
||||
errs.append(f"loop {lid!r} is missing the REQUIRED step_cost_model")
|
||||
|
||||
# 4) required_for / optional_for tasks must be disjoint per component
|
||||
for cid, comp_spec in self.components.items():
|
||||
overlap = comp_spec.required_for & comp_spec.optional_for
|
||||
if overlap:
|
||||
errs.append(f"component {cid!r} lists tasks {overlap} as both required and optional")
|
||||
|
||||
if errs:
|
||||
raise CardValidationError(f"ModelCard {self.model_id!r} failed validation:\n - " + "\n - ".join(errs))
|
||||
return self
|
||||
|
||||
def loops_sharing(self, component_id: str) -> list[str]:
|
||||
"""The loops that bind a given component — the MoT 'many loops, one instance' set."""
|
||||
return [lid for lid, lp in self.loops.items() if component_id in lp.shared_weight_components]
|
||||
@@ -0,0 +1,16 @@
|
||||
"""Deployment & fleet plane — DeploymentCard, our own LocalFleet, Dynamo adapter.
|
||||
|
||||
The engine exports a DeploymentCard; our LocalFleet routes over it (so we don't rely on Dynamo),
|
||||
and the DynamoWorkerAdapter exports the same card so Dynamo can front us too — one object, two
|
||||
consumers.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.deploy.card import DeploymentCard, HealthSchema, SLOSchema, build_deployment_card
|
||||
from v2.deploy.dynamo import DynamoWorkerAdapter, FakeDynamoRuntime
|
||||
from v2.deploy.fleet import LocalFleet, NoWorkerAvailable, Worker
|
||||
|
||||
__all__ = [
|
||||
"DeploymentCard", "HealthSchema", "SLOSchema", "build_deployment_card", "LocalFleet", "Worker", "NoWorkerAvailable",
|
||||
"DynamoWorkerAdapter", "FakeDynamoRuntime"
|
||||
]
|
||||
@@ -0,0 +1,77 @@
|
||||
"""DeploymentCard — what an engine exports to a fleet.
|
||||
|
||||
The engine exports a ``DeploymentCard`` and lets a fleet orchestrator route. The cost model is the
|
||||
SAME object the scheduler budgets with — one object, two consumers (the scheduler's budget and the
|
||||
fleet's routing input).
|
||||
|
||||
This is the contract our OWN fleet (``deploy/fleet.py``) consumes AND the Dynamo adapter
|
||||
(``deploy/dynamo.py``) exports — so we are frontable by Dynamo without depending on it.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from v2._enums import Capability
|
||||
from v2.card import CostModel
|
||||
|
||||
|
||||
@dataclass
|
||||
class HealthSchema:
|
||||
status: str = "healthy" # "healthy" | "draining" | "unhealthy"
|
||||
in_flight: int = 0
|
||||
queue_depth: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class SLOSchema:
|
||||
slo_class: str = "standard" # "latency" | "throughput" | "cost"
|
||||
max_concurrent: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class DeploymentCard:
|
||||
engine_id: str
|
||||
model_cards: list[str] = field(default_factory=list)
|
||||
capabilities: frozenset[Capability] = frozenset()
|
||||
role_pools: list = field(default_factory=list) # list[RolePoolSpec]
|
||||
supported_programs: list[str] = field(default_factory=list)
|
||||
cost_model: CostModel | None = None # the SAME cost model the scheduler uses
|
||||
health: HealthSchema = field(default_factory=HealthSchema)
|
||||
slo: SLOSchema = field(default_factory=SLOSchema)
|
||||
|
||||
def serves(self, model_id: str) -> bool:
|
||||
return model_id in self.model_cards
|
||||
|
||||
|
||||
def build_deployment_card(engine_id: str,
|
||||
model_cards: list,
|
||||
*,
|
||||
max_concurrent: int = 8,
|
||||
slo_class: str = "standard",
|
||||
role_pools: list | None = None,
|
||||
supported_programs: list[str] | None = None) -> DeploymentCard:
|
||||
"""Export a DeploymentCard from the model cards an engine serves.
|
||||
|
||||
Picks a representative ``step_cost_model`` so the fleet/Dynamo route on the SAME cost object the
|
||||
scheduler budgets with."""
|
||||
caps: set = set()
|
||||
cost_model = None
|
||||
ids: list[str] = []
|
||||
for c in model_cards:
|
||||
ids.append(c.model_id)
|
||||
caps |= set(c.capabilities.capabilities)
|
||||
if cost_model is None:
|
||||
for lp in c.loops.values():
|
||||
if lp.step_cost_model is not None:
|
||||
# own copy so online calibration of one card's cost doesn't alias another replica's
|
||||
cost_model = dataclasses.replace(lp.step_cost_model,
|
||||
coefficients=dict(lp.step_cost_model.coefficients))
|
||||
break
|
||||
return DeploymentCard(engine_id=engine_id,
|
||||
model_cards=ids,
|
||||
capabilities=frozenset(caps),
|
||||
role_pools=role_pools or [],
|
||||
supported_programs=supported_programs or [],
|
||||
cost_model=cost_model,
|
||||
slo=SLOSchema(slo_class=slo_class, max_concurrent=max_concurrent))
|
||||
@@ -0,0 +1,95 @@
|
||||
"""Dynamo worker adapter — Dynamo as an *option*, not a dependency.
|
||||
|
||||
NVIDIA Dynamo is the named first-class partner for the fleet layer; the engine's job is to be a good
|
||||
Dynamo citizen: registration, health/drain, cost metrics, affinity events.
|
||||
|
||||
This adapter exposes exactly that contract over an AsyncEngine + DeploymentCard, so Dynamo CAN front
|
||||
this engine. It consumes the same DeploymentCard + cost model as our own LocalFleet — one object, two
|
||||
consumers — so choosing Dynamo vs. our fleet is a deployment decision, not a rewrite.
|
||||
``FakeDynamoRuntime`` proves the contract is satisfiable end-to-end without importing Dynamo.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.request.artifacts import Output
|
||||
from v2.deploy.card import DeploymentCard
|
||||
|
||||
|
||||
class DynamoWorkerAdapter:
|
||||
"""Implements the Dynamo worker surface: registration, metrics, affinity, handle."""
|
||||
|
||||
def __init__(self, engine: Any, card: DeploymentCard, *, worker_type: str = "Aggregated"):
|
||||
self.engine = engine
|
||||
self.card = card
|
||||
self.worker_type = worker_type
|
||||
self.registered = False
|
||||
self.draining = False
|
||||
|
||||
# 1) worker surface — registration payload (ModelType.Videos|Images, roles, endpoints)
|
||||
def registration(self) -> dict[str, Any]:
|
||||
self.registered = True
|
||||
return {
|
||||
"engine_id":
|
||||
self.card.engine_id,
|
||||
"model_type": ["Videos", "Images"] + (["Chat"] if any("omni" in m or "cosmos" in m or "bagel" in m
|
||||
for m in self.card.model_cards) else []),
|
||||
"worker_type":
|
||||
self.worker_type,
|
||||
"models":
|
||||
list(self.card.model_cards),
|
||||
"capabilities":
|
||||
sorted(c.value for c in self.card.capabilities),
|
||||
"supported_programs":
|
||||
list(self.card.supported_programs),
|
||||
}
|
||||
|
||||
# 2) metrics for routing + the SLA Planner (the SAME cost model the scheduler uses)
|
||||
def metrics(self) -> dict[str, Any]:
|
||||
return {"in_flight": self.engine.in_flight, "queue_depth": self.engine.queue_depth, "draining": self.draining}
|
||||
|
||||
def cost_estimate(self, request: Any) -> float:
|
||||
cm = self.card.cost_model
|
||||
steps = max(1, int(getattr(request.diffusion, "num_steps", 1) or 1))
|
||||
work = max(1, int(getattr(request.diffusion, "height", 1)) * int(getattr(request.diffusion, "width", 1)))
|
||||
return steps * (cm.predict(work) if cm is not None else 1e-3)
|
||||
|
||||
# health / graceful drain wired to the engine
|
||||
def health(self) -> dict[str, Any]:
|
||||
return {"status": "draining" if self.draining else "healthy", **self.metrics()}
|
||||
|
||||
def drain(self) -> None:
|
||||
self.draining = True
|
||||
|
||||
# 3) affinity / cache events (KvCacheEventData shape) — checkpoint/session residency
|
||||
def cache_event(self, kind: str, key: str) -> dict[str, Any]:
|
||||
return {"event": kind, "engine_id": self.card.engine_id, "key": key}
|
||||
|
||||
# the worker entrypoint Dynamo's router calls
|
||||
async def handle(self, request: Any) -> Output:
|
||||
if self.draining:
|
||||
raise RuntimeError(f"worker {self.card.engine_id} is draining")
|
||||
return await self.engine.generate(request)
|
||||
|
||||
|
||||
class FakeDynamoRuntime:
|
||||
"""A minimal stand-in for Dynamo: registers worker adapters and routes by the published cost model
|
||||
+ health. Demonstrates the engine satisfies the Dynamo contract WITHOUT a Dynamo dependency."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.workers: list[DynamoWorkerAdapter] = []
|
||||
self.registry: list[dict] = []
|
||||
|
||||
def register_worker(self, adapter: DynamoWorkerAdapter) -> None:
|
||||
self.registry.append(adapter.registration())
|
||||
self.workers.append(adapter)
|
||||
|
||||
def _route(self, request: Any) -> DynamoWorkerAdapter:
|
||||
cands = [w for w in self.workers if not w.draining and request.model_id in w.card.model_cards]
|
||||
if not cands:
|
||||
raise RuntimeError(f"no Dynamo worker serves {request.model_id!r}")
|
||||
# the SLA-planner-style choice: cheapest predicted cost, tie-broken by least in-flight
|
||||
return min(cands, key=lambda w: (w.cost_estimate(request), w.engine.in_flight))
|
||||
|
||||
async def generate(self, request: Any) -> Output:
|
||||
return await self._route(request).handle(request)
|
||||
@@ -0,0 +1,110 @@
|
||||
"""LocalFleet — OUR OWN fleet router.
|
||||
|
||||
Dynamo is the first-class fleet partner, but every Dynamo ask has a first-class fallback — and this
|
||||
is it: a self-contained fleet that does discovery, health/drain, and routing (least-loaded /
|
||||
cost-model / affinity) over multiple engine workers, so we are never *reliant* on Dynamo. The
|
||||
router's cost input is the SAME cost model the scheduler uses (one object, two consumers). Affinity
|
||||
routing is sticky-by-key for checkpoint/session residency (least-loaded + engine redirects).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from collections.abc import AsyncIterator
|
||||
|
||||
from v2.request.artifacts import Output
|
||||
from v2.deploy.card import DeploymentCard, HealthSchema
|
||||
|
||||
|
||||
class NoWorkerAvailable(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class Worker:
|
||||
"""A registered engine worker (an AsyncEngine + its exported DeploymentCard)."""
|
||||
|
||||
def __init__(self, worker_id: str, engine: Any, card: DeploymentCard):
|
||||
self.worker_id = worker_id
|
||||
self.engine = engine
|
||||
self.card = card
|
||||
self.draining = False
|
||||
|
||||
@property
|
||||
def load(self) -> float:
|
||||
return self.engine.in_flight / max(1, self.card.slo.max_concurrent)
|
||||
|
||||
@property
|
||||
def healthy(self) -> bool:
|
||||
return not self.draining and self.engine.in_flight < self.card.slo.max_concurrent * 4
|
||||
|
||||
def serves(self, model_id: str) -> bool:
|
||||
return self.engine.serves(model_id) or self.card.serves(model_id)
|
||||
|
||||
def cost_estimate(self, request: Any) -> float:
|
||||
"""Predicted GPU-time for this request — the cost model as the fleet's routing input."""
|
||||
cm = self.card.cost_model
|
||||
steps = max(1, int(getattr(request.diffusion, "num_steps", 1) or 1))
|
||||
work = max(1, int(getattr(request.diffusion, "height", 1)) * int(getattr(request.diffusion, "width", 1)))
|
||||
per = cm.predict(work) if cm is not None else 1e-3
|
||||
return steps * per
|
||||
|
||||
def health(self) -> HealthSchema:
|
||||
return HealthSchema(status=("draining" if self.draining else "healthy"),
|
||||
in_flight=self.engine.in_flight,
|
||||
queue_depth=self.engine.queue_depth)
|
||||
|
||||
|
||||
class LocalFleet:
|
||||
|
||||
def __init__(self, policy: str = "least_loaded", *, max_affinity: int = 100_000):
|
||||
assert policy in ("least_loaded", "cost", "affinity")
|
||||
self.policy = policy
|
||||
self.max_affinity = max_affinity
|
||||
self.workers: dict[str, Worker] = {}
|
||||
self._affinity: dict[str, str] = {} # affinity key -> worker_id (sticky), FIFO-bounded
|
||||
|
||||
# --- discovery / health (what Dynamo's registry + planner would do) ------ #
|
||||
def register(self, worker_id: str, engine: Any, card: DeploymentCard) -> Worker:
|
||||
w = Worker(worker_id, engine, card)
|
||||
self.workers[worker_id] = w
|
||||
return w
|
||||
|
||||
def deregister(self, worker_id: str) -> None:
|
||||
self.workers.pop(worker_id, None)
|
||||
|
||||
def drain(self, worker_id: str) -> None:
|
||||
if worker_id in self.workers:
|
||||
self.workers[worker_id].draining = True
|
||||
|
||||
def health(self) -> dict[str, HealthSchema]:
|
||||
return {wid: w.health() for wid, w in self.workers.items()}
|
||||
|
||||
# --- routing (least-loaded / cost / affinity) ---------------------------- #
|
||||
def _candidates(self, request: Any) -> list[Worker]:
|
||||
return [w for w in self.workers.values() if w.serves(request.model_id) and w.healthy]
|
||||
|
||||
def route(self, request: Any, *, affinity_key: str | None = None) -> Worker:
|
||||
cands = self._candidates(request)
|
||||
if not cands:
|
||||
raise NoWorkerAvailable(f"no healthy worker serves model {request.model_id!r}")
|
||||
if self.policy == "affinity":
|
||||
key = affinity_key or request.model_id
|
||||
wid = self._affinity.get(key)
|
||||
if wid in self.workers and self.workers[wid] in cands:
|
||||
return self.workers[wid]
|
||||
chosen = min(cands, key=lambda w: w.load) # cold key → least-loaded, then pin
|
||||
if len(self._affinity) >= self.max_affinity: # FIFO-bound the sticky map
|
||||
self._affinity.pop(next(iter(self._affinity)), None)
|
||||
self._affinity[key] = chosen.worker_id
|
||||
return chosen
|
||||
if self.policy == "cost":
|
||||
return min(cands, key=lambda w: w.cost_estimate(request) * (1.0 + w.load))
|
||||
return min(cands, key=lambda w: w.load) # least_loaded (default)
|
||||
|
||||
# --- serving (delegates to the chosen worker's engine) ------------------- #
|
||||
async def generate(self, request: Any, *, affinity_key: str | None = None) -> Output:
|
||||
return await self.route(request, affinity_key=affinity_key).engine.generate(request)
|
||||
|
||||
async def submit(self, request: Any, *, affinity_key: str | None = None) -> AsyncIterator:
|
||||
worker = self.route(request, affinity_key=affinity_key)
|
||||
async for ev in worker.engine.submit(request):
|
||||
yield ev
|
||||
@@ -0,0 +1,4 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""``distributed`` facade — single-GPU dist-init the loaders need (1x1 device mesh). Re-exported so v2
|
||||
code imports ``v2.distributed``; a vendored cutover copies parallel_state getters + communication_op."""
|
||||
from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel # noqa: F401
|
||||
+166
@@ -0,0 +1,166 @@
|
||||
"""Worked examples, runnable as a demo:
|
||||
|
||||
python3 -m v2.examples
|
||||
|
||||
Each prints what it demonstrates. This doubles as living documentation of the API.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.recipes import build_default_engine, build_omni_engine
|
||||
from v2.recipes.wan21 import build_wan21_card
|
||||
from v2.parity import assert_interleave_parity
|
||||
from v2.request import DiffusionParams, OutputSpec, Request, SamplingParams, TaskType, make_request
|
||||
from v2.training import build_diffusion_nft
|
||||
|
||||
|
||||
def _t2v(mid: str, prompt: str, seed: int, steps: int = 4, **kw: Any) -> Request:
|
||||
return make_request(TaskType.T2V, mid, prompt, diffusion=DiffusionParams(num_steps=steps, seed=seed), **kw)
|
||||
|
||||
|
||||
def example_a_text_to_video(eng) -> None:
|
||||
print("\n(a) Text → video, one instance (Wan2.1-1.3B)")
|
||||
out = eng.run(_t2v("wan2.1-1.3b", "a cat surfing a wave", 7))
|
||||
print(f" video {out.artifacts['video'].frames.shape} "
|
||||
f"denoise_steps={out.metrics['denoise_steps']:.0f} gpu_s={out.metrics['gpu_seconds']:.2e}")
|
||||
|
||||
|
||||
def example_b_ltx2_two_stage(eng) -> None:
|
||||
print("\n(b) LTX-2 two-stage distilled (base 8-step → upsample → refine 3-step), shared transformer")
|
||||
out = eng.run(_t2v("ltx2-2stage-distilled", "a neon city at night", 2))
|
||||
print(f" video {out.artifacts['video'].frames.shape} "
|
||||
f"base={out.metrics['base_steps']:.0f} refine={out.metrics['refine_steps']:.0f}")
|
||||
|
||||
|
||||
def example_c_causal_streaming(eng) -> None:
|
||||
print("\n(c) Causal streaming (Wan-causal): chunk rollout + slab-KV, streamable by chunk")
|
||||
out = eng.run(
|
||||
_t2v("wan-causal-sf-1.3b", "a drone flight over mountains", 3, outputs=OutputSpec(stream={"video": True})))
|
||||
print(f" latents {out.artifacts['latents'].latent.shape} chunks={out.metrics['chunks']:.0f} "
|
||||
f"streamed_chunks={out.metrics.get('stream_chunks', 0)}")
|
||||
|
||||
|
||||
def example_c2_interleave_gate(eng) -> None:
|
||||
print("\n(c2) Interleave parity gate — serial == interleaved, bit-identical (the §9.3 obligation)")
|
||||
reqs = [_t2v("wan2.1-1.3b", "alpha", 11), _t2v("wan2.1-1.3b", "beta", 22)]
|
||||
divs = assert_interleave_parity(eng, reqs)
|
||||
print(f" divergences: {divs or 'NONE — gate PASSES ✓'}")
|
||||
|
||||
|
||||
def example_d_rl_rollout() -> None:
|
||||
print("\n(d) RL rollout (DiffusionNFT): the SAME denoise loop + behavior capture (train ≡ serve)")
|
||||
nft = build_diffusion_nft(build_wan21_card(), num_video_per_prompt=4, num_inner_timesteps=2)
|
||||
loss, m = nft.managed_train_step({"prompts": ["a red car", "a blue boat"], "seeds": [1, 2]}, 0)
|
||||
fc = nft.old.caches.stats()["feature"]
|
||||
print(f" policy_loss={loss['policy_loss']:.3f} kl={loss['kl_div_loss']:.5f} "
|
||||
f"reward_mean={m['reward_mean']:.3f} consistency={nft.consistency_level().value} (likelihood-free)")
|
||||
print(f" shared-prompt feature-cache reuse: {fc['hits']} hits / {fc['misses']} misses "
|
||||
f"(K samples encode the prompt once — the 24× reduction)")
|
||||
|
||||
|
||||
def example_g_omni_mot() -> None:
|
||||
print("\n(g) Omni / MoT (§16): ONE resident instance runs AR + diffusion loops on shared weights")
|
||||
eng = build_omni_engine()
|
||||
o = eng.run(
|
||||
make_request(TaskType.T2V,
|
||||
"cosmos3-vfm",
|
||||
"a phoenix",
|
||||
sampling=SamplingParams(max_tokens=6, seed=1),
|
||||
diffusion=DiffusionParams(num_steps=4, seed=1)))
|
||||
print(f" Cosmos3 (reason→joint denoise): text={o.artifacts['text'].text} "
|
||||
f"video={o.artifacts['video'].frames.shape}")
|
||||
o2 = eng.run(
|
||||
make_request(TaskType.T2I,
|
||||
"bagel-mot",
|
||||
"a teapot",
|
||||
sampling=SamplingParams(max_tokens=6, seed=2),
|
||||
diffusion=DiffusionParams(num_steps=4, seed=2)))
|
||||
print(f" BAGEL (generate_text→generate_image): text={o2.artifacts['text'].text} "
|
||||
f"image={o2.artifacts['image'].tensor.shape}")
|
||||
print(f" scheduler priced BOTH WorkUnit kinds (runtime-visible, not one opaque stage): "
|
||||
f"{dict(eng.admission.metrics.by_kind)}")
|
||||
|
||||
|
||||
async def _serving_demo() -> None:
|
||||
import asyncio
|
||||
|
||||
from v2.deploy import DynamoWorkerAdapter, FakeDynamoRuntime, LocalFleet, build_deployment_card
|
||||
from v2.recipes.wan21 import build_wan21_card, build_wan_t2v_program
|
||||
from v2.runtime import AsyncEngine, PoolSet, wan_t2v_disaggregated
|
||||
from v2.serving import OmniOpenAIServer
|
||||
|
||||
eng = build_default_engine()
|
||||
build_omni_engine(eng)
|
||||
ae = AsyncEngine(eng)
|
||||
|
||||
# disaggregated pools: encoder → denoiser → decoder
|
||||
card = build_wan21_card()
|
||||
pools = PoolSet(wan_t2v_disaggregated(), card)
|
||||
pools.warmup()
|
||||
ae.register_disaggregated("wan-disagg", pools, build_wan_t2v_program())
|
||||
out = await ae.generate(
|
||||
make_request(TaskType.T2V, "wan-disagg", "a wave", diffusion=DiffusionParams(num_steps=4, seed=1)))
|
||||
print(f" disaggregated T2V (enc→den→dec): video={out.artifacts['video'].frames.shape} "
|
||||
f"cross-pool transfers={out.metrics['transfers']:.0f}")
|
||||
|
||||
# our own OpenAI server over a real socket
|
||||
server = OmniOpenAIServer(ae, engine_id="worker-0")
|
||||
host, port = await server.serve(port=0)
|
||||
|
||||
async def http(method: str, path: str, body: bytes = b"") -> str:
|
||||
r, w = await asyncio.open_connection(host, port)
|
||||
w.write(f"{method} {path} HTTP/1.1\r\nHost: x\r\nContent-Length: {len(body)}\r\n\r\n".encode() + body)
|
||||
await w.drain()
|
||||
data = await r.read()
|
||||
w.close()
|
||||
return data.decode("utf-8", "replace")
|
||||
|
||||
health = await http("GET", "/health")
|
||||
sse = await http("POST", "/v1/chat/completions",
|
||||
b'{"model":"cosmos3-vfm","messages":[{"role":"user","content":"a comet"}],"stream":true}')
|
||||
n_chunks = sse.count("data: ")
|
||||
print(
|
||||
f" OpenAI server: /health ok={'healthy' in health}; chat SSE streamed {n_chunks} chunks (omni reason→denoise)"
|
||||
)
|
||||
await server.close()
|
||||
|
||||
# our own fleet router + Dynamo adapter (frontable, not relied upon) — same DeploymentCard
|
||||
dcard = build_deployment_card("worker-0", [card])
|
||||
fleet = LocalFleet("least_loaded")
|
||||
fleet.register("worker-0", ae, dcard)
|
||||
routed = fleet.route(make_request(TaskType.T2V, "wan2.1-1.3b", "x", diffusion=DiffusionParams(num_steps=2)))
|
||||
dyn = FakeDynamoRuntime()
|
||||
dyn.register_worker(DynamoWorkerAdapter(ae, dcard))
|
||||
print(f" LocalFleet routes to '{routed.worker_id}'; Dynamo adapter registered "
|
||||
f"({len(dyn.registry)} worker) — both consume the SAME DeploymentCard")
|
||||
|
||||
|
||||
def example_h_serving_and_fleet() -> None:
|
||||
import asyncio
|
||||
print("\n(h) Serving + fleet (OUR OWN — Dynamo-optional): async engine, role pools, OpenAI server")
|
||||
asyncio.run(_serving_demo())
|
||||
|
||||
|
||||
def main() -> None:
|
||||
print("=" * 78)
|
||||
print("v2 — worked examples. One runtime, many loops.")
|
||||
print("=" * 78)
|
||||
eng = build_default_engine()
|
||||
print(f"registered (recipe, runtime) cards: {list(eng._registry)}")
|
||||
example_a_text_to_video(eng)
|
||||
example_b_ltx2_two_stage(eng)
|
||||
example_c_causal_streaming(eng)
|
||||
example_c2_interleave_gate(eng)
|
||||
example_d_rl_rollout()
|
||||
example_g_omni_mot()
|
||||
example_h_serving_and_fleet()
|
||||
print("\n" + "=" * 78)
|
||||
print("All examples ran on CPU with numpy toy components. The architecture (cards, driven")
|
||||
print("loops, scheduler, caches, parity, training-on-shared-loops) is real; the neural")
|
||||
print("forwards are toys. On a GPU box, swap ComponentSpec.factory for the torch adapters.")
|
||||
print("=" * 78)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,12 @@
|
||||
"""Extension plane — observers (read-only) and interceptors (compute-altering)."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.extend.base import Interceptor, InterceptorChain, InterceptorConflict, Observer, ObserverBus
|
||||
from v2.extend.interceptors import ResidualSkipInterceptor
|
||||
from v2.extend.observers import NaNWatch, Profiler
|
||||
from v2.extend.registry import PluginRegistry
|
||||
|
||||
__all__ = [
|
||||
"Observer", "Interceptor", "ObserverBus", "InterceptorChain", "InterceptorConflict", "Profiler", "NaNWatch",
|
||||
"ResidualSkipInterceptor", "PluginRegistry"
|
||||
]
|
||||
@@ -0,0 +1,92 @@
|
||||
"""Observers and interceptors — the optimization/debug/parity surface.
|
||||
|
||||
They compose with the loop cleanly: the hooks wrap ``ctx.execute(plan)``.
|
||||
* Observers (read-only): cannot mutate state. An unused hook is *literally absent* from
|
||||
the hot path (the bus only iterates when observers are attached).
|
||||
* Interceptors (compute-altering): ``before_step`` may supply a cached prediction (skip);
|
||||
``after_step`` updates calibration state. State lives in ``LoopState.plugin_state[id]``,
|
||||
keyed per request AND per CFG branch — the structural fix for module-global residual
|
||||
state that corrupts cache-dit/TeaCache forks under concurrency.
|
||||
|
||||
Trust boundary: plugins are enabled at deploy scope only (registry), never via a per-request
|
||||
``plugins=[...]`` field. Requests only *parameterize* pre-enabled plugins.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Observer(Protocol):
|
||||
|
||||
def observe(self, event: str, **kw) -> None:
|
||||
...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Interceptor(Protocol):
|
||||
plugin_id: str
|
||||
distribution_altering: bool
|
||||
graph_safe: bool
|
||||
|
||||
def before_step(self, plan: Any, state: Any) -> Any | None:
|
||||
... # return override output to skip, or None
|
||||
|
||||
def after_step(self, plan: Any, state: Any, result: Any) -> None:
|
||||
...
|
||||
|
||||
|
||||
class ObserverBus:
|
||||
"""Read-only event fan-out. Cheap when empty (the absent-hook rule)."""
|
||||
|
||||
def __init__(self, observers: list[Observer] | None = None):
|
||||
self._observers = list(observers or [])
|
||||
|
||||
def add(self, observer: Observer) -> None:
|
||||
self._observers.append(observer)
|
||||
|
||||
@property
|
||||
def active(self) -> bool:
|
||||
return bool(self._observers)
|
||||
|
||||
def emit(self, event: str, **kw) -> None:
|
||||
if not self._observers: # absent from the hot path when off
|
||||
return
|
||||
for obs in self._observers:
|
||||
obs.observe(event, **kw)
|
||||
|
||||
|
||||
class InterceptorConflict(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
class InterceptorChain:
|
||||
"""Ordered interceptor chain; conflicting interceptors rejected pre-flight."""
|
||||
|
||||
def __init__(self, interceptors: list[Interceptor] | None = None, *, exact_mode: bool = False):
|
||||
self._chain = list(interceptors or [])
|
||||
self._validate(exact_mode)
|
||||
|
||||
def _validate(self, exact_mode: bool) -> None:
|
||||
skippers = [i for i in self._chain if getattr(i, "distribution_altering", False)]
|
||||
if len(skippers) > 1:
|
||||
raise InterceptorConflict(f"multiple distribution-altering interceptors {[i.plugin_id for i in skippers]} "
|
||||
"conflict — only one step-skipper allowed")
|
||||
if exact_mode and skippers:
|
||||
raise InterceptorConflict(
|
||||
f"exact-mode rejects distribution_altering interceptors {[i.plugin_id for i in skippers]}")
|
||||
|
||||
@property
|
||||
def active(self) -> bool:
|
||||
return bool(self._chain)
|
||||
|
||||
def before(self, plan: Any, state: Any) -> Any | None:
|
||||
for i in self._chain:
|
||||
override = i.before_step(plan, state)
|
||||
if override is not None:
|
||||
return override
|
||||
return None
|
||||
|
||||
def after(self, plan: Any, state: Any, result: Any) -> None:
|
||||
for i in self._chain:
|
||||
i.after_step(plan, state, result)
|
||||
@@ -0,0 +1,53 @@
|
||||
"""cache-dit-style interceptors.
|
||||
|
||||
``ResidualSkipInterceptor`` is the reference step-skip integration. The load-bearing
|
||||
correctness property: its per-step state lives in ``LoopState.plugin_state[id][branch]``,
|
||||
keyed per request AND per CFG branch — NOT a module global. This is exactly why the
|
||||
interleave gate passes with it on: two interleaved requests have disjoint plugin
|
||||
state, so neither smears the other's cached prediction.
|
||||
|
||||
A 4-step distilled card *rejects* this interceptor (capability negotiation) rather than
|
||||
producing garbage — handled by InterceptorChain validation + per-card opt-in.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
|
||||
class ResidualSkipInterceptor:
|
||||
"""Skip the model forward on cadence, reusing the previous step's prediction.
|
||||
|
||||
A deliberately simple stand-in for DBCache/FBCache/TaylorSeer: every ``interval``-th
|
||||
step is recomputed; intermediate steps reuse the cached output. Demonstrates the
|
||||
contract, not the algorithm.
|
||||
"""
|
||||
plugin_id = "residual_skip"
|
||||
distribution_altering = True
|
||||
graph_safe = False
|
||||
|
||||
def __init__(self, interval: int = 2):
|
||||
self.interval = max(2, interval)
|
||||
|
||||
def _branch_state(self, state: Any, branch: str) -> dict:
|
||||
ns = state.plugin_state.setdefault(self.plugin_id, {})
|
||||
return ns.setdefault(branch, {})
|
||||
|
||||
def before_step(self, plan: Any, state: Any) -> Any | None:
|
||||
"""Return a cached forward output to skip the model forward on non-cadence steps.
|
||||
|
||||
The step body still runs the cheap solver step with this prediction, so only the
|
||||
expensive forward is skipped. Cache lives per request (LoopState.plugin_state) AND
|
||||
per branch — interleaving two requests cannot smear caches."""
|
||||
branch = str(plan.payload.get("branch", "combined"))
|
||||
bs = self._branch_state(state, branch)
|
||||
cached = bs.get("last_output")
|
||||
if cached is not None and (state.step_idx % self.interval) != 0:
|
||||
bs["skipped"] = bs.get("skipped", 0) + 1
|
||||
return {"noise_pred": cached}
|
||||
return None
|
||||
|
||||
def after_step(self, plan: Any, state: Any, result: Any) -> None:
|
||||
branch = str(plan.payload.get("branch", "combined"))
|
||||
pred = result.output.get("noise_pred")
|
||||
if pred is not None:
|
||||
self._branch_state(state, branch)["last_output"] = pred
|
||||
@@ -0,0 +1,61 @@
|
||||
"""Built-in read-only observers: Profiler, NaNWatch.
|
||||
|
||||
ParityAligner lives in its own ``parity/`` package but also implements the Observer ``observe``
|
||||
protocol so it attaches to the same bus.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
class Profiler:
|
||||
"""Per-step wall/CUDA timing that calibrates the cost model.
|
||||
|
||||
Accumulates (work_units, actual_seconds) samples and fits a card's CostModel
|
||||
coefficients — the online calibration that refines the conservative baseline.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.samples: list[tuple[str, float, float]] = [] # (batch_key_repr, work_units, seconds)
|
||||
|
||||
def observe(self, event: str, **kw) -> None:
|
||||
if event == "step_complete":
|
||||
plan, result = kw.get("plan"), kw.get("result")
|
||||
if plan is not None and result is not None:
|
||||
self.samples.append(
|
||||
(repr(plan.shape_sig.batch_key), float(plan.shape_sig.work_units), float(result.actual_seconds)))
|
||||
|
||||
def calibrate(self, cost_model) -> None:
|
||||
"""Fit base + per_unit seconds from observed samples (least squares)."""
|
||||
if len(self.samples) < 2:
|
||||
return
|
||||
x = np.array([s[1] for s in self.samples], dtype=np.float64)
|
||||
y = np.array([s[2] for s in self.samples], dtype=np.float64)
|
||||
if np.ptp(x) == 0:
|
||||
cost_model.base_seconds = float(y.mean())
|
||||
return
|
||||
slope, intercept = np.polyfit(x, y, 1)
|
||||
cost_model.per_unit_seconds = float(max(slope, 0.0))
|
||||
cost_model.base_seconds = float(max(intercept, 0.0))
|
||||
|
||||
|
||||
class NaNWatch:
|
||||
"""First-NaN/Inf localization. Request-fatal, triggering an SPMD-consistent abort."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.first: tuple[str, str] | None = None # (tap/output name, plan label)
|
||||
|
||||
def observe(self, event: str, **kw) -> None:
|
||||
if event == "step_complete" and self.first is None:
|
||||
plan, result = kw.get("plan"), kw.get("result")
|
||||
if result is None:
|
||||
return
|
||||
for name, val in result.output.items():
|
||||
arr = np.asarray(val) if hasattr(val, "__array__") or isinstance(val, list | tuple) else None
|
||||
if arr is not None and arr.dtype.kind == "f" and not np.all(np.isfinite(arr)):
|
||||
self.first = (name, getattr(plan, "label", "") or name)
|
||||
return
|
||||
|
||||
@property
|
||||
def tripped(self) -> bool:
|
||||
return self.first is not None
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Plugin registry + trust boundary.
|
||||
|
||||
Plugins are enabled at deploy scope only (never a per-request ``plugins=[...]`` field that
|
||||
would wire third-party code selection into the public API); requests only *parameterize*
|
||||
pre-enabled plugins through validated schemas.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
|
||||
from v2.extend.base import InterceptorChain, Observer
|
||||
|
||||
|
||||
class PluginRegistry:
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._interceptors: dict[str, Callable[..., Any]] = {}
|
||||
self._observers: dict[str, Callable[..., Observer]] = {}
|
||||
self._enabled: set[str] = set() # deploy-scope enablement
|
||||
|
||||
def register_interceptor(self, plugin_id: str, factory: Callable[..., Any]) -> None:
|
||||
self._interceptors[plugin_id] = factory
|
||||
|
||||
def register_observer(self, plugin_id: str, factory: Callable[..., Observer]) -> None:
|
||||
self._observers[plugin_id] = factory
|
||||
|
||||
def enable(self, plugin_id: str) -> None:
|
||||
if plugin_id not in self._interceptors and plugin_id not in self._observers:
|
||||
raise KeyError(f"plugin {plugin_id!r} not registered")
|
||||
self._enabled.add(plugin_id)
|
||||
|
||||
def build_chain(self, params: dict[str, dict] | None = None, *, exact_mode: bool = False) -> InterceptorChain:
|
||||
params = params or {}
|
||||
chain = []
|
||||
for pid in self._enabled:
|
||||
if pid in self._interceptors:
|
||||
chain.append(self._interceptors[pid](**params.get(pid, {})))
|
||||
return InterceptorChain(chain, exact_mode=exact_mode)
|
||||
@@ -0,0 +1,4 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""``FastVideoArgs`` facade — the args object the component loaders consume. Re-exported so v2 code imports
|
||||
``v2.fastvideo_args``; a vendored cutover slims it to inference-only args."""
|
||||
from fastvideo.fastvideo_args import FastVideoArgs # noqa: F401
|
||||
@@ -0,0 +1,6 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""``forward_context`` facade. The fastvideo attention layer reads a thread-local ForwardContext via
|
||||
``get_forward_context()`` (asserts it's set), so every torch forward in the backend runs inside
|
||||
``set_forward_context(...)``. Re-exported here so v2 code imports ``v2.forward_context``; a vendored
|
||||
cutover (which also requires forking the attention layer to drop the global context) replaces this body."""
|
||||
from fastvideo.forward_context import get_forward_context, set_forward_context # noqa: F401
|
||||
@@ -0,0 +1,6 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""``loader`` facade — the v2-owned component-loading seam. ``component_loader`` exposes the loader
|
||||
classes (TransformerLoader / VAELoader / TextEncoderLoader / TokenizerLoader / UpsamplerLoader /
|
||||
AudioDecoderLoader / VocoderLoader) that build real modules from a checkpoint. Re-exported so v2 code
|
||||
imports ``v2.loader``; a vendored cutover replaces this with a slimmed v2-native loader (no caller change)."""
|
||||
from fastvideo.models.loader import component_loader # noqa: F401
|
||||
@@ -0,0 +1,44 @@
|
||||
"""Loop plane — driven loops + policies."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import LoopKind, WorkUnitKind
|
||||
from v2.loop.contracts import (
|
||||
CacheOp,
|
||||
CachePlan,
|
||||
Done,
|
||||
Loop,
|
||||
LoopContext,
|
||||
LoopResult,
|
||||
LoopState,
|
||||
PlacementHint,
|
||||
ResourceRequest,
|
||||
ShapeSignature,
|
||||
StepContext,
|
||||
StepResult,
|
||||
WorkPlan,
|
||||
)
|
||||
from v2.loop.driver import LoopRunner
|
||||
from v2.loop.sampler import add_noise, build_flow_sigmas, flow_match_euler_step, x0_from_velocity
|
||||
|
||||
__all__ = [
|
||||
"Loop",
|
||||
"LoopContext",
|
||||
"LoopState",
|
||||
"LoopResult",
|
||||
"WorkPlan",
|
||||
"Done",
|
||||
"StepResult",
|
||||
"StepContext",
|
||||
"ShapeSignature",
|
||||
"ResourceRequest",
|
||||
"CachePlan",
|
||||
"CacheOp",
|
||||
"PlacementHint",
|
||||
"LoopRunner",
|
||||
"LoopKind",
|
||||
"WorkUnitKind",
|
||||
"build_flow_sigmas",
|
||||
"flow_match_euler_step",
|
||||
"x0_from_velocity",
|
||||
"add_noise",
|
||||
]
|
||||
@@ -0,0 +1,233 @@
|
||||
"""The driven-loop contract — a serializable state machine the runtime drives:
|
||||
|
||||
state = loop.init(req, model_state, ctx)
|
||||
while True:
|
||||
plan = loop.next(state) # describe next step; NO GPU kernels
|
||||
if isinstance(plan, Done): break
|
||||
result = ctx.execute(plan) # ← THE INVERSION POINT (runtime owns it)
|
||||
state = loop.advance(state, result) # fold result in; decide what's next
|
||||
for chunk in plan.emits: ctx.emit(chunk)
|
||||
return loop.finalize(state)
|
||||
|
||||
Two properties this contract buys:
|
||||
* content-adaptive steps are natural — ``next`` reads ``state``, which already folded
|
||||
in the previous ``StepResult`` via ``advance``;
|
||||
* cross-request state safety is *structural* — all per-request mutable state lives in
|
||||
``LoopState`` (incl. ``plugin_state`` per request/CFG-branch), never module globals,
|
||||
so interleaving requests through one ModelInstance cannot smear state.
|
||||
|
||||
This module is pure stdlib (tensors typed as ``TensorLike``) so it imports no backend.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
from v2._enums import ExecutionProfile, WorkUnitKind
|
||||
from v2._types import TensorLike
|
||||
from v2.request.streams import StreamChunk
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Shape / resources / cache plan / placement #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass(frozen=True)
|
||||
class ShapeSignature:
|
||||
"""Batch-compatibility + graph-capture key."""
|
||||
kind: WorkUnitKind
|
||||
dims: tuple[int, ...] = ()
|
||||
dtype: str = "float32"
|
||||
extra: tuple[tuple[str, Any], ...] = () # e.g. (("cfg","classic"),("expert","e0"))
|
||||
|
||||
@property
|
||||
def work_units(self) -> int:
|
||||
u = 1
|
||||
for d in self.dims:
|
||||
u *= max(int(d), 1)
|
||||
return u
|
||||
|
||||
@property
|
||||
def batch_key(self) -> tuple:
|
||||
"""Two WorkPlans batch together iff their batch_keys are equal."""
|
||||
return (self.kind, self.dims, self.dtype, self.extra)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ResourceRequest:
|
||||
"""Everything admission must reserve."""
|
||||
compute_seconds: float = 0.0
|
||||
resident_bytes: int = 0
|
||||
peak_activation_bytes: int = 0
|
||||
cache_blocks: dict[str, int] = field(default_factory=dict)
|
||||
transfer_bytes: int = 0
|
||||
graph_capture_size: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class CacheOp:
|
||||
cache_class: str
|
||||
key: Any # CacheKey
|
||||
nbytes: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class CachePlan:
|
||||
reads: list[CacheOp] = field(default_factory=list)
|
||||
writes: list[CacheOp] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PlacementHint:
|
||||
pool: str = "default"
|
||||
role: str = "denoise"
|
||||
device: str = "cpu"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# WorkPlan / Done / StepResult / LoopResult #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class WorkPlan:
|
||||
"""A typed description of the next step.
|
||||
|
||||
``run`` is the kernel thunk built by ``next()`` but NOT called there (kernel-free
|
||||
planning). The runtime's ``ctx.execute`` calls it — possibly after batching with
|
||||
other compatible plans — preserving 'the runtime owns iteration'.
|
||||
"""
|
||||
loop_id: str
|
||||
instance_id: str
|
||||
kind: WorkUnitKind
|
||||
shape_sig: ShapeSignature
|
||||
resources: ResourceRequest = field(default_factory=ResourceRequest)
|
||||
cache: CachePlan = field(default_factory=CachePlan)
|
||||
placement: PlacementHint = field(default_factory=PlacementHint)
|
||||
emits: list[StreamChunk] = field(default_factory=list)
|
||||
payload: dict[str, Any] = field(default_factory=dict) # inspectable inputs (debug/parity)
|
||||
# the kernel thunk: run(model_instance, override=None) -> StepResult | dict.
|
||||
# Built by next() but NOT called there (kernel-free planning). ``override`` is an optional
|
||||
# interceptor-supplied forward result (e.g. a cached prediction); the step body still runs
|
||||
# the cheap solver step with it.
|
||||
run: Any = None
|
||||
label: str = ""
|
||||
# --- piecewise CUDA-graph capture ------------------------------------------------------------ #
|
||||
# ``capturable``: the model's static declaration that this step is safe to capture/replay — i.e.
|
||||
# no host RNG / data-dependent control flow inside the captured region. A stochastic step (the
|
||||
# FlowGRPO SDE rollout) sets this False, forcing the runtime to eager-break it.
|
||||
capturable: bool = True
|
||||
# ``graph_key``: extra op-structure discriminators (active CFG branch set, expert id) that change
|
||||
# the captured graph's *shape of computation*. Part of the capture key so a step with a different
|
||||
# branch set / expert never replays an incompatible graph. Empty ⇒ structure fixed by shape alone.
|
||||
graph_key: tuple = ()
|
||||
# The static-buffer capture form. A capturable step exposes its deterministic op
|
||||
# structure as ``graph_fn(model, workspace) -> StepResult`` reading EVERY per-step-varying input
|
||||
# (latent, sigmas, conditioning) from ``workspace`` — never from closure — plus ``graph_inputs``,
|
||||
# the dict of those current values. The runtime allocates address-stable buffers once per key and
|
||||
# rebinds ``graph_inputs`` into them in place each step (modeling CUDA static I/O buffers). Loops
|
||||
# that don't provide both stay on the eager path (the runtime eager-breaks them). ``run`` remains
|
||||
# the eager thunk (override / stochastic paths, and the no-capture baseline).
|
||||
graph_fn: Any = None
|
||||
graph_inputs: dict[str, Any] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Done:
|
||||
"""Sentinel returned by ``next`` when the loop is finished. ``finalize`` produces the
|
||||
actual LoopResult; ``result`` here is optional (kept for loops that want to carry it)."""
|
||||
result: LoopResult | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepResult:
|
||||
output: dict[str, Any] = field(default_factory=dict) # typed per-loop (e.g. {"noise_pred": ...})
|
||||
actual_seconds: float = 0.0
|
||||
cache_writes: list[CacheOp] = field(default_factory=list)
|
||||
behavior: Any = None # BehaviorRecord slice (rollout profile)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoopResult:
|
||||
outputs: dict[str, Any] = field(default_factory=dict) # final latents / tokens / artifacts
|
||||
metrics: dict[str, float] = field(default_factory=dict)
|
||||
behavior: Any = None # full BehaviorRecord (rollout)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# LoopState — the per-request mutable container (NEVER module globals) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class LoopState:
|
||||
"""All per-request mutable state lives here (structural cross-request safety)."""
|
||||
loop_id: str
|
||||
instance_id: str
|
||||
request_id: str
|
||||
profile: ExecutionProfile = ExecutionProfile.SERVE
|
||||
step_idx: int = 0
|
||||
done: bool = False
|
||||
rng: Any = None # seeded numpy Generator (per request)
|
||||
seed: int | None = None
|
||||
# common typed fields (resolved at init)
|
||||
latents: dict[str, TensorLike] = field(default_factory=dict)
|
||||
cond: dict[str, Any] = field(default_factory=dict) # conditioning written by ConditioningInjector
|
||||
timesteps: list[float] = field(default_factory=list)
|
||||
sigmas: list[float] = field(default_factory=list)
|
||||
# per-model extension (Cosmos3PackedSeq, MatrixGameState, ...) — typed by LoopSpec.extension_schema
|
||||
extension: Any = None
|
||||
# interceptor/policy state, keyed per plugin id AND per CFG branch
|
||||
plugin_state: dict[str, Any] = field(default_factory=dict)
|
||||
cache_handles: dict[str, Any] = field(default_factory=dict)
|
||||
# capture buffers (rollout): trajectory of per-step records
|
||||
trajectory: list[Any] = field(default_factory=list)
|
||||
scratch: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# StepContext — what policies read each step #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class StepContext:
|
||||
step_idx: int
|
||||
timestep: float
|
||||
sigma: float
|
||||
branch: str = "cond" # current guidance branch
|
||||
active_expert_id: str | None = None # set by ExpertRouting; observed by AdaptiveGateCFG
|
||||
sampler_coeffs: dict[str, float] = field(default_factory=dict)
|
||||
extra: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Protocols: LoopContext (the runtime seam) and Loop (the model contract) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
@runtime_checkable
|
||||
class LoopContext(Protocol):
|
||||
"""The single seam the runtime exposes to a loop. The model never sees the
|
||||
scheduler; the scheduler never sees the model's math."""
|
||||
profile: ExecutionProfile
|
||||
|
||||
def execute(self, plan: WorkPlan) -> StepResult:
|
||||
... # THE INVERSION POINT
|
||||
|
||||
def emit(self, chunk: StreamChunk) -> None:
|
||||
...
|
||||
|
||||
def check_cancel(self) -> None:
|
||||
... # raises request.Cancelled at boundary
|
||||
|
||||
def observe(self, event: str, **kw) -> None:
|
||||
... # observer bus hook
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class Loop(Protocol):
|
||||
"""The model-owned control flow. Four methods; ``next`` is kernel-free."""
|
||||
|
||||
def init(self, req: Any, model: Any, ctx: LoopContext) -> LoopState:
|
||||
...
|
||||
|
||||
def next(self, state: LoopState) -> WorkPlan | Done:
|
||||
...
|
||||
|
||||
def advance(self, state: LoopState, result: StepResult) -> LoopState:
|
||||
...
|
||||
|
||||
def finalize(self, state: LoopState) -> LoopResult:
|
||||
...
|
||||
@@ -0,0 +1,68 @@
|
||||
"""LoopRunner — the only place iteration lives.
|
||||
|
||||
The runtime's driver, factored so the engine can either run it to completion
|
||||
(``run``) or advance it one step at a time (``peek`` + ``step``) to **interleave** the steps
|
||||
of concurrent requests. Per-request state is entirely in ``LoopState``; the runner holds no
|
||||
hidden iteration state beyond a cached pending plan, so interleaving is safe by construction.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.loop.contracts import Done, LoopContext, LoopResult, LoopState, WorkPlan
|
||||
|
||||
|
||||
class LoopRunner:
|
||||
|
||||
def __init__(self, loop: Any, ctx: LoopContext, request: Any, model: Any):
|
||||
self.loop = loop
|
||||
self.ctx = ctx
|
||||
self.state: LoopState = loop.init(request, model, ctx)
|
||||
self._done = False
|
||||
self._result: LoopResult | None = None
|
||||
self._pending: WorkPlan | None = None
|
||||
|
||||
@property
|
||||
def done(self) -> bool:
|
||||
return self._done
|
||||
|
||||
@property
|
||||
def result(self) -> LoopResult | None:
|
||||
return self._result
|
||||
|
||||
def peek(self) -> WorkPlan | None:
|
||||
"""Compute (and cache) the next WorkPlan. ``next`` is kernel-free, so this is cheap.
|
||||
Returns None when the loop is finished (and runs ``finalize`` exactly once)."""
|
||||
if self._done:
|
||||
return None
|
||||
if self._pending is None:
|
||||
self.ctx.check_cancel()
|
||||
nxt = self.loop.next(self.state)
|
||||
if isinstance(nxt, Done):
|
||||
self._result = self.loop.finalize(self.state)
|
||||
self._done = True
|
||||
return None
|
||||
self._pending = nxt
|
||||
return self._pending
|
||||
|
||||
def step(self) -> bool:
|
||||
"""Execute the pending plan (the inversion point) and fold the result in.
|
||||
Returns True when the loop has just finished."""
|
||||
plan = self.peek()
|
||||
if self._done or plan is None:
|
||||
return True
|
||||
# the engine binds the current state onto ctx so interceptors see per-request state
|
||||
if hasattr(self.ctx, "bind_state"):
|
||||
self.ctx.bind_state(self.state)
|
||||
result = self.ctx.execute(plan)
|
||||
for chunk in plan.emits:
|
||||
self.ctx.emit(chunk)
|
||||
self.state = self.loop.advance(self.state, result)
|
||||
self._pending = None
|
||||
return self._done
|
||||
|
||||
def run(self) -> LoopResult:
|
||||
while not self._done:
|
||||
self.step()
|
||||
assert self._result is not None
|
||||
return self._result
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Policies — the default step decomposition."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.loop.policies.base import (
|
||||
BoundaryTimestepRouting,
|
||||
ConditioningInjector,
|
||||
ExpertRouting,
|
||||
FlowShiftPolicy,
|
||||
NoRouting,
|
||||
PassthroughConditioning,
|
||||
PrecisionPolicy,
|
||||
)
|
||||
from v2.loop.policies.cfg import (
|
||||
AdaptiveGateCFG,
|
||||
BatchedCFG,
|
||||
CFGPolicy,
|
||||
ClassicCFG,
|
||||
EmbeddedGuidance,
|
||||
PerModalityCFG,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CFGPolicy",
|
||||
"ClassicCFG",
|
||||
"BatchedCFG",
|
||||
"EmbeddedGuidance",
|
||||
"AdaptiveGateCFG",
|
||||
"PerModalityCFG",
|
||||
"FlowShiftPolicy",
|
||||
"PrecisionPolicy",
|
||||
"ExpertRouting",
|
||||
"NoRouting",
|
||||
"BoundaryTimestepRouting",
|
||||
"ConditioningInjector",
|
||||
"PassthroughConditioning",
|
||||
]
|
||||
@@ -0,0 +1,113 @@
|
||||
"""Policy base classes — the *default* step decomposition.
|
||||
|
||||
Policies (CFG, expert routing, precision, flow-shift, conditioning) delete duplication for
|
||||
the families that fit; they are never required. A family whose math is braided ships a custom
|
||||
``next``/``advance`` and uses these as a library (LTX-2's multi-pass guidance does exactly
|
||||
that). Policy *bindings* resolve at build; policy *state* is per-request in ``LoopState``
|
||||
(the adaptive-gate cached delta).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2.loop.contracts import StepContext
|
||||
from v2.loop.sampler import build_flow_sigmas
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# FlowShiftPolicy — resolution-bucket shift + sigma schedule #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class FlowShiftPolicy:
|
||||
"""config-driven flow-shift lookup (e.g. Wan 480p shift=3.0, 720p=5.0)."""
|
||||
|
||||
def __init__(self, shift: float = 3.0, bucket_lookup: dict[int, float] | None = None):
|
||||
self.shift = shift
|
||||
self.bucket_lookup = bucket_lookup or {}
|
||||
|
||||
def shift_for(self, height: int = 0, width: int = 0) -> float:
|
||||
return self.bucket_lookup.get(height * width, self.shift)
|
||||
|
||||
def build_schedule(self,
|
||||
num_steps: int,
|
||||
height: int = 0,
|
||||
width: int = 0,
|
||||
sigmas: list[float] | None = None) -> np.ndarray:
|
||||
if sigmas is not None: # explicit distilled schedule (LTX-2)
|
||||
return np.asarray(sigmas, dtype=np.float64)
|
||||
return build_flow_sigmas(num_steps, shift=self.shift_for(height, width))
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# PrecisionPolicy — autocast / scheduler-step dtype control #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class PrecisionPolicy:
|
||||
"""Replaces ``prefix=='Flux'`` autocast hacks + ``scheduler_step_in_fp32``."""
|
||||
|
||||
def __init__(self, compute_dtype: str = "float32", scheduler_step_in_fp32: bool = True):
|
||||
self.compute_dtype = compute_dtype
|
||||
self.scheduler_step_in_fp32 = scheduler_step_in_fp32
|
||||
|
||||
def cast(self, arr: Any) -> Any:
|
||||
# Array-preserving: a device (torch) tensor is cast in place on its device — never pulled to
|
||||
# host. numpy on CPU is unchanged. (torch is imported lazily, only on the tensor path.)
|
||||
if isinstance(arr, np.ndarray):
|
||||
return np.asarray(arr, dtype=np.dtype(self.compute_dtype))
|
||||
import torch
|
||||
return arr.to(getattr(torch, self.compute_dtype, torch.float32))
|
||||
|
||||
@property
|
||||
def scheduler_dtype(self):
|
||||
return np.float32 if self.scheduler_step_in_fp32 else np.dtype(self.compute_dtype)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# ExpertRouting — Wan2.2 boundary-timestep transformer switch #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class ExpertRouting(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def expert_for(self, ctx: StepContext) -> str:
|
||||
...
|
||||
|
||||
|
||||
class NoRouting(ExpertRouting):
|
||||
"""Single-expert models (Wan2.1 1.3B): always the same component."""
|
||||
|
||||
def __init__(self, component_id: str = "transformer"):
|
||||
self.component_id = component_id
|
||||
|
||||
def expert_for(self, ctx: StepContext) -> str:
|
||||
return self.component_id
|
||||
|
||||
|
||||
class BoundaryTimestepRouting(ExpertRouting):
|
||||
"""Wan2.2 ``boundary_ratio`` switch between two transformers."""
|
||||
|
||||
def __init__(self, high_noise: str, low_noise: str, boundary: float = 0.5):
|
||||
self.high_noise = high_noise
|
||||
self.low_noise = low_noise
|
||||
self.boundary = boundary
|
||||
|
||||
def expert_for(self, ctx: StepContext) -> str:
|
||||
return self.high_noise if ctx.sigma >= self.boundary else self.low_noise
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# ConditioningInjector — writes RequestState.cond; the loop stays agnostic #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class ConditioningInjector(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def inject(self, state: Any, request: Any) -> None:
|
||||
...
|
||||
|
||||
|
||||
class PassthroughConditioning(ConditioningInjector):
|
||||
"""Copies precomputed encoder outputs (text/image embeds) into ``state.cond``."""
|
||||
|
||||
def inject(self, state: Any, request: Any) -> None:
|
||||
# encoder ComponentNodes write into state.scratch["cond"] upstream of the loop
|
||||
state.cond.update(state.scratch.get("cond", {}))
|
||||
@@ -0,0 +1,103 @@
|
||||
"""CFGPolicy — one taxonomy over one shared denoise body.
|
||||
|
||||
Batched-vs-two-forward is a dispatch detail inside one policy, not a separate mechanism.
|
||||
|
||||
The step body asks ``branches_this_step(ctx, state)`` which forwards to run, runs them, then
|
||||
calls ``combine(preds, scale, ctx, state)``. Per-request mutable state (the adaptive-gate
|
||||
cached delta) lives in the ``state`` dict, which the step body slices out of
|
||||
``LoopState.plugin_state`` — never a module global.
|
||||
|
||||
cfg-parallel is a *parallelism axis* (not a policy); companions are an *orchestrator pattern*
|
||||
(not in the loop). Both compose with any policy here.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2.loop.contracts import StepContext
|
||||
|
||||
|
||||
class CFGPolicy(ABC):
|
||||
batched: bool = False # dispatch detail: stack branches into one forward
|
||||
branch_vocabulary: list[str] = ["cond", "uncond"] # full declared set (for per-branch plugin state)
|
||||
|
||||
def branches_this_step(self, ctx: StepContext, state: dict) -> list[str]:
|
||||
return list(self.branch_vocabulary)
|
||||
|
||||
@abstractmethod
|
||||
def combine(self, preds: dict[str, np.ndarray], guidance_scale: float, ctx: StepContext, state: dict) -> np.ndarray:
|
||||
...
|
||||
|
||||
|
||||
class ClassicCFG(CFGPolicy):
|
||||
"""Sequential two-forward CFG: ``uncond + s·(cond − uncond)`` (Wan default)."""
|
||||
|
||||
def combine(self, preds, guidance_scale, ctx, state):
|
||||
cond, uncond = preds["cond"], preds["uncond"]
|
||||
return uncond + guidance_scale * (cond - uncond)
|
||||
|
||||
|
||||
class BatchedCFG(ClassicCFG):
|
||||
"""Same algebra; the two branches are stacked into one batched forward (dispatch detail).
|
||||
|
||||
Identical output to ClassicCFG — verified by a parity test — which is the whole point:
|
||||
batched-vs-2-forward is not a separate mechanism.
|
||||
"""
|
||||
batched = True
|
||||
|
||||
|
||||
class EmbeddedGuidance(CFGPolicy):
|
||||
"""Degenerate single-branch identity-combine (Flux): guidance rides in the forward kwarg.
|
||||
|
||||
This is *not* 'no CFG' — it is kept inside the same abstraction."""
|
||||
branch_vocabulary = ["cond"]
|
||||
|
||||
def combine(self, preds, guidance_scale, ctx, state):
|
||||
return preds["cond"]
|
||||
|
||||
|
||||
class AdaptiveGateCFG(CFGPolicy):
|
||||
"""Cached-delta reuse with expert-switch self-invalidation.
|
||||
|
||||
On reuse steps it runs ONLY the cond branch and reuses the cached delta:
|
||||
``out = cond + (s−1)·delta`` (algebraically identical to ``uncond + s·(cond−uncond)``).
|
||||
The cached delta is invalidated when ``ExpertRouting`` switches the active expert.
|
||||
"""
|
||||
|
||||
def __init__(self, interval: int = 2):
|
||||
self.interval = max(1, interval)
|
||||
|
||||
def _recompute(self, ctx, state) -> bool:
|
||||
return (ctx.step_idx % self.interval == 0 or "delta" not in state
|
||||
or state.get("expert") != ctx.active_expert_id)
|
||||
|
||||
def branches_this_step(self, ctx, state):
|
||||
return ["cond", "uncond"] if self._recompute(ctx, state) else ["cond"]
|
||||
|
||||
def combine(self, preds, guidance_scale, ctx, state):
|
||||
if "uncond" in preds: # recompute path
|
||||
delta = preds["cond"] - preds["uncond"]
|
||||
state["delta"] = delta
|
||||
state["expert"] = ctx.active_expert_id
|
||||
return preds["uncond"] + guidance_scale * delta
|
||||
delta = state["delta"] # reuse path (skipped the uncond forward)
|
||||
return preds["cond"] + (guidance_scale - 1.0) * delta
|
||||
|
||||
|
||||
class PerModalityCFG(CFGPolicy):
|
||||
"""Joint A/V per-modality scales + interval gating (LTX-2, Cosmos3 t2vs — phase 2).
|
||||
|
||||
``combine`` reads the active modality from ``ctx.extra['modality']`` and applies its scale.
|
||||
"""
|
||||
|
||||
def __init__(self, scales: dict[str, float] | None = None, interval_gated: bool = False):
|
||||
self.scales = scales or {"video": 5.0}
|
||||
self.interval_gated = interval_gated
|
||||
|
||||
def combine(self, preds, guidance_scale, ctx, state):
|
||||
modality = ctx.extra.get("modality", "video")
|
||||
scale = self.scales.get(modality, guidance_scale)
|
||||
cond, uncond = preds["cond"], preds["uncond"]
|
||||
return uncond + scale * (cond - uncond)
|
||||
@@ -0,0 +1,110 @@
|
||||
"""Samplers as a small library; model families compose them.
|
||||
|
||||
Flow-match Euler is shared by the Wan and LTX-2 denoise step bodies. On GPU the real UniPC
|
||||
multistep / distilled schedulers are wrapped by the torch component adapters; this numpy form
|
||||
is the CPU-testable, bit-reproducible reference used by the loop tests.
|
||||
|
||||
Flow-match interpolation: ``x_t = (1-σ_t)·x0 + σ_t·ε``; the model predicts velocity
|
||||
``v = ε - x0``. A deterministic Euler step to σ_next is ``x_next = x_t + (σ_next-σ_t)·v``
|
||||
(exactly LTX-2's ``velocity=(latents-denoised)/σ; latents += velocity·dt``).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def build_flow_sigmas(num_steps: int, shift: float = 1.0, terminal: float = 0.0) -> np.ndarray:
|
||||
"""σ schedule from 1→terminal over ``num_steps+1`` points, with flow-shift applied.
|
||||
|
||||
Flow-shift: ``σ' = shift·σ / (1 + (shift-1)·σ)`` (Wan/LTX). shift=1 is the identity.
|
||||
"""
|
||||
base = np.linspace(1.0, terminal, num_steps + 1, dtype=np.float64)
|
||||
if shift != 1.0:
|
||||
base = shift * base / (1.0 + (shift - 1.0) * base)
|
||||
return base
|
||||
|
||||
|
||||
def x0_from_velocity(x_t: np.ndarray, velocity: np.ndarray, sigma_t: float) -> np.ndarray:
|
||||
"""x0 = x_t - σ_t·v (flow prediction → clean sample; Wan ``x0 = x_t - σ·model_output``)."""
|
||||
return x_t - sigma_t * velocity
|
||||
|
||||
|
||||
def build_karras_sigmas(num_steps: int,
|
||||
sigma_max: float = 80.0,
|
||||
sigma_min: float = 0.002,
|
||||
rho: float = 7.0) -> np.ndarray:
|
||||
"""Karras et al. (2022) ρ-interpolated σ schedule + the terminal σ, as Cosmos/EDM models use it.
|
||||
|
||||
Reproduces ``FlowMatchEulerDiscreteScheduler(use_karras_sigmas=True)`` the Cosmos pipeline configures
|
||||
(``sigma_max=80, sigma_min=0.002, rho=7``): a length-``num_steps`` ramp ``σ_max→σ_min`` via
|
||||
``σ_i = (max_inv_rho + i/(n-1)·(min_inv_rho-max_inv_rho))^ρ``, then ONE appended terminal value. The
|
||||
scheduler appends ``0.0`` and, with ``final_sigmas_type='sigma_min'``, the Cosmos stage overwrites it
|
||||
with ``σ[-2]`` to avoid a divide-by-zero in the EDM coeffs / velocity — so the terminal here is
|
||||
``σ_min`` (the last ramp value), giving ``num_steps+1`` points the EDM loop integrates pairwise.
|
||||
"""
|
||||
ramp = np.linspace(0.0, 1.0, num_steps, dtype=np.float64)
|
||||
min_inv_rho = sigma_min**(1.0 / rho)
|
||||
max_inv_rho = sigma_max**(1.0 / rho)
|
||||
sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho))**rho # σ_max -> σ_min
|
||||
return np.concatenate([sigmas, sigmas[-1:]]).astype(np.float64) # + sigma_min terminal clamp
|
||||
|
||||
|
||||
def flow_match_euler_step(x_t: np.ndarray, velocity: np.ndarray, sigma_t: float, sigma_next: float) -> np.ndarray:
|
||||
"""One deterministic Euler step along the flow ODE."""
|
||||
return x_t + (sigma_next - sigma_t) * velocity
|
||||
|
||||
|
||||
def add_noise(x0: np.ndarray, noise: np.ndarray, sigma: float) -> np.ndarray:
|
||||
"""Forward flow-match interpolation: x_σ = (1-σ)·x0 + σ·noise."""
|
||||
return (1.0 - sigma) * x0 + sigma * noise
|
||||
|
||||
|
||||
def flow_sde_step_with_logprob(x_t,
|
||||
velocity,
|
||||
sigma_t: float,
|
||||
sigma_next: float,
|
||||
*,
|
||||
noise=None,
|
||||
prev_sample=None,
|
||||
noise_scale: float = 0.7):
|
||||
"""FlowGRPO-style stochastic step + Gaussian log-prob.
|
||||
|
||||
The RL *rollout* sampler (vs the deterministic ODE ``flow_match_euler_step`` used at serve time).
|
||||
Injects noise for exploration and returns the log-prob of the realized sample under the step's
|
||||
Gaussian — the quantity the FlowGRPO PPO ratio uses. ``prev_sample=None`` samples a new step;
|
||||
passing ``prev_sample`` (the rollout sample) recomputes its log-prob under the *current* velocity
|
||||
(the ratio's numerator at update time). Returns ``(prev_sample, log_prob, mean, eff_std)``.
|
||||
"""
|
||||
x_t = np.asarray(x_t, dtype=np.float64)
|
||||
velocity = np.asarray(velocity, dtype=np.float64)
|
||||
s = min(float(sigma_t), 0.9999)
|
||||
dt = float(sigma_next) - float(sigma_t) # negative (σ decreases)
|
||||
std = float(np.sqrt(s / (1.0 - s)) * noise_scale)
|
||||
denom = 2.0 * max(s, 1e-6)
|
||||
mean = x_t * (1.0 + std**2 / denom * dt) + velocity * (1.0 + std**2 * (1.0 - s) / denom) * dt
|
||||
eff_std = max(std * np.sqrt(max(-dt, 1e-12)), 1e-6)
|
||||
if prev_sample is None:
|
||||
n = noise if noise is not None else np.zeros_like(x_t)
|
||||
prev_sample = mean + eff_std * np.asarray(n, dtype=np.float64)
|
||||
prev_sample = np.asarray(prev_sample, dtype=np.float64)
|
||||
var = eff_std**2
|
||||
log_prob = float(np.mean(-((prev_sample - mean)**2) / (2.0 * var) - np.log(eff_std) - 0.5 * np.log(2.0 * np.pi)))
|
||||
return prev_sample.astype("float32"), log_prob, mean.astype("float32"), float(eff_std)
|
||||
|
||||
|
||||
def flow_sde_ml_velocity(x_t, sample, sigma_t: float, sigma_next: float, *, noise_scale: float = 0.7):
|
||||
"""The velocity whose deterministic SDE mean lands exactly on ``sample`` (the max-likelihood
|
||||
velocity for a realized FlowGRPO sample). Moving the policy toward it, advantage-weighted, is the
|
||||
policy-gradient direction — and is nonzero even at ratio==1, so it is the correct FlowGRPO update
|
||||
surrogate (nudging toward the velocity the model already produced would be a no-op).
|
||||
"""
|
||||
x_t = np.asarray(x_t, dtype=np.float64)
|
||||
sample = np.asarray(sample, dtype=np.float64)
|
||||
s = min(float(sigma_t), 0.9999)
|
||||
dt = float(sigma_next) - float(sigma_t)
|
||||
std = float(np.sqrt(s / (1.0 - s)) * noise_scale)
|
||||
denom = 2.0 * max(s, 1e-6)
|
||||
a = 1.0 + std**2 / denom * dt
|
||||
b = (1.0 + std**2 * (1.0 - s) / denom) * dt
|
||||
b = b if abs(b) > 1e-8 else (1e-8 if b >= 0.0 else -1e-8)
|
||||
return ((sample - x_t * a) / b).astype("float32")
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Memory plane — reservation before admission, sleep/wake by tag."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.memory.allocator import MemoryManager, OutOfMemory, Reservation
|
||||
|
||||
__all__ = ["MemoryManager", "OutOfMemory", "Reservation"]
|
||||
@@ -0,0 +1,80 @@
|
||||
"""Tagged memory pools with reservation-before-admission.
|
||||
|
||||
Reservation must be pre-flight: a scheduler that admits a diffusion step and then discovers it
|
||||
cannot allocate the VAE tile is wrong.
|
||||
|
||||
Sleep/wake is component-granular (tags are component names) for RL: drop DiT + caches,
|
||||
keep VAE/text-encoder resident (CuMem-style).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import itertools
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
class OutOfMemory(Exception):
|
||||
pass
|
||||
|
||||
|
||||
_res_ctr = itertools.count(1)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Reservation:
|
||||
res_id: int
|
||||
tag: str
|
||||
nbytes: int
|
||||
active: bool = True
|
||||
|
||||
|
||||
class MemoryManager:
|
||||
|
||||
def __init__(self, total_bytes: int = 1 << 40, per_tag_budget: dict[str, int] | None = None):
|
||||
self.total_bytes = total_bytes
|
||||
self.per_tag_budget = dict(per_tag_budget or {})
|
||||
self.reserved = 0
|
||||
self._by_tag: dict[str, int] = {}
|
||||
self._reservations: dict[int, Reservation] = {}
|
||||
self._asleep: set[str] = set()
|
||||
|
||||
@property
|
||||
def available(self) -> int:
|
||||
return self.total_bytes - self.reserved
|
||||
|
||||
def can_reserve(self, tag: str, nbytes: int) -> bool:
|
||||
if nbytes > self.available:
|
||||
return False
|
||||
budget = self.per_tag_budget.get(tag)
|
||||
return budget is None or self._by_tag.get(tag, 0) + nbytes <= budget
|
||||
|
||||
def reserve(self, tag: str, nbytes: int) -> Reservation:
|
||||
if not self.can_reserve(tag, nbytes):
|
||||
raise OutOfMemory(f"cannot reserve {nbytes} bytes for tag {tag!r} "
|
||||
f"(available={self.available}, used={self._by_tag.get(tag, 0)})")
|
||||
res = Reservation(next(_res_ctr), tag, nbytes)
|
||||
self._reservations[res.res_id] = res
|
||||
self.reserved += nbytes
|
||||
self._by_tag[tag] = self._by_tag.get(tag, 0) + nbytes
|
||||
return res
|
||||
|
||||
def release(self, res: Reservation) -> None:
|
||||
if not res.active:
|
||||
return
|
||||
res.active = False
|
||||
self._reservations.pop(res.res_id, None)
|
||||
self.reserved -= res.nbytes
|
||||
self._by_tag[res.tag] = max(0, self._by_tag.get(res.tag, 0) - res.nbytes)
|
||||
|
||||
# component-granular sleep/wake (tags = component names) ------------------- #
|
||||
def sleep(self, tags: list[str]) -> int:
|
||||
freed = 0
|
||||
for tag in tags:
|
||||
self._asleep.add(tag)
|
||||
for res in [r for r in self._reservations.values() if r.tag == tag]:
|
||||
freed += res.nbytes
|
||||
self.release(res)
|
||||
return freed
|
||||
|
||||
def wake(self, tags: list[str]) -> None:
|
||||
for tag in tags:
|
||||
self._asleep.discard(tag)
|
||||
@@ -0,0 +1,5 @@
|
||||
"""v2 model architectures (vendored from fastvideo/models/). Mirrors fastvideo's layout — ``dits/``,
|
||||
``vaes/``, ``encoders/``, ``audio/``, ``upsamplers/`` — so a per-module cutover is a mechanical
|
||||
``cp`` + ``sed 'fastvideo.'->'v2.'``. Submodule bodies currently start as re-export STUBS backed by
|
||||
fastvideo (see memory: v2-vendoring-approach); nothing is imported eagerly here so ``import v2`` stays
|
||||
torch-free — the stubs (which import fastvideo/torch) load only when the GPU backend references them."""
|
||||
@@ -0,0 +1 @@
|
||||
"""Audio architectures (vendored, mirrors fastvideo/models/audio/). Stub re-exports for now."""
|
||||
@@ -0,0 +1,4 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""LTX-2 audio VAE facade. The backend uses ``AudioLatentShape`` (audio token-count helper); the
|
||||
AudioDecoder / Vocoder classes are constructed by the loaders from the card's load_id, not imported here."""
|
||||
from fastvideo.models.audio.ltx2_audio_vae import AudioLatentShape # noqa: F401
|
||||
@@ -0,0 +1 @@
|
||||
"""DiT architectures (vendored, mirrors fastvideo/models/dits/). Stub re-exports for now."""
|
||||
@@ -0,0 +1,10 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""LingBot-World DiT facade.
|
||||
|
||||
The transformer class is constructed by the loader from the card's ``load_id``, so it is not imported
|
||||
here. This stub only re-exports the camera/Plucker embedding builder ``prepare_camera_embedding``
|
||||
(poses.npy + intrinsics.npy -> ``c2ws_plucker_emb [B, 6*s^2, F, H, W]``) so the program's camera node
|
||||
can build the tensor without reaching into the model package.
|
||||
"""
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import ( # noqa: F401
|
||||
prepare_camera_embedding, )
|
||||
@@ -0,0 +1,4 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""LTX-2 DiT facade. The backend uses ``VideoLatentShape`` (token-count helper); the transformer class
|
||||
itself is constructed by the loader from the card's load_id, not imported here."""
|
||||
from fastvideo.models.dits.ltx2 import VideoLatentShape # noqa: F401
|
||||
@@ -0,0 +1,5 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""SD3 (Stable Diffusion 3.5) DiT facade. The transformer class itself is constructed by the loader
|
||||
from the card's ``load_id``; this stub only re-exports the output dataclass so the SD3 torch adapter can
|
||||
type-check / unwrap the forward result symbolically (mirrors ``v2/models/dits/ltx2.py``)."""
|
||||
from fastvideo.models.dits.sd3 import SD3Transformer2DModelOutput # noqa: F401
|
||||
@@ -0,0 +1 @@
|
||||
"""Upsampler architectures (vendored, mirrors fastvideo/models/upsamplers/). Stub re-exports for now."""
|
||||
@@ -0,0 +1,4 @@
|
||||
# STUB: re-exports fastvideo until vendored (see memory: v2-vendoring-approach).
|
||||
"""LTX-2 latent upsampler facade. The backend uses ``upsample_video`` (un_normalize -> learned 2x ->
|
||||
normalize); the LTX2LatentUpsampler module is constructed by the loader from the card's load_id."""
|
||||
from fastvideo.models.upsamplers.ltx2_upsampler import upsample_video # noqa: F401
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Parallelism plane — named axes -> validated mesh, part of the cache key."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.parallel.mesh import FakeDeviceMesh, build_mesh
|
||||
from v2.parallel.plan import AXIS_NAMES, ParallelPlan
|
||||
from v2.parallel.validation import ParallelValidationError, validate_plan
|
||||
|
||||
__all__ = ["ParallelPlan", "AXIS_NAMES", "FakeDeviceMesh", "build_mesh", "validate_plan", "ParallelValidationError"]
|
||||
@@ -0,0 +1,56 @@
|
||||
"""FakeDeviceMesh — a torchtitan-style ParallelDims builder, CPU-testable.
|
||||
|
||||
On a GPU box this compiles to ``torch.distributed.device_mesh.DeviceMesh``. Here it is a
|
||||
pure-Python mesh of rank tuples so topology logic is unit-tested without GPUs.
|
||||
Degree-one axes exist as trivial groups so component code needs no special cases.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from itertools import product
|
||||
|
||||
from v2.parallel.plan import ParallelPlan
|
||||
from v2.parallel.validation import validate_plan
|
||||
|
||||
|
||||
class FakeDeviceMesh:
|
||||
|
||||
def __init__(self, plan: ParallelPlan):
|
||||
self.plan = plan
|
||||
self.order = plan.mesh_order or list(plan.axes.keys())
|
||||
self.shape = tuple(plan.degree(a) for a in self.order)
|
||||
self.world_size = plan.world_size()
|
||||
|
||||
def group_ranks(self, axis: str) -> list[list[int]]:
|
||||
"""All collective groups along one axis (each is a list of global ranks)."""
|
||||
if axis not in self.order:
|
||||
return [[0]] # degree-one trivial group
|
||||
axis_idx = self.order.index(axis)
|
||||
groups: list[list[int]] = []
|
||||
ranges = [range(d) for d in self.shape]
|
||||
seen: set[tuple] = set()
|
||||
for coord in product(*ranges):
|
||||
key = coord[:axis_idx] + coord[axis_idx + 1:]
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
group = []
|
||||
for i in range(self.shape[axis_idx]):
|
||||
c = list(coord)
|
||||
c[axis_idx] = i
|
||||
group.append(self._coord_to_rank(tuple(c)))
|
||||
groups.append(group)
|
||||
return groups
|
||||
|
||||
def _coord_to_rank(self, coord: tuple[int, ...]) -> int:
|
||||
rank = 0
|
||||
for c, d in zip(coord, self.shape, strict=False):
|
||||
rank = rank * d + c
|
||||
return rank
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"FakeDeviceMesh(order={self.order}, shape={self.shape}, world={self.world_size})"
|
||||
|
||||
|
||||
def build_mesh(plan: ParallelPlan, card=None, **kw) -> FakeDeviceMesh:
|
||||
validate_plan(plan, card=card, **kw)
|
||||
return FakeDeviceMesh(plan)
|
||||
@@ -0,0 +1,46 @@
|
||||
"""ParallelPlan — parallelism as a model contract.
|
||||
|
||||
Parallelism is not a launch flag; it affects cache keys, scheduling, transport,
|
||||
capture, and parity, so it lives on the card. Declarative, validated, compiled to a
|
||||
mesh via a ParallelDims-style builder (``parallel/mesh.py``). This module is a pure
|
||||
leaf (no card/runtime imports) so ``card/`` can hold plans without a cycle.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
# Canonical axis names. cfgp is <=2.
|
||||
AXIS_NAMES = ("dp", "tp", "sp", "cp", "cfgp", "pp_patch", "vae", "ep", "fsdp", "role", "replica")
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParallelPlan:
|
||||
axes: dict[str, int] = field(default_factory=dict) # e.g. {"tp": 2, "sp": 4, "cfgp": 2}
|
||||
mesh_order: list[str] = field(default_factory=list)
|
||||
placement: str = "colocated"
|
||||
communication: dict[str, str] = field(default_factory=dict)
|
||||
# applicability conditions travel with axes
|
||||
applicability: dict[str, dict] = field(default_factory=dict)
|
||||
per_axis_communication: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
def degree(self, axis: str) -> int:
|
||||
return int(self.axes.get(axis, 1))
|
||||
|
||||
def world_size(self) -> int:
|
||||
w = 1
|
||||
for v in self.axes.values():
|
||||
w *= int(v)
|
||||
return max(w, 1)
|
||||
|
||||
@property
|
||||
def hash(self) -> str:
|
||||
"""Stable hash — part of the CacheKey (parallel_plan_hash)."""
|
||||
payload = json.dumps({"axes": self.axes, "order": self.mesh_order}, sort_keys=True)
|
||||
return hashlib.sha256(payload.encode()).hexdigest()[:16]
|
||||
|
||||
@classmethod
|
||||
def single(cls) -> ParallelPlan:
|
||||
"""The default deployment: one device, all degree-one trivial groups."""
|
||||
return cls(axes={"dp": 1}, mesh_order=["dp"])
|
||||
@@ -0,0 +1,59 @@
|
||||
"""Pre-flight parallelism validation.
|
||||
|
||||
Validate at load or fail, never halfway. Ownership conflicts are build errors, and
|
||||
applicability conditions travel with axes (e.g. ``pp_patch`` is invalid for causal/AR).
|
||||
All checks are CPU-testable on a fake mesh.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import LoopKind
|
||||
from v2.parallel.plan import AXIS_NAMES, ParallelPlan
|
||||
|
||||
|
||||
class ParallelValidationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
# loop kinds whose causality is broken by PipeFusion displaced-patch pipelining
|
||||
_CAUSAL_KINDS = {LoopKind.CHUNK_ROLLOUT, LoopKind.AR_DECODE}
|
||||
|
||||
|
||||
def validate_plan(plan: ParallelPlan,
|
||||
card=None,
|
||||
*,
|
||||
world_size: int | None = None,
|
||||
cfg_policy_batched: bool = False) -> ParallelPlan:
|
||||
"""Validate a plan, optionally against a card. Returns the plan or raises.
|
||||
|
||||
Checks:
|
||||
1. axis names are known; degrees >= 1
|
||||
2. cfgp <= 2
|
||||
3. product of degrees matches world_size (when given)
|
||||
4. pp_patch invalid for any causal/AR loop on the card
|
||||
5. ownership conflict: cfgp>1 AND a batched CFGPolicy is rejected
|
||||
"""
|
||||
errs: list[str] = []
|
||||
|
||||
for name, deg in plan.axes.items():
|
||||
if name not in AXIS_NAMES:
|
||||
errs.append(f"unknown parallel axis {name!r} (known: {AXIS_NAMES})")
|
||||
if int(deg) < 1:
|
||||
errs.append(f"axis {name!r} has degree {deg} < 1")
|
||||
|
||||
if plan.degree("cfgp") > 2:
|
||||
errs.append(f"cfgp degree {plan.degree('cfgp')} > 2 (CFG has at most 2 branches)")
|
||||
|
||||
if world_size is not None and plan.world_size() != world_size:
|
||||
errs.append(f"product of degrees {plan.world_size()} != world_size {world_size}")
|
||||
|
||||
if plan.degree("cfgp") > 1 and cfg_policy_batched:
|
||||
errs.append("ownership conflict: a request owns a BatchedCFG policy OR a cfgp group, never both")
|
||||
|
||||
if card is not None and plan.degree("pp_patch") > 1:
|
||||
causal = [lid for lid, lp in card.loops.items() if lp.kind in _CAUSAL_KINDS]
|
||||
if causal:
|
||||
errs.append(f"pp_patch is invalid for causal/AR loops {causal} (stale KV breaks causality, §8)")
|
||||
|
||||
if errs:
|
||||
raise ParallelValidationError("ParallelPlan failed validation:\n - " + "\n - ".join(errs))
|
||||
return plan
|
||||
@@ -0,0 +1,16 @@
|
||||
"""Parity plane — a first-class package, not a test folder.
|
||||
|
||||
It is how the (recipe, runtime) pair is kept honest: parity is measured by ParityAligner,
|
||||
the consistency ladder is typed, and the interleave gate is non-negotiable.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import ConsistencyLevel, ExecutionProfile
|
||||
from v2.parity.aligner import ParityAligner
|
||||
from v2.parity.interleave_gate import assert_interleave_parity, compare_outputs
|
||||
from v2.parity.ladder import Divergence, array_diff, bit_identical, within
|
||||
|
||||
__all__ = [
|
||||
"ConsistencyLevel", "ExecutionProfile", "ParityAligner", "Divergence", "array_diff", "bit_identical", "within",
|
||||
"assert_interleave_parity", "compare_outputs"
|
||||
]
|
||||
@@ -0,0 +1,52 @@
|
||||
"""ParityAligner — parity is measured, never assumed.
|
||||
|
||||
A read-only observer: *record mode* dumps named taps per step/block from a reference
|
||||
(the official framework, or a pre-change build); *compare mode* replays with fixed seeds
|
||||
and reports the first divergence beyond per-tap tolerance. This is the engine behind
|
||||
the "old loop vs new loop, bit-identical" gate and the standing instrument for every
|
||||
port, precision change, and kernel swap.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.parity.ladder import ConsistencyLevel, Divergence, array_diff
|
||||
|
||||
|
||||
class ParityAligner:
|
||||
"""Observer that records named taps and compares two runs."""
|
||||
|
||||
def __init__(self, name: str = "parity", default_atol: float = 0.0, default_rtol: float = 0.0):
|
||||
self.name = name
|
||||
self.default_atol = default_atol
|
||||
self.default_rtol = default_rtol
|
||||
self.taps: dict[tuple[int, str], Any] = {} # (step, tap_name) -> value
|
||||
self.tolerances: dict[str, tuple[float, float]] = {} # tap_name -> (atol, rtol)
|
||||
|
||||
def set_tolerance(self, tap_name: str, atol: float = 0.0, rtol: float = 0.0) -> None:
|
||||
self.tolerances[tap_name] = (atol, rtol)
|
||||
|
||||
# observer hook (the loop/engine calls this) ------------------------------ #
|
||||
def record_tap(self, step: int, tap_name: str, value: Any) -> None:
|
||||
self.taps[(step, tap_name)] = value
|
||||
|
||||
def observe(self, event: str, **kw) -> None:
|
||||
if event == "tap":
|
||||
self.record_tap(kw.get("step", 0), kw["name"], kw["value"])
|
||||
|
||||
def first_divergence(self,
|
||||
reference: ParityAligner,
|
||||
level: ConsistencyLevel = ConsistencyLevel.C1) -> Divergence | None:
|
||||
"""Compare this (current) run against a reference, reporting the FIRST divergence
|
||||
in step order beyond the per-tap tolerance."""
|
||||
for (step, tap_name) in sorted(reference.taps, key=lambda k: (k[0], k[1])):
|
||||
ref_val = reference.taps[(step, tap_name)]
|
||||
if (step, tap_name) not in self.taps:
|
||||
return Divergence(f"{tap_name}@{step}", level, float("inf"), float("inf"), "tap missing in current run")
|
||||
cur_val = self.taps[(step, tap_name)]
|
||||
atol, rtol = self.tolerances.get(tap_name, (self.default_atol, self.default_rtol))
|
||||
abs_d, rel_d = array_diff(ref_val, cur_val)
|
||||
if not (abs_d <= atol or rel_d <= rtol): # diverged: outside BOTH tolerances
|
||||
return Divergence(f"{tap_name}@{step}", level, abs_d, rel_d,
|
||||
f"abs={abs_d:.3e} rel={rel_d:.3e} > (atol={atol},rtol={rtol})")
|
||||
return None
|
||||
@@ -0,0 +1,62 @@
|
||||
"""The batch-of-N interleave parity gate — required, not optional.
|
||||
|
||||
Loop inversion's real hazard is cross-request state smearing under interleaving, which a
|
||||
batch-of-1 parity gate cannot detect. So this gate runs two or more concurrent requests
|
||||
interleaved at step granularity and requires bit-identical output versus running them serially.
|
||||
|
||||
The gate is pure and duck-typed: it takes any engine exposing ``run_serial`` and ``run_interleaved``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.request.artifacts import (
|
||||
Output, )
|
||||
from v2.parity.ladder import ConsistencyLevel, Divergence, bit_identical
|
||||
|
||||
|
||||
def _artifact_payload(art: Any) -> Any:
|
||||
for attr in ("frames", "samples", "tensor", "latent", "token_ids", "text"):
|
||||
if hasattr(art, attr) and getattr(art, attr) is not None:
|
||||
return getattr(art, attr)
|
||||
return None
|
||||
|
||||
|
||||
def compare_outputs(serial: dict[str, Output], interleaved: dict[str, Output]) -> list[Divergence]:
|
||||
"""Bitwise-compare per-request outputs from serial vs interleaved runs."""
|
||||
divs: list[Divergence] = []
|
||||
if set(serial) != set(interleaved):
|
||||
divs.append(
|
||||
Divergence("request_set", ConsistencyLevel.C1, float("inf"), float("inf"),
|
||||
f"serial reqs {set(serial)} != interleaved {set(interleaved)}"))
|
||||
return divs
|
||||
for rid in sorted(serial):
|
||||
so, io = serial[rid], interleaved[rid]
|
||||
if set(so.artifacts) != set(io.artifacts):
|
||||
divs.append(
|
||||
Divergence(f"{rid}:artifacts", ConsistencyLevel.C1, float("inf"), float("inf"),
|
||||
f"artifact names differ: {set(so.artifacts)} vs {set(io.artifacts)}"))
|
||||
continue
|
||||
for name in sorted(so.artifacts):
|
||||
a, b = _artifact_payload(so.artifacts[name]), _artifact_payload(io.artifacts[name])
|
||||
if a is None and b is None:
|
||||
# both empty is NOT parity — a deferred/aborted request must not pass the gate vacuously
|
||||
divs.append(
|
||||
Divergence(
|
||||
f"{rid}:{name}", ConsistencyLevel.C1, float("inf"), float("inf"),
|
||||
"both serial and interleaved produced EMPTY output for a declared "
|
||||
"artifact (deferred/aborted?) — suspicious, not parity"))
|
||||
continue
|
||||
if not bit_identical(a, b):
|
||||
divs.append(
|
||||
Divergence(f"{rid}:{name}", ConsistencyLevel.C1, float("nan"), float("nan"),
|
||||
"serial vs interleaved output not bit-identical "
|
||||
"(cross-request state smearing!)"))
|
||||
return divs
|
||||
|
||||
|
||||
def assert_interleave_parity(engine: Any, requests: list[Any]) -> list[Divergence]:
|
||||
"""Drive ``engine`` both ways and return divergences (empty list == gate PASSES)."""
|
||||
serial = engine.run_serial(requests)
|
||||
interleaved = engine.run_interleaved(requests)
|
||||
return compare_outputs(serial, interleaved)
|
||||
@@ -0,0 +1,59 @@
|
||||
"""The consistency ladder + numeric comparison helpers.
|
||||
|
||||
C0 component · C1 loop · C2 behavioral · C3 distribution · C4 artifact.
|
||||
|
||||
The C2 split is load-bearing: likelihood-based methods compare per-step log-probs;
|
||||
likelihood-free methods (DiffusionNFT) compare seeded final-sample + prediction-space
|
||||
identity (old_deviate / ref-MSE) — there are NO log-probs to match.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2._enums import ConsistencyLevel # re-exported for the package
|
||||
|
||||
|
||||
@dataclass
|
||||
class Divergence:
|
||||
where: str # tap name / artifact path
|
||||
level: ConsistencyLevel
|
||||
max_abs_diff: float
|
||||
max_rel_diff: float
|
||||
message: str = ""
|
||||
|
||||
def __bool__(self) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def array_diff(a: Any, b: Any) -> tuple[float, float]:
|
||||
"""Return (max_abs, max_rel) difference between two array-likes."""
|
||||
a = np.asarray(a, dtype=np.float64)
|
||||
b = np.asarray(b, dtype=np.float64)
|
||||
if a.shape != b.shape:
|
||||
return (float("inf"), float("inf"))
|
||||
if a.size == 0:
|
||||
return (0.0, 0.0)
|
||||
diff = np.abs(a - b)
|
||||
max_abs = float(diff.max())
|
||||
denom = np.maximum(np.abs(a), np.abs(b))
|
||||
with np.errstate(divide="ignore", invalid="ignore"):
|
||||
rel = np.where(denom > 0, diff / denom, 0.0)
|
||||
return (max_abs, float(np.nanmax(rel)) if rel.size else 0.0)
|
||||
|
||||
|
||||
def within(a: Any, b: Any, rtol: float = 0.0, atol: float = 0.0) -> bool:
|
||||
"""Close enough if within the absolute OR the relative tolerance (allclose-style).
|
||||
With atol=rtol=0 this reduces to exact equality (the bit-identical case)."""
|
||||
abs_d, rel_d = array_diff(a, b)
|
||||
return abs_d <= atol or rel_d <= rtol
|
||||
|
||||
|
||||
def bit_identical(a: Any, b: Any) -> bool:
|
||||
"""C1/C2 bit-identical check (fixed seed, same kernels) — array_equal on raw values."""
|
||||
try:
|
||||
return bool(np.array_equal(np.asarray(a), np.asarray(b)))
|
||||
except Exception:
|
||||
return a == b
|
||||
@@ -0,0 +1,42 @@
|
||||
"""The multi-backend dispatch substrate.
|
||||
|
||||
Two tuple-keyed registries (``COMPONENTS``, ``KERNELS``) + a ``Platform`` that detects ``(device,
|
||||
arch)`` and resolves both, with the numpy reference as the terminal fallback rung and parity oracle.
|
||||
Lets CPU, GPU, and other backends coexist: a model's loops, policies, scheduler, caches, and
|
||||
training code never name a device; they call through the component/kernel seams and the resolved
|
||||
``Platform`` decides the implementation.
|
||||
|
||||
Importing this package is cheap and cycle-free (registry + platform classes only); the concrete
|
||||
backend registrations in ``backends/`` are imported lazily on first ``Platform`` use.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.platform.platform import (
|
||||
KernelTable,
|
||||
Platform,
|
||||
component_matrix,
|
||||
ensure_backends_loaded,
|
||||
kernel_matrix,
|
||||
)
|
||||
from v2.platform.registry import (
|
||||
COMPONENTS,
|
||||
FLOW_MATCH_STEP,
|
||||
FLOW_SDE_STEP,
|
||||
KERNELS,
|
||||
register_component,
|
||||
register_kernel,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Platform",
|
||||
"KernelTable",
|
||||
"COMPONENTS",
|
||||
"KERNELS",
|
||||
"register_component",
|
||||
"register_kernel",
|
||||
"ensure_backends_loaded",
|
||||
"kernel_matrix",
|
||||
"component_matrix",
|
||||
"FLOW_MATCH_STEP",
|
||||
"FLOW_SDE_STEP",
|
||||
]
|
||||
@@ -0,0 +1,69 @@
|
||||
"""Array namespace — numpy on a CPU box, torch-on-device on a GPU box.
|
||||
|
||||
The denoise loop's math is already array-agnostic: CFG combine (``uncond + s·(cond−uncond)``) and the
|
||||
flow-match Euler step (``x + (σ_next−σ_t)·v``) are pure arithmetic that runs identically on numpy and
|
||||
torch. This namespace supplies the few NON-arithmetic helpers the loop still needs — birthing the
|
||||
latent, dtype casts, and the single host round-trip at the request boundary — so on a GPU box the
|
||||
latent stays resident on-device for the whole loop (no per-step host<->device copy) while the CPU/toy
|
||||
path stays pure numpy (torch-free; the parity mini is unchanged).
|
||||
|
||||
The latent is still *seeded* with numpy (``rng.standard_normal``) and uploaded once via ``from_host``,
|
||||
so a GPU run is bit-identical to the pre-on-device path (the noise values are the same; upload is
|
||||
lossless). ``Platform.xp`` picks the namespace from the platform's device.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
class NumpyNS:
|
||||
"""The CPU/toy namespace: every op is a numpy identity (no torch dependency)."""
|
||||
|
||||
device = "cpu"
|
||||
|
||||
def from_host(self, a: Any) -> np.ndarray:
|
||||
return np.asarray(a, dtype=np.float32)
|
||||
|
||||
def to_host(self, a: Any) -> np.ndarray:
|
||||
return np.asarray(a)
|
||||
|
||||
def to_f32(self, a: Any) -> np.ndarray:
|
||||
return np.asarray(a, dtype=np.float32)
|
||||
|
||||
def is_native(self, a: Any) -> bool:
|
||||
return isinstance(a, np.ndarray)
|
||||
|
||||
|
||||
class TorchNS:
|
||||
"""The GPU namespace: keeps arrays as device tensors; marshals host<->device only on demand."""
|
||||
|
||||
def __init__(self, device: str = "cuda") -> None:
|
||||
import torch
|
||||
self._torch = torch
|
||||
self.device = device
|
||||
|
||||
def from_host(self, a: Any) -> Any:
|
||||
t = self._torch
|
||||
if t.is_tensor(a):
|
||||
return a.to(self.device)
|
||||
return t.as_tensor(np.asarray(a, dtype=np.float32), device=self.device)
|
||||
|
||||
def to_host(self, a: Any) -> np.ndarray:
|
||||
t = self._torch
|
||||
if t.is_tensor(a):
|
||||
return a.detach().to("cpu", t.float32).numpy()
|
||||
return np.asarray(a)
|
||||
|
||||
def to_f32(self, a: Any) -> Any:
|
||||
t = self._torch
|
||||
return a.float() if t.is_tensor(a) else np.asarray(a, dtype=np.float32)
|
||||
|
||||
def is_native(self, a: Any) -> bool:
|
||||
return bool(self._torch.is_tensor(a))
|
||||
|
||||
|
||||
def get_array_ns(device: str) -> Any:
|
||||
"""The array namespace for a platform device: torch-on-device for cuda, numpy otherwise."""
|
||||
return TorchNS(device) if device == "cuda" else NumpyNS()
|
||||
@@ -0,0 +1,79 @@
|
||||
# GPU bring-up — the real torch/CUDA backend
|
||||
|
||||
The `cuda` backend (`torch_cuda.py` + `torch_adapters.py` + `torch_kernels.py`) is **written but not
|
||||
run**: this repo's dev environment has no GPU and no torch, so the torch code is grounded in the
|
||||
verbatim real `fastvideo` APIs but **could not be executed or verified here**. Every spot the real
|
||||
API must be confirmed on a box is marked `# BRINGUP` in the source and listed as a risk below.
|
||||
|
||||
This doc is the checklist for a human on a GPU box to take it from "resolves" to "generates".
|
||||
|
||||
## What it does
|
||||
|
||||
On a box with torch + CUDA + the parent `fastvideo` package installed, `Platform.detect()` returns a
|
||||
`cuda` platform. The existing v2 loops/policies/scheduler/training are **unchanged**; only the
|
||||
resolved implementations differ:
|
||||
|
||||
- **Components** resolve to torch adapters that wrap the real module named by the card's `load_id`
|
||||
(`fastvideo.models.dits.wanvideo:WanTransformer3DModel`, `…vaes.wanvae:AutoencoderKLWan`,
|
||||
`…encoders.t5`), loading weights from `ComponentSpec.checkpoint`. Each adapter bridges the real
|
||||
forward to the mini's duck-typed surface (`dit(latent,text,sigma)->velocity`, `vae.decode/encode`,
|
||||
`text_encoder.encode`), marshalling numpy↔torch at its boundary (the loop math stays numpy fp32).
|
||||
- **Solver ops** (`flow_match_step`/`flow_sde_step`) resolve to plain torch elementwise — there is
|
||||
**no fused solver kernel** in fastvideo-kernel (it ships only attention/norm/quant primitives), so
|
||||
the honest source is "torch elementwise", registered at arch `generic`.
|
||||
|
||||
## Already verified on CPU (see `test_torch_backend.py`)
|
||||
|
||||
1. All three cuda components (`dit`, `vae`, `text_encoder`) + both solver ops are registered, with
|
||||
honest sources (no claimed `fastvideo-kernel:flow_*`), `available=False` here.
|
||||
2. Importing the backends **never imports torch** (`torch_adapters`/`torch_kernels` load only inside
|
||||
builder bodies) — the CPU mini stays green.
|
||||
3. On a (forced-available) cuda platform, resolution picks the **real cuda cells, not a silent toy**.
|
||||
4. Building without torch **fails loudly** (not a quiet numpy toy decoding real latents).
|
||||
|
||||
## Ordered bring-up checklist (on the GPU box)
|
||||
|
||||
1. **Set weights + args.** Fill `ComponentSpec.checkpoint` for each component (HF id or local path),
|
||||
and provide the `FastVideoArgs` the real loaders need (`_fastvideo_args` builds a minimal one from
|
||||
the path — confirm its required fields/precision). *(Risk A — without these, builders raise.)*
|
||||
2. **Detect.** `Platform.detect()` returns `cuda(smXY)`; confirm the arch string. `component_matrix()`
|
||||
shows the three cuda components `available=True`.
|
||||
3. **Build in isolation.** Build each component via the FastVideo loaders (`TransformerLoader` /
|
||||
`VAELoader` / `TextEncoderLoader` + `TokenizerLoader`); assert type + a single forward's output
|
||||
shape (no full denoise yet). *(A)*
|
||||
4. **One DiT step.** `dit(x, pe, sigma) -> velocity`; check finite + shape `[C,T,H,W]`. The
|
||||
`timestep = sigma*1000` convention and the velocity (noise−clean) semantics were cross-checked as
|
||||
MATCHING the real `forward`/scheduler — confirm numerically. *(B, C)*
|
||||
5. **One solver step.** `flow_match_step` finite + right shape (math mirrors `loop/sampler.py`).
|
||||
6. **VAE.** Normalization is now applied in-adapter (`(z-mean)*inv_std` on encode, inverse on decode,
|
||||
`latents_std` as reciprocal). Confirm the `shift_factor` placement/sign and dtype on the box. *(D)*
|
||||
7. **Text.** Class resolution (UMT5 vs T5, from config) and `set_forward_context(...)` are now wired.
|
||||
Confirm the exact tokenizer kwargs (`text_len`/`max_length`, special tokens) from the model config. *(E)*
|
||||
8. **End-to-end.** Full t2v denoise → VAE decode → compare to a known-good fastvideo generation
|
||||
(SSIM / the ssim regression harness).
|
||||
9. **RL / SDE path.** `flow_sde_step` returns a finite `(prev, log_prob, mean, eff_std)`; run a
|
||||
rollout. Note: the **training** weight-surface (`mse_grad_step`) is NOT implemented on the cuda
|
||||
rung — RL/distill on GPU is a separate workstream. *(F)*
|
||||
10. **CUDA-graph capture.** Last. The wan21 loop declares `breakable_cudagraph`; capturing a real
|
||||
torch graph (vs the numpy `StaticWorkspace` model) is GPU-only work.
|
||||
|
||||
## Risks / open unknowns (the `# BRINGUP` points)
|
||||
|
||||
Cross-checked against the real source (`crosscheck-gpu-adapters`): the **interface contracts matched**
|
||||
(DiT returns a bare velocity tensor; `timestep=sigma*1000`; `encode().mode()` + bare `decode`;
|
||||
`.last_hidden_state`; no fused solver kernel). The **construction layer was wrong and is now fixed**
|
||||
in code (real loaders instead of the nonexistent `from_pretrained`; UMT5-vs-T5 resolved from config;
|
||||
`set_forward_context` wired; latent normalization applied). What remains is genuinely box-dependent:
|
||||
|
||||
| | Risk | Failure mode if wrong |
|
||||
|---|---|---|
|
||||
| **A** | `FastVideoArgs` construction the real loaders need (exact fields/precision); checkpoint path | builder raises (blocking) |
|
||||
| **B** | sigma→timestep scaling — cross-checked as `sigma*1000`; confirm numerically | silent garbage, not a crash |
|
||||
| **C** | DiT output velocity (noise−clean) — cross-checked as matching; confirm sign on box | denoises backward |
|
||||
| **D** | `shift_factor` placement/sign + dtype of the (now-applied) latent normalization | washed-out / saturated video |
|
||||
| **E** | exact tokenizer kwargs (`text_len`/`max_length`, special tokens) from config — class + `set_forward_context` now fixed in code | wrong/empty conditioning |
|
||||
| **F** | torch training surface (`mse_grad_step`) not implemented | RL/distill on cuda is a separate workstream |
|
||||
| **G** | numpy↔torch marshalling per DiT call (H2D/D2H + dtype round-trip) | correct but slow; torch-native surface is the perf follow-up |
|
||||
|
||||
Items A–E are correctness; F–G are scope/perf. None block the **CPU mini** (all gated `available=False`).
|
||||
Multi-GPU FSDP sharding via the loaders also needs on-box verification (single-GPU bring-up first).
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Backend registration modules. Imported for side effects by ``Platform.ensure_backends_loaded``.
|
||||
|
||||
Each module calls ``register_component`` / ``register_kernel`` at import time. They live behind the
|
||||
lazy loader (not imported by ``platform/__init__`` or ``card/``) so the registry stays a pure leaf
|
||||
and there are no import cycles. On disk these mirror the kernel-colocation answer: ``cpu`` is the
|
||||
unified numpy reference; a real ``torch_cuda`` backend would register both unified primitives and
|
||||
model-co-located fusion *variants* through this same API.
|
||||
"""
|
||||
@@ -0,0 +1,111 @@
|
||||
"""``accel`` — a pure-python stand-in accelerator backend.
|
||||
|
||||
There is no GPU in this environment, so to prove the dispatch substrate is genuinely
|
||||
device-generic (not a single-backend abstraction with a CPU special case) this backend registers a
|
||||
second device, ``accel``, entirely in CPU-resident Python. It exercises everything a real GPU
|
||||
backend would:
|
||||
|
||||
* a **component** override — ``AccelDiT`` is built via ``COMPONENTS[("dit", "accel", …)]`` instead
|
||||
of the card's numpy factory, while components it does NOT register (text_encoder, vae) fall back
|
||||
over the device chain ``accel → cpu`` to the numpy factory;
|
||||
* **kernels at different arch rungs** — ``flow_sde_step`` is registered at ``sm90`` directly, while
|
||||
``flow_match_step`` is registered only at ``generic``, so resolving it on an ``sm90`` accel
|
||||
platform must walk the arch fallback ``sm90 → sm80 → generic`` to find it;
|
||||
* the **parity oracle** — the ODE solver op (``_accel_flow_match_step``) is an independently
|
||||
written kernel that matches the numpy reference bit-for-bit, so a denoise run on the accel
|
||||
platform equals the cpu run. That independent-impl-vs-reference equality is the oracle a real
|
||||
backend must pass. (The SDE op below deliberately reuses the reference impl — a backend may
|
||||
legitimately not specialize every op — so the SDE test exercises the dispatch path, not an
|
||||
independent implementation.)
|
||||
|
||||
On a real box this file is where a torch/cuda backend would register its adapters and fused kernels
|
||||
(unified primitives + model-co-located fusion variants) through the same two functions.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.platform.backends.toy import ToyDiT
|
||||
from v2.platform.registry import FLOW_MATCH_STEP, FLOW_SDE_STEP, register_component, register_kernel
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Component: an accel-resident DiT (stand-in for a torch module placed on the #
|
||||
# device). Wraps the same toy compute so the parity oracle is bit-identical. #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class AccelDiT:
|
||||
"""A device-tagged DiT wrapper. On a real box this would hold a torch module on the accelerator;
|
||||
here it delegates to the same numpy toy so accel ≡ cpu numerically (the oracle)."""
|
||||
|
||||
def __init__(self, inner: Any):
|
||||
self._inner = inner
|
||||
self.device = "accel"
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
return self._inner(*args, **kwargs)
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
# delegate the trainable surface (clone/copy_from/blend_from/mse_grad_step) to the inner toy
|
||||
return getattr(self._inner, name)
|
||||
|
||||
|
||||
def _build_accel_dit(spec: Any, instance: Any, platform: Any) -> AccelDiT:
|
||||
# Build the same component the cpu rung would (via the card's factory), then place it "on device".
|
||||
inner = spec.factory(instance) if getattr(spec, "factory", None) is not None else ToyDiT()
|
||||
return AccelDiT(inner)
|
||||
|
||||
|
||||
register_component("dit", _build_accel_dit, device="accel", source="accel(stand-in):AccelDiT")
|
||||
|
||||
|
||||
class AccelComponent:
|
||||
"""Generic device-tagged wrapper for non-callable component kinds (vae, audio_vae, …). Delegates
|
||||
the full call surface (``decode``/``encode``/…) to the inner toy, so accel ≡ cpu numerically."""
|
||||
|
||||
def __init__(self, inner: Any):
|
||||
self._inner = inner
|
||||
self.device = "accel"
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
return getattr(self._inner, name)
|
||||
|
||||
|
||||
def _build_accel_component(spec: Any, instance: Any, platform: Any) -> AccelComponent:
|
||||
if getattr(spec, "factory", None) is None:
|
||||
raise RuntimeError(f"accel: no factory to wrap for kind={spec.kind!r}")
|
||||
return AccelComponent(spec.factory(instance))
|
||||
|
||||
|
||||
# Override more than one component kind on accel (proving the seam is kind-generic, not dit-special).
|
||||
# text_encoder is deliberately left UNregistered so the device→cpu fallback path stays demonstrated.
|
||||
for _kind in ("vae", "audio_vae"):
|
||||
register_component(_kind, _build_accel_component, device="accel", source="accel(stand-in):AccelComponent")
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Kernels: independent accel implementations that match the numpy reference. #
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _accel_flow_match_step(x_t, velocity, sigma_t: float, sigma_next: float):
|
||||
"""Stand-in 'fused' accel solver step. Identical math to the numpy reference (the oracle)."""
|
||||
return x_t + (sigma_next - sigma_t) * velocity
|
||||
|
||||
|
||||
# flow_match_step lives only at the `generic` arch rung → resolving on sm90 must fall back through
|
||||
# the arch chain (sm90 → sm80 → generic). This proves the arch-fallback walk.
|
||||
register_kernel(FLOW_MATCH_STEP,
|
||||
_accel_flow_match_step,
|
||||
device="accel",
|
||||
arch="generic",
|
||||
source="accel(stand-in):fused_flow_match")
|
||||
|
||||
# flow_sde_step is registered at sm90 directly. This backend does NOT specialize the SDE op — it
|
||||
# reuses the numpy reference, so the accel SDE path is the same function as cpu (the SDE parity test
|
||||
# therefore checks dispatch-path equivalence, not an independent impl; the ODE op above is the real
|
||||
# independent oracle). A device that DID specialize SDE would be checked the same way the ODE op is.
|
||||
from v2.loop.sampler import flow_sde_step_with_logprob # noqa: E402
|
||||
|
||||
register_kernel(FLOW_SDE_STEP,
|
||||
flow_sde_step_with_logprob,
|
||||
device="accel",
|
||||
arch="sm90",
|
||||
source="accel(stand-in):sde_logprob")
|
||||
@@ -0,0 +1,28 @@
|
||||
"""CPU / numpy backend — the terminal rung and parity oracle.
|
||||
|
||||
This backend registers the numpy reference kernels. It is the bottom of every device fallback
|
||||
chain: any op a richer backend hasn't implemented resolves here, and every other backend's output
|
||||
is checked against this one (the consistency ladder's C0/C1 oracle).
|
||||
|
||||
Components are NOT registered here: the card's ``ComponentSpec.factory`` *is* the cpu/numpy
|
||||
component rung (``Platform.build_component`` falls back to it when no ``(kind, device, …)`` cell
|
||||
matches). That keeps every existing card working with zero registration.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.loop.sampler import flow_match_euler_step, flow_sde_step_with_logprob
|
||||
from v2.platform.registry import FLOW_MATCH_STEP, FLOW_SDE_STEP, register_kernel
|
||||
|
||||
# The numpy solver primitives, registered as the terminal (cpu, numpy) kernels. These are the exact
|
||||
# same functions the loops used to call directly — so dispatching through the registry on a CPU box
|
||||
# is bit-for-bit identical to the pre-registry behavior.
|
||||
register_kernel(FLOW_MATCH_STEP,
|
||||
flow_match_euler_step,
|
||||
device="cpu",
|
||||
arch="numpy",
|
||||
source="loop.sampler:flow_match_euler_step")
|
||||
register_kernel(FLOW_SDE_STEP,
|
||||
flow_sde_step_with_logprob,
|
||||
device="cpu",
|
||||
arch="numpy",
|
||||
source="loop.sampler:flow_sde_step_with_logprob")
|
||||
@@ -0,0 +1,646 @@
|
||||
"""Real torch component adapters for the GPU backend. See v2/README.md ("Running the real models").
|
||||
|
||||
A ``TorchComponent`` base centralizes the shared mechanics — ``.to(device,dtype).eval()``, the
|
||||
numpy<->torch marshalling at the loop boundary, the ``set_forward_context`` wrap every fastvideo
|
||||
forward needs, and the weight surface (copy_from / blend_from / clone). Thin per-model subclasses
|
||||
carry only the forward semantics (Wan ``sigma*1000`` -> velocity; LTX-2 per-token sigma,
|
||||
x0 -> ``(x_t-x0)/sigma``, joint A/V; VAE normalization; T5 padding; Gemma dual-projection; upsampler;
|
||||
audio decode->vocoder). A single ``build_component(spec, instance, platform)`` dispatches by
|
||||
``spec.kind`` through ``_MAKERS``.
|
||||
|
||||
The loop math (CFG combine, flow-match Euler) is array-agnostic, so the boundary is the ONLY place
|
||||
host<->device crossing happens (``TorchComponent._t``/``_out``). When the card sets ``device_io`` (a GPU
|
||||
box), egress keeps tensors on-device and the latent stays resident for the whole denoise loop — no
|
||||
per-step round-trip; otherwise the default marshals to host numpy (un-migrated families / CPU toy). See
|
||||
``v2/platform/array_ns`` for the loop-side namespace. Component construction goes through
|
||||
``load_component``, the v2-owned loader seam (currently delegates to fastvideo's component_loader; a
|
||||
per-module vendored cutover replaces it later, no caller changes).
|
||||
|
||||
Imported ONLY lazily — inside ``torch_cuda.py``'s registered builder — so importing the platform package
|
||||
on a CPU box never imports torch and the mini stays green. Wan2.1 + LTX-2 A/V GPU-verified.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
# Flow-match timestep convention: the loop hands the raw sigma (1->0); the diffusers/FastVideo Wan
|
||||
# convention embeds ``timestep = sigma * num_train_timesteps`` (BRINGUP risk B, confirmed on box).
|
||||
NUM_TRAIN_TIMESTEPS = 1000
|
||||
|
||||
_RUNTIME_READY = False
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Construction helpers (FastVideoArgs / loaders / dist-init / marshalling) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _ensure_fastvideo_runtime() -> None:
|
||||
"""Single-process distributed env the real fastvideo loaders require (mirrors
|
||||
``fastvideo/worker/gpu_worker.py:init_device``): the loaders call ``get_local_torch_device()`` and
|
||||
build a 1x1 device mesh, which needs an initialized process group. Idempotent. BRINGUP risk A."""
|
||||
global _RUNTIME_READY
|
||||
if _RUNTIME_READY:
|
||||
return
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29500")
|
||||
os.environ.setdefault("RANK", "0")
|
||||
os.environ.setdefault("WORLD_SIZE", "1")
|
||||
os.environ.setdefault("LOCAL_RANK", "0")
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
|
||||
from v2.distributed import maybe_init_distributed_environment_and_model_parallel
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=1)
|
||||
_RUNTIME_READY = True
|
||||
|
||||
|
||||
def _require_checkpoint(spec) -> str:
|
||||
ckpt = getattr(spec, "checkpoint", "") or ""
|
||||
if not ckpt:
|
||||
raise RuntimeError(f"component {spec.component_id!r}: set ComponentSpec.checkpoint to the weights path / HF id "
|
||||
f"for the GPU backend (load_id={spec.load_id!r}). See GPU_BRINGUP.md, risk A.")
|
||||
return ckpt
|
||||
|
||||
|
||||
def _model_root(spec) -> str:
|
||||
"""Model root = the dir holding ``model_index.json`` + component subfolders. ``spec.checkpoint`` is
|
||||
the component subfolder (e.g. ``<root>/transformer``); its parent is the registry root. BRINGUP A."""
|
||||
return os.path.dirname(os.path.normpath(_require_checkpoint(spec)))
|
||||
|
||||
|
||||
def _fastvideo_args(spec: Any) -> Any:
|
||||
"""Build the FastVideoArgs the real loaders need (BRINGUP risk A). ``from_kwargs`` populates
|
||||
``pipeline_config`` (dit/vae/text-encoder configs + precisions) from the model root. Single-GPU,
|
||||
all offload/FSDP OFF — weights resident on one device, the simplest correct bring-up."""
|
||||
from v2.fastvideo_args import FastVideoArgs
|
||||
return FastVideoArgs.from_kwargs(model_path=_model_root(spec),
|
||||
num_gpus=1,
|
||||
tp_size=1,
|
||||
sp_size=1,
|
||||
dit_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
image_encoder_cpu_offload=False,
|
||||
dit_layerwise_offload=False,
|
||||
use_fsdp_inference=False,
|
||||
pin_cpu_memory=False)
|
||||
|
||||
|
||||
def load_component(loader_attr: str, path: str, args):
|
||||
"""v2-owned loader seam (``v2/loader``). Constructs a real component from a checkpoint via the
|
||||
component loaders. ``WanTransformer3DModel`` / ``AutoencoderKLWan`` have NO ``from_pretrained``; the
|
||||
loader reads the checkpoint config, resolves the class (UMT5 vs T5 from config), and loads weights.
|
||||
Currently delegates to fastvideo's loader; a vendored cutover swaps the body, not the callers."""
|
||||
from v2.loader import component_loader as _cl
|
||||
return getattr(_cl, loader_attr)().load(path, args)
|
||||
|
||||
|
||||
def _device(platform) -> str:
|
||||
return "cuda" if platform.device == "cuda" else platform.device
|
||||
|
||||
|
||||
def _native_dtype(module: Any) -> Any:
|
||||
"""Keep each component at the precision its loader produced (Wan DiT bf16, VAE/text fp32) rather than
|
||||
the card's uniform fp32 — faithful to fastvideo and ~2x faster/smaller for the DiT."""
|
||||
try:
|
||||
return next(module.parameters()).dtype
|
||||
except StopIteration:
|
||||
return torch.float32
|
||||
|
||||
|
||||
def _to_torch(a: Any, *, device: Any, dtype: Any) -> torch.Tensor | None:
|
||||
return None if a is None else torch.as_tensor(np.asarray(a), dtype=dtype, device=device)
|
||||
|
||||
|
||||
def _to_numpy(t: Any) -> np.ndarray:
|
||||
if hasattr(t, "sample"): # some heads wrap output in an object with .sample
|
||||
t = t.sample
|
||||
return t.detach().to("cpu", torch.float32).numpy()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# TorchComponent — the shared base (eval/marshalling/forward-context/weights) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class TorchComponent:
|
||||
"""Wraps a real ``fastvideo.models.*`` module to the mini's numpy duck-typed surface. Subclasses
|
||||
override only the forward semantics (``__call__`` / ``encode`` / ``decode`` / ``upsample`` / ...)."""
|
||||
|
||||
def __init__(self, module: Any, *, device: Any, dtype: Any, eager: bool = True) -> None:
|
||||
self.device, self.dtype = device, dtype
|
||||
self.module = module.to(device=device, dtype=dtype).eval() if eager else module
|
||||
# When True (set by build_component from the card's device_io flag on a GPU box), the loop
|
||||
# boundary keeps tensors on-device — no per-step host<->device copy. Default False = today's
|
||||
# numpy in/out, so un-migrated families and the CPU toy are unchanged.
|
||||
self.device_io = False
|
||||
|
||||
# host<->device marshalling at the loop boundary (ONE place) ------------------------------------- #
|
||||
def _t(self, a: Any, *, batch: bool = True) -> Any:
|
||||
if a is None:
|
||||
return None
|
||||
# Ingress: an already-on-device tensor (the loop kept it resident) is moved to our dtype, never
|
||||
# round-tripped through numpy; a host array is uploaded here.
|
||||
t = a.to(device=self.device, dtype=self.dtype) if torch.is_tensor(a) else _to_torch(
|
||||
a, device=self.device, dtype=self.dtype)
|
||||
return t.unsqueeze(0) if batch else t
|
||||
|
||||
def _out(self, t: Any) -> Any:
|
||||
"""Egress: keep the tensor on-device when ``device_io`` (it stays resident for the next step),
|
||||
else marshal to host numpy (the default — un-migrated families / the CPU toy are unchanged).
|
||||
Either way it lands in fp32 — matching the host path's ``_to_numpy`` (``.to(cpu, float32)``) so
|
||||
the loop's CFG combine runs in fp32 regardless of the module's native dtype (e.g. Wan DiT bf16)."""
|
||||
if not self.device_io:
|
||||
return _to_numpy(t)
|
||||
if hasattr(t, "sample"): # unwrap heads that wrap output in an object with .sample
|
||||
t = t.sample
|
||||
return t.float()
|
||||
|
||||
def _n(self, t: Any) -> Any:
|
||||
return self._out(t.squeeze(0))
|
||||
|
||||
@staticmethod
|
||||
def _ctx(current_timestep: Any = 0) -> Any:
|
||||
"""The FastVideo attention layer reads attn_metadata via get_forward_context(), so every forward
|
||||
runs inside set_forward_context(...). attn_metadata=None selects the dense SDPA path."""
|
||||
from v2.forward_context import set_forward_context
|
||||
return set_forward_context(current_timestep=current_timestep, attn_metadata=None)
|
||||
|
||||
# weight surface used by serving weight-sync + training ----------------------------------------- #
|
||||
def copy_from(self, other) -> None:
|
||||
self.module.load_state_dict(other.module.state_dict())
|
||||
|
||||
def blend_from(self, other, decay: float) -> None: # EMA / decayed-old-policy
|
||||
with torch.no_grad():
|
||||
for p, q in zip(self.module.parameters(), other.module.parameters(), strict=False):
|
||||
p.mul_(decay).add_(q, alpha=1.0 - decay)
|
||||
|
||||
def clone(self) -> TorchComponent:
|
||||
import copy
|
||||
c = self.__class__.__new__(self.__class__)
|
||||
c.__dict__.update(self.__dict__) # tokenizer / shapes / sibling refs shared
|
||||
c.module = copy.deepcopy(self.module) # independent weights
|
||||
return c
|
||||
|
||||
def mse_grad_step(self, *a, **k):
|
||||
raise NotImplementedError(
|
||||
"GPU training surface (mse_grad_step) is a separate workstream — see GPU_BRINGUP.md risk F")
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# DiT adapters #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class WanDiT(TorchComponent):
|
||||
"""dit(latent[C,T,H,W], text_embed[seq,dim], sigma) -> velocity[C,T,H,W]. Real forward
|
||||
(wanvideo.py): forward(hidden_states[B,C,T,H,W], encoder_hidden_states, timestep,
|
||||
encoder_hidden_states_image=None) -> velocity (bare tensor)."""
|
||||
|
||||
def __init__(self,
|
||||
module: Any,
|
||||
*,
|
||||
device: Any,
|
||||
dtype: Any,
|
||||
offload_group: Any = None,
|
||||
component_id: str = "transformer") -> None:
|
||||
# Wan2.2 MoE (A14B): two 14B experts don't both fit one 80GB GPU. With ``offload_group`` set, keep
|
||||
# this expert on CPU and bring only the *active* one onto the GPU on demand (single swap at the
|
||||
# boundary, not per-step thrash). Single-expert Wan stays resident (offload_group=None).
|
||||
self.offload_group = offload_group
|
||||
self.component_id = component_id
|
||||
super().__init__(module, device=device, dtype=dtype, eager=(offload_group is None))
|
||||
if offload_group is not None:
|
||||
self.module = self.module.to(device="cpu", dtype=dtype).eval()
|
||||
self._on_gpu = False
|
||||
offload_group[component_id] = self
|
||||
else:
|
||||
self._on_gpu = True
|
||||
# CausalWanTransformer3DModel (self-forcing student) conditions across chunks via an internal
|
||||
# kv_cache, not a forward arg, so the chunk_rollout loop's latent ``context`` is ignored here.
|
||||
self.causal = "Causal" in type(module).__name__
|
||||
|
||||
def _ensure_resident(self) -> None:
|
||||
if self.offload_group is None or self._on_gpu:
|
||||
return
|
||||
for other in self.offload_group.values():
|
||||
if other is not self and other._on_gpu:
|
||||
other.module.to("cpu")
|
||||
other._on_gpu = False
|
||||
torch.cuda.empty_cache()
|
||||
self.module.to(self.device)
|
||||
self._on_gpu = True
|
||||
|
||||
def alloc_causal_caches(self, frame_seqlen: int, *, batch: int = 1, max_text_len: int = 512) -> Any:
|
||||
"""Allocate the persistent per-block KV + cross-attn caches the causal student needs for
|
||||
cross-chunk conditioning (mirrors CausalDenoisingStage._initialize_kv_cache). The chunk_rollout
|
||||
loop owns these in its LoopState (per request, never on this shared adapter) and threads them
|
||||
through __call__ so the model routes to _forward_inference."""
|
||||
self._ensure_resident()
|
||||
m = self.module
|
||||
nblocks = len(m.blocks)
|
||||
nheads, hdim = int(m.num_attention_heads), int(m.attention_head_dim)
|
||||
local = int(getattr(m, "local_attn_size", -1) or -1)
|
||||
if local != -1:
|
||||
kv_size = local * frame_seqlen
|
||||
else: # global window: model keeps the last GLOBAL_ATTN_COMPAT_MAX_LATENT_FRAMES (21) frames
|
||||
arch = getattr(getattr(m, "config", None), "arch_config", None)
|
||||
sliding = int(getattr(arch, "sliding_window_num_frames", 0) or 0)
|
||||
kv_size = max(sliding, 21) * frame_seqlen
|
||||
dev, dt = self.device, self.dtype
|
||||
kv = [{
|
||||
"k": torch.zeros(batch, kv_size, nheads, hdim, device=dev, dtype=dt),
|
||||
"v": torch.zeros(batch, kv_size, nheads, hdim, device=dev, dtype=dt),
|
||||
"global_end_index": torch.tensor([0], dtype=torch.long, device=dev),
|
||||
"local_end_index": torch.tensor([0], dtype=torch.long, device=dev),
|
||||
} for _ in range(nblocks)]
|
||||
ca = [{
|
||||
"k": torch.zeros(batch, max_text_len, nheads, hdim, device=dev, dtype=dt),
|
||||
"v": torch.zeros(batch, max_text_len, nheads, hdim, device=dev, dtype=dt),
|
||||
"is_init": False,
|
||||
} for _ in range(nblocks)]
|
||||
return {"kv": kv, "crossattn": ca}
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self,
|
||||
latent,
|
||||
text_embed,
|
||||
sigma,
|
||||
context=None,
|
||||
*,
|
||||
cond=None,
|
||||
kv_cache=None,
|
||||
crossattn_cache=None,
|
||||
current_start=0,
|
||||
start_frame=0,
|
||||
frame_seqlen=1560):
|
||||
self._ensure_resident()
|
||||
hs = self._t(latent)
|
||||
if cond is not None: # i2v: concat [noise (16ch) ; mask+cond_latent (20ch)] -> 36ch DiT input
|
||||
hs = torch.cat([hs, self._t(cond)], dim=1)
|
||||
ehs = self._t(text_embed)
|
||||
# timestep = sigma*1000 (BRINGUP risk B). Causal model needs per-latent-frame [B, num_frames].
|
||||
ts = float(sigma) * NUM_TRAIN_TIMESTEPS
|
||||
timestep = (torch.full(
|
||||
(1, hs.shape[2]), ts, device=self.device) if self.causal else torch.tensor([ts], device=self.device))
|
||||
img = None if self.causal else self._t(context) # i2v image embedding for standard Wan
|
||||
fwd = dict(hidden_states=hs, encoder_hidden_states=ehs, timestep=timestep, encoder_hidden_states_image=img)
|
||||
if self.causal and kv_cache is not None:
|
||||
# Pass the persistent caches + frame offset so the model runs _forward_inference (CausVid
|
||||
# Algorithm 2) and conditions this chunk on prior clean chunks. Without it the model falls
|
||||
# through to _forward_train (no cross-chunk KV) -> chunks denoise blind -> discontinuity.
|
||||
fwd.update(kv_cache=kv_cache,
|
||||
crossattn_cache=crossattn_cache,
|
||||
current_start=int(current_start),
|
||||
cache_start=int(current_start),
|
||||
start_frame=int(start_frame),
|
||||
frame_seqlen=int(frame_seqlen))
|
||||
with self._ctx():
|
||||
velocity = self.module(**fwd)
|
||||
return self._n(velocity) # rectified-flow velocity (BRINGUP risk C)
|
||||
|
||||
def clone(self) -> WanDiT:
|
||||
c = super().clone()
|
||||
assert isinstance(c, WanDiT)
|
||||
c.offload_group, c._on_gpu = None, True # standalone resident copy
|
||||
return c
|
||||
|
||||
|
||||
class LTX2DiT(TorchComponent):
|
||||
"""LTX-2 DiT: patchifies internally, takes a PER-TOKEN timestep ``ones(B, token_count, 1) * sigma``
|
||||
(sigma direct, 0..1) + a per-sample ``video_sigma``; predicts ``denoised`` (x0), so the adapter
|
||||
returns the flow-match velocity ``(x_t - x0)/sigma`` the loop integrates. Joint A/V (audio_latent
|
||||
given) -> (video_velocity, audio_velocity) in one forward."""
|
||||
|
||||
def __init__(self, module: Any, *, device: Any, dtype: Any) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
from v2.models.dits.ltx2 import VideoLatentShape
|
||||
self._VideoLatentShape = VideoLatentShape
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self, latent, text_embed, sigma, context=None, *, audio_latent=None, audio_text=None):
|
||||
hs = self._t(latent)
|
||||
ehs = self._t(text_embed)
|
||||
s = float(sigma)
|
||||
token_count = self.module.patchifier.get_token_count(self._VideoLatentShape.from_torch_shape(tuple(hs.shape)))
|
||||
timestep = torch.full((1, token_count, 1), s, device=self.device, dtype=torch.float32)
|
||||
video_sigma = torch.tensor([s], device=self.device, dtype=torch.float32)
|
||||
av_kwargs: dict = {}
|
||||
au = None
|
||||
if audio_latent is not None: # joint A/V: audio latent [1,8,T,16] + audio text
|
||||
from v2.models.audio.ltx2_audio_vae import AudioLatentShape
|
||||
au = self._t(audio_latent)
|
||||
aeh = self._t(audio_text)
|
||||
atok = self.module.audio_patchifier.get_token_count(AudioLatentShape.from_torch_shape(tuple(au.shape)))
|
||||
av_kwargs = dict(audio_hidden_states=au,
|
||||
audio_encoder_hidden_states=aeh,
|
||||
audio_timestep=torch.full((1, atok, 1), s, device=self.device, dtype=torch.float32),
|
||||
audio_sigma=torch.tensor([s], device=self.device, dtype=torch.float32))
|
||||
with self._ctx(current_timestep=s):
|
||||
out = self.module(hidden_states=hs,
|
||||
encoder_hidden_states=ehs,
|
||||
timestep=timestep,
|
||||
video_sigma=video_sigma,
|
||||
encoder_attention_mask=None,
|
||||
**av_kwargs)
|
||||
if au is not None: # LTX-2 DiT predicts x0; integrate velocity
|
||||
denoised_v, denoised_a = out
|
||||
vel_v = (hs.float() - denoised_v.float()) / max(s, 1e-6)
|
||||
vel_a = (au.float() - denoised_a.float()) / max(s, 1e-6)
|
||||
return self._n(vel_v), self._n(vel_a)
|
||||
denoised = out[0] if isinstance(out, tuple) else out
|
||||
return self._n((hs.float() - denoised.float()) / max(s, 1e-6))
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# VAE adapters #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class WanVAE(TorchComponent):
|
||||
"""The DiT operates in NORMALIZED latent space, so encode applies ``(z - mean)/std`` and decode
|
||||
inverts it; ``AutoencoderKLWan.encode/decode`` operate in RAW latent space (BRINGUP risk D)."""
|
||||
|
||||
def _mean_invstd(self, like: Any) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
mean = torch.tensor(self.module.latents_mean, device=like.device, dtype=torch.float32).view(1, -1, 1, 1, 1)
|
||||
inv_std = (1.0 / torch.tensor(self.module.latents_std, device=like.device, dtype=torch.float32)).view(
|
||||
1, -1, 1, 1, 1)
|
||||
return mean, inv_std
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, video):
|
||||
x = self._t(video)
|
||||
dist = self.module.encode(x)
|
||||
z = dist.mode() if hasattr(dist, "mode") else (dist.sample() if hasattr(dist, "sample") else dist)
|
||||
mean, inv_std = self._mean_invstd(z)
|
||||
return self._n((z.float() - mean) * inv_std) # -> the normalized latent the DiT expects
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latent):
|
||||
z = self._t(latent).float()
|
||||
mean, inv_std = self._mean_invstd(z)
|
||||
z = z / inv_std + mean # invert encode normalization -> raw latent
|
||||
video = self.module.decode(z.to(self.dtype)) # -> video [B,3,T,H,W] in [-1,1]
|
||||
return self._n(video)
|
||||
|
||||
|
||||
class LTX2VAE(TorchComponent):
|
||||
"""LTX-2 VideoDecoder un-normalizes internally (per-channel stats); the adapter just marshals."""
|
||||
|
||||
def __init__(self, module: Any, *, device: Any, dtype: Any) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
# Tiled decode (spatial tiles concatenated): the LTX-2 CausalVideoAutoencoder decode of a full-res,
|
||||
# full-length (121f / 1536x1024) latent allocates >80GB of conv activations otherwise. fastvideo
|
||||
# defaults vae_tiling=True for LTX-2; mirror it so the full clip fits on one GPU.
|
||||
if hasattr(self.module, "enable_tiling"):
|
||||
self.module.enable_tiling()
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latent):
|
||||
video = self.module.decode(self._t(latent))
|
||||
if hasattr(video, "sample"):
|
||||
video = video.sample
|
||||
return self._n(video)
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, video):
|
||||
dist = self.module.encode(self._t(video))
|
||||
z = dist.mode() if hasattr(dist, "mode") else (dist.sample() if hasattr(dist, "sample") else dist)
|
||||
return self._n(z)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Text encoders #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class T5Encoder(TorchComponent):
|
||||
"""(U)MT5 text_encoder.encode(text) -> embedding[text_len, dim] (numpy out)."""
|
||||
|
||||
def __init__(self, module: Any, tokenizer: Any, *, device: Any, dtype: Any, max_length: int = 512) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, text):
|
||||
toks = self.tokenizer(text or "", return_tensors="pt", max_length=self.max_length, truncation=True)
|
||||
ids = toks.input_ids.to(self.device)
|
||||
mask = toks.attention_mask.to(self.device)
|
||||
with self._ctx():
|
||||
out = self.module(input_ids=ids, attention_mask=mask)
|
||||
hidden = (out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0]).squeeze(0)
|
||||
# Wan convention: the DiT cross-attends over a FIXED-length text sequence — real-token rows then
|
||||
# a ZERO-padded tail to text_len (a short unpadded sequence mis-conditions the DiT).
|
||||
if hidden.shape[0] < self.max_length:
|
||||
pad = hidden.new_zeros(self.max_length - hidden.shape[0], hidden.shape[1])
|
||||
hidden = torch.cat([hidden, pad], dim=0)
|
||||
return self._out(hidden)
|
||||
|
||||
|
||||
class Gemma(TorchComponent):
|
||||
"""LTX-2 Gemma text encoder. ``encode`` -> video projection; ``encode_av`` -> (video, audio): the 2.3
|
||||
connector emits SEPARATE projections (audio in ``hidden_states[0]`` only when output_hidden_states)."""
|
||||
|
||||
def __init__(self, module: Any, tokenizer: Any, *, device: Any, dtype: Any, max_length: int = 256) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
self.tokenizer = tokenizer
|
||||
self.max_length = max_length
|
||||
|
||||
def _tok(self, text: Any) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
toks = self.tokenizer(text or "",
|
||||
return_tensors="pt",
|
||||
max_length=self.max_length,
|
||||
truncation=True,
|
||||
padding="max_length")
|
||||
return toks.input_ids.to(self.device), toks.attention_mask.to(self.device)
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, text):
|
||||
ids, mask = self._tok(text)
|
||||
with self._ctx():
|
||||
out = self.module(input_ids=ids, attention_mask=mask)
|
||||
hidden = out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0]
|
||||
return _to_numpy(hidden.squeeze(0))
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_av(self, text):
|
||||
ids, mask = self._tok(text)
|
||||
with self._ctx():
|
||||
out = self.module(input_ids=ids, attention_mask=mask, output_hidden_states=True)
|
||||
video = out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0]
|
||||
hs = getattr(out, "hidden_states", None)
|
||||
audio = hs[0] if hs else video # separate audio connector projection
|
||||
return _to_numpy(video.squeeze(0)), _to_numpy(audio.squeeze(0))
|
||||
|
||||
|
||||
class CLIPImageEncoder(TorchComponent):
|
||||
"""CLIP-vision image encoder for Wan i2v: ``encode_image(image) -> image_embeds`` (the DiT's
|
||||
``encoder_hidden_states_image``). Mirrors fastvideo's ImageEncodingStage — the HF image processor
|
||||
preprocesses, the encoder returns ``last_hidden_state``. BRINGUP: written-not-run; GPU-verify the
|
||||
processor/encoder subfolders + dtype against a real i2v checkpoint (see GPU_BRINGUP.md)."""
|
||||
|
||||
def __init__(self, module: Any, processor: Any, *, device: Any, dtype: Any) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
self.processor = processor
|
||||
|
||||
@torch.no_grad()
|
||||
def encode_image(self, image):
|
||||
inputs = self.processor(images=image, return_tensors="pt").to(self.device)
|
||||
with self._ctx():
|
||||
out = self.module(**inputs)
|
||||
embeds = out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0]
|
||||
return self._out(embeds.squeeze(0))
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Upsampler + audio (LTX-2) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
class LTX2Upsampler(TorchComponent):
|
||||
"""upsample(latent[C,T,H,W]) -> [C,T,2H,2W]. Applies the repo's ``upsample_video``: un_normalize (via
|
||||
the video VAE's per_channel_statistics) -> learned 2x upsample -> normalize."""
|
||||
|
||||
def __init__(self, module: Any, stats_owner: Any, *, device: Any, dtype: Any) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
self.stats_owner = stats_owner # exposes per_channel_statistics (the VAE decoder/encoder)
|
||||
|
||||
@torch.no_grad()
|
||||
def upsample(self, latent):
|
||||
from v2.models.upsamplers.ltx2_upsampler import upsample_video
|
||||
out = upsample_video(self._t(latent), self.stats_owner, self.module)
|
||||
return self._n(out)
|
||||
|
||||
|
||||
class LTX2Vocoder(TorchComponent):
|
||||
"""Thin wrapper over the LTX-2 Vocoder (mel -> waveform @24kHz); chained by LTX2AudioVAE."""
|
||||
|
||||
|
||||
class LTX2AudioVAE(TorchComponent):
|
||||
"""audio_vae.decode(audio_latent[8,T,16]) -> waveform. Runs AudioDecoder (latent->mel) then the
|
||||
Vocoder (mel->waveform). ``sample_rate`` = the vocoder's output rate (so a saver writes the right wav)."""
|
||||
|
||||
def __init__(self, module: Any, vocoder: Any, *, device: Any, dtype: Any) -> None:
|
||||
super().__init__(module, device=device, dtype=dtype)
|
||||
self.vocoder = vocoder # LTX2Vocoder | None
|
||||
self.sample_rate = int(getattr(getattr(vocoder, "module", None), "output_sample_rate", 24000))
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, audio_latent):
|
||||
mel = self.module(self._t(audio_latent))
|
||||
if hasattr(mel, "sample"):
|
||||
mel = mel.sample
|
||||
if self.vocoder is not None:
|
||||
wav = self.vocoder.module(mel)
|
||||
if hasattr(wav, "sample"):
|
||||
wav = wav.sample
|
||||
else:
|
||||
wav = mel
|
||||
return self._n(wav)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# build_component — one dispatch (replaces the six build_torch_* + trampolines) #
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _explicit_adapter(spec: Any, module: Any, device: Any, dtype: Any, *extra: Any) -> Any:
|
||||
"""Build the card's explicitly-declared adapter (``ComponentSpec.adapter='module:Class'``) if set —
|
||||
the seam that lets a NEW architecture's recipe carry its own TorchComponent subclass without editing
|
||||
the dispatch here (so a port is a self-contained recipe package). Constructed as
|
||||
``cls(module, *extra, device=device, dtype=dtype)``. Returns None when unset (-> built-in dispatch)."""
|
||||
ref = getattr(spec, "adapter", "") or ""
|
||||
if not ref:
|
||||
return None
|
||||
import importlib
|
||||
mod, _, cls = ref.partition(":")
|
||||
return getattr(importlib.import_module(mod), cls)(module, *extra, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def _make_dit(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("TransformerLoader", spec.checkpoint, args)
|
||||
device, dtype = _device(platform), _native_dtype(module)
|
||||
explicit = _explicit_adapter(spec, module, device, dtype)
|
||||
if explicit is not None:
|
||||
return explicit
|
||||
if "LTX2" in type(module).__name__:
|
||||
return LTX2DiT(module, device=device, dtype=dtype)
|
||||
# Wan2.2 MoE: >1 DiT expert -> CPU-offload all but the active one (swapped at the boundary).
|
||||
n_experts = sum(1 for c in instance.card.components.values() if getattr(c, "kind", None) == "dit")
|
||||
grp = None
|
||||
if n_experts > 1:
|
||||
grp = getattr(instance, "_dit_offload_group", None)
|
||||
if grp is None:
|
||||
grp = instance._dit_offload_group = {}
|
||||
return WanDiT(module, device=device, dtype=dtype, offload_group=grp, component_id=spec.component_id)
|
||||
|
||||
|
||||
def _make_vae(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("VAELoader", spec.checkpoint, args)
|
||||
device, dtype = _device(platform), _native_dtype(module)
|
||||
explicit = _explicit_adapter(spec, module, device, dtype)
|
||||
if explicit is not None:
|
||||
return explicit
|
||||
cls = LTX2VAE if "LTX2" in type(module).__name__ else WanVAE
|
||||
return cls(module, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def _make_text_encoder(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("TextEncoderLoader", spec.checkpoint, args)
|
||||
tokenizer = load_component("TokenizerLoader", os.path.join(_model_root(spec), "tokenizer"), args)
|
||||
device, dtype = _device(platform), _native_dtype(module)
|
||||
explicit = _explicit_adapter(spec, module, device, dtype, tokenizer) # adapter cls(module, tokenizer, ...)
|
||||
if explicit is not None:
|
||||
return explicit
|
||||
cls = Gemma if "Gemma" in type(module).__name__ else T5Encoder
|
||||
return cls(module, tokenizer, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def _make_image_encoder(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("ImageEncoderLoader", spec.checkpoint, args) # CLIP vision
|
||||
# the HF image processor is a sibling subfolder (BRINGUP: confirm path on a real i2v checkpoint)
|
||||
processor = load_component("ImageProcessorLoader", os.path.join(_model_root(spec), "image_processor"), args)
|
||||
return CLIPImageEncoder(module, processor, device=_device(platform), dtype=_native_dtype(module))
|
||||
|
||||
|
||||
def _make_upsampler(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("UpsamplerLoader", spec.checkpoint, args) # LTX2LatentUpsampler
|
||||
vae = instance.component("vae")
|
||||
dtype = getattr(vae, "dtype", _native_dtype(module))
|
||||
vae_module = getattr(vae, "module", None)
|
||||
# per_channel_statistics lives on the AE's decoder/encoder sub-module, not the top-level autoencoder.
|
||||
stats_owner = getattr(vae_module, "decoder", None) or getattr(vae_module, "encoder", None) or vae_module
|
||||
return LTX2Upsampler(module, stats_owner, device=_device(platform), dtype=dtype)
|
||||
|
||||
|
||||
def _make_audio_vae(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("AudioDecoderLoader", spec.checkpoint, args) # LTX2AudioDecoder
|
||||
voc = instance.component("vocoder") # chained: decoder -> vocoder
|
||||
return LTX2AudioVAE(module, voc, device=_device(platform), dtype=_native_dtype(module))
|
||||
|
||||
|
||||
def _make_vocoder(spec: Any, instance: Any, platform: Any, args: Any) -> TorchComponent:
|
||||
module = load_component("VocoderLoader", spec.checkpoint, args) # LTX2Vocoder
|
||||
return LTX2Vocoder(module, device=_device(platform), dtype=_native_dtype(module))
|
||||
|
||||
|
||||
_MAKERS = {
|
||||
"dit": _make_dit,
|
||||
"vae": _make_vae,
|
||||
"text_encoder": _make_text_encoder,
|
||||
"image_encoder": _make_image_encoder,
|
||||
"upsampler": _make_upsampler,
|
||||
"audio_vae": _make_audio_vae,
|
||||
"vocoder": _make_vocoder,
|
||||
}
|
||||
|
||||
|
||||
def build_component(spec: Any, instance: Any, platform: Any) -> TorchComponent:
|
||||
"""The single cuda component builder (registered for every kind in ``torch_cuda.py``). Shared prefix
|
||||
— checkpoint check, dist-init, FastVideoArgs — then dispatch by ``spec.kind`` to its maker."""
|
||||
_require_checkpoint(spec) # fail fast on a mis-stamped card, before any dist init
|
||||
_ensure_fastvideo_runtime()
|
||||
args = _fastvideo_args(spec)
|
||||
try:
|
||||
maker = _MAKERS[spec.kind]
|
||||
except KeyError:
|
||||
raise RuntimeError(f"torch backend: no builder for component kind {spec.kind!r} "
|
||||
f"(have {sorted(_MAKERS)})") from None
|
||||
comp = maker(spec, instance, platform, args)
|
||||
# On-device I/O opt-in: keep this component's loop boundary on-device when the card declares it
|
||||
# (and we're actually on cuda). Off by default -> numpy in/out (un-migrated families unchanged).
|
||||
if isinstance(comp, TorchComponent):
|
||||
comp.device_io = bool(platform.device == "cuda"
|
||||
and getattr(getattr(instance, "card", None), "device_io", False))
|
||||
return comp
|
||||
@@ -0,0 +1,77 @@
|
||||
"""``cuda`` registration — the real torch/GPU backend (see v2/README.md "Running the real models").
|
||||
|
||||
Torch-free at import (stdlib only): gates every cell on ``available=_cuda_available`` (a ``find_spec``
|
||||
check that never imports torch) and registers ONE lazy component builder for all cuda component kinds +
|
||||
the two solver ops. So ``ensure_backends_loaded()`` runs on every box without importing torch, the matrix
|
||||
stays dumpable (these rows ``available=False`` on a CPU box), and ``resolve_first`` skips them.
|
||||
|
||||
On a real box ``available`` flips True and these resolve: components -> the single generic builder in
|
||||
``torch_backend.py`` (which dispatches by ``spec.kind``, wraps the real ``fastvideo.models.*`` module the
|
||||
card's load_id names — weights from ``ComponentSpec.checkpoint`` via the ``v2.loader`` seam); the solver
|
||||
ops -> plain torch elementwise in ``torch_kernels.py`` (no fused solver kernel exists in fastvideo-kernel).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
|
||||
from v2.platform.registry import FLOW_MATCH_STEP, FLOW_SDE_STEP, register_component, register_kernel
|
||||
|
||||
|
||||
def _torch_present() -> bool:
|
||||
return importlib.util.find_spec("torch") is not None
|
||||
|
||||
|
||||
def _cuda_available() -> bool:
|
||||
if not _torch_present():
|
||||
return False
|
||||
try:
|
||||
import torch # type: ignore
|
||||
return bool(torch.cuda.is_available())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# --- ONE lazy builder for every cuda component kind: imports torch only when called (on a GPU box) --- #
|
||||
def _build_component(spec, instance, platform):
|
||||
from v2.platform.backends.torch_backend import build_component
|
||||
return build_component(spec, instance, platform)
|
||||
|
||||
|
||||
def _flow_match_cuda(*args, **kwargs):
|
||||
from v2.platform.backends.torch_kernels import flow_match_step
|
||||
return flow_match_step(*args, **kwargs)
|
||||
|
||||
|
||||
def _flow_sde_cuda(*args, **kwargs):
|
||||
from v2.platform.backends.torch_kernels import flow_sde_step
|
||||
return flow_sde_step(*args, **kwargs)
|
||||
|
||||
|
||||
# --- components: one generic builder; torch_backend.build_component dispatches by spec.kind ---------- #
|
||||
_B = "v2.platform.backends.torch_backend"
|
||||
_KIND_SOURCE = {
|
||||
"dit": "WanDiT/LTX2DiT",
|
||||
"vae": "WanVAE/LTX2VAE",
|
||||
"text_encoder": "T5Encoder/Gemma",
|
||||
"image_encoder": "CLIPImageEncoder",
|
||||
"upsampler": "LTX2Upsampler",
|
||||
"audio_vae": "LTX2AudioVAE",
|
||||
"vocoder": "LTX2Vocoder",
|
||||
}
|
||||
for _kind, _cls in _KIND_SOURCE.items():
|
||||
register_component(_kind, _build_component, device="cuda", available=_cuda_available, source=f"{_B}:{_cls}")
|
||||
del _kind, _cls
|
||||
|
||||
# --- solver ops: plain torch elementwise (no fused solver kernel in fastvideo-kernel) ---------------- #
|
||||
register_kernel(FLOW_MATCH_STEP,
|
||||
_flow_match_cuda,
|
||||
device="cuda",
|
||||
arch="generic",
|
||||
available=_cuda_available,
|
||||
source="torch elementwise (no fused solver kernel)")
|
||||
register_kernel(FLOW_SDE_STEP,
|
||||
_flow_sde_cuda,
|
||||
device="cuda",
|
||||
arch="generic",
|
||||
available=_cuda_available,
|
||||
source="torch elementwise (no fused solver kernel)")
|
||||
@@ -0,0 +1,58 @@
|
||||
"""Torch solver-step kernels for the GPU backend.
|
||||
|
||||
There is NO fused flow-match / SDE *solver* kernel in fastvideo-kernel — it ships only primitives
|
||||
(sparse/tiled attention, RMS/LayerNorm, int8 GEMM/quant). So the GPU solver step is a plain torch
|
||||
elementwise op with the SAME math as ``loop/sampler.py``'s numpy reference (which it matches bit-for-bit).
|
||||
|
||||
Imports torch; loaded only lazily from ``torch_cuda.py``. Under the current numpy loop surface these
|
||||
marshal numpy<->torch at the boundary (a torch-native surface would keep the latent on-device through
|
||||
forward->combine->solver — a perf follow-up).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def _as_torch(a: Any) -> tuple[Any, bool]:
|
||||
"""numpy → cuda torch (passthrough if already torch). Returns (tensor, was_numpy)."""
|
||||
if torch.is_tensor(a):
|
||||
return a, False
|
||||
return torch.as_tensor(np.asarray(a), device="cuda"), True
|
||||
|
||||
|
||||
def _back(t, like_numpy: bool):
|
||||
return t.detach().to("cpu", torch.float32).numpy() if like_numpy else t
|
||||
|
||||
|
||||
def flow_match_step(x_t, velocity, sigma_t, sigma_next):
|
||||
"""One deterministic flow-match Euler step — identical to sampler.flow_match_euler_step."""
|
||||
xt, was_np = _as_torch(x_t)
|
||||
v, _ = _as_torch(velocity)
|
||||
return _back(xt + (float(sigma_next) - float(sigma_t)) * v, was_np)
|
||||
|
||||
|
||||
def flow_sde_step(x_t, velocity, sigma_t, sigma_next, *, noise=None, prev_sample=None, noise_scale=0.7):
|
||||
"""FlowGRPO stochastic step + Gaussian log-prob — the torch port of
|
||||
sampler.flow_sde_step_with_logprob. Returns (prev_sample, log_prob, mean, eff_std)."""
|
||||
xt, was_np = _as_torch(x_t)
|
||||
v, _ = _as_torch(velocity)
|
||||
xt = xt.double()
|
||||
v = v.double()
|
||||
s = min(float(sigma_t), 0.9999)
|
||||
dt = float(sigma_next) - float(sigma_t)
|
||||
std = float((s / (1.0 - s))**0.5 * noise_scale)
|
||||
denom = 2.0 * max(s, 1e-6)
|
||||
mean = xt * (1.0 + std**2 / denom * dt) + v * (1.0 + std**2 * (1.0 - s) / denom) * dt
|
||||
eff_std = max(std * max(-dt, 1e-12)**0.5, 1e-6)
|
||||
if prev_sample is None:
|
||||
n = _as_torch(noise)[0].double() if noise is not None else torch.zeros_like(xt)
|
||||
prev = mean + eff_std * n
|
||||
else:
|
||||
prev = _as_torch(prev_sample)[0].double()
|
||||
var = eff_std**2
|
||||
log_prob = float(
|
||||
(-((prev - mean)**2) / (2.0 * var) - float(np.log(eff_std)) - 0.5 * float(np.log(2.0 * np.pi))).mean().item())
|
||||
return (_back(prev.float(), was_np), log_prob, _back(mean.float(), was_np), float(eff_std))
|
||||
@@ -0,0 +1,476 @@
|
||||
"""Toy numpy components — the CPU-testable backend.
|
||||
|
||||
There is no GPU/torch/weights in this environment, so the heavy Wan/LTX neural forwards are
|
||||
represented by small, deterministic numpy components. They exercise the *real* loop control
|
||||
flow, CFG/flow-shift policies, scheduler steps, cache reuse, parity gates, and training-method
|
||||
math with real numbers — just with a toy network instead of a 1.3B DiT.
|
||||
|
||||
The contract these toys honor is what matters:
|
||||
* deterministic given (weights, inputs) → bit-reproducible (the interleave/parity gates rely on this);
|
||||
* same prompt → same text embedding → feature-cache reuse is correct;
|
||||
* a velocity ("flow") prediction the flow-match sampler consumes (Wan/LTX prediction_type).
|
||||
|
||||
On a GPU box, ``ComponentSpec.factory`` swaps these for lazy torch adapters wrapping the real
|
||||
``fastvideo.models`` modules + weights (see each model's ``components.py``); the loops, policies,
|
||||
scheduler, caches, parity, and training code are unchanged. That is the whole point of the
|
||||
(recipe, runtime) separation.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
|
||||
import numpy as np
|
||||
|
||||
LATENT_CHANNELS = 4
|
||||
TEXT_SEQ = 4
|
||||
TEXT_DIM = 8
|
||||
|
||||
|
||||
def _seed_from(s: str) -> int:
|
||||
return int(hashlib.sha256(s.encode()).hexdigest()[:8], 16)
|
||||
|
||||
|
||||
class ToyTextEncoder:
|
||||
"""Deterministic text→embedding (same text ⇒ same embedding ⇒ feature-cache reuse works)."""
|
||||
|
||||
def __init__(self, seq: int = TEXT_SEQ, dim: int = TEXT_DIM):
|
||||
self.seq, self.dim = seq, dim
|
||||
|
||||
def encode(self, text: str) -> np.ndarray:
|
||||
rng = np.random.default_rng(_seed_from("txt:" + (text or "<empty>")))
|
||||
return (rng.standard_normal((self.seq, self.dim)) * 0.1).astype("float32")
|
||||
|
||||
def encode_av(self, text: str):
|
||||
"""LTX-2.3's connector emits SEPARATE video + audio text projections; toy returns two distinct
|
||||
embeddings (the real Gemma.encode_av returns last_hidden_state + hidden_states[0])."""
|
||||
return self.encode(text), self.encode((text or "") + "\x00audio")
|
||||
|
||||
|
||||
class ToyImageEncoder:
|
||||
"""Deterministic image→embedding stand-in for the CLIP vision encoder (Wan i2v conditioning context)."""
|
||||
|
||||
def __init__(self, seq: int = TEXT_SEQ, dim: int = TEXT_DIM):
|
||||
self.seq, self.dim = seq, dim
|
||||
|
||||
def encode_image(self, image) -> np.ndarray:
|
||||
a = np.asarray(image, dtype="float32")
|
||||
rng = np.random.default_rng(_seed_from(f"img:{a.size}:{(float(a.mean()) if a.size else 0.0):.4f}"))
|
||||
return (rng.standard_normal((self.seq, self.dim)) * 0.1).astype("float32")
|
||||
|
||||
|
||||
class ToyDiT:
|
||||
"""A tiny deterministic velocity predictor (stands in for the 1.3B DiT).
|
||||
|
||||
velocity = tanh( Wx·latent + Wt·σ + s·mean(text_embed) + c·mean(context) )
|
||||
Channel-mixing via ``Wx`` makes the trajectory non-trivial; ``context`` carries causal
|
||||
chunk history for the self-forcing chunk_rollout loop.
|
||||
"""
|
||||
|
||||
def __init__(self, channels: int = LATENT_CHANNELS, seed: int = 0):
|
||||
self.C = channels
|
||||
rng = np.random.default_rng(seed)
|
||||
self.w_x = (rng.standard_normal((channels, channels)) * 0.15).astype("float32")
|
||||
self.w_t = (rng.standard_normal(channels) * 0.1).astype("float32")
|
||||
self.s_text = 0.5
|
||||
self.s_ctx = 0.3
|
||||
|
||||
def _pre_tanh(self, latent: np.ndarray, text_embed, sigma: float, context) -> np.ndarray:
|
||||
mixed = np.tensordot(self.w_x, latent, axes=([1], [0])) # [C,...] channel mix
|
||||
mixed = mixed + self.w_t.reshape((self.C, ) + (1, ) * (latent.ndim - 1)) * float(sigma)
|
||||
cond = float(np.mean(text_embed)) if text_embed is not None else 0.0
|
||||
mixed = mixed + self.s_text * cond
|
||||
if context is not None:
|
||||
mixed = mixed + self.s_ctx * float(np.mean(context))
|
||||
return mixed
|
||||
|
||||
def __call__(self,
|
||||
latent: np.ndarray,
|
||||
text_embed: np.ndarray | None,
|
||||
sigma: float,
|
||||
context: np.ndarray | None = None,
|
||||
*,
|
||||
audio_latent: np.ndarray | None = None,
|
||||
audio_text: np.ndarray | None = None,
|
||||
cond: np.ndarray | None = None):
|
||||
# ``cond`` (i2v mask+cond latent) is the real WanDiT's 36ch concat; the toy denoises the noise
|
||||
# channels only (image-conditioning is a GPU-path concern), so it's accepted and ignored here.
|
||||
latent = np.asarray(latent, dtype=np.float32)
|
||||
video = np.tanh(self._pre_tanh(latent, text_embed, sigma, context)).astype("float32")
|
||||
if audio_latent is None:
|
||||
return video
|
||||
# toy joint A/V (LTX-2.3): a simple element-wise audio velocity (the real DiT cross-attends
|
||||
# video<->audio in one forward and returns (video_vel, audio_vel)).
|
||||
au = np.asarray(audio_latent, dtype=np.float32)
|
||||
a_cond = float(np.mean(audio_text)) if audio_text is not None else 0.0
|
||||
audio = np.tanh(0.9 * au + 0.1 * float(sigma) + 0.3 * a_cond).astype("float32")
|
||||
return video, audio
|
||||
|
||||
# --- minimal trainable surface (so training methods do real optimizer steps) --------- #
|
||||
def clone(self) -> ToyDiT:
|
||||
c = ToyDiT.__new__(ToyDiT)
|
||||
c.C, c.s_text, c.s_ctx = self.C, self.s_text, self.s_ctx
|
||||
c.w_x, c.w_t = self.w_x.copy(), self.w_t.copy()
|
||||
return c
|
||||
|
||||
def blend_from(self, other: ToyDiT, decay: float) -> None:
|
||||
"""EMA / decay-blended-old-policy update: self ← decay·self + (1-decay)·other."""
|
||||
self.w_x = (decay * self.w_x + (1.0 - decay) * other.w_x).astype("float32")
|
||||
self.w_t = (decay * self.w_t + (1.0 - decay) * other.w_t).astype("float32")
|
||||
|
||||
def copy_from(self, other: ToyDiT) -> None:
|
||||
self.w_x, self.w_t = other.w_x.copy(), other.w_t.copy()
|
||||
|
||||
def mse_grad_step(self,
|
||||
latent: np.ndarray,
|
||||
text_embed,
|
||||
sigma: float,
|
||||
target: np.ndarray,
|
||||
lr: float,
|
||||
context=None,
|
||||
weight: float = 1.0) -> tuple[float, float]:
|
||||
"""One exact SGD step minimizing ``weight·MSE(forward(latent), target)`` w.r.t. ``w_x``.
|
||||
|
||||
Returns (loss, grad_norm). The grad_norm is the #1396-style per-method regression signal.
|
||||
This is a *real* gradient (chain rule through tanh), so the toy student actually learns —
|
||||
the structural stand-in for the FSDP/optimizer step a GPU trainer would run."""
|
||||
latent = np.asarray(latent, dtype=np.float32)
|
||||
target = np.asarray(target, dtype=np.float32)
|
||||
z = self._pre_tanh(latent, text_embed, sigma, context)
|
||||
pred = np.tanh(z)
|
||||
err = pred - target
|
||||
loss = float(weight * np.mean(err**2))
|
||||
n = err.size
|
||||
g = weight * (2.0 * err / n) * (1.0 - pred**2)
|
||||
axes = (list(range(1, g.ndim)), list(range(1, latent.ndim)))
|
||||
grad_wx = np.tensordot(g, latent, axes=axes) # [C_out, C_in]
|
||||
grad_norm = float(np.linalg.norm(grad_wx))
|
||||
self.w_x = (self.w_x - lr * grad_wx).astype("float32")
|
||||
return loss, grad_norm
|
||||
|
||||
|
||||
class ToyTokenizer:
|
||||
"""Deterministic byte-ish tokenizer for the omni AR pathway (phase 2)."""
|
||||
EOS = 0
|
||||
VOCAB = 256
|
||||
|
||||
def encode(self, text: str) -> list[int]:
|
||||
toks = [(ord(c) % (self.VOCAB - 1)) + 1 for c in (text or "")[:16]]
|
||||
return toks or [1]
|
||||
|
||||
def decode(self, tokens) -> str:
|
||||
return "".join(chr(32 + (int(t) % 90)) for t in tokens)
|
||||
|
||||
|
||||
class ToyMoTDiT(ToyDiT):
|
||||
"""Mixture-of-Transformers stand-in: ONE resident module that runs BOTH an AR (understanding)
|
||||
pathway and a diffusion (generation) pathway on shared weights.
|
||||
|
||||
``ar_forward`` is the und pathway (next-token); ``__call__`` (inherited) is the gen pathway
|
||||
(velocity). Binding both the ar_decode and diffusion_denoise loops to one instance of this
|
||||
component is the MoT requirement no DAG-of-engines can express.
|
||||
"""
|
||||
VOCAB = 256
|
||||
EOS = 0
|
||||
|
||||
def ar_forward(self, tokens) -> int:
|
||||
# deterministic next-token from the context (greedy ⇒ trivially interleave-safe)
|
||||
ctx = sum(int(t) for t in tokens)
|
||||
nxt = (ctx * 7 + 13) % self.VOCAB
|
||||
return int(nxt)
|
||||
|
||||
def reasoner_embed(self, tokens) -> np.ndarray:
|
||||
"""Pack the und-pathway tokens into a conditioning embed the gen pathway consumes
|
||||
(the Cosmos3 'prompt upsampling before diffusion in the same request')."""
|
||||
rng = np.random.default_rng((sum(int(t) for t in tokens) + 1) % (1 << 31))
|
||||
return (rng.standard_normal((4, 8)) * 0.1).astype("float32")
|
||||
|
||||
|
||||
class ToyPromptRefiner:
|
||||
"""Toy LM prompt-refiner — UniRL/PromptRL's Qwen role (a *separate expert*, not MoT-shared).
|
||||
|
||||
A categorical policy over a small set of refinement "actions" (stand-in for the LM emitting a
|
||||
refined prompt). Sampling an action picks an embedding offset that shifts the diffusion
|
||||
conditioning; the realized reward then drives the policy by REINFORCE. This is the *second
|
||||
trainable expert* in the unified-RL stress test: one reward → a token-policy-gradient update
|
||||
here AND a FlowGRPO update on the DiT, under one advantage (PromptRL §6).
|
||||
|
||||
Real gradient: loss = −A·logπ(a); ∇_logits = −A·(onehot(a) − softmax). So the toy LM actually
|
||||
learns to prefer the reward-favored action — the LM-side analogue of ToyDiT.mse_grad_step.
|
||||
"""
|
||||
|
||||
def __init__(self, n_actions: int = 8, dim: int = TEXT_DIM, seq: int = TEXT_SEQ, seed: int = 0):
|
||||
self.n = int(n_actions)
|
||||
self.dim, self.seq = int(dim), int(seq)
|
||||
self.logits: np.ndarray = np.zeros(self.n, dtype="float64") # the trainable policy params
|
||||
rng = np.random.default_rng(seed)
|
||||
# each action ⇒ a fixed conditioning offset (the "content" of that refined prompt)
|
||||
self._emb = (rng.standard_normal((self.n, self.seq, self.dim)) * 0.1).astype("float32")
|
||||
|
||||
def _probs(self) -> np.ndarray:
|
||||
z = self.logits - self.logits.max()
|
||||
e = np.exp(z)
|
||||
return e / e.sum()
|
||||
|
||||
def sample_refinement(self, rng) -> tuple[int, float]:
|
||||
"""Sample an action (exploration) and return (action, logπ(action))."""
|
||||
p = self._probs()
|
||||
a = int(rng.choice(self.n, p=p))
|
||||
return a, float(np.log(p[a] + 1e-12))
|
||||
|
||||
def argmax_refinement(self) -> tuple[int, float]:
|
||||
a = int(np.argmax(self.logits))
|
||||
return a, float(np.log(self._probs()[a] + 1e-12))
|
||||
|
||||
def action_logprob(self, action: int) -> float:
|
||||
return float(np.log(self._probs()[int(action)] + 1e-12))
|
||||
|
||||
def ar_forward(self, tokens) -> int:
|
||||
"""und-pathway hook so the declared ``ar_decode`` loop is runnable (greedy ⇒ interleave-safe)."""
|
||||
return int(np.argmax(self.logits))
|
||||
|
||||
def refined_embed(self, base_embed, action: int) -> np.ndarray:
|
||||
"""Apply the chosen refinement to the base text embedding → the diffusion conditioning."""
|
||||
base = np.zeros((self.seq, self.dim), dtype="float32") if base_embed is None \
|
||||
else np.asarray(base_embed, dtype="float32")
|
||||
return (base + self._emb[int(action)]).astype("float32")
|
||||
|
||||
def pg_step(self, action: int, advantage: float, lr: float) -> float:
|
||||
"""REINFORCE on the categorical policy. Returns grad_norm (the per-method regression signal)."""
|
||||
p = self._probs()
|
||||
onehot: np.ndarray = np.zeros(self.n, dtype="float64")
|
||||
onehot[int(action)] = 1.0
|
||||
grad = -float(advantage) * (onehot - p) # ∇_logits (−A·logπ(a))
|
||||
self.logits = self.logits - float(lr) * grad # descent on −A·logπ ⇒ ascent on A·logπ
|
||||
return float(np.linalg.norm(grad))
|
||||
|
||||
def kl_to(self, other: ToyPromptRefiner) -> float:
|
||||
"""KL(π_self ‖ π_other) over the categorical — the LM-PG reference-KL penalty term."""
|
||||
p, q = self._probs(), other._probs()
|
||||
return float(np.sum(p * (np.log(p + 1e-12) - np.log(q + 1e-12))))
|
||||
|
||||
def clone(self) -> ToyPromptRefiner:
|
||||
c = ToyPromptRefiner.__new__(ToyPromptRefiner)
|
||||
c.n, c.dim, c.seq = self.n, self.dim, self.seq
|
||||
c.logits, c._emb = self.logits.copy(), self._emb.copy()
|
||||
return c
|
||||
|
||||
def copy_from(self, other: ToyPromptRefiner) -> None:
|
||||
self.logits = other.logits.copy()
|
||||
|
||||
def blend_from(self, other: ToyPromptRefiner, decay: float) -> None:
|
||||
self.logits = decay * self.logits + (1.0 - decay) * other.logits
|
||||
|
||||
|
||||
class ToyTalker(ToyMoTDiT):
|
||||
"""Toy Talker — Qwen-Omni stage 1 (AR over a speech-codec vocab, conditioned on the Thinker).
|
||||
|
||||
A *separate expert* from the thinker (not weight-shared): its next-token depends on its OWN
|
||||
weights (a seed salt) AND the prefilled thinker payload (tokens + hidden state). So the cascade
|
||||
is real — change the thinker and the talker's tokens change; change the talker's weights and they
|
||||
change too. The vllm-omni Talker is likewise a distinct model conditioned on Thinker hidden states.
|
||||
"""
|
||||
|
||||
def __init__(self, channels: int = LATENT_CHANNELS, seed: int = 0):
|
||||
super().__init__(channels=channels, seed=seed)
|
||||
self.salt = (seed % 97) + 1
|
||||
|
||||
def ar_forward(self, tokens) -> int:
|
||||
ctx = sum(int(t) for t in tokens)
|
||||
return int((ctx * 7 + 13 + self.salt) % self.VOCAB) # weight-dependent (salt) + context
|
||||
|
||||
|
||||
class ToyVocoder:
|
||||
"""Toy streaming code2wav vocoder — Qwen-Omni stage 2 (BigVGAN/Code2Wav role).
|
||||
|
||||
Speech codec tokens → audio waveform, synthesized in chunks. Deterministic given the tokens so
|
||||
the ``audio_decode`` loop is interleave-safe. Each token contributes a short, timbre-stamped
|
||||
waveform segment — enough to exercise chunked synthesis, streaming emits, and the AudioArtifact.
|
||||
"""
|
||||
|
||||
def __init__(self, samples_per_token: int = 16, vocab: int = 256, seed: int = 7):
|
||||
self.spt = int(samples_per_token)
|
||||
rng = np.random.default_rng(seed)
|
||||
self.bank = (rng.standard_normal(vocab) * 0.3).astype("float32") # per-token timbre offset
|
||||
|
||||
def synthesize(self, token_chunk) -> np.ndarray:
|
||||
"""One chunk of codec tokens → a waveform segment in [-1, 1], length spt·len(chunk)."""
|
||||
segs = []
|
||||
for t in token_chunk:
|
||||
t = int(t) % self.bank.size
|
||||
phase = np.linspace(0.0, np.pi * (1 + (t % 8)), self.spt, dtype="float32")
|
||||
segs.append(np.tanh(self.bank[t] + np.sin(phase)).astype("float32"))
|
||||
return np.concatenate(segs) if segs else np.zeros(0, dtype="float32")
|
||||
|
||||
|
||||
class ToyLoRA:
|
||||
"""Toy LoRA adapter — a low-rank velocity delta on a base DiT, selected and applied per request.
|
||||
|
||||
Many of these are served over ONE resident base (the adapter is tiny: a rank-r channel delta);
|
||||
a request picks which to apply. Swappable/versioned independently (the cache key's
|
||||
``adapter_versions``). Stand-in for a real LoRA / DoRA / IP-Adapter."""
|
||||
kind = "lora"
|
||||
|
||||
def __init__(self,
|
||||
adapter_id: str,
|
||||
channels: int = LATENT_CHANNELS,
|
||||
scale: float = 0.6,
|
||||
rank: int = 2,
|
||||
seed: int = 0):
|
||||
self.adapter_id = adapter_id
|
||||
self.scale = float(scale)
|
||||
rng = np.random.default_rng(seed)
|
||||
self.a = (rng.standard_normal((channels, rank)) * 0.4).astype("float32")
|
||||
self.b = (rng.standard_normal((rank, channels)) * 0.4).astype("float32")
|
||||
|
||||
def delta(self, latent, control=None) -> np.ndarray:
|
||||
w = self.a @ self.b # [C, C] low-rank weight delta
|
||||
d = np.tensordot(w, np.asarray(latent, dtype="float32"), axes=([1], [0])) # [C, ...]
|
||||
return (self.scale * np.tanh(d)).astype("float32")
|
||||
|
||||
def update(self, seed: int) -> None: # hot-swap: replace the adapter's weights
|
||||
rng = np.random.default_rng(seed)
|
||||
self.a = (rng.standard_normal(self.a.shape) * 0.4).astype("float32")
|
||||
|
||||
|
||||
class ToyControlNet:
|
||||
"""Toy ControlNet adapter — conditions the velocity on a control signal (pose/depth/edge stand-in).
|
||||
Different control ⇒ different generation; selected per request like a LoRA."""
|
||||
kind = "controlnet"
|
||||
|
||||
def __init__(self, adapter_id: str, channels: int = LATENT_CHANNELS, scale: float = 0.8, seed: int = 1):
|
||||
self.adapter_id = adapter_id
|
||||
self.scale = float(scale)
|
||||
rng = np.random.default_rng(seed)
|
||||
self.proj = (rng.standard_normal(channels) * 0.3).astype("float32")
|
||||
|
||||
def delta(self, latent, control=None) -> np.ndarray:
|
||||
latent = np.asarray(latent, dtype="float32")
|
||||
c = float(np.mean(control)) if control is not None else 0.0 # the control image's signal
|
||||
shape = (latent.shape[0], ) + (1, ) * (latent.ndim - 1)
|
||||
return (self.scale * c * self.proj.reshape(shape)).astype("float32")
|
||||
|
||||
def update(self, seed: int) -> None:
|
||||
rng = np.random.default_rng(seed)
|
||||
self.proj = (rng.standard_normal(self.proj.shape) * 0.3).astype("float32")
|
||||
|
||||
|
||||
def _spec_target_next(tokens) -> int:
|
||||
"""The target AR model's greedy next-token. Length-dependent so the sequence stays varied (no
|
||||
trivial absorbing fixed point) — the substrate for a meaningful speculative-decoding accept rate."""
|
||||
s = sum(int(t) for t in tokens)
|
||||
n = len(tokens)
|
||||
return int((s * 31 + 7 * n + 13) % 256)
|
||||
|
||||
|
||||
class ToyTargetModel:
|
||||
"""Toy speculative-decoding TARGET (the large, ground-truth AR model)."""
|
||||
VOCAB = 256
|
||||
EOS = 0
|
||||
|
||||
def __init__(self, seed: int = 0):
|
||||
self.seed = seed
|
||||
|
||||
def ar_forward(self, tokens) -> int:
|
||||
return _spec_target_next(tokens)
|
||||
|
||||
|
||||
class ToyDraftModel:
|
||||
"""Toy speculative-decoding DRAFT model — a cheap approximation of the target.
|
||||
|
||||
It replicates the target on a context-and-position-varying subset (~``agree`` fraction → those tokens
|
||||
are accepted) and diverges otherwise (→ a rejection + a target correction). The accept rate varies
|
||||
*within* a decode, so different rounds accept different lengths (a ragged loop). The demonstration is
|
||||
exact regardless of draft quality: speculative decoding accepts only tokens the target would have
|
||||
produced, so the output equals the target's own greedy decode."""
|
||||
VOCAB = 256
|
||||
EOS = 0
|
||||
|
||||
def __init__(self, agree: float = 0.7, seed: int = 0):
|
||||
self.agree = float(agree)
|
||||
|
||||
def ar_forward(self, tokens) -> int:
|
||||
target = _spec_target_next(tokens)
|
||||
s, n = sum(int(t) for t in tokens), len(tokens)
|
||||
if ((s * 13 + n * 7) % 10) < int(round(self.agree * 10)): # agree on a varying subset
|
||||
return int(target)
|
||||
return int((target + 1) % self.VOCAB) # a cheap wrong guess (rejected)
|
||||
|
||||
|
||||
class ToyRewardModel:
|
||||
"""Toy learned reward model — a served scorer/verifier (PickScore / CLIP / a VLM judge stand-in).
|
||||
|
||||
Maps a media latent → a scalar reward in [-1, 1] via a fixed projection over per-channel features.
|
||||
The point isn't the math: it's that the reward is computed by a *model* (served, scheduled as
|
||||
``REWARD_BATCH`` work units, place-able on its own pool), not a numpy heuristic bolted onto the
|
||||
method — the reward plane composing with the serving plane."""
|
||||
|
||||
def __init__(self, dim: int = LATENT_CHANNELS, seed: int = 5):
|
||||
rng = np.random.default_rng(seed)
|
||||
self.w = (rng.standard_normal(dim) * 0.5).astype("float32")
|
||||
|
||||
def score(self, media) -> float:
|
||||
a = np.asarray(media, dtype="float64")
|
||||
feat = a.reshape(a.shape[0], -1).mean(axis=1) if a.ndim > 1 else np.atleast_1d(a)
|
||||
n = self.w.size
|
||||
feat = feat[:n] if feat.size >= n else np.pad(feat, (0, n - feat.size))
|
||||
return float(np.tanh(float(np.dot(self.w, feat))))
|
||||
|
||||
|
||||
class ToyAudioVAE:
|
||||
"""Toy audio VAE — mel-like audio latent → mono waveform (the LTX-2 / Cosmos3 audio branch).
|
||||
|
||||
Decodes a small ``[C, A, 1, 1]`` audio latent (denoised jointly with the video latent) into a 1-D
|
||||
waveform: channel-collapse via a fixed projection, then upsample along time. Deterministic, so the
|
||||
joint A/V denoise loop stays interleave-safe. The video VAE and this share nothing — two decoders,
|
||||
one synchronized latent pair."""
|
||||
|
||||
def __init__(self, samples_per_frame: int = 16, seed: int = 2):
|
||||
self.spf = int(samples_per_frame)
|
||||
rng = np.random.default_rng(seed)
|
||||
self.proj = (rng.standard_normal(LATENT_CHANNELS) * 0.3).astype("float32")
|
||||
|
||||
def decode(self, audio_latent: np.ndarray) -> np.ndarray:
|
||||
a = np.asarray(audio_latent, dtype="float32")
|
||||
flat = a.reshape(a.shape[0], -1) # [C, A·...]
|
||||
# project over the actual channel count (LTX-2.3 audio latent is 8ch; the 2-stage toy is
|
||||
# LATENT_CHANNELS) — np.resize is identity when they already match, so existing T2VS is unchanged.
|
||||
proj = np.resize(self.proj, a.shape[0]).astype("float32")
|
||||
mono = np.tensordot(proj, flat, axes=([0], [0])) # [A·...]
|
||||
return np.tanh(np.repeat(mono, self.spf)).astype("float32") # [A·spf] mono waveform
|
||||
|
||||
|
||||
class ToyVAE:
|
||||
"""Tiny deterministic VAE. encode: video→latent (mean-pool + channel proj); decode: latent→video."""
|
||||
|
||||
def __init__(self, channels: int = LATENT_CHANNELS, spatial: int = 8, temporal: int = 4, seed: int = 1):
|
||||
self.C = channels
|
||||
self.spatial = spatial
|
||||
self.temporal = temporal
|
||||
rng = np.random.default_rng(seed)
|
||||
self.dec_proj = (rng.standard_normal((3, channels)) * 0.2).astype("float32")
|
||||
|
||||
def decode(self, latent: np.ndarray) -> np.ndarray:
|
||||
"""latent [C,T,H,W] -> video [3, T, H*spatial, W*spatial] (deterministic upsample)."""
|
||||
latent = np.asarray(latent, dtype=np.float32)
|
||||
rgb = np.tensordot(self.dec_proj, latent, axes=([1], [0])) # [3,T,H,W]
|
||||
up = np.repeat(np.repeat(rgb, self.spatial, axis=2), self.spatial, axis=3)
|
||||
return np.tanh(up).astype("float32")
|
||||
|
||||
def encode(self, video: np.ndarray) -> np.ndarray:
|
||||
video = np.asarray(video, dtype=np.float32)
|
||||
pooled = video[:self.C] # crude: take C channels
|
||||
return pooled.astype("float32")
|
||||
|
||||
|
||||
class ToyUpsampler:
|
||||
"""Toy latent upsampler — CPU stand-in for the learned LTX-2 spatial upsampler. Nearest-neighbor 2x
|
||||
spatial repeat on a ``[C,T,H,W]`` latent. The real ``LTX2LatentUpsampler`` learns this super-res; the
|
||||
GPU backend swaps in that module (torch_backend.LTX2Upsampler) via the ``upsampler`` component kind,
|
||||
so the program calls ``component('spatial_upsampler').upsample(...)`` on both backends (no device branch)."""
|
||||
|
||||
def __init__(self, scale: int = 2):
|
||||
self.scale = int(scale)
|
||||
|
||||
def upsample(self, latent: np.ndarray) -> np.ndarray:
|
||||
z = np.asarray(latent, dtype="float32")
|
||||
return np.repeat(np.repeat(z, self.scale, axis=-2), self.scale, axis=-1).astype("float32")
|
||||
@@ -0,0 +1,204 @@
|
||||
"""Platform — the detected ``(device, arch)`` that resolves both registries.
|
||||
|
||||
One ``Platform`` per pool. It owns:
|
||||
* the **device fallback chain** (e.g. ``accel → cpu``) — so a component/kernel a backend hasn't
|
||||
implemented falls back to the numpy reference (the terminal rung + parity oracle);
|
||||
* the **arch fallback chain** per device (e.g. cuda ``sm100→sm90→sm80→ptx→generic``) — so a kernel
|
||||
built for an older arch still resolves on a newer card;
|
||||
* ``build_component(spec, instance)`` — the single materialization seam (replaces a bare
|
||||
``spec.factory`` call); and
|
||||
* ``kernels`` — a per-platform ``KernelTable`` that resolves an op once and caches the callable.
|
||||
|
||||
``Platform.detect()`` returns a CPU/numpy platform unless torch+CUDA are actually present. The
|
||||
(device, arch) machinery is exercised on a CPU box by the pure-python ``accel`` stand-in backend,
|
||||
proving the dispatch is generic without a real GPU.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
from collections.abc import Iterator
|
||||
|
||||
from v2.platform.registry import COMPONENTS, KERNELS
|
||||
|
||||
_BACKENDS_LOADED = False
|
||||
|
||||
|
||||
def ensure_backends_loaded() -> None:
|
||||
"""Import the backend modules once so they self-register. Lazy → import-cycle-free."""
|
||||
global _BACKENDS_LOADED
|
||||
if _BACKENDS_LOADED:
|
||||
return
|
||||
_BACKENDS_LOADED = True
|
||||
# imported for side effects (each module calls register_component/register_kernel at import)
|
||||
from v2.platform.backends import accel, cpu, torch_cuda # noqa: F401
|
||||
|
||||
|
||||
# Arch fallback chains (newest → oldest → portable terminal). Resolution walks these in order.
|
||||
_CUDA_ARCH_FALLBACK = ("sm100", "sm90", "sm80", "ptx", "generic")
|
||||
_ACCEL_ARCH_FALLBACK = ("sm90", "sm80", "generic")
|
||||
_PORTABLE_ARCHS = ("ptx", "generic") # vendor-neutral terminals — valid on any device
|
||||
|
||||
|
||||
def _arch_rank(arch: str) -> int:
|
||||
"""``smXY`` → integer ``XY`` for comparison; portable terminals (ptx/generic/unknown) → -1."""
|
||||
if arch.startswith("sm") and arch[2:].isdigit():
|
||||
return int(arch[2:])
|
||||
return -1
|
||||
|
||||
|
||||
def _suffix_from(chain: tuple[str, ...], arch: str) -> tuple[str, ...]:
|
||||
"""The arch fallback tail for ``arch``: itself, then ONLY older real archs from ``chain``, then
|
||||
portable terminals. Never falls back to a *newer* arch — a kernel built for sm90 is binary-
|
||||
incompatible on an sm70 device, so resolution must only ever degrade to older/more-portable code.
|
||||
(A newer device, e.g. sm120, may still run older-arch kernels — those are kept.)
|
||||
"""
|
||||
rank = _arch_rank(arch)
|
||||
tail: list[str] = [arch]
|
||||
for a in chain:
|
||||
if a in tail:
|
||||
continue
|
||||
ar = _arch_rank(a)
|
||||
if ar < 0 or rank >= 0 and ar < rank: # portable terminal (ptx/generic): always valid
|
||||
tail.append(a)
|
||||
# else: newer-or-equal real arch → skip (would be binary-incompatible on this device)
|
||||
return tuple(tail)
|
||||
|
||||
|
||||
class KernelTable:
|
||||
"""A platform's resolved op→callable map. Resolves each op once (over the fallback chain),
|
||||
then caches — so the hot loop pays the registry walk at most once per (op, variant)."""
|
||||
|
||||
def __init__(self, platform: Platform):
|
||||
self.platform = platform
|
||||
self._cache: dict[tuple[str, str], Any] = {}
|
||||
|
||||
def get(self, op: str, variant: str = "default") -> Any:
|
||||
ck = (op, variant)
|
||||
if ck not in self._cache:
|
||||
reg = self.platform.resolve_kernel(op, variant)
|
||||
if reg is None:
|
||||
raise KeyError(f"no kernel for op={op!r} variant={variant!r} on device chain "
|
||||
f"{self.platform.device_chain} (and no numpy terminal registered)")
|
||||
self._cache[ck] = reg.fn
|
||||
return self._cache[ck]
|
||||
|
||||
def resolved_source(self, op: str, variant: str = "default") -> str:
|
||||
"""Which registered cell an op resolves to — used by parity/diagnostics tests."""
|
||||
reg = self.platform.resolve_kernel(op, variant)
|
||||
return "" if reg is None else f"{reg.key}@{reg.source}"
|
||||
|
||||
|
||||
class Platform:
|
||||
|
||||
def __init__(self, device: str, arch: str, device_chain: tuple[str, ...], arch_chains: dict[str, tuple[str, ...]]):
|
||||
self.device = device
|
||||
self.arch = arch
|
||||
self.device_chain = tuple(device_chain)
|
||||
self.arch_chains = dict(arch_chains)
|
||||
self._kernels: KernelTable | None = None
|
||||
self._xp: Any = None
|
||||
|
||||
# --- the kernel table (lazy, cached) ------------------------------------- #
|
||||
@property
|
||||
def kernels(self) -> KernelTable:
|
||||
if self._kernels is None:
|
||||
ensure_backends_loaded()
|
||||
self._kernels = KernelTable(self)
|
||||
return self._kernels
|
||||
|
||||
# --- the array namespace (numpy on CPU, torch-on-device on a GPU box) ----- #
|
||||
@property
|
||||
def xp(self) -> Any:
|
||||
"""Birth/cast/marshal helpers so the denoise loop is array-agnostic: numpy on CPU (torch-free),
|
||||
torch-on-device on cuda (the latent stays resident — no per-step host<->device round-trip)."""
|
||||
if self._xp is None:
|
||||
from v2.platform.array_ns import get_array_ns
|
||||
self._xp = get_array_ns(self.device)
|
||||
return self._xp
|
||||
|
||||
def arch_chain(self, device: str) -> tuple[str, ...]:
|
||||
return self.arch_chains.get(device, ("generic", ))
|
||||
|
||||
# --- candidate-key generators (the fallback chains, in priority order) --- #
|
||||
def _component_keys(self, kind: str, variant: str) -> Iterator[tuple]:
|
||||
for dev in self.device_chain:
|
||||
yield (kind, dev, variant)
|
||||
if variant != "default":
|
||||
yield (kind, dev, "default")
|
||||
|
||||
def _kernel_keys(self, op: str, variant: str) -> Iterator[tuple]:
|
||||
for dev in self.device_chain:
|
||||
for arch in self.arch_chain(dev):
|
||||
yield (op, dev, arch, variant)
|
||||
if variant != "default":
|
||||
yield (op, dev, arch, "default")
|
||||
|
||||
# --- the two resolution seams -------------------------------------------- #
|
||||
def build_component(self, spec: Any, instance: Any, variant: str = "default") -> Any:
|
||||
"""Materialize a component. Registry first (device chain), then ``spec.factory`` as the
|
||||
terminal numpy rung — so existing cards (factory-only) are unchanged, while a GPU/accel
|
||||
backend overrides by registering ``(kind, device, …)``."""
|
||||
ensure_backends_loaded()
|
||||
reg = COMPONENTS.resolve_first(self._component_keys(spec.kind, variant))
|
||||
if reg is not None:
|
||||
return reg.fn(spec, instance, self)
|
||||
if getattr(spec, "factory", None) is not None:
|
||||
return spec.factory(instance) # the (kind, "cpu", "numpy") terminal rung
|
||||
raise RuntimeError(f"component kind={spec.kind!r} ({getattr(spec, 'component_id', '?')!r}) has no impl on "
|
||||
f"device chain {self.device_chain} and no factory terminal")
|
||||
|
||||
def resolve_kernel(self, op: str, variant: str = "default"):
|
||||
ensure_backends_loaded()
|
||||
return KERNELS.resolve_first(self._kernel_keys(op, variant))
|
||||
|
||||
# --- constructors -------------------------------------------------------- #
|
||||
@classmethod
|
||||
def cpu(cls) -> Platform:
|
||||
return cls("cpu", "numpy", ("cpu", ), {"cpu": ("numpy", )})
|
||||
|
||||
@classmethod
|
||||
def accel(cls, arch: str = "sm90") -> Platform:
|
||||
"""A pure-python stand-in accelerator (CPU-resident) used to prove cross-device dispatch,
|
||||
arch fallback, and the parity oracle without a real GPU."""
|
||||
return cls("accel", arch, ("accel", "cpu"), {
|
||||
"accel": _suffix_from(_ACCEL_ARCH_FALLBACK, arch),
|
||||
"cpu": ("numpy", )
|
||||
})
|
||||
|
||||
@classmethod
|
||||
def cuda(cls, arch: str = "sm90") -> Platform:
|
||||
"""A real CUDA platform (resolves torch adapters + fastvideo-kernel cells). Only usable on a
|
||||
box where those cells are ``available``; declared here so detection has a target."""
|
||||
return cls("cuda", arch, ("cuda", "cpu"), {"cuda": _suffix_from(_CUDA_ARCH_FALLBACK, arch), "cpu": ("numpy", )})
|
||||
|
||||
@classmethod
|
||||
def detect(cls) -> Platform:
|
||||
"""Honest detection: CUDA iff torch+CUDA are actually importable, else CPU/numpy."""
|
||||
ensure_backends_loaded()
|
||||
import importlib.util
|
||||
if importlib.util.find_spec("torch") is not None:
|
||||
try:
|
||||
import torch # type: ignore
|
||||
if torch.cuda.is_available():
|
||||
major, minor = torch.cuda.get_device_capability()
|
||||
return cls.cuda(f"sm{major}{minor}")
|
||||
except Exception:
|
||||
pass
|
||||
return cls.cpu()
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"Platform(device={self.device!r}, arch={self.arch!r}, chain={self.device_chain})"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Matrix dumps — enumerate the full (op/kind × device × arch × variant) grid, #
|
||||
# including unavailable backends, without importing torch or touching a GPU. #
|
||||
# --------------------------------------------------------------------------- #
|
||||
def kernel_matrix() -> list[dict[str, Any]]:
|
||||
ensure_backends_loaded()
|
||||
return KERNELS.manifest()
|
||||
|
||||
|
||||
def component_matrix() -> list[dict[str, Any]]:
|
||||
ensure_backends_loaded()
|
||||
return COMPONENTS.manifest()
|
||||
@@ -0,0 +1,157 @@
|
||||
"""The two tuple-keyed backend registries — the dispatch membrane.
|
||||
|
||||
Everything device/arch-specific resolves through one of two registries, keyed by a tuple:
|
||||
|
||||
* ``COMPONENTS`` — keyed ``(kind, device, variant)``. Resolves which *weight-bearing* component
|
||||
implementation to materialize (the toy numpy module, a torch adapter on a GPU box, …). This is
|
||||
the generalization of ``ComponentSpec.factory``: the card's ``factory`` is simply the terminal
|
||||
``(kind, "cpu", "numpy")`` rung that every other backend falls back to.
|
||||
* ``KERNELS`` — keyed ``(op, device, arch, variant)``. Resolves which *stateless primitive* kernel
|
||||
implements an op (the flow-match solver step, an SDE step, a gemm/norm on a real box). Arch is a
|
||||
fallback axis (sm100→sm90→…→generic), with the numpy reference as the terminal rung.
|
||||
|
||||
Two properties this module guarantees (the kernel-colocation answer made concrete):
|
||||
|
||||
1. **Location is decoupled from dispatch.** A registration carries a key, not a path. Where the .cu
|
||||
(or .py) lives on disk — unified ``fastvideo-kernel/`` or co-located with a model — is invisible
|
||||
here; both call the same ``register_*``.
|
||||
2. **The matrix is enumerable without a GPU and without importing torch.** ``manifest()`` lists
|
||||
every declared cell — including ones whose backend is *unavailable* here (e.g. a torch/cuda cell
|
||||
on a CPU box, gated by a ``find_spec``-based ``available`` predicate that never imports torch) —
|
||||
so the parity/fallback matrix can be dumped on a laptop. ``resolve`` skips unavailable cells;
|
||||
``manifest`` lists them.
|
||||
|
||||
Scope caveat (honest): in this in-tree mini the cells are populated by a fixed backend import
|
||||
list (``Platform.ensure_backends_loaded`` imports ``backends/{cpu,accel,torch_cuda}``), not by
|
||||
setuptools entry points. So the matrix is complete for the backends in that list, but an
|
||||
*out-of-tree* or *model-co-located* backend would not self-enumerate until added to it — full
|
||||
entry-point discovery is the real-package mechanism and is deliberately not wired here.
|
||||
|
||||
This is a pure leaf module: stdlib only. The actual backend registrations live in ``backends/`` and
|
||||
are imported lazily (``Platform.ensure_backends_loaded``) to keep this import-cycle-free.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from collections.abc import Callable, Iterable
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Op names — the stable identifiers a loop dispatches on (not free functions). #
|
||||
# --------------------------------------------------------------------------- #
|
||||
FLOW_MATCH_STEP = "flow_match_step" # deterministic flow-match Euler solver step (ODE serve)
|
||||
FLOW_SDE_STEP = "flow_sde_step" # FlowGRPO stochastic step + log-prob (RL rollout)
|
||||
|
||||
|
||||
def _always_available() -> bool:
|
||||
return True
|
||||
|
||||
|
||||
@dataclass
|
||||
class Registration:
|
||||
"""One cell in a registry: a key, the implementation, and an availability predicate.
|
||||
|
||||
``available`` lets a cell be *declared* (so it shows up in ``manifest()``) while being
|
||||
*unresolvable* in this environment — e.g. a torch/cuda kernel on a CPU box declares
|
||||
``available=lambda: torch.cuda.is_available()`` and is listed-but-skipped here.
|
||||
"""
|
||||
key: tuple
|
||||
fn: Callable[..., Any]
|
||||
available: Callable[[], bool] = _always_available
|
||||
source: str = ""
|
||||
meta: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def is_available(self) -> bool:
|
||||
try:
|
||||
return bool(self.available())
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
class TupleRegistry:
|
||||
"""A registry keyed by an ordered field tuple, with availability-aware resolution."""
|
||||
|
||||
def __init__(self, name: str, key_fields: tuple[str, ...]):
|
||||
self.name = name
|
||||
self.key_fields = key_fields
|
||||
self._reg: dict[tuple, Registration] = {}
|
||||
|
||||
def put(self,
|
||||
key: tuple,
|
||||
fn: Callable[..., Any],
|
||||
*,
|
||||
available: Callable[[], bool] | None = None,
|
||||
source: str = "",
|
||||
meta: dict[str, Any] | None = None) -> None:
|
||||
if len(key) != len(self.key_fields):
|
||||
raise ValueError(f"{self.name}: key {key} does not match fields {self.key_fields}")
|
||||
self._reg[key] = Registration(key=key,
|
||||
fn=fn,
|
||||
available=available or _always_available,
|
||||
source=source,
|
||||
meta=dict(meta or {}))
|
||||
|
||||
def lookup(self, key: tuple) -> Registration | None:
|
||||
"""Exact lookup, ignoring availability (used by ``manifest``/diagnostics)."""
|
||||
return self._reg.get(key)
|
||||
|
||||
def resolve_first(self, candidate_keys: Iterable[tuple]) -> Registration | None:
|
||||
"""First candidate key that is both registered AND available wins (the fallback chain)."""
|
||||
for key in candidate_keys:
|
||||
reg = self._reg.get(key)
|
||||
if reg is not None and reg.is_available():
|
||||
return reg
|
||||
return None
|
||||
|
||||
def manifest(self) -> list[dict[str, Any]]:
|
||||
"""Every declared cell as a flat dict — enumerable without importing/calling the kernels.
|
||||
|
||||
This is what makes the parity matrix dumpable on a CPU box: it lists torch/cuda cells too,
|
||||
each with ``available=False`` here, so coverage is never silently truncated to 'what imported'.
|
||||
"""
|
||||
rows: list[dict[str, Any]] = []
|
||||
for key, reg in sorted(self._reg.items(), key=lambda kv: tuple(map(str, kv[0]))):
|
||||
row = dict(zip(self.key_fields, key, strict=False))
|
||||
row["available"] = reg.is_available()
|
||||
row["source"] = reg.source
|
||||
row.update(reg.meta) # surface declared metadata (e.g. workspace_bytes)
|
||||
rows.append(row)
|
||||
return rows
|
||||
|
||||
|
||||
# The two registries, populated by ``backends/`` modules (lazily imported).
|
||||
COMPONENTS = TupleRegistry("components", ("kind", "device", "variant"))
|
||||
KERNELS = TupleRegistry("kernels", ("op", "device", "arch", "variant"))
|
||||
|
||||
|
||||
def register_component(kind: str,
|
||||
fn: Callable[..., Any],
|
||||
*,
|
||||
device: str,
|
||||
variant: str = "default",
|
||||
available: Callable[[], bool] | None = None,
|
||||
source: str = "") -> None:
|
||||
"""Register a component builder. ``fn(spec, instance, platform) -> live component``."""
|
||||
COMPONENTS.put((kind, device, variant), fn, available=available, source=source)
|
||||
|
||||
|
||||
def register_kernel(op: str,
|
||||
fn: Callable[..., Any],
|
||||
*,
|
||||
device: str,
|
||||
arch: str,
|
||||
variant: str = "default",
|
||||
available: Callable[[], bool] | None = None,
|
||||
source: str = "",
|
||||
workspace_bytes: int = 0) -> None:
|
||||
"""Register a stateless kernel. ``fn(*args, **kwargs)`` — same signature as its numpy reference.
|
||||
|
||||
``workspace_bytes`` is the kernel's static scratch requirement — the capture-safety contract: a
|
||||
kernel used inside a captured CUDA graph must draw its workspace from the pool-provided static
|
||||
buffer, not malloc per call, or replay corrupts. Declared here so the requirement is enumerable
|
||||
in the matrix; the numpy reference needs none (0)."""
|
||||
KERNELS.put((op, device, arch, variant),
|
||||
fn,
|
||||
available=available,
|
||||
source=source,
|
||||
meta={"workspace_bytes": int(workspace_bytes)})
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Program plane — compose a card's loops into a task."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.program.specs import (
|
||||
ComponentNode,
|
||||
Edge,
|
||||
EdgeKind,
|
||||
ModelLoopNode,
|
||||
Program,
|
||||
ProgramKind,
|
||||
ProgramNode,
|
||||
always,
|
||||
when_opt,
|
||||
when_task,
|
||||
)
|
||||
from v2.program.workflow import (
|
||||
BestOfNWorkflow,
|
||||
ParallelWorkflow,
|
||||
Workflow,
|
||||
WorkflowRegistry,
|
||||
WorkflowStage,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Program", "ProgramKind", "ProgramNode", "ComponentNode", "ModelLoopNode", "Edge", "EdgeKind", "always",
|
||||
"when_task", "when_opt", "Workflow", "WorkflowStage", "WorkflowRegistry", "ParallelWorkflow", "BestOfNWorkflow"
|
||||
]
|
||||
@@ -0,0 +1,110 @@
|
||||
"""Programs compose a card's loops into a task.
|
||||
|
||||
> The card says what loops *exist*, the program says how to *run* them for this request.
|
||||
|
||||
Kinds: InlineProgram (many loops, one resident instance — the omni default), TrainingProgram,
|
||||
etc. Nodes are typed (ModelLoopNode / ComponentNode / ...); edges are typed and carry *named*
|
||||
artifacts (not a god-batch). Linear pipelines are the degenerate case; ``when=`` predicates and
|
||||
multiple producers give branches/fan-out (LTX-2 base→upsample→refine→decode; A/V fan-out).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
|
||||
from v2.request.tasks import TaskType
|
||||
|
||||
|
||||
class ProgramKind(str, Enum):
|
||||
INLINE = "inline"
|
||||
DISAGGREGATED = "disaggregated"
|
||||
TRAINING = "training"
|
||||
REALTIME = "realtime"
|
||||
WORKFLOW = "workflow"
|
||||
|
||||
|
||||
# --- node predicates (declarative, replacing geometry heuristics) ------------- #
|
||||
def always(_request: Any) -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def when_task(*tasks: TaskType) -> Callable[[Any], bool]:
|
||||
s = set(tasks)
|
||||
return lambda req: req.task in s
|
||||
|
||||
|
||||
def when_opt(node_id: str, key: str) -> Callable[[Any], bool]:
|
||||
return lambda req: bool(req.node_override(node_id).get(key, False))
|
||||
|
||||
|
||||
# --- nodes -------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class ProgramNode:
|
||||
node_id: str
|
||||
when: Callable[[Any], bool] = always
|
||||
reads: tuple[str, ...] = () # input slot names (typed edges in)
|
||||
writes: tuple[str, ...] = () # output slot names (typed edges out)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ComponentNode(ProgramNode):
|
||||
"""A one-shot transform: text/vision encode, VAE decode, latent upsample, vocoder.
|
||||
|
||||
``fn(instance, slots, request, ctx)`` reads its ``reads`` slots and writes its ``writes``
|
||||
slots. Runs to completion in a single engine tick (not step-interleaved)."""
|
||||
fn: Callable[..., None] | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelLoopNode(ProgramNode):
|
||||
"""Binds one of the card's loops. The engine drives it via a LoopRunner, one step per tick,
|
||||
so its steps interleave with other requests' steps."""
|
||||
loop_id: str = ""
|
||||
output_slot: str = "latents"
|
||||
|
||||
|
||||
# --- typed edges (graph IR; validated, data flows via named slots) ------------ #
|
||||
class EdgeKind(str, Enum):
|
||||
TENSOR = "tensor"
|
||||
ARTIFACT = "artifact"
|
||||
STREAM = "stream"
|
||||
CONTROL = "control"
|
||||
CACHE = "cache"
|
||||
BEHAVIOR = "behavior"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Edge:
|
||||
src: str # slot or node id
|
||||
dst: str
|
||||
kind: EdgeKind = EdgeKind.TENSOR
|
||||
|
||||
|
||||
# --- program ------------------------------------------------------------------ #
|
||||
@dataclass
|
||||
class Program:
|
||||
program_id: str
|
||||
kind: ProgramKind
|
||||
nodes: list[ProgramNode] = field(default_factory=list)
|
||||
edges: list[Edge] = field(default_factory=list)
|
||||
# artifact name -> slot that holds its value at the end
|
||||
output_artifacts: dict[str, str] = field(default_factory=dict)
|
||||
|
||||
def active_nodes(self, request: Any) -> list[ProgramNode]:
|
||||
return [n for n in self.nodes if n.when(request)]
|
||||
|
||||
def validate(self) -> Program:
|
||||
produced: set[str] = set()
|
||||
for n in self.nodes:
|
||||
for r in n.reads:
|
||||
if r not in produced:
|
||||
# not fatal in mini (some reads come from the request itself), but flag dangling
|
||||
pass
|
||||
produced.update(n.writes)
|
||||
for name, slot in self.output_artifacts.items():
|
||||
if slot not in produced:
|
||||
raise ValueError(f"program {self.program_id!r} output {name!r} reads slot "
|
||||
f"{slot!r} that no node writes")
|
||||
return self
|
||||
@@ -0,0 +1,187 @@
|
||||
"""Workflow — multi-MODEL composition above the engine (``ProgramKind.WORKFLOW``).
|
||||
|
||||
A ``Program`` composes the loops of ONE resident model instance (the omni/MoT default — every
|
||||
``ModelLoopNode`` resolves on the program's single instance). A ``Workflow`` composes *across model
|
||||
instances*: each stage is a full ``engine.run`` on a (possibly different) registered card, with typed
|
||||
artifacts threaded stage→stage. This is the realistic "T2I model **followed by** I2V model" pipeline
|
||||
(FLUX→Wan): two distinct cards, two distinct weight sets, chained.
|
||||
|
||||
The split mirrors the doc's two layers exactly:
|
||||
* **Program** — single-model, step-interleaved loop composition (the hot path; unchanged).
|
||||
* **Workflow** — multi-model orchestration; a thin layer that calls ``engine.run`` per stage and
|
||||
passes artifacts forward. No change to the engine's single-instance program runner — so cross-
|
||||
model chaining costs nothing on the per-step scheduling path. The engine already holds every
|
||||
instance in its registry; the Workflow just selects which card each stage runs on.
|
||||
|
||||
Why not put cross-model in ``Program``? Because a step-interleaved program shares one resident
|
||||
instance, one cache scope, one parity gate. Crossing instances is a *boundary* (different weights,
|
||||
different caches, an artifact hand-off) — exactly a Workflow edge, not a loop step. Keeping it out of
|
||||
the program runner is what preserves the interleave-parity guarantee for the single-model case.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
|
||||
|
||||
@dataclass
|
||||
class WorkflowStage:
|
||||
"""One model invocation. ``make_request(state)`` builds this stage's Request from the accumulated
|
||||
``state`` (the initial inputs plus every prior stage's artifacts, keyed ``"<stage_label>:<name>"``
|
||||
and also under ``"prev"`` for the immediately-preceding stage's artifact dict)."""
|
||||
model_id: str
|
||||
make_request: Callable[[dict[str, Any]], Any]
|
||||
label: str = ""
|
||||
|
||||
|
||||
def _engine_serves(engine: Any, model_id: str) -> bool:
|
||||
if hasattr(engine, "serves"):
|
||||
return engine.serves(model_id) # AsyncEngine
|
||||
return model_id in getattr(engine, "_registry", {}) # Engine
|
||||
|
||||
|
||||
@dataclass
|
||||
class Workflow:
|
||||
"""A named, registrable cross-model pipeline. ``workflow_id`` is its servable name — it lives in
|
||||
the SAME namespace as a card's ``model_id`` (so the engine, server, and fleet route to it the same
|
||||
way), and by convention is dotted/namespaced to avoid colliding with a card id (e.g.
|
||||
``"image_video.t2i_i2v"``). ``requires`` is the set of cards it composes — declared, validated,
|
||||
never inferred."""
|
||||
workflow_id: str
|
||||
stages: list[WorkflowStage] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def requires(self) -> list[str]:
|
||||
"""The model_ids this workflow composes (deduped, in first-use order). The engine must serve
|
||||
all of them before the workflow can run."""
|
||||
seen: set[str] = set()
|
||||
out: list[str] = []
|
||||
for s in self.stages:
|
||||
if s.model_id not in seen:
|
||||
seen.add(s.model_id)
|
||||
out.append(s.model_id)
|
||||
return out
|
||||
|
||||
def validate(self, engine: Any) -> Workflow:
|
||||
"""Fail-fast: every required card must already be registered on ``engine`` (no silent skips)."""
|
||||
if not self.stages:
|
||||
raise ValueError(f"workflow {self.workflow_id!r} has no stages")
|
||||
missing = [m for m in self.requires if not _engine_serves(engine, m)]
|
||||
if missing:
|
||||
raise ValueError(f"workflow {self.workflow_id!r} requires unregistered models {missing} "
|
||||
f"(register those cards before the workflow)")
|
||||
return self
|
||||
|
||||
def run(self, engine: Any, **initial: Any) -> Any:
|
||||
"""Execute the stages in order on ``engine``; return the final stage's Output.
|
||||
|
||||
Each stage's artifacts are merged into ``state`` so downstream stages can read them
|
||||
(e.g. the I2V stage reads the T2I stage's ``image`` artifact)."""
|
||||
if not self.stages:
|
||||
raise ValueError(f"workflow {self.workflow_id!r} has no stages")
|
||||
state: dict[str, Any] = dict(initial)
|
||||
out = None
|
||||
for stage in self.stages:
|
||||
label = stage.label or stage.model_id
|
||||
req = stage.make_request(state)
|
||||
if req.model_id != stage.model_id:
|
||||
raise ValueError(f"workflow {self.workflow_id!r} stage {label!r}: request model "
|
||||
f"{req.model_id!r} != stage model {stage.model_id!r}")
|
||||
out = engine.run(req)
|
||||
state["prev"] = out.artifacts
|
||||
for name, art in out.artifacts.items():
|
||||
state[f"{label}:{name}"] = art
|
||||
return out
|
||||
|
||||
|
||||
def _payload(artifact: Any) -> Any:
|
||||
for attr in ("frames", "samples", "latent", "tensor", "text"):
|
||||
v = getattr(artifact, attr, None)
|
||||
if v is not None:
|
||||
return v
|
||||
return None
|
||||
|
||||
|
||||
class ParallelWorkflow(Workflow):
|
||||
"""Fan-out: run every stage on the SAME initial input (conceptually in parallel — the engine can
|
||||
interleave their steps), then merge their artifacts into one Output, namespaced by stage label.
|
||||
The non-linear shape a linear chain can't express: one prompt → N models → a merged result (e.g.
|
||||
variants, or video+audio+upscale in parallel)."""
|
||||
|
||||
def run(self, engine: Any, **initial: Any) -> Any:
|
||||
if not self.stages:
|
||||
raise ValueError(f"workflow {self.workflow_id!r} has no stages")
|
||||
from v2.request.artifacts import Output
|
||||
merged: dict[str, Any] = {}
|
||||
metrics: dict[str, float] = {}
|
||||
rid = self.workflow_id
|
||||
for stage in self.stages:
|
||||
label = stage.label or stage.model_id
|
||||
req = stage.make_request(dict(initial)) # each branch sees the ORIGINAL input
|
||||
if req.model_id != stage.model_id:
|
||||
raise ValueError(f"{self.workflow_id!r} stage {label!r}: model mismatch")
|
||||
out = engine.run(req)
|
||||
rid = out.request_id
|
||||
for name, art in out.artifacts.items():
|
||||
merged[f"{label}:{name}"] = art # namespaced — branches don't collide
|
||||
for k, v in out.metrics.items():
|
||||
metrics[f"{label}:{k}"] = v
|
||||
return Output(request_id=rid, artifacts=merged, metrics=metrics)
|
||||
|
||||
|
||||
class BestOfNWorkflow(Workflow):
|
||||
"""Inference-time scaling / rejection sampling: generate N candidates (varying seed), score each with
|
||||
a reward scorer, return the best. A feedback loop across models — generator + reward (e.g. the served
|
||||
REWARD_BATCH card) — the other non-linear shape."""
|
||||
|
||||
def __init__(self, workflow_id, generator_stage, *, scorer, n: int = 4, score_key: str = "latents"):
|
||||
super().__init__(workflow_id, [generator_stage])
|
||||
self.scorer = scorer # any .score(media, prompts) -> {"avg": ...}
|
||||
self.n = int(n)
|
||||
self.score_key = score_key
|
||||
|
||||
def run(self, engine: Any, **initial: Any) -> Any:
|
||||
import numpy as np
|
||||
stage = self.stages[0]
|
||||
base_seed = int(initial.get("seed", 0))
|
||||
cands = []
|
||||
for i in range(self.n):
|
||||
req = stage.make_request({**initial, "seed": base_seed * 100 + i})
|
||||
if req.model_id != stage.model_id:
|
||||
raise ValueError(f"{self.workflow_id!r}: model mismatch")
|
||||
cands.append(engine.run(req))
|
||||
media = [_payload(c.artifacts[self.score_key]) for c in cands]
|
||||
scores = np.asarray(self.scorer.score(media, [initial.get("prompt", "")] * len(cands))["avg"])
|
||||
best = int(np.argmax(scores))
|
||||
out = cands[best]
|
||||
out.metrics["best_of_n"] = float(self.n)
|
||||
out.metrics["best_index"] = float(best)
|
||||
out.metrics["best_score"] = float(scores[best])
|
||||
return out
|
||||
|
||||
|
||||
class WorkflowRegistry:
|
||||
"""Declarative ``workflow_id → builder`` catalog — the cross-model analog of the card builders in
|
||||
``models/__init__.py`` (cf. vllm-omni's ``pipeline_registry``). Adding a custom pipeline is one
|
||||
``register`` call; the builder is a zero/kw-arg factory returning a ``Workflow``."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._builders: dict[str, Any] = {}
|
||||
|
||||
def register(self, workflow_id: str, builder: Any) -> Any:
|
||||
if workflow_id in self._builders:
|
||||
raise ValueError(f"workflow {workflow_id!r} already registered")
|
||||
self._builders[workflow_id] = builder
|
||||
return builder
|
||||
|
||||
def build(self, workflow_id: str, **kw: Any) -> Workflow:
|
||||
if workflow_id not in self._builders:
|
||||
raise KeyError(f"no workflow {workflow_id!r} (have {list(self._builders)})")
|
||||
return self._builders[workflow_id](**kw)
|
||||
|
||||
def names(self) -> list[str]:
|
||||
return list(self._builders)
|
||||
|
||||
def __contains__(self, workflow_id: str) -> bool:
|
||||
return workflow_id in self._builders
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Model definitions — concrete (recipe, runtime) cards.
|
||||
|
||||
Phase 1: Wan2.1-1.3B (T2V), LTX2.3 (2-stage distilled), Wan-causal (self-forcing student).
|
||||
``build_default_engine`` loads all three onto one engine (one resident instance per card).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.recipes.adaptive import build_adaptive_card
|
||||
from v2.recipes.adapters import build_adapter_card, build_adapter_program
|
||||
from v2.recipes.bagel import build_bagel_card, build_bagel_program
|
||||
from v2.recipes.cosmos3 import build_cosmos3_card, build_cosmos3_program
|
||||
from v2.recipes.image_video import (
|
||||
build_flux_t2i_card,
|
||||
build_flux_t2i_program,
|
||||
build_t2i_i2v_extend_workflow,
|
||||
build_t2i_then_i2v_workflow,
|
||||
build_wan_i2v_card,
|
||||
build_wan_i2v_program,
|
||||
)
|
||||
from v2.recipes.ltx2 import build_ltx2_av_program, build_ltx2_card, build_ltx2_program
|
||||
from v2.recipes.qwen_omni import build_qwen_omni_card, build_qwen_omni_program
|
||||
from v2.recipes.reward import build_reward_card
|
||||
from v2.recipes.speculative import build_speculative_card, build_speculative_program
|
||||
from v2.recipes.tiled import build_tiled_card, build_tiled_program
|
||||
from v2.recipes.unified import build_unified_card, build_unified_program
|
||||
from v2.recipes.wan21 import build_wan21_card, build_wan_t2v_program
|
||||
from v2.recipes.wan_causal import build_wan_causal_card, build_wan_causal_program
|
||||
|
||||
__all__ = [
|
||||
"build_wan21_card",
|
||||
"build_wan_t2v_program",
|
||||
"build_ltx2_card",
|
||||
"build_ltx2_program",
|
||||
"build_ltx2_av_program",
|
||||
"build_wan_causal_card",
|
||||
"build_wan_causal_program",
|
||||
"build_cosmos3_card",
|
||||
"build_cosmos3_program",
|
||||
"build_bagel_card",
|
||||
"build_bagel_program",
|
||||
"build_qwen_omni_card",
|
||||
"build_qwen_omni_program",
|
||||
"build_reward_card",
|
||||
"build_speculative_card",
|
||||
"build_speculative_program",
|
||||
"build_unified_card",
|
||||
"build_unified_program",
|
||||
"build_flux_t2i_card",
|
||||
"build_flux_t2i_program",
|
||||
"build_wan_i2v_card",
|
||||
"build_wan_i2v_program",
|
||||
"build_t2i_then_i2v_workflow",
|
||||
"build_t2i_i2v_extend_workflow",
|
||||
"register_workflows",
|
||||
"build_tiled_card",
|
||||
"build_tiled_program",
|
||||
"build_adaptive_card",
|
||||
"build_adapter_card",
|
||||
"build_adapter_program",
|
||||
"build_default_engine",
|
||||
"build_omni_engine",
|
||||
"build_unified_engine",
|
||||
"build_image_video_engine",
|
||||
"build_tiled_engine",
|
||||
]
|
||||
|
||||
_BUILDERS = [
|
||||
(build_wan21_card, build_wan_t2v_program),
|
||||
(build_ltx2_card, build_ltx2_program),
|
||||
(build_wan_causal_card, build_wan_causal_program),
|
||||
]
|
||||
|
||||
# Phase-2 omni cards: MoT shared-weight (Cosmos3/BAGEL) + the cascaded thinker→talker→vocoder
|
||||
# (Qwen-Omni, three disjoint experts / three loop types in one request).
|
||||
_OMNI_BUILDERS = [
|
||||
(build_cosmos3_card, build_cosmos3_program),
|
||||
(build_bagel_card, build_bagel_program),
|
||||
(build_qwen_omni_card, build_qwen_omni_program),
|
||||
]
|
||||
|
||||
|
||||
def build_default_engine(engine: Any = None) -> Any:
|
||||
"""Register Wan2.1, LTX2.3, and Wan-causal onto one engine (one resident instance each)."""
|
||||
from v2.cache import CacheManager
|
||||
from v2.card import load_card
|
||||
from v2.runtime import Engine
|
||||
eng = engine if engine is not None else Engine()
|
||||
for build_card, build_program in _BUILDERS:
|
||||
card = build_card()
|
||||
inst = load_card(card, cache_manager=CacheManager.from_card(card))
|
||||
eng.register(card.model_id, inst, build_program())
|
||||
return eng
|
||||
|
||||
|
||||
def build_omni_engine(engine: Any = None) -> Any:
|
||||
"""Register the phase-2 omni cards (Cosmos3 + BAGEL/lance) onto one engine.
|
||||
|
||||
Each is ONE resident MoT instance whose ``transformer`` runs BOTH an ar_decode loop and a
|
||||
diffusion_denoise loop (shared weights) — true omni/MoT serving.
|
||||
"""
|
||||
from v2.cache import CacheManager
|
||||
from v2.card import load_card
|
||||
from v2.runtime import Engine
|
||||
eng = engine if engine is not None else Engine()
|
||||
for build_card, build_program in _OMNI_BUILDERS:
|
||||
card = build_card()
|
||||
inst = load_card(card, cache_manager=CacheManager.from_card(card))
|
||||
eng.register(card.model_id, inst, build_program())
|
||||
return eng
|
||||
|
||||
|
||||
# Declarative catalog of cross-model workflows: workflow_id -> (builder, required model cards).
|
||||
# Adding a custom pipeline is one line here (the cross-model analog of _BUILDERS / _OMNI_BUILDERS;
|
||||
# cf. vllm-omni's pipeline_registry). ``register_workflows`` registers each whose cards are present.
|
||||
_WORKFLOWS: dict[str, tuple] = {
|
||||
"image_video.t2i_i2v": (build_t2i_then_i2v_workflow, [(build_flux_t2i_card, build_flux_t2i_program),
|
||||
(build_wan_i2v_card, build_wan_i2v_program)]),
|
||||
}
|
||||
|
||||
|
||||
def register_workflows(engine: Any, *, only: list[str] | None = None) -> Any:
|
||||
"""Register catalog workflows (and the cards they require) onto ``engine``. ``only`` selects a
|
||||
subset by workflow_id; default registers all whose cards aren't yet present."""
|
||||
from v2.cache import CacheManager
|
||||
from v2.card import load_card
|
||||
names = only if only is not None else list(_WORKFLOWS)
|
||||
for wf_id in names:
|
||||
build_workflow, card_builders = _WORKFLOWS[wf_id]
|
||||
for build_card, build_program in card_builders:
|
||||
card = build_card()
|
||||
if not engine.serves(card.model_id):
|
||||
inst = load_card(card, cache_manager=CacheManager.from_card(card))
|
||||
engine.register(card.model_id, inst, build_program())
|
||||
engine.register_workflow(build_workflow())
|
||||
return engine
|
||||
|
||||
|
||||
def build_image_video_engine(engine: Any = None) -> Any:
|
||||
"""Register the T2I (``flux-t2i``) and I2V (``wan-i2v``) cards plus the ``image_video.t2i_i2v``
|
||||
workflow on one engine — two *separate* models chained by a cross-model workflow, addressable by
|
||||
its workflow_id like any servable."""
|
||||
from v2.runtime import Engine
|
||||
eng = engine if engine is not None else Engine()
|
||||
return register_workflows(eng, only=["image_video.t2i_i2v"])
|
||||
|
||||
|
||||
def build_tiled_engine(engine: Any = None) -> Any:
|
||||
"""Register the tiled-decode card (``wan-tiled``, whose VAE decode emits ``VAE_TILE`` work units)
|
||||
alongside plain Wan2.1 — so an interleaved batch mixes ``VAE_TILE`` and ``DIFFUSION_STEP`` units
|
||||
through one admission/scheduler budget (the heterogeneous co-scheduling probe)."""
|
||||
from v2.cache import CacheManager
|
||||
from v2.card import load_card
|
||||
from v2.runtime import Engine
|
||||
eng = engine if engine is not None else Engine()
|
||||
for build_card, build_program in [(build_tiled_card, build_tiled_program),
|
||||
(build_wan21_card, build_wan_t2v_program)]:
|
||||
card = build_card()
|
||||
inst = load_card(card, cache_manager=CacheManager.from_card(card))
|
||||
eng.register(card.model_id, inst, build_program())
|
||||
return eng
|
||||
|
||||
|
||||
def build_unified_engine(engine: Any = None) -> Any:
|
||||
"""Register the unified LM+generator card (UniRL/PromptRL): a prompt-refiner ``llm`` expert + a
|
||||
flow ``transformer`` expert on the SAME engine, each driving its own loop. Two separate experts,
|
||||
one request, both trainable under one RL reward (the joint-RL stress test)."""
|
||||
from v2.cache import CacheManager
|
||||
from v2.card import load_card
|
||||
from v2.runtime import Engine
|
||||
eng = engine if engine is not None else Engine()
|
||||
card = build_unified_card()
|
||||
inst = load_card(card, cache_manager=CacheManager.from_card(card))
|
||||
eng.register(card.model_id, inst, build_unified_program())
|
||||
return eng
|
||||
@@ -0,0 +1,47 @@
|
||||
"""Default negative prompts per model family — copied verbatim from fastvideo's pipeline presets
|
||||
(`fastvideo/pipelines/basic/{wan,ltx2}/presets.py`). These are stable recipe DATA, not model code,
|
||||
so v2 keeps its own copy rather than importing private symbols from fastvideo's pipelines package."""
|
||||
from __future__ import annotations
|
||||
|
||||
# Wan family — English (Wan2.1 T2V) and Chinese (Wan2.2 / self-forcing) negative prompts.
|
||||
WAN_NEG_EN = ("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")
|
||||
|
||||
WAN_NEG_CN = ("色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,"
|
||||
"静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,"
|
||||
"多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,"
|
||||
"形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,"
|
||||
"背景人很多,倒着走")
|
||||
|
||||
# Cosmos-Predict2 family negative prompt (verbatim from the Cosmos2 pipeline preset).
|
||||
COSMOS_NEG = ("The video captures a series of frames showing ugly scenes, static with no motion, motion"
|
||||
" blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images,"
|
||||
" poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out"
|
||||
" colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding,"
|
||||
" unnatural transitions, outdated special effects, fake elements, unconvincing visuals,"
|
||||
" poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of"
|
||||
" poor quality.")
|
||||
|
||||
# LTX-2 family negative prompt (base / 2.3); the distilled few-step presets use "" instead.
|
||||
LTX2_NEG = ("blurry, out of focus, overexposed, underexposed, low contrast, "
|
||||
"washed out colors, excessive noise, grainy texture, poor lighting, "
|
||||
"flickering, motion blur, distorted proportions, unnatural skin "
|
||||
"tones, deformed facial features, asymmetrical face, missing facial "
|
||||
"features, extra limbs, disfigured hands, wrong hand count, "
|
||||
"artifacts around text, inconsistent perspective, camera shake, "
|
||||
"incorrect depth of field, background too sharp, background clutter, "
|
||||
"distracting reflections, harsh shadows, inconsistent lighting "
|
||||
"direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
|
||||
"wrong gender, exaggerated expressions, wrong gaze direction, "
|
||||
"mismatched lip sync, silent or muted audio, distorted voice, "
|
||||
"robotic voice, echo, background noise, off-sync audio, incorrect "
|
||||
"dialogue, added dialogue, repetitive speech, jittery movement, "
|
||||
"awkward pauses, incorrect timing, unnatural transitions, "
|
||||
"inconsistent framing, tilted camera, flat lighting, inconsistent "
|
||||
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Adapter plane: one base + swappable LoRA/ControlNet adapters, selected per request."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.recipes.adapters.card import ADAPTERS, build_adapter_card
|
||||
from v2.recipes.adapters.loop import AdapterDenoiseLoop
|
||||
from v2.recipes.adapters.program import build_adapter_program
|
||||
|
||||
__all__ = ["build_adapter_card", "build_adapter_program", "AdapterDenoiseLoop", "ADAPTERS"]
|
||||
@@ -0,0 +1,98 @@
|
||||
"""Adapter-serving card — one base DiT + a set of swappable LoRA/ControlNet adapters.
|
||||
|
||||
Many adapters are declared as components over one resident ``transformer``; a request picks which to
|
||||
apply (``DiffusionParams.adapters``). Adapters are versioned independently (the cache key's
|
||||
``adapter_versions``) and hot-swappable.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import Capability, ConsistencyLevel, LoopKind, WorkUnitKind
|
||||
from v2.card import (
|
||||
CacheContract,
|
||||
CapabilityMatrix,
|
||||
ComponentSpec,
|
||||
CostModel,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ParallelismContract,
|
||||
ParitySpec,
|
||||
PrecisionContract,
|
||||
RecipeSpec,
|
||||
)
|
||||
from v2.loop.policies import ClassicCFG, FlowShiftPolicy, NoRouting, PrecisionPolicy
|
||||
from v2.parallel import ParallelPlan
|
||||
from v2.platform.backends.toy import ToyControlNet, ToyDiT, ToyLoRA, ToyTextEncoder, ToyVAE, _seed_from
|
||||
from v2.recipes.adapters.loop import AdapterDenoiseLoop
|
||||
|
||||
ADAPTERS = ("lora_anime", "lora_realistic", "control_pose") # the served adapter library
|
||||
|
||||
|
||||
def build_adapter_card(model_id: str = "wan-adapters") -> ModelCard:
|
||||
seed = _seed_from(model_id)
|
||||
cost = CostModel(kind=WorkUnitKind.DIFFUSION_STEP, base_seconds=1e-4, per_unit_seconds=1e-7)
|
||||
cfg, flow = ClassicCFG(), FlowShiftPolicy(shift=3.0)
|
||||
precision, expert = PrecisionPolicy(), NoRouting("transformer")
|
||||
|
||||
def loop_factory():
|
||||
return AdapterDenoiseLoop(loop_id="diffusion_denoise",
|
||||
cfg=cfg,
|
||||
flow_shift=flow,
|
||||
precision=precision,
|
||||
expert=expert,
|
||||
cost=cost)
|
||||
|
||||
components = {
|
||||
"text_encoder":
|
||||
ComponentSpec("text_encoder", kind="text_encoder", factory=lambda inst: ToyTextEncoder(), required_for={"t2v"}),
|
||||
"transformer":
|
||||
ComponentSpec("transformer",
|
||||
kind="dit",
|
||||
factory=lambda inst: ToyDiT(seed=seed),
|
||||
resident_for=["diffusion_denoise"],
|
||||
required_for={"t2v"}),
|
||||
"vae":
|
||||
ComponentSpec("vae", kind="vae", factory=lambda inst: ToyVAE(), required_for={"t2v"}),
|
||||
# the adapter library — lightweight, optional, applied per request:
|
||||
"lora_anime":
|
||||
ComponentSpec("lora_anime",
|
||||
kind="adapter",
|
||||
factory=lambda inst: ToyLoRA("lora_anime", scale=0.6, seed=seed + 1),
|
||||
optional_for={"t2v"}),
|
||||
"lora_realistic":
|
||||
ComponentSpec("lora_realistic",
|
||||
kind="adapter",
|
||||
factory=lambda inst: ToyLoRA("lora_realistic", scale=0.6, seed=seed + 2),
|
||||
optional_for={"t2v"}),
|
||||
"control_pose":
|
||||
ComponentSpec("control_pose",
|
||||
kind="adapter",
|
||||
factory=lambda inst: ToyControlNet("control_pose", scale=0.8, seed=seed + 3),
|
||||
optional_for={"t2v"}),
|
||||
}
|
||||
loops = {
|
||||
"diffusion_denoise":
|
||||
LoopSpec("diffusion_denoise",
|
||||
kind=LoopKind.DIFFUSION_DENOISE,
|
||||
work_unit_kind=WorkUnitKind.DIFFUSION_STEP,
|
||||
step_cost_model=cost,
|
||||
shared_weight_components=["transformer"],
|
||||
cache_policy=["feature"],
|
||||
loop_factory=loop_factory),
|
||||
}
|
||||
return ModelCard(
|
||||
model_id=model_id,
|
||||
family="wan",
|
||||
components=components,
|
||||
loops=loops,
|
||||
capabilities=CapabilityMatrix.of(Capability.TEXT_TO_VIDEO, Capability.VAE_DECODE),
|
||||
recipe=RecipeSpec(method="base",
|
||||
assumes_loop="diffusion_denoise",
|
||||
assumes_precision="float32",
|
||||
consistency_required=ConsistencyLevel.C1),
|
||||
parity=ParitySpec(consistency_levels=[ConsistencyLevel.C1], interleave_required=True),
|
||||
caches={
|
||||
"feature": CacheContract("feature", max_bytes=1 << 24, reuse_across_requests=True)
|
||||
},
|
||||
precision=PrecisionContract(default_dtype="float32", training_precision="float32"),
|
||||
parallelism=ParallelismContract(valid_plans=[ParallelPlan.single()], default_plan=ParallelPlan.single()),
|
||||
).validate()
|
||||
@@ -0,0 +1,49 @@
|
||||
"""AdapterDenoiseLoop — per-request LoRA / ControlNet over one resident base.
|
||||
|
||||
A request selects adapters (``DiffusionParams.adapters``); the loop applies each active adapter's
|
||||
velocity delta on top of the base DiT's prediction, then re-integrates. Many adapters are served over
|
||||
ONE resident base; the active set lives in the request/``LoopState`` (never global), so two requests
|
||||
using different adapters interleave without smearing. An empty adapter set is identical to the base
|
||||
``WanDenoiseLoop``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2.loop.contracts import Done
|
||||
from v2.platform import FLOW_MATCH_STEP
|
||||
from v2.recipes.wan21.loop import WanDenoiseLoop
|
||||
|
||||
|
||||
class AdapterDenoiseLoop(WanDenoiseLoop):
|
||||
|
||||
def init(self, req, model, ctx):
|
||||
st = super().init(req, model, ctx)
|
||||
st.scratch["adapters"] = tuple(getattr(req.diffusion, "adapters", ()) or ())
|
||||
st.scratch["control"] = ctx.slots.get("control_signal") # ControlNet conditioning (optional)
|
||||
return st
|
||||
|
||||
def next(self, st):
|
||||
plan = super().next(st)
|
||||
adapters = st.scratch.get("adapters") or ()
|
||||
if isinstance(plan, Done) or not adapters:
|
||||
return plan
|
||||
base_run = plan.run
|
||||
x = st.latents["video"]
|
||||
control = st.scratch.get("control")
|
||||
sigma_t, sigma_next = st.sigmas[st.step_idx], st.sigmas[st.step_idx + 1]
|
||||
prec = self.precision
|
||||
|
||||
def run(model, override=None):
|
||||
res = base_run(model, override) # base velocity (CFG combine)
|
||||
v = np.asarray(res.output["noise_pred"], dtype="float32")
|
||||
for aid in adapters: # apply each active adapter's delta
|
||||
v = v + model.component(aid).delta(x, control)
|
||||
res.output["noise_pred"] = v
|
||||
fm = model.platform.kernels.get(FLOW_MATCH_STEP) # solver dispatched per (device, arch)
|
||||
res.output["latents"] = fm(prec.cast(x), v, sigma_t, sigma_next).astype("float32")
|
||||
return res
|
||||
|
||||
plan.run = run
|
||||
plan.label = f"{plan.label}.adapt[{'+'.join(adapters)}]"
|
||||
return plan
|
||||
@@ -0,0 +1,38 @@
|
||||
"""Adapter-serving program: control_signal → text_encode → diffusion_denoise(+adapters) → vae_decode."""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2.program import ComponentNode, ModelLoopNode, Program, ProgramKind
|
||||
from v2.recipes.common import text_encode_node_fn as _text_encode
|
||||
|
||||
|
||||
def _control_signal(instance, slots, request, ctx) -> None:
|
||||
"""A ControlNet's conditioning image (pose/depth/edge) from the request → a slot the loop reads."""
|
||||
img = request.image()
|
||||
slots["control_signal"] = None if img is None else np.asarray(img.pixels, dtype="float32")
|
||||
|
||||
|
||||
def _vae_decode(instance, slots, request, ctx) -> None:
|
||||
slots["video"] = instance.component("vae").decode(slots["denoise_out"]["latents"])
|
||||
|
||||
|
||||
def build_adapter_program() -> Program:
|
||||
return Program(
|
||||
program_id="wan.adapters",
|
||||
kind=ProgramKind.INLINE,
|
||||
nodes=[
|
||||
ComponentNode("control_signal", fn=_control_signal, writes=("control_signal", )),
|
||||
ComponentNode("text_encode", fn=_text_encode, writes=("text_embeds", "neg_text_embeds")),
|
||||
ModelLoopNode("denoise",
|
||||
loop_id="diffusion_denoise",
|
||||
output_slot="denoise_out",
|
||||
reads=("text_embeds", "control_signal"),
|
||||
writes=("denoise_out", )),
|
||||
ComponentNode("vae_decode", fn=_vae_decode, reads=("denoise_out", ), writes=("video", )),
|
||||
],
|
||||
output_artifacts={
|
||||
"video": "video",
|
||||
"latents": "denoise_out"
|
||||
},
|
||||
).validate()
|
||||
@@ -0,0 +1,7 @@
|
||||
"""Adaptive-compute denoise: the loop owns cache-dit skip + early-exit."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.recipes.adaptive.card import build_adaptive_card
|
||||
from v2.recipes.adaptive.loop import CacheDiTDenoiseLoop
|
||||
|
||||
__all__ = ["build_adaptive_card", "CacheDiTDenoiseLoop"]
|
||||
@@ -0,0 +1,82 @@
|
||||
"""Adaptive-compute card — a Wan-like T2V model whose denoise loop owns content-adaptive control flow
|
||||
(cache-dit skip + early-exit). Reuses the Wan T2V program (loop_id ``diffusion_denoise``); only the
|
||||
loop differs."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import Capability, ConsistencyLevel, LoopKind, WorkUnitKind
|
||||
from v2.card import (
|
||||
CacheContract,
|
||||
CapabilityMatrix,
|
||||
ComponentSpec,
|
||||
CostModel,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ParallelismContract,
|
||||
ParitySpec,
|
||||
PrecisionContract,
|
||||
RecipeSpec,
|
||||
)
|
||||
from v2.loop.policies import ClassicCFG, FlowShiftPolicy, NoRouting, PrecisionPolicy
|
||||
from v2.parallel import ParallelPlan
|
||||
from v2.platform.backends.toy import ToyDiT, ToyTextEncoder, ToyVAE, _seed_from
|
||||
from v2.recipes.adaptive.loop import CacheDiTDenoiseLoop
|
||||
|
||||
|
||||
def build_adaptive_card(model_id: str = "wan-adaptive",
|
||||
*,
|
||||
cache_threshold: float = 0.02,
|
||||
exit_threshold: float = 0.0) -> ModelCard:
|
||||
seed = _seed_from(model_id)
|
||||
cost = CostModel(kind=WorkUnitKind.DIFFUSION_STEP, base_seconds=1e-4, per_unit_seconds=1e-7)
|
||||
cfg, flow = ClassicCFG(), FlowShiftPolicy(shift=3.0)
|
||||
precision, expert = PrecisionPolicy(), NoRouting("transformer")
|
||||
|
||||
def loop_factory():
|
||||
return CacheDiTDenoiseLoop(loop_id="diffusion_denoise",
|
||||
cfg=cfg,
|
||||
flow_shift=flow,
|
||||
precision=precision,
|
||||
expert=expert,
|
||||
cost=cost,
|
||||
cache_threshold=cache_threshold,
|
||||
exit_threshold=exit_threshold)
|
||||
|
||||
components = {
|
||||
"text_encoder":
|
||||
ComponentSpec("text_encoder", kind="text_encoder", factory=lambda inst: ToyTextEncoder(), required_for={"t2v"}),
|
||||
"transformer":
|
||||
ComponentSpec("transformer",
|
||||
kind="dit",
|
||||
factory=lambda inst: ToyDiT(seed=seed),
|
||||
resident_for=["diffusion_denoise"],
|
||||
required_for={"t2v"}),
|
||||
"vae":
|
||||
ComponentSpec("vae", kind="vae", factory=lambda inst: ToyVAE(), required_for={"t2v"}),
|
||||
}
|
||||
loops = {
|
||||
"diffusion_denoise":
|
||||
LoopSpec("diffusion_denoise",
|
||||
kind=LoopKind.DIFFUSION_DENOISE,
|
||||
work_unit_kind=WorkUnitKind.DIFFUSION_STEP,
|
||||
step_cost_model=cost,
|
||||
shared_weight_components=["transformer"],
|
||||
cache_policy=["feature"],
|
||||
loop_factory=loop_factory),
|
||||
}
|
||||
return ModelCard(
|
||||
model_id=model_id,
|
||||
family="wan",
|
||||
components=components,
|
||||
loops=loops,
|
||||
capabilities=CapabilityMatrix.of(Capability.TEXT_TO_VIDEO, Capability.VAE_DECODE),
|
||||
recipe=RecipeSpec(method="base",
|
||||
assumes_loop="diffusion_denoise",
|
||||
assumes_precision="float32",
|
||||
consistency_required=ConsistencyLevel.C1),
|
||||
parity=ParitySpec(consistency_levels=[ConsistencyLevel.C1], interleave_required=True),
|
||||
caches={
|
||||
"feature": CacheContract("feature", max_bytes=1 << 24, reuse_across_requests=True)
|
||||
},
|
||||
precision=PrecisionContract(default_dtype="float32", training_precision="float32"),
|
||||
parallelism=ParallelismContract(valid_plans=[ParallelPlan.single()], default_plan=ParallelPlan.single()),
|
||||
).validate()
|
||||
@@ -0,0 +1,78 @@
|
||||
"""CacheDiTDenoiseLoop — the loop owns content-adaptive control flow.
|
||||
|
||||
Because the loop owns control flow, each step can decide whether to:
|
||||
|
||||
* skip the forward and reuse the cached velocity when consecutive predictions barely change
|
||||
(cache-dit / TeaCache / Δ-DiT), or
|
||||
* early-exit once the latent has converged.
|
||||
|
||||
This yields a variable, content-dependent step count. Subclass of ``WanDenoiseLoop`` (threshold 0 ⇒
|
||||
identical behavior) and interleave-safe: different requests skip different steps without smearing.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2.loop.contracts import Done, StepResult
|
||||
from v2.platform import FLOW_MATCH_STEP
|
||||
from v2.recipes.wan21.loop import WanDenoiseLoop
|
||||
|
||||
|
||||
class CacheDiTDenoiseLoop(WanDenoiseLoop):
|
||||
|
||||
def __init__(self, *, cache_threshold: float = 0.0, exit_threshold: float = 0.0, **kw):
|
||||
super().__init__(**kw)
|
||||
self.cache_threshold = cache_threshold # rel velocity change below which we reuse (skip the DiT)
|
||||
self.exit_threshold = exit_threshold # rel latent change below which we stop early
|
||||
|
||||
def init(self, req, model, ctx):
|
||||
st = super().init(req, model, ctx)
|
||||
st.scratch.update(last_v=None, skip_next=False, skipped=0, exited=False)
|
||||
return st
|
||||
|
||||
def next(self, st):
|
||||
if st.scratch.get("exited"):
|
||||
return Done() # converged early — content-adaptive termination
|
||||
plan = super().next(st) # the normal full-forward step plan
|
||||
last_v = st.scratch.get("last_v")
|
||||
if isinstance(plan, Done) or not self.cache_threshold or not st.scratch.get("skip_next") \
|
||||
or last_v is None:
|
||||
return plan
|
||||
# SKIP: reuse the cached velocity, no DiT forward (the cache-dit win) — wrap the run thunk
|
||||
x, prec = st.latents["video"], self.precision
|
||||
sigma_t, sigma_next = st.sigmas[st.step_idx], st.sigmas[st.step_idx + 1]
|
||||
v = np.asarray(last_v, dtype="float32")
|
||||
|
||||
def run(model, override=None):
|
||||
x_next = model.platform.kernels.get(FLOW_MATCH_STEP)(prec.cast(x), v, sigma_t, sigma_next)
|
||||
return StepResult(output={"noise_pred": v, "latents": x_next.astype("float32"), "skipped": True})
|
||||
|
||||
plan.run = run
|
||||
plan.label = f"{plan.label}.skip"
|
||||
return plan
|
||||
|
||||
def advance(self, st, result):
|
||||
x_prev = np.asarray(st.latents["video"], dtype="float64")
|
||||
v = np.asarray(result.output["noise_pred"], dtype="float64")
|
||||
prev_v = st.scratch.get("last_v")
|
||||
st = super().advance(st, result) # folds latents, step_idx, (rollout) trajectory
|
||||
if result.output.get("skipped"):
|
||||
st.scratch["skipped"] += 1
|
||||
st.scratch["skip_next"] = False # don't chain skips — re-check with a full forward
|
||||
elif self.cache_threshold and prev_v is not None:
|
||||
rel = float(np.linalg.norm(v - np.asarray(prev_v, dtype="float64")) / (np.linalg.norm(prev_v) + 1e-8))
|
||||
st.scratch["skip_next"] = rel < self.cache_threshold # next prediction will barely change
|
||||
st.scratch["last_v"] = result.output["noise_pred"]
|
||||
half = max(1, (len(st.sigmas) - 1) // 2) # only consider exiting in the low-noise tail
|
||||
if self.exit_threshold and st.step_idx >= half:
|
||||
x_next = np.asarray(st.latents["video"], dtype="float64")
|
||||
dl = float(np.linalg.norm(x_next - x_prev) / (np.linalg.norm(x_prev) + 1e-8))
|
||||
if dl < self.exit_threshold:
|
||||
st.scratch["exited"] = True # latent converged → stop next tick
|
||||
return st
|
||||
|
||||
def finalize(self, st):
|
||||
res = super().finalize(st)
|
||||
res.metrics["skipped_steps"] = float(st.scratch.get("skipped", 0))
|
||||
res.metrics["early_exited"] = 1.0 if st.scratch.get("exited") else 0.0
|
||||
return res
|
||||
@@ -0,0 +1,7 @@
|
||||
"""BAGEL/lance — canonical vllm-omni MoT (AR generate_text + diffusion generate_image). Phase 2."""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.recipes.bagel.card import build_bagel_card
|
||||
from v2.recipes.bagel.program import build_bagel_program
|
||||
|
||||
__all__ = ["build_bagel_card", "build_bagel_program"]
|
||||
@@ -0,0 +1,104 @@
|
||||
"""BAGEL/lance-style MoT ModelCard — canonical vllm-omni omni model.
|
||||
|
||||
vllm-omni's ``bagel_single_stage``/``lance`` run one resident MoT instance doing both AR
|
||||
``generate_text`` and diffusion ``generate_image`` on co-resident experts in a single request, but
|
||||
bury the interleaving inside one opaque ``DIFFUSION`` stage the scheduler can't see into
|
||||
(``max_num_running_reqs=1``). This card expresses the same shared-weight MoT with both loops
|
||||
runtime-visible, step-scheduled, and batchable: ``generate_text`` (ar_decode) and ``generate_image``
|
||||
(diffusion_denoise) both bind the one resident ``transformer``, and every AR token and denoise step
|
||||
is a WorkUnit the scheduler can interleave and price.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import Capability, ConsistencyLevel, LoopKind, WorkUnitKind
|
||||
from v2.card import (
|
||||
CacheContract,
|
||||
CapabilityMatrix,
|
||||
ComponentSpec,
|
||||
CostModel,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ParallelismContract,
|
||||
ParitySpec,
|
||||
PrecisionContract,
|
||||
RecipeSpec,
|
||||
)
|
||||
from v2.loop.policies import ClassicCFG, FlowShiftPolicy, NoRouting, PrecisionPolicy
|
||||
from v2.parallel import ParallelPlan
|
||||
from v2.platform.backends.toy import ToyMoTDiT, ToyTokenizer, ToyVAE, _seed_from
|
||||
from v2.recipes.omni import ARDecodeLoop
|
||||
from v2.recipes.wan21.loop import WanDenoiseLoop
|
||||
|
||||
|
||||
def build_bagel_card(model_id: str = "bagel-mot") -> ModelCard:
|
||||
seed = _seed_from(model_id)
|
||||
ar_cost = CostModel(kind=WorkUnitKind.AR_TOKEN, base_seconds=5e-5, per_unit_seconds=1e-7)
|
||||
dn_cost = CostModel(kind=WorkUnitKind.DIFFUSION_STEP, base_seconds=1e-4, per_unit_seconds=1e-7)
|
||||
cfg, flow = ClassicCFG(), FlowShiftPolicy(shift=3.0)
|
||||
precision, expert = PrecisionPolicy(), NoRouting("transformer")
|
||||
|
||||
def text_factory():
|
||||
return ARDecodeLoop(loop_id="generate_text", transformer_id="transformer", cost=ar_cost, max_tokens=6)
|
||||
|
||||
def image_factory():
|
||||
return WanDenoiseLoop(loop_id="generate_image",
|
||||
cfg=cfg,
|
||||
flow_shift=flow,
|
||||
precision=precision,
|
||||
expert=expert,
|
||||
cost=dn_cost)
|
||||
|
||||
components = {
|
||||
"tokenizer":
|
||||
ComponentSpec("tokenizer",
|
||||
kind="tokenizer",
|
||||
factory=lambda inst: ToyTokenizer(),
|
||||
required_for={"reason", "t2i"}),
|
||||
"transformer":
|
||||
ComponentSpec(
|
||||
"transformer",
|
||||
kind="dit",
|
||||
load_id="vllm_omni.diffusion.models.bagel:BagelTransformer",
|
||||
factory=lambda inst: ToyMoTDiT(seed=seed),
|
||||
resident_for=["generate_text", "generate_image"], # one resident copy for BOTH loops
|
||||
required_for={"reason", "t2i"}),
|
||||
"vae":
|
||||
ComponentSpec("vae", kind="vae", factory=lambda inst: ToyVAE(), required_for={"t2i"}),
|
||||
}
|
||||
loops = {
|
||||
"generate_text":
|
||||
LoopSpec("generate_text",
|
||||
kind=LoopKind.AR_DECODE,
|
||||
work_unit_kind=WorkUnitKind.AR_TOKEN,
|
||||
step_cost_model=ar_cost,
|
||||
shared_weight_components=["transformer"],
|
||||
cache_policy=["paged_kv"],
|
||||
loop_factory=text_factory),
|
||||
"generate_image":
|
||||
LoopSpec("generate_image",
|
||||
kind=LoopKind.DIFFUSION_DENOISE,
|
||||
work_unit_kind=WorkUnitKind.DIFFUSION_STEP,
|
||||
step_cost_model=dn_cost,
|
||||
shared_weight_components=["transformer"],
|
||||
cache_policy=["feature"],
|
||||
loop_factory=image_factory),
|
||||
}
|
||||
card = ModelCard(
|
||||
model_id=model_id,
|
||||
family="bagel",
|
||||
components=components,
|
||||
loops=loops,
|
||||
capabilities=CapabilityMatrix.of(Capability.TEXT_TO_IMAGE, Capability.REASONING_TEXT, Capability.VAE_DECODE),
|
||||
recipe=RecipeSpec(method="base",
|
||||
assumes_loop="generate_image",
|
||||
assumes_precision="float32",
|
||||
consistency_required=ConsistencyLevel.C1),
|
||||
parity=ParitySpec(consistency_levels=[ConsistencyLevel.C1], interleave_required=True),
|
||||
caches={
|
||||
"feature": CacheContract("feature", max_bytes=1 << 24, reuse_across_requests=True),
|
||||
"paged_kv": CacheContract("paged_kv", max_bytes=1 << 24, block_bytes=1 << 12, reuse_across_requests=False),
|
||||
},
|
||||
precision=PrecisionContract(default_dtype="float32", training_precision="float32"),
|
||||
parallelism=ParallelismContract(valid_plans=[ParallelPlan.single()], default_plan=ParallelPlan.single()),
|
||||
)
|
||||
return card.validate()
|
||||
@@ -0,0 +1,47 @@
|
||||
"""BAGEL/lance program: one resident MoT instance, two runtime-visible loops.
|
||||
|
||||
tokenize → generate_text (ar_decode) → pack(text→cond) → generate_image (diffusion) → vae_decode
|
||||
|
||||
generate_text and generate_image hit the SAME resident weights; their steps are WorkUnits the
|
||||
scheduler interleaves.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.program import ComponentNode, ModelLoopNode, Program, ProgramKind
|
||||
from v2.recipes.omni import emit_text_node, pack_cond_from_tokens, tokenize_node, vae_decode_node
|
||||
|
||||
|
||||
def build_bagel_program() -> Program:
|
||||
return Program(
|
||||
program_id="bagel.mot",
|
||||
kind=ProgramKind.INLINE,
|
||||
nodes=[
|
||||
ComponentNode("tokenize", fn=tokenize_node, writes=("prompt_tokens", )),
|
||||
ModelLoopNode("generate_text",
|
||||
loop_id="generate_text",
|
||||
output_slot="gen_text_out",
|
||||
reads=("prompt_tokens", ),
|
||||
writes=("gen_text_out", )),
|
||||
ComponentNode("emit_text",
|
||||
fn=emit_text_node("gen_text_out", "text"),
|
||||
reads=("gen_text_out", ),
|
||||
writes=("text", )),
|
||||
ComponentNode("pack",
|
||||
fn=pack_cond_from_tokens("gen_text_out"),
|
||||
reads=("gen_text_out", ),
|
||||
writes=("text_embeds", "neg_text_embeds")),
|
||||
ModelLoopNode("generate_image",
|
||||
loop_id="generate_image",
|
||||
output_slot="gen_image_out",
|
||||
reads=("text_embeds", ),
|
||||
writes=("gen_image_out", )),
|
||||
ComponentNode("vae_decode",
|
||||
fn=vae_decode_node("gen_image_out", "image"),
|
||||
reads=("gen_image_out", ),
|
||||
writes=("image", )),
|
||||
],
|
||||
output_artifacts={
|
||||
"text": "text",
|
||||
"image": "image"
|
||||
},
|
||||
).validate()
|
||||
@@ -0,0 +1,38 @@
|
||||
"""Shared component-node helpers (text encode with feature cache, etc.).
|
||||
|
||||
The content-hash feature cache lets a K-sample RL group encode its shared prompt once instead of
|
||||
K times.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from v2.cache.keys import CacheKey, content_hash
|
||||
|
||||
|
||||
def cached_text_encode(instance: Any, text: str) -> Any:
|
||||
enc = instance.component("text_encoder")
|
||||
cache = instance.caches
|
||||
if cache is not None and cache.has("feature"):
|
||||
# Every output-semantic field is in the key — adapter stack + precision partition it, and the
|
||||
# text-encoder's OWN version (not the instance's) so a transformer weight sync does not
|
||||
# invalidate the (frozen) text encoder's embeddings.
|
||||
key = CacheKey(model_id=instance.card.model_id,
|
||||
component_id="text_encoder",
|
||||
weights_version=instance.version_of("text_encoder"),
|
||||
adapter_versions=CacheKey.adapters(instance.adapter_versions),
|
||||
precision=instance.card.precision.dtype_for("text_encoder"),
|
||||
input_hashes=(("text", content_hash(text)), ))
|
||||
hit = cache.pool("feature").get(key)
|
||||
if hit is not None:
|
||||
return hit
|
||||
emb = enc.encode(text)
|
||||
cache.pool("feature").put(key, emb)
|
||||
return emb
|
||||
return enc.encode(text)
|
||||
|
||||
|
||||
def text_encode_node_fn(instance, slots, request, ctx) -> None:
|
||||
"""ComponentNode fn: prompt + negative prompt → cached text embeddings in slots."""
|
||||
slots["text_embeds"] = cached_text_encode(instance, request.prompt())
|
||||
slots["neg_text_embeds"] = cached_text_encode(instance, request.diffusion.negative_prompt)
|
||||
@@ -0,0 +1,13 @@
|
||||
"""Cosmos-Predict2 (Video2World) EDM-Karras denoiser recipe.
|
||||
|
||||
Self-contained recipe: the card declares torch adapters (``CosmosDiT``/``CosmosT5Encoder`` in
|
||||
``v2/recipes/cosmos2/adapter.py``) plus ``CosmosDenoiseLoop`` (EDM preconditioning folded into
|
||||
flow-match Euler), reusing the Wan VAE adapter + T5 + ``stamp_wan21_checkpoints``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.recipes.cosmos2.card import build_cosmos2_card
|
||||
from v2.recipes.cosmos2.loop import CosmosDenoiseLoop
|
||||
from v2.recipes.cosmos2.program import build_cosmos2_program
|
||||
|
||||
__all__ = ["build_cosmos2_card", "build_cosmos2_program", "CosmosDenoiseLoop"]
|
||||
@@ -0,0 +1,61 @@
|
||||
"""Cosmos-Predict2 torch adapters (GPU backend) — declared on the card via ``ComponentSpec.adapter``
|
||||
so the Cosmos recipe is self-contained (no edit to the shared ``_make_dit``/``_make_text_encoder``
|
||||
dispatch in ``torch_backend.py``). Imported lazily by ``_explicit_adapter`` only on a GPU box.
|
||||
|
||||
* ``CosmosDiT`` — the EDM denoiser ``F_θ``. The loop hands the *already EDM-input-scaled* model input
|
||||
(``x·c_in``) and the model timestep ``t = σ·1000``; this adapter returns the **raw** transformer
|
||||
output (the EDM ``c_skip``/``c_out`` → x0 reconstruction + x0-space CFG live in ``CosmosDenoiseLoop``,
|
||||
NOT here). It builds the mandatory zero ``condition_mask`` / ``padding_mask`` and ``fps`` the Cosmos
|
||||
forward requires (the model concats the masks internally → 18ch patch_embed input). Faithful to
|
||||
``CosmosDenoisingStage._run_transformer`` (fastvideo/pipelines/stages/denoising.py).
|
||||
* ``CosmosT5Encoder`` — T5-Large (1024-dim). Cosmos uses the **raw** last_hidden_state (NaN→0, no
|
||||
fixed-length zero-pad), unlike the Wan T5 convention that zero-pads to 512.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from v2.platform.backends.torch_backend import T5Encoder, TorchComponent, _to_numpy
|
||||
|
||||
|
||||
class CosmosDiT(TorchComponent):
|
||||
"""``dit(model_input[C,T,h,w], text_embed[seq,1024], timestep) -> raw noise_pred[C,T,h,w]``.
|
||||
|
||||
``model_input`` is pre-scaled by the loop (``x·c_in``); ``timestep`` is ``σ·1000`` (1D [B] float)."""
|
||||
|
||||
@torch.no_grad()
|
||||
def __call__(self, model_input, text_embed, timestep, context=None, *, cond=None):
|
||||
hs = self._t(model_input) # [1, C, T, h, w]
|
||||
ehs = self._t(text_embed)
|
||||
b, _c, t, h, w = hs.shape
|
||||
ts = torch.tensor([float(timestep)], device=self.device, dtype=self.dtype).expand(b)
|
||||
condition_mask = torch.zeros(b, 1, t, h, w, device=self.device, dtype=self.dtype) # t2v: zeros
|
||||
padding_mask = torch.zeros(1, 1, h, w, device=self.device, dtype=self.dtype) # resized→[h,w]
|
||||
with self._ctx():
|
||||
out = self.module(hidden_states=hs,
|
||||
timestep=ts,
|
||||
encoder_hidden_states=ehs,
|
||||
fps=24,
|
||||
condition_mask=condition_mask,
|
||||
padding_mask=padding_mask,
|
||||
return_dict=False)[0]
|
||||
return self._n(out) # RAW EDM output (loop reconstructs x0)
|
||||
|
||||
|
||||
class CosmosT5Encoder(T5Encoder):
|
||||
"""Cosmos T5-Large: raw last_hidden_state (NaN→0), NO zero-pad-to-max (the Wan convention would
|
||||
mis-condition Cosmos). Truncates to ``max_length`` but returns only the real-token rows."""
|
||||
|
||||
@torch.no_grad()
|
||||
def encode(self, text):
|
||||
toks = self.tokenizer(text or "", return_tensors="pt", max_length=self.max_length, truncation=True)
|
||||
ids = toks.input_ids.to(self.device)
|
||||
mask = toks.attention_mask.to(self.device)
|
||||
with self._ctx():
|
||||
out = self.module(input_ids=ids, attention_mask=mask)
|
||||
hidden = (out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0]).squeeze(0)
|
||||
hidden = torch.nan_to_num(hidden, nan=0.0)
|
||||
n = int(mask.sum().item()) # zero any row beyond real-token length
|
||||
if n < hidden.shape[0]:
|
||||
hidden[n:] = 0.0
|
||||
return _to_numpy(hidden)
|
||||
@@ -0,0 +1,121 @@
|
||||
"""Cosmos-Predict2-2B-Video2World ModelCard — text->video (the registered preset; video2world
|
||||
conditioning is a later capability the loop already threads).
|
||||
|
||||
Architecture deltas vs Wan (declared on the card so the recipe is self-contained):
|
||||
* DiT ``fastvideo.models.dits.cosmos:CosmosTransformer3DModel`` — an EDM denoiser (not flow-match); the
|
||||
``CosmosDiT`` adapter returns the raw network output and ``CosmosDenoiseLoop`` does the EDM
|
||||
``c_in/c_skip/c_out`` -> x0 reconstruction + x0-space CFG + the Karras (rho=7, sigma 80->0.002) schedule.
|
||||
* VAE ``fastvideo.models.vaes.wanvae:AutoencoderKLWan`` — Wan-style (z=16, 8x/4x); reuses the v2
|
||||
``WanVAE`` adapter unchanged (sigma_data=1.0 makes the Cosmos sigma_data factor a no-op).
|
||||
* Text ``fastvideo.models.encoders.t5:T5EncoderModel`` (T5-Large, 1024-dim) via ``CosmosT5Encoder``
|
||||
(raw last_hidden_state, no Wan zero-pad).
|
||||
``stamp_wan21_checkpoints`` applies (diffusers transformer/vae/text_encoder subfolder layout).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2._enums import Capability, ConsistencyLevel, LoopKind, WorkUnitKind
|
||||
from v2.card import (
|
||||
CacheContract,
|
||||
CapabilityMatrix,
|
||||
ComponentSpec,
|
||||
CostModel,
|
||||
LoopSpec,
|
||||
ModelCard,
|
||||
ParallelismContract,
|
||||
ParitySpec,
|
||||
ParityTestSpec,
|
||||
PrecisionContract,
|
||||
RecipeSpec,
|
||||
SamplingDefaults,
|
||||
)
|
||||
from v2.loop.policies import ClassicCFG, NoRouting, PrecisionPolicy
|
||||
from v2.parallel import ParallelPlan
|
||||
from v2.platform.backends.toy import ToyDiT, ToyTextEncoder, ToyVAE, _seed_from
|
||||
from v2.recipes._prompts import COSMOS_NEG
|
||||
from v2.recipes.cosmos2.loop import CosmosDenoiseLoop
|
||||
from v2.recipes.wan21.card import stamp_wan21_checkpoints
|
||||
|
||||
_COSMOS_DIT = "v2.recipes.cosmos2.adapter:CosmosDiT"
|
||||
_COSMOS_T5 = "v2.recipes.cosmos2.adapter:CosmosT5Encoder"
|
||||
|
||||
|
||||
def build_cosmos2_card(model_id: str = "cosmos-predict2-2b-video2world",
|
||||
*,
|
||||
checkpoint_root: str | None = None,
|
||||
sampling_defaults: SamplingDefaults | None = None) -> ModelCard:
|
||||
seed = _seed_from(model_id)
|
||||
cost = CostModel(kind=WorkUnitKind.DIFFUSION_STEP, base_seconds=1e-4, per_unit_seconds=1e-7)
|
||||
cfg = ClassicCFG()
|
||||
precision = PrecisionPolicy(compute_dtype="float32", scheduler_step_in_fp32=True)
|
||||
expert = NoRouting("transformer")
|
||||
|
||||
def loop_factory():
|
||||
return CosmosDenoiseLoop(loop_id="diffusion_denoise",
|
||||
cfg=cfg,
|
||||
precision=precision,
|
||||
expert=expert,
|
||||
cost=cost,
|
||||
sigma_max=80.0,
|
||||
sigma_min=0.002,
|
||||
sigma_data=1.0,
|
||||
rho=7.0,
|
||||
augment_sigma=0.001)
|
||||
|
||||
components = {
|
||||
"text_encoder":
|
||||
ComponentSpec(component_id="text_encoder",
|
||||
kind="text_encoder",
|
||||
load_id="fastvideo.models.encoders.t5:T5EncoderModel",
|
||||
adapter=_COSMOS_T5,
|
||||
factory=lambda inst: ToyTextEncoder(),
|
||||
required_for={"t2v"}),
|
||||
"vae":
|
||||
ComponentSpec(component_id="vae",
|
||||
kind="vae",
|
||||
load_id="fastvideo.models.vaes.wanvae:AutoencoderKLWan",
|
||||
factory=lambda inst: ToyVAE(),
|
||||
required_for={"t2v"}),
|
||||
"transformer":
|
||||
ComponentSpec(component_id="transformer",
|
||||
kind="dit",
|
||||
load_id="fastvideo.models.dits.cosmos:CosmosTransformer3DModel",
|
||||
adapter=_COSMOS_DIT,
|
||||
factory=lambda inst: ToyDiT(seed=seed),
|
||||
resident_for=["diffusion_denoise"],
|
||||
required_for={"t2v"}),
|
||||
}
|
||||
loops = {
|
||||
"diffusion_denoise":
|
||||
LoopSpec(loop_id="diffusion_denoise",
|
||||
kind=LoopKind.DIFFUSION_DENOISE,
|
||||
work_unit_kind=WorkUnitKind.DIFFUSION_STEP,
|
||||
step_cost_model=cost,
|
||||
shared_weight_components=["transformer"],
|
||||
cache_policy=["feature"],
|
||||
graph_capture="breakable_cudagraph",
|
||||
loop_factory=loop_factory),
|
||||
}
|
||||
card = ModelCard(
|
||||
model_id=model_id,
|
||||
family="cosmos",
|
||||
components=components,
|
||||
loops=loops,
|
||||
capabilities=CapabilityMatrix.of(Capability.TEXT_TO_VIDEO, Capability.VAE_DECODE, Capability.POLICY_ROLLOUT),
|
||||
recipe=RecipeSpec(method="base",
|
||||
assumes_loop="diffusion_denoise",
|
||||
assumes_precision="float32",
|
||||
consistency_required=ConsistencyLevel.C1),
|
||||
parity=ParitySpec(consistency_levels=[ConsistencyLevel.C1],
|
||||
interleave_required=True,
|
||||
tests=[ParityTestSpec(name="denoise_trajectory", level=ConsistencyLevel.C1, tap="latents")]),
|
||||
caches={"feature": CacheContract(cache_class="feature", max_bytes=1 << 24, reuse_across_requests=True)},
|
||||
precision=PrecisionContract(default_dtype="float32", training_precision="float32"),
|
||||
parallelism=ParallelismContract(valid_plans=[ParallelPlan.single()], default_plan=ParallelPlan.single()),
|
||||
sampling_defaults=sampling_defaults or SamplingDefaults(
|
||||
num_steps=35, guidance_scale=7.0, height=704, width=1280, num_frames=93, fps=16,
|
||||
negative_prompt=COSMOS_NEG),
|
||||
)
|
||||
card.validate()
|
||||
if checkpoint_root:
|
||||
stamp_wan21_checkpoints(card, checkpoint_root)
|
||||
return card
|
||||
@@ -0,0 +1,198 @@
|
||||
"""CosmosDenoiseLoop — EDM-Karras preconditioning folded into a flow-match Euler integrator.
|
||||
|
||||
Cosmos-Predict2 is not a plain flow-match model despite the pipeline wrapping a
|
||||
``FlowMatchEulerDiscreteScheduler``. The network is an EDM denoiser ``F_θ`` whose preconditioned output
|
||||
reconstructs ``x0`` (``x0 = c_skip·x + c_out·F_θ(x·c_in)``, ``sigma_data=1``); CFG combines in x0 space;
|
||||
then ``x0`` is converted to a flow-match velocity ``(x - x0)/σ`` for the Euler update
|
||||
``x_next = x + (σ_next - σ)·v``. The σ schedule is Karras (ρ=7, σ_max=80 -> σ_min=0.002), not the
|
||||
flow-shift linspace, and the latent starts at ``randn·σ_max``. Port of
|
||||
``fastvideo/pipelines/stages/denoising.py:CosmosDenoisingStage`` (see that file for the exact math).
|
||||
|
||||
video2world (frame_replace) conditioning is threaded (``conditioning_latents``/``cond_indicator`` from
|
||||
slots) but is ``None`` for the registered t2v preset -> the gated injection is inert; the loop degrades
|
||||
to pure t2v exactly as the fastvideo pipeline does when no image/video is given.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from v2._enums import ExecutionProfile, WorkUnitKind
|
||||
from v2.loop.contracts import (
|
||||
Done,
|
||||
LoopResult,
|
||||
LoopState,
|
||||
ResourceRequest,
|
||||
ShapeSignature,
|
||||
StepContext,
|
||||
StepResult,
|
||||
WorkPlan,
|
||||
)
|
||||
from v2.loop.sampler import build_karras_sigmas
|
||||
from v2.platform import FLOW_MATCH_STEP
|
||||
from v2.recipes.wan21.loop import latent_shape
|
||||
|
||||
COSMOS_LATENT_CHANNELS = 16 # config.in_channels(17) - 1 (the condition_mask channel)
|
||||
COSMOS_TEMPORAL_RATIO = 4
|
||||
COSMOS_SPATIAL_RATIO = 8
|
||||
|
||||
|
||||
class CosmosDenoiseLoop:
|
||||
|
||||
def __init__(self,
|
||||
*,
|
||||
loop_id,
|
||||
cfg,
|
||||
precision,
|
||||
expert,
|
||||
cost,
|
||||
sigma_max: float = 80.0,
|
||||
sigma_min: float = 0.002,
|
||||
sigma_data: float = 1.0,
|
||||
rho: float = 7.0,
|
||||
augment_sigma: float = 0.001,
|
||||
latent_channels: int = COSMOS_LATENT_CHANNELS,
|
||||
spatial_ratio: int = COSMOS_SPATIAL_RATIO,
|
||||
temporal_ratio: int = COSMOS_TEMPORAL_RATIO):
|
||||
self.loop_id = loop_id
|
||||
self.cfg = cfg # carried for the WorkPlan op-structure key; x0-space CFG is done here
|
||||
self.precision = precision
|
||||
self.expert = expert
|
||||
self.cost = cost
|
||||
self.sigma_max = sigma_max
|
||||
self.sigma_min = sigma_min
|
||||
self.sigma_data = sigma_data
|
||||
self.rho = rho
|
||||
self.augment_sigma = augment_sigma
|
||||
self.latent_channels = latent_channels
|
||||
self.spatial_ratio = spatial_ratio
|
||||
self.temporal_ratio = temporal_ratio
|
||||
|
||||
def init(self, req, model, ctx) -> LoopState:
|
||||
seed = req.diffusion.seed if req.diffusion.seed is not None else 0
|
||||
rng = np.random.default_rng(seed)
|
||||
sig = build_karras_sigmas(req.diffusion.num_steps,
|
||||
sigma_max=self.sigma_max,
|
||||
sigma_min=self.sigma_min,
|
||||
rho=self.rho)
|
||||
shape = latent_shape(req,
|
||||
model,
|
||||
channels=self.latent_channels,
|
||||
spatial_ratio=self.spatial_ratio,
|
||||
temporal_ratio=self.temporal_ratio)
|
||||
x = (rng.standard_normal(shape) * float(self.sigma_max)).astype("float32") # EDM: randn·σ_max
|
||||
st = LoopState(loop_id=self.loop_id,
|
||||
instance_id=model.card.model_id,
|
||||
request_id=req.request_id,
|
||||
profile=ctx.profile,
|
||||
rng=rng,
|
||||
seed=seed,
|
||||
latents={"video": x},
|
||||
sigmas=[float(s) for s in sig],
|
||||
timesteps=[float(s) * 1000.0 for s in sig]) # model timestep = σ·1000 (large for EDM)
|
||||
st.cond["prompt_embeds"] = ctx.slots.get("text_embeds")
|
||||
st.cond["negative_prompt_embeds"] = ctx.slots.get("neg_text_embeds")
|
||||
st.scratch["guidance_scale"] = float(req.diffusion.guidance_scale)
|
||||
st.scratch["stream_video"] = bool(req.outputs.stream.get("video"))
|
||||
# video2world (frame_replace): VAE-encoded conditioning latents + cond/uncond indicators. None for
|
||||
# the t2v preset -> the gated injection below is skipped (pure t2v, matching the fastvideo pipeline).
|
||||
st.scratch["conditioning_latents"] = ctx.slots.get("conditioning_latents")
|
||||
st.scratch["cond_indicator"] = ctx.slots.get("cond_indicator")
|
||||
st.scratch["uncond_indicator"] = ctx.slots.get("uncond_indicator")
|
||||
st.plugin_state["cfg"] = {}
|
||||
return st
|
||||
|
||||
def next(self, st: LoopState):
|
||||
i = st.step_idx
|
||||
if i >= len(st.sigmas) - 1:
|
||||
return Done()
|
||||
sigma_t, sigma_next = st.sigmas[i], st.sigmas[i + 1]
|
||||
t = st.timesteps[i] # σ·1000, the model timestep
|
||||
expert_id = self.expert.expert_for(StepContext(i, t, sigma_t))
|
||||
x = st.latents["video"]
|
||||
pe, ne = st.cond["prompt_embeds"], st.cond["negative_prompt_embeds"]
|
||||
scale = st.scratch["guidance_scale"]
|
||||
precision = self.precision
|
||||
sd = float(self.sigma_data)
|
||||
do_cfg = scale != 1.0 and ne is not None
|
||||
cond_lat = st.scratch.get("conditioning_latents")
|
||||
cond_ind = st.scratch.get("cond_indicator")
|
||||
uncond_ind = st.scratch.get("uncond_indicator")
|
||||
aug = float(self.augment_sigma)
|
||||
|
||||
def _x0(model: Any, dit: Any, latent_in: np.ndarray, text_embed: Any, ind: Any) -> np.ndarray:
|
||||
# EDM preconditioning + (gated) frame-replace conditioning, returning the x0 prediction for one
|
||||
# CFG branch. ``ind`` is the (un)cond indicator; None for t2v -> no frame injection.
|
||||
s = float(sigma_t)
|
||||
c_in = 1.0 / (s**2 + sd**2)**0.5
|
||||
c_skip = sd**2 / (s**2 + sd**2)
|
||||
c_out = s * sd / (s**2 + sd**2)**0.5
|
||||
cur_ind = (ind * 0 if (ind is not None and aug >= s) else ind)
|
||||
lat = np.array(latent_in, dtype="float32")
|
||||
if cur_ind is not None and cond_lat is not None: # video2world frame injection (inert for t2v)
|
||||
c_in_aug = 1.0 / (aug**2 + sd**2)**0.5
|
||||
cn = st.rng.standard_normal(lat.shape).astype("float32")
|
||||
cf = (cond_lat + cn * aug) * c_in_aug / c_in
|
||||
lat = cur_ind * cf + (1 - cur_ind) * lat
|
||||
model_input = (lat * c_in).astype("float32")
|
||||
np_pred = np.asarray(dit(model_input, text_embed, t), dtype="float32")
|
||||
x0 = c_skip * x + c_out * np_pred
|
||||
if cur_ind is not None and cond_lat is not None:
|
||||
x0 = cur_ind * cond_lat + (1 - cur_ind) * x0
|
||||
return x0
|
||||
|
||||
def run(model, override=None):
|
||||
dit = model.component(expert_id)
|
||||
if override is not None and "noise_pred" in override:
|
||||
velocity = precision.cast(np.asarray(override["noise_pred"], dtype="float32"))
|
||||
else:
|
||||
cond_x0 = _x0(model, dit, x, pe, cond_ind)
|
||||
if do_cfg:
|
||||
uncond_x0 = _x0(model, dit, x, ne, uncond_ind)
|
||||
final_x0 = cond_x0 + scale * (cond_x0 - uncond_x0)
|
||||
else:
|
||||
final_x0 = cond_x0
|
||||
velocity = precision.cast((x - final_x0) / max(float(sigma_t), 1e-6))
|
||||
x_next = model.platform.kernels.get(FLOW_MATCH_STEP)(precision.cast(x), velocity, sigma_t, sigma_next)
|
||||
return StepResult(output={
|
||||
"noise_pred": np.asarray(velocity, dtype="float32"),
|
||||
"latents": x_next.astype("float32")
|
||||
})
|
||||
|
||||
cond_bytes = sum(int(np.asarray(e).nbytes) for e in (pe, ne) if e is not None)
|
||||
res = ResourceRequest(compute_seconds=self.cost.predict(int(np.prod(x.shape)), 2.0 if do_cfg else 1.0),
|
||||
resident_bytes=int(x.nbytes) + cond_bytes,
|
||||
peak_activation_bytes=int(x.nbytes))
|
||||
return WorkPlan(loop_id=self.loop_id,
|
||||
instance_id=st.instance_id,
|
||||
kind=WorkUnitKind.DIFFUSION_STEP,
|
||||
shape_sig=ShapeSignature(WorkUnitKind.DIFFUSION_STEP,
|
||||
dims=tuple(x.shape),
|
||||
dtype=precision.compute_dtype,
|
||||
extra=(("cfg", type(self.cfg).__name__), ("edm", True))),
|
||||
resources=res,
|
||||
payload={
|
||||
"branch": "edm",
|
||||
"step": i
|
||||
},
|
||||
run=run,
|
||||
label=f"cosmos.denoise.{i}",
|
||||
capturable=False) # EDM x0-space CFG + (gated) host-RNG frame injection -> eager path
|
||||
|
||||
def advance(self, st: LoopState, result: StepResult) -> LoopState:
|
||||
st.latents["video"] = result.output["latents"]
|
||||
if st.profile == ExecutionProfile.ROLLOUT:
|
||||
st.trajectory.append({
|
||||
"step": st.step_idx,
|
||||
"sigma": st.sigmas[st.step_idx],
|
||||
"velocity": np.asarray(result.output["noise_pred"]).copy(),
|
||||
"latents": np.asarray(st.latents["video"]).copy()
|
||||
})
|
||||
st.step_idx += 1
|
||||
return st
|
||||
|
||||
def finalize(self, st: LoopState) -> LoopResult:
|
||||
return LoopResult(outputs={"latents": st.latents["video"]},
|
||||
metrics={"denoise_steps": float(st.step_idx)},
|
||||
behavior=st.trajectory or None)
|
||||
@@ -0,0 +1,34 @@
|
||||
"""Cosmos-Predict2 t2v program: text_encode → diffusion_denoise (EDM) → vae_decode.
|
||||
|
||||
Same inline shape as the Wan t2v program — the EDM specifics live in ``CosmosDenoiseLoop`` and the
|
||||
``CosmosDiT`` adapter, so the node graph is unchanged. (video2world conditioning would add a VAE-encode
|
||||
node writing ``conditioning_latents``/``cond_indicator`` into slots; the loop already reads them.)
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from v2.program import ComponentNode, ModelLoopNode, Program, ProgramKind
|
||||
from v2.recipes.common import text_encode_node_fn as _text_encode
|
||||
|
||||
|
||||
def _vae_decode(instance, slots, request, ctx) -> None:
|
||||
slots["video"] = instance.component("vae").decode(slots["denoise_out"]["latents"])
|
||||
|
||||
|
||||
def build_cosmos2_program() -> Program:
|
||||
return Program(
|
||||
program_id="cosmos2.t2v.inline",
|
||||
kind=ProgramKind.INLINE,
|
||||
nodes=[
|
||||
ComponentNode("text_encode", fn=_text_encode, writes=("text_embeds", "neg_text_embeds")),
|
||||
ModelLoopNode("denoise",
|
||||
loop_id="diffusion_denoise",
|
||||
output_slot="denoise_out",
|
||||
reads=("text_embeds", "neg_text_embeds"),
|
||||
writes=("denoise_out", )),
|
||||
ComponentNode("vae_decode", fn=_vae_decode, reads=("denoise_out", ), writes=("video", )),
|
||||
],
|
||||
output_artifacts={
|
||||
"video": "video",
|
||||
"latents": "denoise_out"
|
||||
},
|
||||
).validate()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user