Compare commits
83
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 | ||
|
|
87f98c9b8b | ||
|
|
e60601df7f | ||
|
|
eed9c4bfbf | ||
|
|
1dee77f4a4 | ||
|
|
b80148819c | ||
|
|
c3b971488e | ||
|
|
88e753f281 | ||
|
|
77832059cc |
@@ -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,331 @@
|
||||
# Performance Dashboard Memory
|
||||
|
||||
Date: 2026-06-16
|
||||
Branch: `ci/dashboard`
|
||||
|
||||
## Purpose
|
||||
|
||||
This branch adds a local live dashboard for FastVideo performance benchmark
|
||||
history. It is intended for maintainer/operator use: inspect latest benchmark
|
||||
status, compare current values with recent baseline context, and view trends
|
||||
from the Hugging Face performance-tracking dataset.
|
||||
|
||||
The dashboard is a FastAPI + React app. It is separate from the existing
|
||||
Svelte `ui/` app.
|
||||
|
||||
## Main Files Added Or Changed
|
||||
|
||||
Backend:
|
||||
|
||||
- `fastvideo/performance_dashboard/__init__.py`
|
||||
- `fastvideo/performance_dashboard/__main__.py`
|
||||
- `fastvideo/performance_dashboard/api.py`
|
||||
- `fastvideo/performance_dashboard/metrics.py`
|
||||
- `fastvideo/performance_dashboard/service.py`
|
||||
|
||||
Frontend:
|
||||
|
||||
- `performance_dashboard/frontend/package.json`
|
||||
- `performance_dashboard/frontend/package-lock.json`
|
||||
- `performance_dashboard/frontend/tsconfig.json`
|
||||
- `performance_dashboard/frontend/vite.config.ts`
|
||||
- `performance_dashboard/frontend/index.html`
|
||||
- `performance_dashboard/frontend/scripts/build.mjs`
|
||||
- `performance_dashboard/frontend/src/main.tsx`
|
||||
- `performance_dashboard/frontend/src/api.ts`
|
||||
- `performance_dashboard/frontend/src/App.tsx`
|
||||
- `performance_dashboard/frontend/src/styles.css`
|
||||
|
||||
Docs/tests:
|
||||
|
||||
- `performance_dashboard/README.md`
|
||||
- `docs/contributing/performance_benchmarks.md`
|
||||
- `fastvideo/tests/performance/test_dashboard_service.py`
|
||||
- `fastvideo/tests/performance/test_dashboard_api.py`
|
||||
|
||||
Shared HF utility change:
|
||||
|
||||
- `fastvideo/tests/performance/hf_store.py`
|
||||
|
||||
## Data Source
|
||||
|
||||
The source of truth remains the Hugging Face dataset repo used by existing
|
||||
performance CI:
|
||||
|
||||
```text
|
||||
HF_REPO_ID=FastVideo/performance-tracking
|
||||
```
|
||||
|
||||
The dataset stores normalized JSON records emitted by
|
||||
`fastvideo/tests/performance/compare_baseline.py`. The current v1 normalized
|
||||
schema includes:
|
||||
|
||||
- `model_id`
|
||||
- `timestamp`
|
||||
- `commit_sha`
|
||||
- `gpu_type`
|
||||
- `latency`
|
||||
- `throughput`
|
||||
- `memory`
|
||||
- `text_encoder_time_s`
|
||||
- `dit_time_s`
|
||||
- `vae_decode_time_s`
|
||||
- `success`
|
||||
|
||||
Records are grouped by `(model_id, gpu_type)` for v1 dashboard behavior.
|
||||
|
||||
## Local Cache
|
||||
|
||||
The backend syncs the HF dataset to a local cache directory:
|
||||
|
||||
```text
|
||||
PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard
|
||||
```
|
||||
|
||||
If `PERFORMANCE_TRACKING_ROOT` is not set, the dashboard defaults to:
|
||||
|
||||
```text
|
||||
/tmp/fastvideo-perf-dashboard
|
||||
```
|
||||
|
||||
The sync is performed through the existing helper:
|
||||
|
||||
```python
|
||||
fastvideo.tests.performance.hf_store.sync_from_hf(...)
|
||||
```
|
||||
|
||||
The dashboard then loads JSON files from the local cache through:
|
||||
|
||||
```python
|
||||
fastvideo.tests.performance.hf_store.load_records(...)
|
||||
```
|
||||
|
||||
## Authentication
|
||||
|
||||
Originally `hf_store.py` only read `HF_API_KEY`. This caused local dashboard
|
||||
runs to fail when users had standard Hugging Face token variables set.
|
||||
|
||||
`hf_store.py` now resolves tokens from the first available variable in:
|
||||
|
||||
```text
|
||||
HF_API_KEY
|
||||
HUGGINGFACE_HUB_TOKEN
|
||||
HF_TOKEN
|
||||
```
|
||||
|
||||
For local use:
|
||||
|
||||
```bash
|
||||
export HF_TOKEN=hf_...
|
||||
```
|
||||
|
||||
If the HF repo is private or gated, the token must have dataset read access.
|
||||
|
||||
## Backend API
|
||||
|
||||
The FastAPI app is created by:
|
||||
|
||||
```python
|
||||
fastvideo.performance_dashboard.api:create_app
|
||||
```
|
||||
|
||||
The module-level app is:
|
||||
|
||||
```python
|
||||
fastvideo.performance_dashboard.api:app
|
||||
```
|
||||
|
||||
Endpoints:
|
||||
|
||||
- `GET /api/performance/health`
|
||||
- `POST /api/performance/refresh`
|
||||
- `GET /api/performance/records?days=90`
|
||||
- `GET /api/performance/summary?days=90`
|
||||
- `GET /api/performance/trends?days=90`
|
||||
|
||||
`POST /api/performance/refresh` forces a fresh HF sync.
|
||||
|
||||
## Status Semantics
|
||||
|
||||
Important: the dashboard intentionally separates stored CI status from
|
||||
recomputed context.
|
||||
|
||||
Stored status:
|
||||
|
||||
- Comes directly from the latest JSON record's `success` field.
|
||||
- This is what the dashboard displays as `Stored Status`.
|
||||
- This is the primary latest status.
|
||||
|
||||
Recomputed status:
|
||||
|
||||
- Calculated locally from cached records for explanatory context.
|
||||
- Uses the latest record's metric values compared to the median of the latest
|
||||
five previous successful records in the same `(model_id, gpu_type)` group.
|
||||
- Displayed separately as `Recomputed`.
|
||||
- Does not override the stored JSON `success` status.
|
||||
|
||||
This distinction was added after observing that recomputing pass/fail from the
|
||||
local cache can disagree with the status originally uploaded by CI.
|
||||
|
||||
## Time Window Behavior
|
||||
|
||||
The default dashboard time window is 90 days.
|
||||
|
||||
The selected `days` value affects:
|
||||
|
||||
- trend charts
|
||||
- record browsing/filtering
|
||||
|
||||
The selected `days` value does not affect:
|
||||
|
||||
- latest stored status
|
||||
- latest summary baseline context
|
||||
|
||||
Reason: latest status should not change when users widen or narrow the trend
|
||||
window. The API keeps `days` on `/summary` only for shared frontend filter
|
||||
state, but summary loading uses all cached records.
|
||||
|
||||
This fixed a bug where changing from roughly 35 days to 42 days could change
|
||||
the latest status from pass to fail because older records entered the local
|
||||
baseline window.
|
||||
|
||||
## Metric Logic
|
||||
|
||||
Dashboard metric definitions live in:
|
||||
|
||||
```text
|
||||
fastvideo/performance_dashboard/metrics.py
|
||||
```
|
||||
|
||||
Tracked metrics:
|
||||
|
||||
- `latency` lower is better
|
||||
- `throughput` higher is better
|
||||
- `memory` lower is better
|
||||
- `text_encoder_time_s` lower is better
|
||||
- `dit_time_s` lower is better
|
||||
- `vae_decode_time_s` lower is better
|
||||
|
||||
Baseline context uses the median of up to five previous successful records for
|
||||
the same `(model_id, gpu_type)`.
|
||||
|
||||
## Frontend Behavior
|
||||
|
||||
The React app:
|
||||
|
||||
- fetches `/api/performance/summary`
|
||||
- fetches `/api/performance/trends`
|
||||
- displays summary cards
|
||||
- displays latest rows by model/GPU
|
||||
- displays native SVG trend charts
|
||||
- has model/GPU/day filters
|
||||
- includes a refresh button
|
||||
- auto-refreshes every five minutes
|
||||
|
||||
The UI is implemented without a charting library. Trend charts are native SVG
|
||||
in `performance_dashboard/frontend/src/App.tsx`.
|
||||
|
||||
The production frontend build uses `esbuild` through
|
||||
`performance_dashboard/frontend/scripts/build.mjs`. Vite is still used for the
|
||||
dev server and `/api` proxy.
|
||||
|
||||
Why esbuild for production build:
|
||||
|
||||
- Vite/Rollup hit a local macOS native optional dependency code-signing issue
|
||||
in this environment.
|
||||
- Direct esbuild worked reliably and is sufficient for this small dashboard.
|
||||
|
||||
## Static Serving
|
||||
|
||||
After frontend build, the FastAPI server serves:
|
||||
|
||||
- static JS/CSS from `performance_dashboard/frontend/dist/assets`
|
||||
- `performance_dashboard/frontend/dist/index.html` for the dashboard page
|
||||
|
||||
This allows a single local port to serve both the API and UI.
|
||||
|
||||
## Local Run Workflow
|
||||
|
||||
Build frontend:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
conda run -n fastvideo env PATH=/Applications/Codex.app/Contents/Resources/cua_node/bin:/usr/local/bin:/usr/bin:/bin \
|
||||
/Applications/Codex.app/Contents/Resources/cua_node/bin/npm install
|
||||
conda run -n fastvideo env PATH=/Applications/Codex.app/Contents/Resources/cua_node/bin:/usr/local/bin:/usr/bin:/bin \
|
||||
/Applications/Codex.app/Contents/Resources/cua_node/bin/npm run build
|
||||
```
|
||||
|
||||
Run dashboard:
|
||||
|
||||
```bash
|
||||
export HF_TOKEN=hf_...
|
||||
python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000
|
||||
```
|
||||
|
||||
Open locally:
|
||||
|
||||
```text
|
||||
http://127.0.0.1:8000
|
||||
```
|
||||
|
||||
## ngrok Workflow
|
||||
|
||||
`python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000`
|
||||
starts the actual local dashboard server.
|
||||
|
||||
`ngrok http 8000` does not start the dashboard. It exposes the already-running
|
||||
local server through a temporary public URL.
|
||||
|
||||
Typical flow:
|
||||
|
||||
```bash
|
||||
python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000
|
||||
ngrok http 8000
|
||||
```
|
||||
|
||||
Use the HTTPS URL printed by ngrok to view the dashboard remotely.
|
||||
|
||||
## Verification Commands
|
||||
|
||||
Backend tests:
|
||||
|
||||
```bash
|
||||
conda run -n fastvideo python -m pytest \
|
||||
fastvideo/tests/performance/test_dashboard_service.py \
|
||||
fastvideo/tests/performance/test_dashboard_api.py \
|
||||
-q
|
||||
```
|
||||
|
||||
Expected after latest changes:
|
||||
|
||||
```text
|
||||
8 passed
|
||||
```
|
||||
|
||||
Frontend build:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
conda run -n fastvideo env PATH=/Applications/Codex.app/Contents/Resources/cua_node/bin:/usr/local/bin:/usr/bin:/bin \
|
||||
/Applications/Codex.app/Contents/Resources/cua_node/bin/npm run build
|
||||
```
|
||||
|
||||
Expected:
|
||||
|
||||
```text
|
||||
tsc && node scripts/build.mjs
|
||||
```
|
||||
|
||||
with exit code 0.
|
||||
|
||||
## Known Notes
|
||||
|
||||
- `performance_dashboard/frontend/node_modules/` and
|
||||
`performance_dashboard/frontend/dist/` are ignored by git.
|
||||
- `npm install` reported two high-severity audit findings in dependency tree.
|
||||
`npm audit fix --force` was not run because it can introduce breaking
|
||||
dependency upgrades.
|
||||
- Existing `fastvideo` package imports may emit platform warnings such as NPU
|
||||
or macOS torch distributed messages. These are not dashboard-specific errors.
|
||||
|
||||
+46
-61
@@ -1,71 +1,56 @@
|
||||
# FastVideo Next-Gen Runtime — Two-Page Summary
|
||||
# FastVideo — Design Philosophy
|
||||
|
||||
**Companion to** `design.md` (v19, 2026-06-12) · **Status:** draft for discussion · **Ask:** read this, then dive into the sections you own.
|
||||
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.
|
||||
|
||||
---
|
||||
|
||||
## The problem
|
||||
**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.
|
||||
|
||||
FastVideo's pipeline abstraction has been outgrown by its own model zoo. Four facts, all on `main` today:
|
||||
**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 denoise/sampling loop exists in four copies** — inference stages (`pipelines/stages/denoising.py`, a 1,381-line file), `train/` distillation methods, legacy `training/` monoliths, and the just-landed RL work. The fourth copy documents the cause in its own docstring: `DiffusionSampler` (PR #1450) *"intentionally does not call FastVideo's full inference pipelines"* because the only consumable units are family-bound pipeline classes. Every new post-training method must pick between a wrong dependency and a private loop.
|
||||
- **There is no serving runtime.** One request at a time, no queue, no cross-request batching. Dreamverse — our shipping product — hand-rolls its own GPU pool, queue, warmup, and streaming relay at a cost of **one B200 per user session**.
|
||||
- **Cosmos3 outgrows both the stage abstraction and the alternatives.** Its AR text reasoner and multimodal diffusion denoiser are *the same resident weights* driven by two loop types within one request — our full port (incl. action modality, on `feat/cosmos3-reasoning`) runs it only by bypassing stages with one monolithic block. Multi-engine DAG stacks (vllm-omni, sglang-omni) compose separable stages with disjoint weights; none can express a mixture-of-transformers model. This is the forcing function.
|
||||
- **RL has arrived and pays the tax already** — likelihood-free DiffusionNFT for Wan (#1450): in-process rollouts with zero serving-grade optimizations (no CFG, dense attention, full 25-step ODE), a vendored sampler, and a parallel validation path built *because* inference pipelines aren't consumable as a library.
|
||||
**The 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.
|
||||
|
||||
## The design
|
||||
**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.
|
||||
|
||||
```
|
||||
Request plane OpenAI-compatible server (videos/images/audio/chat) · AsyncEngine
|
||||
OmniRequest ─► admission ─► queue ─► OmniOutput (typed modality parts)
|
||||
Pipeline plane PipelineSpec: declarative graph per family
|
||||
nodes: Stage | LoopStage (DenoiseLoop, ARDecodeLoop) · typed Artifacts
|
||||
policies: CFG, ExpertRouting, AttnMetadata, Precision, FlowShift
|
||||
Execution plane StepScheduler (multiplexes denoise + AR steps) · worker pools (TP/SP + CFG-parallel)
|
||||
CacheManager (paged text-KV ▪ chunked causal-video KV ▪ feature caches) · connectors
|
||||
```
|
||||
**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`**.
|
||||
|
||||
Default deployment is exactly today's: one SPMD pool, co-located nodes, synchronous call. Serving is additive configuration, not a different code path. Five load-bearing decisions:
|
||||
**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**.
|
||||
|
||||
1. **Loop inversion.** Loops become `LoopStage` nodes exposing `init / step / finalize`; the runtime owns iteration, families own step bodies (with a custom-step escape hatch — the runtime never dictates step factoring). This is the enabler for step-level scheduling, streaming, MoT interleaving, and one shared loop across inference, distillation, and RL. Stated honestly: **no surveyed system does this at scheduler granularity** — the risk is retired by Phase-1 bit-identical parity gates and a measured falsifier, not borrowed validation.
|
||||
2. **Cost-model scheduling.** Denoise steps and AR tokens are incommensurable (bidirectional attention is O(L²) per step with zero KV amortization; steps differ ~1000×) — the budget currency is **predicted GPU-time** from a per-(model, phase, shape) cost model, calibrated by the profiler and published to Dynamo's router/Planner as the same artifact.
|
||||
3. **One substrate for inference, training, and RL** — models, loaders, configs, schedulers, parallel state, loop step bodies — under a strict `engine never imports train` rule. Trainers keep their internals; their embedded sampling paths migrate onto the shared loops.
|
||||
4. **Consistency is a declared, measured contract**, not a hope: **C0** corrected (profiles differ, TIS/MIS fixes it) / **C1** kernel-pinned (RL default; drift gated in CI) / **C2** bitwise (batch-invariant kernels + Behavior Record, for goldens and MoE parity). One repo ≠ automatic parity — the ladder is what makes the single-runtime bet honest.
|
||||
5. **Extensions, never monkeypatching**: read-only observers (ParityAligner, ActivationTrace, Profiler, NaNWatch) and compute-altering interceptors at declared points — **cache-dit** is the first interceptor. **Dynamo is the fleet layer** (first-class partner, not a dependency we rebuild): registration/health/cost contract in-engine, seven concrete upstream asks (affinity key spaces, cost interface, media streaming, role-graph disagg, RL weight plane, KVBM generalization, sessions) — each with a fallback.
|
||||
|
||||
## Why this is the moat
|
||||
|
||||
Unlike LLMs — where inference optimization is post-hoc on frozen weights — **a usable video model is itself a post-training artifact**. Every inference capability we ship is a *(recipe, runtime)* pair: step distillation ↔ few-step samplers; self-forcing ↔ causal KV streaming; QAT-NVFP4 ↔ FP4 kernels; VSA ↔ sparse attention backend; RL ↔ samplers + capture. The training loop *embeds* the inference loop, so whoever owns both sides of the pair owns the optimization frontier. The industry's RL pain proves the converse: verl-omni re-implements Wan2.2 inside vLLM-Omni and corrects the numerics afterward; miles' headline features are all mismatch patches for two runtimes with different kernels. We answer with one model definition, one kernel set, one measured ladder — at FastVideo's 1–30B FSDP2 scale, where the bet is viable.
|
||||
|
||||
## What's pulling on it
|
||||
|
||||
| Customer | Pull | Proof point |
|
||||
|---|---|---|
|
||||
| **Cosmos3 / omni** | MoT loops, packed sequences, reasoner KV, world-model rollout | 150-test parity suite on `feat/cosmos3-reasoning` |
|
||||
| **Dreamverse** | Engine-client replaces hand-rolled pool; capacity = duty cycle + cost-model admission + distillation | today 1 B200/session; Phase-3 gate: ≥2 sessions/GPU on a recorded duty-cycle trace, p95 within SLO |
|
||||
| **RL (landed)** | #1450 migrates onto shared loops (Phase 1), engine-client rollouts (Phase 2+); GRPO-class next | the vendored-sampler docstring; C1-by-construction discipline |
|
||||
| **ComfyUI funnel** | embed (nodes) → **compile** (workflow→PipelineSpec, accelerated cloud) → productize (Studio) | tier-1 ~20-node static sublanguage maps onto PipelineSpec |
|
||||
|
||||
## Migration — seven phases, each independently shippable
|
||||
|
||||
| Phase | Ships | Gate |
|
||||
|---|---|---|
|
||||
| **−1** | Merge cosmos3 chain; seed SSIM for uncovered families | baselines exist |
|
||||
| **0** | Typed omni I/O; config freeze (`compat.py` shrinks monotonically to zero) | all SSIM suites unchanged |
|
||||
| **1** | **Loop inversion** + policies + extension core (cache-dit, ParityAligner); RL migrates off its vendored sampler | old vs new loop **bit-identical**; per-method grad-norm refs (#1396) extended |
|
||||
| **2** | AsyncEngine + StepScheduler; LTX-2 linear graph; Dynamo worker (stock); colocated weight sync | ≤2% batch-1 latency regression; Dreamverse single-session parity; RL engine-client parity |
|
||||
| **3** | PipelineSpec graphs, role pools, declarative parallelism, ComfyUI compiler MVP, general WeightSyncPlan | ≥2 Dreamverse sessions/GPU on recorded duty-cycle trace |
|
||||
| **4** | Cosmos3 native; AR continuous batching + paged KV (arriving *with* their workload, per N5); RL hardening (C1/C2, Behavior Record) | Cosmos3 parity suite on new runtime; drift ≈ 0 on a Wan RL run |
|
||||
| **5** | Deletion: legacy `training/` retires, then `ComposedPipelineBase`, legacy loop, `forward_context.py`, `compat.py`, `RayDistributedExecutor` | the deletion diff — **4 loop copies → 1** |
|
||||
|
||||
Enforcement the last freeze lacked (it was broken 19×): `compat.py` frozen from Phase 0; CI path gates + CODEOWNERS once Phase 1 lands; new families land on new abstractions from Phase-1 completion.
|
||||
|
||||
## What we are deliberately not doing
|
||||
|
||||
Datacenter orchestration (Dynamo's job) · trainer internals (frozen `training/`; `train/` is a consumer) · replacing the bit-exact porting methodology · migrating 20+ families at once · **standalone LLM-serving excellence** — AR machinery arrives only at the sophistication omni workloads pull (N5).
|
||||
|
||||
## Decisions we need from this review
|
||||
|
||||
1. **sglang `multimodal_gen` relationship** — upstream, friendly fork, or shared core (decide by Phase 2; drift is a strategic cost either way).
|
||||
2. **Dynamo asks** — green-light proposing the Phase-2 asks (A2 cost interface, A3 media streaming) to the team first, with A5 (RL weight plane) queued behind them?
|
||||
3. **Phase −1 start** — merge the cosmos3 chain and seed SSIM baselines now; it blocks everything else.
|
||||
> 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.
|
||||
|
||||
-830
@@ -1,830 +0,0 @@
|
||||
# FastVideo v3 — A Model-Native Runtime for the (Recipe, Runtime) Era
|
||||
|
||||
**Status:** unconstrained north-star. This document assumes we are free to build a brand-new architecture with no
|
||||
backward-compatibility, no migration tax, and no obligation to the current code. It exists to define the *ceiling*:
|
||||
the system FastVideo should be if nothing held it back. Migration is a separate, later question — deliberately out of
|
||||
scope here.
|
||||
|
||||
**Lineage:** this is the synthesis of `design.md` (the strategic thesis, product pull, and hard-won serving realism)
|
||||
and `designv2.md` (the model-native center and typed contracts), with the open tensions of both resolved rather than
|
||||
hedged.
|
||||
|
||||
---
|
||||
|
||||
## Table of contents
|
||||
|
||||
1. [The thesis](#1-the-thesis)
|
||||
2. [The two signature ideas](#2-the-two-signature-ideas)
|
||||
3. [Planes and their dependency order](#3-planes-and-their-dependency-order)
|
||||
4. [Model Plane — the center](#4-model-plane--the-center)
|
||||
5. [The loop contract — driven loops](#5-the-loop-contract--driven-loops)
|
||||
6. [Runtime and scheduler — one currency, one WorkUnit](#6-runtime-and-scheduler)
|
||||
7. [Memory, cache, transport, compile](#7-memory-cache-transport-compile)
|
||||
8. [Parallelism as a model contract](#8-parallelism)
|
||||
9. [Correctness — parity as a typed gate](#9-correctness)
|
||||
10. [Training and RL on the same loops](#10-training-and-rl)
|
||||
11. [Extensions — observers and interceptors](#11-extensions)
|
||||
12. [Request, session, artifact, stream](#12-request-session-artifact-stream)
|
||||
13. [Programs and workflows](#13-programs-and-workflows)
|
||||
14. [Deployment and fleet](#14-deployment-and-fleet)
|
||||
15. [Worked examples](#15-worked-examples)
|
||||
16. [What this unlocks](#16-what-this-unlocks)
|
||||
17. [Honest unknowns and falsifiers](#17-honest-unknowns-and-falsifiers)
|
||||
18. [Package layout](#18-package-layout)
|
||||
19. [Reference synthesis](#19-reference-synthesis)
|
||||
|
||||
---
|
||||
|
||||
## 1. The thesis
|
||||
|
||||
Three facts about video generation, taken together, dictate the architecture.
|
||||
|
||||
**A deployable video model is a post-training artifact.** Unlike an LLM — where inference optimizes frozen weights
|
||||
post-hoc — a *usable* video model is *created* by training: step distillation is mandatory for usable latency, low
|
||||
precision needs QAT, and causal/world models are made by distillation plus self-forcing. Every inference capability is
|
||||
therefore a **(recipe, runtime) pair**: the weights and the loop that produced-and-assumes them are inseparable. A
|
||||
"4-step NVFP4 FastWan" is not weights plus a flag; it is a distillation recipe, a sampler, a precision path, and a
|
||||
parity contract that are one object.
|
||||
|
||||
**Video systems are loop systems.** Denoise timesteps, AR decode, chunked world-model rollout, VAE tiles, encoder
|
||||
chunks, audio tokens, reward batches, optimizer steps, media chunks — the work is iteration, not a single `forward`. A
|
||||
runtime that reduces everything to `forward()` cannot schedule, batch, cancel, stream, reserve memory for, or capture
|
||||
the behavior of the thing that actually runs.
|
||||
|
||||
**Omni models share weights across loop types within one request.** Cosmos3's text reasoner and multimodal denoiser
|
||||
are the *same resident weights*, driven by an AR loop and then a diffusion loop in one request. This cannot be a
|
||||
DAG of separate engines (that doubles 30B+ of weights and severs the shared KV/denoise state); it must be one resident
|
||||
model instance running many loop types. This is achievable — vllm-omni's `bagel_single_stage`/`lance` already run one
|
||||
resident MoT instance doing AR `generate_text` and diffusion `generate_image` on co-resident experts in a single
|
||||
request — *but* they bury that interleaving inside one opaque `DIFFUSION` stage their scheduler never sees inside,
|
||||
request-scheduled with `max_num_running_reqs` forced to 1. The hard part, and the differentiation, is not *expressing*
|
||||
the shared-weight loops — it is making them **runtime-visible, step-scheduled, batchable, and cost-priced.**
|
||||
|
||||
The architecture that falls out:
|
||||
|
||||
> **The atomic unit is the (recipe, runtime) pair, owned by a typed `ModelCard`.** Everything — serving, training, RL,
|
||||
> products, deployment — is a *view* over that card. The **runtime owns loop *lifecycle*** (admission, scheduling,
|
||||
> batching, caching, cancellation, streaming, behavior capture); the **model owns loop *semantics*** (typed state
|
||||
> transitions and kernel execution). One resident model instance runs many loops; one scheduler schedules the
|
||||
> *steps* of all of them in a single currency; one parity contract binds the train-forward to the serve-forward so
|
||||
> the recipe and the runtime never silently drift apart.
|
||||
|
||||
The single invariant, stated once:
|
||||
|
||||
```text
|
||||
Model cards own components, loops, recipes, and parity.
|
||||
Programs compose loops into tasks.
|
||||
The scheduler executes the steps of loops as WorkUnits under one budget.
|
||||
Caches are correct by key, not by hope.
|
||||
Training records behavior on the same loops it serves.
|
||||
Deployment places and routes; it does not define semantics.
|
||||
Products stream artifacts; they do not reach into the model.
|
||||
```
|
||||
|
||||
Everything below is the elaboration of that invariant.
|
||||
|
||||
---
|
||||
|
||||
## 2. The two signature ideas
|
||||
|
||||
Two ideas do most of the work and are what an unconstrained design can reach that an incremental one cannot.
|
||||
|
||||
### 2.1 The (recipe, runtime) pair is a first-class, versioned, typed object
|
||||
|
||||
A model in v3 is not a checkpoint. It is a `ModelCard` that owns, as one versioned unit:
|
||||
|
||||
- the **components** (weights, loaders, layouts),
|
||||
- the **loops** it can run (the runtime semantics),
|
||||
- the **recipe** that produced the weights (distillation/QAT/RL config, teacher, data contract, the sampler the
|
||||
recipe assumes), and
|
||||
- the **parity contract** asserting that the train-forward and the serve-forward agree to a declared level.
|
||||
|
||||
You cannot ship the weights without the loop they assume, and you cannot change the loop without re-proving parity.
|
||||
This makes design.md's "(recipe, runtime) pair" *literal*: the deployable artifact carries its own provenance and its
|
||||
own correctness obligation. It is the thing that turns "we do training and inference in one repo" from an org chart
|
||||
into a guarantee — and it is essentially un-retrofittable, which is exactly why it belongs in a clean design.
|
||||
|
||||
### 2.2 Driven loops: the model owns control flow, the runtime owns execution
|
||||
|
||||
Loop inversion, done right, is not a heavy `plan_step`/`run_step` contract and not a hidden `for t in timesteps`. It
|
||||
is a **driven loop**: the model describes the *next step it needs*, the runtime *decides when and with whom that step
|
||||
runs*, and the model folds the result back into its own state and decides what to do next. The model keeps its control
|
||||
flow (so content-adaptive decisions — cache-dit skips, EOS, VSA tile selection — are natural); the runtime keeps the
|
||||
`await` (so admission, batching, cancellation, streaming, and behavior capture are universal). Per-request state lives
|
||||
in the loop's own typed `LoopState`, never in module globals, so interleaving requests through one model instance
|
||||
cannot smear state — the failure mode that makes naive loop-inversion dangerous is *structurally* excluded.
|
||||
|
||||
These two ideas are developed in §4–§5. The rest of the system is their consequence.
|
||||
|
||||
---
|
||||
|
||||
## 3. Planes and their dependency order
|
||||
|
||||
```text
|
||||
Products: Python · CLI · OpenAI API · ComfyUI · Dreamverse · RTC · Trainer
|
||||
│ (thin: validate intent, make requests/sessions, subscribe)
|
||||
Request / Session / Artifact / Stream
|
||||
│ (typed runtime objects, cancellation, streaming)
|
||||
Program Plane ← typed loop programs + compiled workflows
|
||||
│
|
||||
┌──────────────── Model Plane (CENTER) ────────────────┐
|
||||
│ ModelCard: components · loops · recipe · parity │
|
||||
│ capabilities · caches · parallelism · precision │
|
||||
└───────────────────────┬──────────────────────────────┘
|
||||
│
|
||||
┌─────────────────────────────┼─────────────────────────────┐
|
||||
│ Runtime / Scheduler │ Training / RL │ (same loops, different capture)
|
||||
│ WorkUnits · GPU-time budget │ rollout · reward · weight-sync│
|
||||
└─────────────────────────────┼─────────────────────────────┘
|
||||
│
|
||||
Memory · Cache · Transport · Compile (typed CacheKey, per-class pools, CuMem sleep/wake)
|
||||
│
|
||||
Parallelism (named axes → DeviceMesh, validated, part of the cache key)
|
||||
│
|
||||
Deployment / Fleet (DeploymentCard → Dynamo; never the core)
|
||||
```
|
||||
|
||||
Dependency rules (enforced at the package boundary, §18):
|
||||
|
||||
- Products do not define model semantics. Workflows do not define model semantics. Deployment does not define model
|
||||
semantics. Training does not redefine model semantics. **All of them reference the Model Plane.**
|
||||
- The runtime *executes* model loops but does not *own* their math. Training *captures* behavior on serving loops but
|
||||
does not *fork* them. Cross-cutting concerns — extensions (§11), parallelism (§8), and parity (§9) — are contracts
|
||||
declared on the card, not features bolted onto the runtime.
|
||||
|
||||
---
|
||||
|
||||
## 4. Model Plane — the center
|
||||
|
||||
### 4.1 ModelCard
|
||||
|
||||
```python
|
||||
class ModelCard:
|
||||
model_id: str # "fastwan-1.3b-nvfp4-4step"
|
||||
family: str # "wan"
|
||||
components: dict[str, ComponentSpec]
|
||||
loops: dict[str, LoopSpec]
|
||||
capabilities: CapabilityMatrix # text_to_video, image_to_video, reasoning_text, vae_decode, ...
|
||||
recipe: RecipeSpec # ← what produced these weights (signature idea §2.1)
|
||||
parity: ParitySpec # ← train-forward ≡ serve-forward, to a declared level (§9)
|
||||
caches: dict[str, CacheContract]
|
||||
parallelism: ParallelismContract
|
||||
precision: PrecisionContract
|
||||
checkpoint: CheckpointManifest # explicit components, layouts, key maps — no name-detector guessing
|
||||
```
|
||||
|
||||
The card is both a **declarative contract** (strict enough to validate before any GPU touches it) and a **runtime
|
||||
factory** (it knows how to instantiate components, bind loops, and resolve caches). It is hub-interchange compatible
|
||||
with diffusers' `modular_model_index.json` / `ComponentSpec` so models published either way load both ways.
|
||||
|
||||
`CheckpointManifest` replaces today's implicit `model_index.json` + name-detector resolution with explicit declared
|
||||
components and `required_for` / `optional_for` task sets (the Cosmos3 lazy-sound-VAE problem becomes a declaration, not
|
||||
an `if env_var` inside `forward`).
|
||||
|
||||
### 4.2 RecipeSpec — the provenance half of the pair
|
||||
|
||||
```python
|
||||
class RecipeSpec:
|
||||
method: str # "dmd2" | "self_forcing" | "attn_qat_nvfp4" | "diffusion_nft" | "base"
|
||||
parents: list[str] # teacher / base model_ids this was distilled or RL'd from
|
||||
data_contract: DataRef # what the recipe trained on (for governance and reproduction)
|
||||
assumes_loop: str # the loop_id this recipe's weights require at serve time
|
||||
assumes_precision: str # the precision the QAT recipe baked in
|
||||
consistency_required: str # the minimum parity level this recipe's outputs are valid under (§9)
|
||||
```
|
||||
|
||||
`assumes_loop` and `assumes_precision` are the teeth: a 4-step distilled model whose `assumes_loop = "ddim_4step"`
|
||||
cannot be served under a 50-step sampler without a typed mismatch error. The recipe and the runtime are bound.
|
||||
|
||||
### 4.3 ComponentSpec and LoopSpec
|
||||
|
||||
```python
|
||||
class ComponentSpec:
|
||||
component_id: str
|
||||
kind: str # dit | vae | text_encoder | reasoner_tower | reward_head | ...
|
||||
load_id: str
|
||||
config_schema: type
|
||||
io_schema: tuple[type, type]
|
||||
precision_policy: PrecisionPolicy
|
||||
placement_policy: PlacementPolicy
|
||||
parallel_constraints: ParallelConstraint
|
||||
parity_tests: list[ParityTestSpec]
|
||||
|
||||
class LoopSpec:
|
||||
loop_id: str # diffusion_denoise | ar_decode | chunk_rollout | vae_tile | ...
|
||||
state_schema: type # the typed LoopState (no dicts)
|
||||
step_schema: type # the typed WorkPlan a step emits
|
||||
result_schema: type # the typed StepResult a step returns
|
||||
behavior_schema: type | None # what to capture for RL (None if not training-relevant)
|
||||
step_cost_model: CostModel # predicted GPU-time per step at (shape, precision, policy) — §6
|
||||
valid_parallel_plans: list[ParallelPlanPattern]
|
||||
graph_capture: GraphCapturePolicy
|
||||
cache_policy: CachePolicy
|
||||
```
|
||||
|
||||
A `ModelInstance` is a resident, loaded card: component instances, model state, caches, compiled graphs, and a
|
||||
parallel plan. **A request may run several of the card's loops against one `ModelInstance`.** That single sentence is
|
||||
the difference between this design and a stage-only design, and it is what makes omni native:
|
||||
|
||||
```text
|
||||
one Cosmos3 ModelInstance, one request:
|
||||
ar_decode(reasoner) → pack → diffusion_denoise(vision[+action][+sound]) → vae_tile_decode → audio_decode
|
||||
└────────────── same resident weights, shared packed state, scheduled as steps ──────────────┘
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 5. The loop contract — driven loops
|
||||
|
||||
### 5.1 The contract
|
||||
|
||||
A loop is a **serializable state machine** the runtime drives:
|
||||
|
||||
```python
|
||||
class Loop(Protocol):
|
||||
def init(self, req: Request, model: ModelState, ctx: LoopContext) -> LoopState: ...
|
||||
def next(self, state: LoopState) -> WorkPlan | Done: ... # describe the next step; NO GPU kernels here
|
||||
def advance(self, state: LoopState, result: StepResult) -> LoopState: ... # fold result in; decide what's next
|
||||
def finalize(self, state: LoopState) -> LoopResult: ...
|
||||
```
|
||||
|
||||
The runtime's driver — the only place iteration lives:
|
||||
|
||||
```python
|
||||
state = loop.init(req, model_state, ctx)
|
||||
while True:
|
||||
plan = loop.next(state) # typed WorkPlan: resources, cache reads/writes, shape, sinks, cancel-scope
|
||||
if isinstance(plan, Done):
|
||||
break
|
||||
result = await ctx.execute(plan) # ← THE INVERSION POINT: runtime admits, batches, places, runs, returns
|
||||
state = loop.advance(state, result) # content-adaptive: next() can branch on everything in state, incl. result
|
||||
for chunk in plan.emits:
|
||||
ctx.emit(chunk) # streaming falls out
|
||||
return loop.finalize(state)
|
||||
```
|
||||
|
||||
Why this is the right contract, point by point against the failure modes:
|
||||
|
||||
- **Content-adaptive steps are natural.** `next()` reads `state`, and `advance()` has already folded in the last
|
||||
`StepResult` — so cache-dit's skip decision (a residual comparison from the prior step), AR's EOS, and VSA's
|
||||
content-dependent tile selection are ordinary control flow in the model. This is the tension `designv2.md`'s
|
||||
"pure plan_step that pre-declares shape" could not resolve; here it dissolves, because planning the *next* step is
|
||||
allowed to depend on the *previous* result. `next()` is still kernel-free (it *describes* work; it does not run it),
|
||||
which is all the scheduler needs.
|
||||
- **Cross-request state safety is structural.** All per-request state is in the loop's typed `LoopState`. There are no
|
||||
module-level residual/KV globals (the bug that silently corrupts cache-dit and TeaCache forks under concurrency).
|
||||
Interleaving requests through one `ModelInstance` cannot smear state because there is no shared mutable state to
|
||||
smear. This is the safety property naive loop inversion lacks, made impossible-to-get-wrong by construction.
|
||||
- **The runtime owns everything it needs and nothing it doesn't.** `execute(plan)` is the single seam for admission,
|
||||
memory reservation, cross-request batching, placement, graph dispatch, cancellation, and behavior capture. The model
|
||||
never sees the scheduler; the scheduler never sees the model's math.
|
||||
- **Serializable, therefore migratable and resumable.** `LoopState` is typed and serializable, so a half-finished
|
||||
1000-step job is a resume point: preempt by stopping the driver, migrate by shipping `LoopState` to another worker,
|
||||
recover from a crash by replaying from the last serialized state. (A coroutine that keeps state in a suspended Python
|
||||
frame — the tempting sugar — cannot do this; the explicit state machine is the price of resumability, and it is
|
||||
worth paying.)
|
||||
|
||||
Custom step bodies are first-class, not an escape hatch with an asterisk: a family whose math is genuinely braided
|
||||
(Cosmos's EDM coefficients consumed inside the CFG branch with an x0-space combine; LTX-2's 1–4 runtime-decided
|
||||
guidance passes) writes `next()`/`advance()` by hand using samplers and CFG utilities as a *library*. The runtime
|
||||
requires only the four methods; *how* a step body is factored is the model's business. Policies (CFG, expert routing,
|
||||
precision, flow-shift, conditioning) are the *default* decomposition that deletes duplication for the families that
|
||||
fit — never an admission requirement.
|
||||
|
||||
### 5.2 Loop granularity
|
||||
|
||||
Chosen by runtime value, not purity:
|
||||
|
||||
- too coarse → cannot cancel/batch/reserve/stream/record at useful points;
|
||||
- too fine → scheduler overhead dominates, graph capture fragments;
|
||||
- good default → one denoise step (or window), one AR decode batch, one encoder chunk, one VAE-tile batch, one
|
||||
reward/logprob batch.
|
||||
|
||||
The runtime may **fuse** adjacent compatible WorkPlans after planning (an optimization); the unfused boundary remains
|
||||
the semantic model, so parity and behavior capture are defined on the unfused loop.
|
||||
|
||||
### 5.3 CFG is a policy over *one* shared denoise body (verified)
|
||||
|
||||
A natural worry: CFG changes the *shape* of the step (one forward vs two vs a batched pair vs a data-parallel split),
|
||||
so can one shared denoise loop really host all of it by swapping a policy? **Yes — and it is proven by existing code,
|
||||
not aspiration.** vllm-omni's `CFGParallelMixin.predict_noise_maybe_with_cfg` + `combine_cfg_noise`
|
||||
(`diffusion/.../cfg_parallel.py:76-212`) already runs sequential-2-forward, batched-1-forward, *and* cfg-parallel
|
||||
through **one** pair, with the loop body unaware of which. The clean cut is three layers:
|
||||
|
||||
- **In-loop `CFGPolicy`** — branch vocabulary (`[cond]`, `[cond, uncond]`, per-modality, STG-perturbed),
|
||||
the combine formula (standard `uncond + s·(cond−uncond)`, CFG-zero `st_star`, `cfg_normalize`/`guidance_rescale`),
|
||||
and **per-request mutable state** (the adaptive-gate cached delta with model-id self-invalidation is the canonical
|
||||
state case — and exactly why state lives in `LoopState`, §5.1). **Batched-vs-two-forward is a *dispatch detail
|
||||
inside one policy*, not a separate mechanism.** This covers classic / batched / adaptive-gate / per-modality.
|
||||
- **`cfg`-parallel is a *parallelism axis*, not a policy** — it shards the policy's branches across ranks and runs the
|
||||
*same rank-invariant `combine`* on every rank. It composes *under* any `CFGPolicy`; you own a `BatchedCFG` policy
|
||||
**or** a `cfg` group, never both (the §9 build-guard).
|
||||
- **Companions are an *orchestrator pattern*, not in the loop** — splitting a request into companion sub-requests
|
||||
upstream of diffusion (the conditioning is precomputed and bundled in; the loop is unchanged).
|
||||
|
||||
Two caveats keep the first pass honest: the `combine` runs in *the step body's* numeric space (Cosmos combines in
|
||||
x0-space after EDM preconditioning, not noise-space — the body fixes the space, the policy fixes the algebra), and
|
||||
embedded-guidance (Flux) is a **degenerate single-branch identity-combine policy** (guidance rides inside the forward
|
||||
kwarg), kept *inside* the same abstraction rather than special-cased as "no CFG." This is the same shared denoise body
|
||||
that RL rollout reuses (§10) — one CFG taxonomy serves both serving and rollout.
|
||||
|
||||
---
|
||||
|
||||
## 6. Runtime and scheduler
|
||||
|
||||
### 6.1 One WorkUnit, one currency
|
||||
|
||||
Every `await ctx.execute(plan)` produces a **WorkUnit**: the smallest schedulable action with a resource reservation
|
||||
and a loop boundary. Kinds: `ar_prefill`, `ar_token`, `diffusion_step`, `diffusion_window`, `chunk_step`,
|
||||
`encoder_chunk`, `vae_tile`, `audio_chunk`, `reward_batch`, `logprob_batch`, `transfer`, `cache_io`, `graph_capture`.
|
||||
Tokens are *one kind*, not the scheduler — this is the generalization of vLLM's token scheduler that diffusion forces.
|
||||
|
||||
**The budget currency is predicted GPU-time, not counts.** A bidirectional denoise step re-attends the full latent at
|
||||
O(L²) with zero KV amortization (every step pays full price); an AR decode step is ~O(context) against a cache; a
|
||||
chunked-causal step sits between. Counting "steps" or "tokens" puts items three orders of magnitude apart in one
|
||||
bucket. So each WorkUnit converts to GPU-seconds via its `LoopSpec.step_cost_model`, calibrated online by the Profiler
|
||||
(§11). **The same cost model is the interface published to the fleet** (§14): the scheduler's internal budget and
|
||||
Dynamo's routing/autoscaling input are one object, built once.
|
||||
|
||||
Two honesty caveats, kept from design.md's contact with reality:
|
||||
|
||||
- **Admission uses the conservative baseline.** The design's own flagship features make realized cost unknowable in
|
||||
advance — cache-dit skips are residual comparisons, VSA tiles are content-dependent, AR length is unbounded
|
||||
(budgeted at the `max_tokens` cap, refunded on early EOS). Telemetry refines calibration; it never licenses
|
||||
admission optimism.
|
||||
- **A denoise step is indivisible.** A 30s-1080p step bounds iteration latency no matter the budget. Mitigations are
|
||||
first-class, not afterthoughts: cost-class pools (jumbo steps don't co-schedule with latency-class work), SP within
|
||||
a node to shrink jumbo wall-time, and admission-time SLO classes so the fleet planner scales pools per class. This
|
||||
is why the scheduler is *cost-aware*, not just *count-aware*.
|
||||
|
||||
### 6.2 WorkPlan and admission
|
||||
|
||||
```python
|
||||
class WorkPlan:
|
||||
loop_id: str
|
||||
instance_id: str
|
||||
kind: str
|
||||
shape_sig: ShapeSignature # for batch compatibility + graph capture key
|
||||
resources: ResourceRequest # compute (GPU-s), resident bytes, peak-activation bytes, cache blocks, xfer bw, sinks
|
||||
cache: CachePlan # typed reads/writes (§7)
|
||||
placement: PlacementHint
|
||||
cancel_scope: CancelScope
|
||||
emits: list[StreamChunk]
|
||||
class Done: result: LoopResult
|
||||
```
|
||||
|
||||
**Admission rule (the soundness condition of multiplexing):** *do not admit a waiting WorkUnit unless every resource
|
||||
it requests can be reserved* — compute budget **and** memory (resident + worst-case peak) **and** cache blocks **and**
|
||||
transfer bandwidth **and** graph-capture shape **and** output sinks. Two requests that fit individually but jointly OOM
|
||||
are rejected at admission, not discovered at step 37. This is vLLM's "token budget is half the story, `allocate_slots`
|
||||
is the other half," generalized.
|
||||
|
||||
### 6.3 The scheduler, in layers (each testable on a fake pool, no GPU)
|
||||
|
||||
1. **RequestScheduler** — accepts requests/sessions, selects programs, starts loop drivers.
|
||||
2. **LoopScheduler** — drives `next()`, collects pending WorkPlans.
|
||||
3. **BatchScheduler** — groups compatible WorkPlans by `(instance, loop_kind, shape_sig, precision, parallel_plan,
|
||||
graph_key)`; image diffusion and AR decode batch across requests, jumbo video stays batch-of-1, Cosmos-style
|
||||
token-budget packing is an opt-in.
|
||||
4. **PlacementScheduler** — worker, role pool, instance, device mesh.
|
||||
5. **TransferScheduler** — tensor / cache / artifact movement as scheduled WorkUnits.
|
||||
6. **AdmissionController** — the reservation gate of §6.2.
|
||||
|
||||
Policies: running loops first (vLLM); preempt only at loop step boundaries; cancel only at declared scopes; prefer
|
||||
cache hits when latency/fairness allow; never starve a long denoise loop behind short AR requests.
|
||||
|
||||
### 6.4 SPMD consistency and failure isolation
|
||||
|
||||
All ranks of a pool must make identical scheduling decisions or NCCL deadlocks. **Rank-0 decides, broadcasts** — the
|
||||
existing discipline, now also the channel for the **abort broadcast** (failure isolation and scheduling share one
|
||||
consistency mechanism). Failure classes: *request-fatal* (NaN flagged by NaNWatch, one request's step error) →
|
||||
SPMD-consistent abort of that request, deliver partial artifacts with a structured error; *pool-fatal* (illegal
|
||||
access, NCCL desync) → pool re-init, invalidate pool caches, resume requests from serialized `LoopState` where one
|
||||
exists. **Cancellation is common-path, not exceptional** — vibe-directing makes abandoning in-flight work the *normal*
|
||||
user action; cancel takes effect at the next step boundary, drops queued WorkUnits, releases `LoopState` and cache
|
||||
handles, reports `cancelled`.
|
||||
|
||||
---
|
||||
|
||||
## 7. Memory, cache, transport, compile
|
||||
|
||||
Video and omni inference are memory systems as much as compute systems. This plane is explicit and typed.
|
||||
|
||||
### 7.1 Cache correctness is a contract
|
||||
|
||||
```python
|
||||
class CacheKey:
|
||||
model_id: str; component_id: str; loop_id: str | None
|
||||
weights_version: str; adapter_versions: dict[str, str]
|
||||
precision: str; parallel_plan_hash: str
|
||||
shape_sig: str; layout_sig: str
|
||||
scheduler_sig: str | None; guidance_sig: str | None; seed: int | None
|
||||
input_hashes: dict[str, str]; step_index: int | None
|
||||
contract_version: str
|
||||
```
|
||||
|
||||
**If a field can change output semantics, it is in the key.** Incorrect reuse is worse than no reuse. The serving
|
||||
hazard this kills: a workflow-cloud request that shares a prompt but differs in te-LoRA stack must not serve stale
|
||||
embeddings — so the key is *partitioned* by `adapter_versions`, not flushed. An RL `update_weights` bumps
|
||||
`weights_version` and invalidates wholesale.
|
||||
|
||||
### 7.2 Per-class pools (the granularity reality)
|
||||
|
||||
There is **no single unified block pool**, because a unified pool requires uniform bytes-per-block and our cache
|
||||
classes differ by 150–500× in natural granularity (a text-KV page ≈ 64 KB/layer; a causal-video latent-chunk slab is
|
||||
9.6–32 MB/layer) and their demand is workload-decoupled. Each class gets a statically budgeted pool behind one
|
||||
`CacheHandle`: paged text-KV (`ar_decode`), slab chunk-KV (`chunk_rollout`, with a declared training mode that
|
||||
disables mid-rollout recycling and keeps grad-aware index snapshots), feature caches (text/vision-encoder, content-hash
|
||||
keyed, reference-counted FIFO), residual caches (cache-dit, scoped per `LoopState`), weight/adapter cache
|
||||
(disk→CPU→GPU LRU for the workflow cloud). MoT falls out: the und pathway draws paged KV, the gen pathway draws slab or
|
||||
nothing — independent budgets, no interference. **KV is the minority case** — a pure bidirectional deployment allocates
|
||||
none of it; the machinery materializes only when a card declares KV-bearing loops.
|
||||
|
||||
### 7.3 Memory, transport, compile
|
||||
|
||||
- **Memory** — tagged pools, sleep/wake by tag (CuMem-style; tags are component names), reservation before admission,
|
||||
per-role budgets, host-pinned staging. Sleep/wake is component-granular for RL (drop DiT + caches, keep
|
||||
VAE/text-encoder resident).
|
||||
- **Transport** — manifest-based and pluggable: in-proc reference → SHM → CUDA IPC → NCCL/UCXX/NIXL/RDMA →
|
||||
object-store. KV/cache-bearing edges speak a `KVConnector`-shaped protocol (scheduler-side query/alloc/finish +
|
||||
worker-side async load/save) so NIXL/LMCache/Mooncake/KVBM implement it directly. Transfers are scheduled WorkUnits,
|
||||
not side effects.
|
||||
- **Compile** — CUDA graphs and `torch.compile` managed by a `CompileCache` keyed on `(model, component, loop,
|
||||
work_kind, shape_sig, precision, parallel_plan, backend)`. **Never full-graph across the engine** (per vLLM's own
|
||||
reversal): per-block compile where it pays, manual fused ops permitted in model code, breakable CUDA graphs as an
|
||||
*optimization tier* over an always-correct eager baseline. Graph capture is planned by the scheduler (padding,
|
||||
bucketing, capture sizes affect admission and batching).
|
||||
|
||||
---
|
||||
|
||||
## 8. Parallelism as a model contract
|
||||
|
||||
Parallelism is not a launch flag; it affects cache keys, scheduling, transport, capture, and parity, so it lives on
|
||||
the card.
|
||||
|
||||
```python
|
||||
class ParallelPlan:
|
||||
axes: dict[str, int] # dp, tp, sp(=ulysses×ring), cp, cfgp(≤2), pp_patch, vae, ep, fsdp, role, replica
|
||||
mesh_order: list[str]
|
||||
placement: PlacementSpec
|
||||
communication: CommunicationSpec
|
||||
```
|
||||
|
||||
Declarative, validated, compiled to a PyTorch `DeviceMesh` via a `ParallelDims`-style builder
|
||||
(product-of-degrees validation, cached submeshes). **Pre-flight or it fails at load, never halfway.** Ownership
|
||||
conflicts are build errors (CFG owned by a `BatchedCFG` *policy* or a `cfgp` *group*, never both). Applicability
|
||||
conditions travel with axes: `pp_patch` (PipeFusion displaced-patch pipelining) is **invalid for causal/AR** (stale KV
|
||||
breaks causality) and the validator enforces it per card. Degree-one axes exist as trivial groups so component code
|
||||
needs no special cases. **Pools are single-node**; multi-node scale is *multiple pools* fronted by the fleet (§14) —
|
||||
the engine never owns a cross-node NCCL mesh inside one pool.
|
||||
|
||||
---
|
||||
|
||||
## 9. Correctness — parity as a typed gate
|
||||
|
||||
This is the section both prior documents needed and neither fully had.
|
||||
|
||||
### 9.1 The parity contract
|
||||
|
||||
Every card carries a `ParitySpec`. Parity is **measured, never assumed**, by a `ParityAligner` observer (§11): record
|
||||
named taps per step/block from a reference (the official framework, or a pre-change build); compare-mode replays with
|
||||
fixed seeds and reports the first divergence beyond per-tap tolerance. This is the engine behind the "old loop vs new
|
||||
loop, bit-identical" gate and the standing instrument for every port, precision change, and kernel swap.
|
||||
|
||||
### 9.2 The consistency ladder (with the rung both prior docs missed)
|
||||
|
||||
```text
|
||||
C0 component parity — VAE, encoder, transformer block, scheduler step in isolation
|
||||
C1 loop parity — full denoise trajectory / AR logits, fixed seed
|
||||
C2 behavioral identity — the train-forward and serve-forward agree on the quantity the RL objective uses:
|
||||
· likelihood-based methods (GRPO-class): per-step log-prob identity
|
||||
· likelihood-free methods (DiffusionNFT-class): seeded final-sample +
|
||||
prediction-space identity (old_deviate / ref-MSE) — there are NO log-probs to match
|
||||
C3 distribution parity — rollout distribution under allowed nondeterminism
|
||||
C4 artifact quality — SSIM-class, reward agreement, human-preference (gates product claims; needs the eval system)
|
||||
```
|
||||
|
||||
The C2 split is load-bearing and is the lesson of the landed RL stack: the shipped Wan DiffusionNFT is
|
||||
**likelihood-free** — it captures only final clean latents and contrasts the student against an implicit negative
|
||||
policy in prediction space, so "log-prob identity" is *undefined* for it. A ladder that assumes log-probs (as both
|
||||
`design.md`'s and `designv2.md`'s early framings did) cannot describe the only RL method actually in the tree. RL
|
||||
methods declare their required level on the `RecipeSpec`.
|
||||
|
||||
### 9.3 The gate that catches what batch-of-1 cannot
|
||||
|
||||
Loop inversion's real hazard is **cross-request state smearing under interleaving** — and a batch-of-1 parity gate is
|
||||
*structurally blind* to it, because the corruption only manifests when two requests share a loop. §5.1 excludes the
|
||||
hazard by construction (state in `LoopState`, never globals), but construction-arguments need a test. So v3 makes a
|
||||
**batch-of-N interleave parity test** a *required* gate: two (or more) concurrent requests, interleaved at step
|
||||
granularity, must be bit-identical to the same requests run serially. This is the test the whole loop-inversion bet
|
||||
lives or dies on, and it is named here as a first-class obligation, not left implicit.
|
||||
|
||||
### 9.4 Three execution profiles, one definition
|
||||
|
||||
Even in one runtime there are three forwards: the **serve** forward (no-grad, graphed, cached, possibly quantized), the
|
||||
**rollout** forward (serve profile + behavior capture), and the **train** forward (grad, checkpointed, FSDP-gathered).
|
||||
They share *one* loop definition; they differ only in grad mode and capture. The ladder measures the gap; the recipe
|
||||
declares the level it needs. "Train BF16, serve FP8" is legal only at C2-corrected with importance-sampling, and the
|
||||
card says so. This is how the (recipe, runtime) pair stays honest: the contract is typed and tested, not trusted.
|
||||
|
||||
---
|
||||
|
||||
## 10. Training and RL on the same loops
|
||||
|
||||
```text
|
||||
serve : request → program → loop → WorkUnits → artifacts
|
||||
rollout : prompt batch → program → loop → WorkUnits → BehaviorRecords → rewards → update
|
||||
```
|
||||
|
||||
The loop kernel is shared; the only difference is output capture and training policy — not a second interpretation of
|
||||
the model. This is design.md's §8 thesis and v2's training plane, with the dependency rule kept absolute: **`training`
|
||||
may require behavior records but must not fork serving loop logic; the engine never imports `training`.** The engine
|
||||
*is* the rollout engine (it already runs the loops); the trainer is a client.
|
||||
|
||||
**This is the moat — and it is the one place a serving-only runtime structurally cannot follow.** vllm-omni proves
|
||||
omni serving can be production-grade, but it has *no* training/RL plane at all; verl-omni and miles prove the
|
||||
alternative — a standalone trainer-side sampler on a *different* runtime than serving — costs the two-runtime tax
|
||||
forever. The whole point of collocation is that **the rollout forward *is* the serve forward plus capture**: same loop,
|
||||
same caches, same batcher, same numerics. Three consequences nothing else gets:
|
||||
|
||||
- **Every serving optimization is automatically a rollout optimization.** Distilled few-step samplers, cache-dit
|
||||
skips, CFG-parallel, paged/feature caches, step batching — the recipe team builds them once for serving and the RL
|
||||
rollout inherits them for free. FastVideo's *own* landed DiffusionNFT is the negative example that proves the
|
||||
point: it vendors a bare-model `for`-loop (`rl/common/sampling.py`, whose docstring says it "intentionally does not
|
||||
call FastVideo's full inference pipelines"), and DMD2 vendors a *second* one (`dmd2.py::_student_rollout`) — so
|
||||
today's rollout runs with **zero serving-grade optimizations** (no CFG, dense attention, full 25-step ODE,
|
||||
one-sample-at-a-time). Collocation deletes both private loops.
|
||||
- **RL rollout is a *better* batching case than open-world serving — not a worse one.** A GRPO/NFT group is K
|
||||
*identical-config* samples of one prompt: same shape, same schedule, same CFG branch. The landed config is K=24
|
||||
(`num_video_per_prompt: 24`), 6 prompts/batch × 48 batches = **288 prompt-slots/GPU/epoch, each a 24-wide homogeneous
|
||||
denoise batch** — zero bucketing required (serving must bucket heterogeneous resolutions/steps/CFG across users; a
|
||||
GRPO group is homogeneous *by construction*). And all K samples share one prompt embedding, so the content-hash
|
||||
feature cache computes the text encoder **once per group instead of 24×**. The vendored sampler captures none of
|
||||
this; it carries the embedding per sample and runs one shape at a time.
|
||||
- **One numerics surface.** Serve-forward and rollout-forward differ only in grad mode and capture (§9.4), so there is
|
||||
no rollout-vs-train kernel gap to patch — the consistency ladder *measures* the gap rather than a correction layer
|
||||
*papering over* it. For the landed likelihood-free NFT, "reuse holds" means it holds at the **C2 behavioral rung**
|
||||
(seeded sample + prediction-space identity), under a `CFGPolicy` that is conditional-only and a `WeightSyncPlan`
|
||||
whose role is the decay-blended old policy — all of which the card already declares.
|
||||
|
||||
- **BehaviorRecord** — captured at generation time (reconstructing later is fragile): seeds, scheduler trajectory,
|
||||
timesteps, latents-or-refs, logprobs *where applicable*, sampled/action tokens, guidance, reward in/out, cache
|
||||
assumptions, precision, parallel plan, attention backend, deterministic flags, `weights_version`. Sized honestly:
|
||||
full MoE-routing capture is GB/sample for Cosmos3-class requests, so it is an **opt-in instrument** for goldens and
|
||||
debugging, not always-on.
|
||||
- **Weight-sync lifecycle** — freeze admission for the affected role/version → drain or boundary-stop in-flight loops →
|
||||
transfer weights/deltas → bump `weights_version` → invalidate incompatible caches and graphs → publish version →
|
||||
resume. A `WeightSyncPlan` is three inputs (mesh specs + per-model layout adapters + transport), validated
|
||||
pre-flight, CPU-testable on fake pools. RL ships a *role*, not "the weights": student / EMA / decay-blended old
|
||||
policy is declared (the landed NFT behavior policy is the *old* copy, not the student — the plan must carry that).
|
||||
- **Roles** (policy, rollout, reference, reward, critic, evaluator, data, coordinator) reference the same cards and
|
||||
loops; they are deployment concerns, scaled by the fleet.
|
||||
- **The industry tax we delete:** verl-omni re-implements Wan inside vLLM-Omni and corrects numerics afterward; miles'
|
||||
headline features (TIS/MIS, bitwise logprobs, R3 routing replay, unified FP8) are all mismatch patches for *two
|
||||
runtimes with different kernels*. One model definition, one kernel set, one measured ladder is the answer — viable
|
||||
at FastVideo's 1–30B FSDP2 scale (the boundary condition: a Megatron-class trainer at 100B+ re-enters the
|
||||
two-runtime world, and the ladder is the fallback there).
|
||||
|
||||
---
|
||||
|
||||
## 11. Extensions — observers and interceptors
|
||||
|
||||
The optimization, debugging, and parity surface, as versioned hook points assembled at loop build (an unused hook is
|
||||
*literally absent* from the hot path). It composes with §5 cleanly: the hooks wrap `ctx.execute(plan)`.
|
||||
|
||||
- **Observers (read-only):** `ParityAligner` (§9), `Profiler` (per-step wall+CUDA, calibrates the cost model),
|
||||
`NaNWatch` (first-NaN localization), `ActivationTrace`. They cannot mutate state.
|
||||
- **Interceptors (compute-altering):** `StepInterceptor` (step-skip / cached-prediction) and `BlockInterceptor`
|
||||
(cache-dit's DBCache/FBCache/TaylorSeer). State lives in `LoopState.plugin_state[id]`, keyed **per request and per
|
||||
CFG branch** — the structural fix for the module-global residual state that silently corrupts cache-dit/TeaCache
|
||||
forks under concurrency. cache-dit is the reference integration (the library sglang's serving already uses);
|
||||
conflicting interceptors are rejected pre-flight; a 4-step distilled card *rejects* step-skip caches rather than
|
||||
producing garbage.
|
||||
|
||||
**Trust boundary:** plugins are enabled at deploy scope only (never a per-request `plugins=[...]` field that would wire
|
||||
third-party code selection into the public API); requests only *parameterize* pre-enabled plugins through validated
|
||||
schemas, and exact-mode requests reject `distribution_altering` parameterization outright.
|
||||
|
||||
---
|
||||
|
||||
## 12. Request, session, artifact, stream
|
||||
|
||||
Typed runtime objects, not IDs in a batch (the Dreamverse/LiveKit lesson):
|
||||
|
||||
- `Request` — one generation, scoring, encoding, training-sample, or conversion job.
|
||||
- `Session` — a long-lived interactive context: prompt memory, media streams, cancellation, partial updates,
|
||||
cross-request chunk-KV that persists for a game/scene session.
|
||||
- `Artifact` — a *named, typed* output with provenance (which node produced it): `VideoArtifact`, `AudioArtifact(
|
||||
sample_rate)`, `TextArtifact(token_ids, text)`, `TensorArtifact`, `LatentArtifact`. This kills the `extra["audio"]`
|
||||
pattern — audio carries its sample rate as a first-class artifact, not a dict passenger.
|
||||
- `Stream` — one ordered event channel for previews, media chunks, progress, logs, finals.
|
||||
- `CancelScope` — structured cancellation target (request / loop / stream / session).
|
||||
|
||||
Typed event taxonomy (`request.*`, `session.*`, `artifact.*`, `media.{init,chunk,complete}`, `trace.*`). A
|
||||
`media.chunk` must know its stream, byte-range or shared-buffer ref, codec/container, timestamp range, and
|
||||
preview-vs-final — invalid combinations are unrepresentable.
|
||||
|
||||
The **request is the only currency crossing the product boundary.** A typed `Request` carries `task: TaskType`
|
||||
(declared, never inferred), `inputs: list[ModalPart]` (Text/Image/Video/Audio/Action/Latent), AR `sampling` vs
|
||||
`diffusion` params, an `OutputSpec` (requested modalities + streaming + capture flags), and per-node overrides. Task is
|
||||
declared; heuristics may only *suggest* a default at the boundary.
|
||||
|
||||
---
|
||||
|
||||
## 13. Programs and workflows
|
||||
|
||||
A **Program** composes a card's loops into a task; the card says what loops *exist*, the program says how to *run* them
|
||||
for this request. Kinds: `InlineProgram` (many loops, one resident instance — the omni default), `DisaggregatedProgram`
|
||||
(encoder→denoiser→decoder role pools), `WorkflowProgram` (compiled from ComfyUI), `TrainingProgram`,
|
||||
`RealtimeProgram`. Nodes: `ModelLoopNode`, `ComponentNode`, `ExternalNode`, `ArtifactNode`, `ControlNode`,
|
||||
`StreamNode`, `TransferNode`. Edges are typed (`TensorEdge`, `ArtifactEdge`, `StreamEdge`, `ControlEdge`, `CacheEdge`,
|
||||
`BehaviorEdge`). Linear pipelines are the degenerate case; branches/fan-out/fan-in are real (video and audio decode in
|
||||
parallel after a joint denoise). A separate deploy config maps nodes → pools/devices/parallelism, defaulting to "one
|
||||
pool, everything colocated."
|
||||
|
||||
**Workflows compile, they are not the runtime.** A ComfyUI workflow's tier-1/tier-2 static sublanguage maps onto a
|
||||
`Program` (`CheckpointLoaderSimple→card`, `KSampler→diffusion_denoise` with sampler/CFG policies,
|
||||
`LoraLoader→adapter hot-swap`, `ControlNetApply→ConditioningInjector`); unknown nodes become `ExternalNode`s or a
|
||||
coverage rejection — never silent wrongness. The moat: an orchestrator can run stock workflows on rented silicon;
|
||||
substituting a *credibly faster* model requires owning the recipe (§2.1) — which an orchestrator structurally cannot
|
||||
do. Equivalence is a quality-metric vs a reference render (C4), never a bit-parity claim.
|
||||
|
||||
---
|
||||
|
||||
## 14. Deployment and fleet
|
||||
|
||||
The engine exports a `DeploymentCard` and lets a fleet orchestrator (Dynamo) route — Dynamo orchestrates engines, it
|
||||
is never the engine core.
|
||||
|
||||
```python
|
||||
class DeploymentCard:
|
||||
engine_id: str; model_cards: list[str]
|
||||
capabilities: CapabilityMatrix; role_pools: list[RolePoolSpec]
|
||||
supported_programs: list[str]; supported_parallel_plans: list[ParallelPlan]
|
||||
cache_events: list[CacheEventSpec]; transfer_endpoints: list[TransferEndpoint]
|
||||
cost_model: CostModel # the SAME §6 cost model — one object, two consumers
|
||||
health: HealthSchema; slo: SLOSchema
|
||||
```
|
||||
|
||||
Clean line: the **fleet** owns global routing, tenant policy, cold start, role-pool scaling, cross-node transfer,
|
||||
placement-by-SLO, health/failover, multi-engine upgrades, global cache routing. The **engine** owns model load, loop
|
||||
execution, local scheduling, local memory/cache, model-specific behavior, parity, WorkUnit batching. The asks of the
|
||||
fleet are concrete and each has a fallback: generic affinity key-spaces (checkpoint/session/lora/weight_version beyond
|
||||
token prefixes), a heterogeneous request-cost interface (the §6 cost model), chunked media streaming through the
|
||||
frontend, role-graph disagg (N roles, not two), an RL weight plane (versioned broadcast + staleness-aware routing),
|
||||
cache-object tiering (KVBM generalized to latent/session caches), and session lifecycle as a routing primitive.
|
||||
|
||||
---
|
||||
|
||||
## 15. Worked examples
|
||||
|
||||
**(a) Text → video, one instance.** `Request(T2V)` → `InlineProgram` → `diffusion_denoise` loop. Driver: `init`
|
||||
builds sigmas/latents; `next` emits a `diffusion_step` WorkPlan (batch-of-1, SP+CFG-parallel); `advance` folds the
|
||||
model output and the CFG combine; cache-dit's `BlockInterceptor` may skip blocks based on the prior residual; at
|
||||
`Done`, a `vae_tile_decode` loop runs; output is a named `VideoArtifact`. Compiles/captures exactly like today's inner
|
||||
loop — nothing tensor-level changes for batch-of-1.
|
||||
|
||||
**(b) Cosmos3 omni, one request, shared weights.** `InlineProgram` over one `ModelInstance`: `ar_decode(reasoner)`
|
||||
yields `ar_token` WorkUnits that join the AR continuous-batching group → `pack` → `diffusion_denoise(vision+action+
|
||||
sound)` yields `diffusion_step` WorkUnits → fan-out `vae_tile_decode` + `audio_decode`. The reasoner's tokens and the
|
||||
denoiser's steps hit the *same resident weights*; the scheduler is the mode multiplexer. AR decode runs data-parallel
|
||||
across the cfg×sp weight-replica axes (decode is sequence-length-1; SP has nothing to shard). This is the workload no
|
||||
DAG-of-engines can express.
|
||||
|
||||
**(c) Image serving at scale.** Many `Request(T2I)` → the `BatchScheduler` groups `diffusion_step` WorkUnits by
|
||||
resolution bucket and batches across requests every step — the case where cross-request batching pays most. The
|
||||
*same* scheduler that runs (a) and (b).
|
||||
|
||||
**(d) RL rollout.** A `TrainingProgram` drives the *same* `diffusion_denoise` loop with `OutputSpec(capture=behavior)`;
|
||||
each step emits a `BehaviorRecord` slice; rollouts run C2 by construction (in-process, trainer kernels, pinned
|
||||
attention). For likelihood-free NFT the behavior is seeded final latents + prediction-space deviations; for a future
|
||||
GRPO-class method it is per-step log-probs — the loop is identical, the capture differs, the ladder rung is declared.
|
||||
|
||||
**(e) Dreamverse session.** A `Session(realtime_video_continue)` holds chunk-KV across 5s segments; `push_text`
|
||||
updates prompt memory; `stream` yields `media.chunk` previews from loop `emit`s; a direction change throws
|
||||
`Cancelled` at the next step boundary and starts a new segment. Capacity comes from duty cycle + cost-model admission +
|
||||
distillation — interleaving is fairness, not throughput.
|
||||
|
||||
**(f) ComfyUI compile.** `workflow.compile(json)` → a `WorkflowProgram` of `ModelLoopNode`/`ComponentNode` over a
|
||||
weight-fleet-cached card, with stacked-LoRA patch/unpatch priced by the §6 weight-transition cost — same runtime, new
|
||||
frontend.
|
||||
|
||||
---
|
||||
|
||||
## 16. What this unlocks (the unconstrained payoff)
|
||||
|
||||
Things no incremental design — and neither prior document — could actually claim:
|
||||
|
||||
- **True omni/MoT serving.** One resident model, many loop types, scheduled at step granularity, in one request. Not a
|
||||
monolith bypassing the abstraction (the Cosmos3 port's necessary hack), not a DAG doubling weights — native.
|
||||
- **Train ≡ serve by construction.** Because rollout and serve are the *same loop*, the (recipe, runtime) flywheel is
|
||||
real and measured, not aspirational: Dreamverse's directing sessions emit preference data → the RL plane → faster
|
||||
distilled cards → a better product, with the ladder guaranteeing the preferences collected under the serving profile
|
||||
transfer into training.
|
||||
- **Real-time interactive omni.** The driven-loop contract + sessions + WebRTC frame/PTS streaming + step-boundary
|
||||
cancellation make the <100ms motion-to-photon interactive world-model loop expressible in the same runtime that
|
||||
serves batch T2V.
|
||||
- **One substrate, three personas.** A research vehicle (new ports land as cards), a product engine (Dreamverse, the
|
||||
workflow cloud), and an RL rollout engine — without three codebases. The dependency rules keep them from fusing into
|
||||
mud.
|
||||
- **Correctness you can sign.** A deployable card is a *(recipe, runtime)* pair with a typed parity obligation; "this
|
||||
fast model is equivalent" is a claim with a test behind it, which is the one thing an orchestrator-without-recipes
|
||||
can never say.
|
||||
|
||||
---
|
||||
|
||||
## 17. Honest unknowns and falsifiers
|
||||
|
||||
An unconstrained design is not an unfalsifiable one. The bets, stated with the experiment that kills each:
|
||||
|
||||
- **The novelty is concentrated and real.** Runtime-owned diffusion iteration has a **narrow, opt-in precedent** —
|
||||
vllm-omni's `SupportsStepExecution` (`prepare_encode/denoise_step/step_scheduler/post_decode`,
|
||||
`diffusion/models/interface.py:44-67`) is exactly runtime-owned diffusion iteration at step granularity, and maps
|
||||
almost 1:1 onto our `init/next/advance/finalize` — but it is Qwen-Image-only and off in every shipped deploy. What
|
||||
is unprecedented is making it the **always-on universal contract** *and* a fully general WorkUnit scheduler over
|
||||
heterogeneous units. The risk is not the loop contract (a state machine is well-understood, and now demonstrably
|
||||
shippable); it is whether step-level cross-request scheduling *pays* for video. **Falsifier:** publish a load profile and targets from a real duty-cycle trace; if step-level scheduling does
|
||||
not beat a request-level baseline (≥2 concurrent sessions/GPU, p95 within SLO), the scheduler degrades to
|
||||
request-level dispatch and the loop contract keeps only its streaming/cancellation/behavior seams — which still
|
||||
justify it. The contract is safe even if the scheduling bet loses; that is the design's insurance.
|
||||
- **The general WorkUnit scheduler may be over-general.** Scheduling VAE tiles, transfers, and graph-captures through
|
||||
the *same* admission machinery as denoise steps is elegant and unproven. **Falsifier:** if, after Phase 2, the
|
||||
non-diffusion/non-AR WorkUnit kinds (tile, transfer, cache_io) gain nothing from unified scheduling over a simple
|
||||
in-loop call, collapse them back to in-loop operations and keep WorkUnits for the step-bearing kinds only.
|
||||
- **Cost-model admission is a modeling bet.** It converges toward cost-class pool routing once the indivisible-step
|
||||
reality is respected — which is close to what request-level pooling + a fleet planner already do. The fine-grained
|
||||
interleave win has a *narrow* window (many small concurrent jobs); it should be argued on that window, measured.
|
||||
- **The clean-slate premise is the elephant.** This document deliberately ignores migration. The org that would build
|
||||
it broke its own freeze 19 times and ships 20+ families, a live product, and a landed RL stack. A clean-slate
|
||||
rebuild is the highest-risk path that exists for *this* org; the responsible realization is to build v3 as a
|
||||
*parallel* engine around one forcing-function card (Cosmos3), prove it on the parity ladder, then migrate families
|
||||
onto it behind an adapter while everything keeps shipping — i.e., reach this architecture incrementally. That plan is
|
||||
out of scope here by request; it is non-optional in reality.
|
||||
- **Quality is unmeasured.** C4 (artifact quality / human preference) and the eval system it needs do not exist yet,
|
||||
and they gate every product claim ("fast mode is equivalent", RL reward validity, distillation comparisons). Named
|
||||
as a required, currently-absent subsystem, not assumed.
|
||||
|
||||
---
|
||||
|
||||
## 18. Package layout
|
||||
|
||||
```text
|
||||
fastvideo/
|
||||
card/ specs, components, loops, recipes, parity, checkpoints, capabilities # the Model Plane
|
||||
loop/ driver, loopstate, workplan, policies (cfg, expert, precision, flowshift, conditioning)
|
||||
runtime/ engine, scheduler/{request,loop,batch,placement,transfer,admission}, workers, events
|
||||
cache/ keys, classes/{paged_kv, slab_kv, feature, residual, weight_fleet}, policies
|
||||
memory/ allocator, sleep_wake, reservations
|
||||
transport/ manifests, backends/{shm, cuda_ipc, nccl, nixl, kvbm}, relay
|
||||
parallel/ plans, mesh, process_groups, validation
|
||||
parity/ aligner, ladder, interleave_gate # §9 is its own home
|
||||
extend/ observers, interceptors, cache_dit, registry, trust
|
||||
program/ specs, compiler, workflows
|
||||
request/ requests, sessions, artifacts, streams, cancel
|
||||
training/ rollout, behavior, rewards, weight_sync, methods # imports card/loop/runtime; never imported by them
|
||||
deploy/ cards, role_pools, dynamo_adapter
|
||||
integrations/ comfyui, dreamverse, livekit, diffusers
|
||||
```
|
||||
|
||||
Enforced boundaries: `card/` imports no product/runtime; `runtime/` executes `card/` loops but defines no semantics;
|
||||
`training/` may require behavior records but forks no loop; `integrations/` adapt external systems into core specs and
|
||||
events, never bypass them. **`parity/` is a first-class package**, not a test folder — it is how the (recipe, runtime)
|
||||
pair is kept honest.
|
||||
|
||||
---
|
||||
|
||||
## 19. Reference synthesis
|
||||
|
||||
| Source | Take | Constrain / reject |
|
||||
|---|---|---|
|
||||
| Cosmos3 (official + port) | Shared model instance across reasoning/diffusion/action/sound; packed multimodal sequences; component+scheduler parity matrices | A strong `ModelCard`, not the framework; no Cosmos-specific branching in global runtime |
|
||||
| vLLM core | Running-first scheduling, reservation-before-admission, model-owned state, encoder/KV cache managers, CuMem sleep/wake, CUDA-graph dispatch, KV-connector split | Token scheduling is one WorkUnit kind; never full-graph compile |
|
||||
| sglang `multimodal_gen` | Role pools, request lifecycle, capacity dispatch, transfer manifests, disagg state machine, cache-dit integration | No large mutable `Req`/`ForwardBatch` as the stable API; not single-item diffusion scheduling |
|
||||
| vLLM-Omni | Frozen pipeline spec separate from deploy YAML (verified, adopt); `OmniConnectorBase` + `chunk_ready` readiness; **`SupportsStepExecution` as loop-inversion prior art** (opt-in, Qwen-Image-only — we generalize to always-on); TP-rank- and CFG-branch-aware KV-copy transfer; **`CFGParallelMixin` proves CFG-as-policy over one shared denoise body** (§5.3); 3 separate cache subsystems confirm per-class pools | Expresses shared-weight MoT (`bagel`/`lance`) only as **one opaque request-scheduled stage** the scheduler never sees inside — no step visibility, no cross-request batching by default; cross-stage KV is a *copy*, not a shared live cache; **no cost model** (per-stage count budgets); readiness-parking, not credit flow; RDMA = Mooncake/Mori/Yuanrong, not NIXL/NCCL |
|
||||
| sglang-omni | The `next/wait_for/merge_fn/stream_to` edge vocabulary; Relay transport + **credit-based flow control** (this is sglang-omni's, not vllm-omni's) | Stages own disjoint weights; hybrid AR+diffusion only as AR-stage → DiT-stage; per-model bootstrap duplication |
|
||||
| Dynamo | Fleet routing, disagg role pools, KV-aware routing, KVBM, SLA planner, ModelExpress cold-start/weight streaming | Orchestrates engines; never the engine core. Export a `DeploymentCard` + cost model to it |
|
||||
| diffusers Modular | `ComponentSpec`/`modular_model_index.json` interchange; Guiders ≈ CFG policies | A Python pipeline interpreter is not the performance boundary; import is lossy |
|
||||
| xDiT | DiT parallelism catalog (USP, ring/ulysses, PipeFusion, CFG-parallel, DistVAE) + world-size validation | Parallelism lives in the runtime + card, not a wrapper-per-model library; `pp_patch` invalid for causal |
|
||||
| TorchTitan | Named mesh axes, `ParallelDims` validation, ModelSpec discipline, TorchStore weight-sync, batch-invariance utils | Adopt the discipline, not the stack; DCP/TorchStore don't reshard — `WeightSyncPlan` owns layout |
|
||||
| verl-omni / miles / cosmos-rl | Rollout adapters, per-step capture, async rewards, group-relative advantage, TIS/MIS, deterministic/batch-invariant modes, per-payload `weight_version`, AIPO/off-policy masking | The two-runtime tax is the thing to delete; capture behavior *in* the serving loop, not after the fact |
|
||||
| ComfyUI | Workflow graph, node-signature cache, model memory management, App-Mode (workflows-as-products) | Compile to `Program`; dynamic node execution is not the serving/training core; GPL hygiene |
|
||||
| Dreamverse | Sessions, prompt memory, typed media IPC, cancellation, the duty-cycle capacity reality, the preference-data flywheel | Product/session behavior is first-class in the request plane, never merged into the model core |
|
||||
| LiveKit | Realtime sessions, push audio/video, interruptions, turn/activity state, frame+PTS streaming | Realtime triggers only when they fire (<100ms interactive); don't force RTC onto offline jobs |
|
||||
| Thinking Machines (batch-invariance) | The C2/C3 mechanism: batch-invariant kernels for bitwise rollout↔train identity | Scoped to goldens; the conservative baseline governs admission |
|
||||
|
||||
---
|
||||
|
||||
## Final position
|
||||
|
||||
```text
|
||||
A model card is a (recipe, runtime) pair with a parity obligation.
|
||||
The model owns loop semantics; the runtime owns loop lifecycle.
|
||||
One resident instance runs many loops; one scheduler runs their steps in one currency.
|
||||
Caches are correct by key; parity is correct by test; the interleave gate is non-negotiable.
|
||||
Training records behavior on the same loops it serves.
|
||||
Deployment places and routes; products stream artifacts; neither defines the model.
|
||||
```
|
||||
|
||||
This is the ceiling: a model-native runtime where omni is native, train and serve are the same loops by construction,
|
||||
correctness is a typed contract you can sign, and the (recipe, runtime) flywheel is real. The constraint we removed to
|
||||
see it was migration. Putting that constraint back is the next document, not this one.
|
||||
-1895
File diff suppressed because it is too large
Load Diff
@@ -41,6 +41,13 @@ when you want local Markdown/normalized-result artifacts from the comparator.
|
||||
`fastvideo/tests/performance/results/`; remove stale result files if you only
|
||||
want to compare the latest local run.
|
||||
|
||||
## Local live dashboard
|
||||
|
||||
For an app-style local dashboard backed by the same HF performance-tracking
|
||||
records, see `performance_dashboard/README.md`. The dashboard provides a
|
||||
FastAPI API plus a React UI and can be exposed with `ngrok` after building the
|
||||
frontend.
|
||||
|
||||
## Architecture
|
||||
|
||||
```
|
||||
|
||||
@@ -24,6 +24,7 @@ This page describes the various options for speeding up generation times in Fast
|
||||
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
|
||||
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
|
||||
- Sage Attention 3: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN_THREE`
|
||||
- Attn-QAT inference (modified SageAttention3 FP4, sm_120/RTX 5090): `FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER`
|
||||
- Video MoBA Attention: `FASTVIDEO_ATTENTION_BACKEND=VMOBA_ATTN`
|
||||
- Sparse Linear Attention: `FASTVIDEO_ATTENTION_BACKEND=SLA_ATTN`
|
||||
- SageSLA Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_SLA_ATTN`
|
||||
@@ -122,6 +123,45 @@ gen.generate_video(prompt="A raccoon in sunflowers", save_video=True)
|
||||
- Per-call cosine similarity vs BF16: ~0.99 (slight quantization error accumulates over denoising steps)
|
||||
- Only supports `headdim >= 128`
|
||||
|
||||
### NVFP4 + Attn-QAT (modified SageAttention3, Blackwell sm_120)
|
||||
|
||||
**`ATTN_QAT_INFER`** with **`transformer_quant=nvfp4_qat`**
|
||||
|
||||
Runs the DiT fully in 4-bit: NVFP4 linear layers (activations quantized on the
|
||||
fly) plus the modified SageAttention3 FP4 attention backend. This is the
|
||||
inference half of the Quantization-Aware Distillation (QAD) recipe and the path
|
||||
used for the RTX 5090 release.
|
||||
|
||||
The `attn_qat_infer` kernel hard-gates on **sm_120 (consumer Blackwell / RTX
|
||||
5090)**; on other GPUs the backend logs a notice and falls back to Flash
|
||||
Attention. See the [Attn-QAT paper](https://arxiv.org/abs/2603.00040).
|
||||
|
||||
Enable both halves — attention via the env var, linear via `transformer_quant`:
|
||||
|
||||
```python
|
||||
import os
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "ATTN_QAT_INFER"
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
# Wan-2.1 uses the nvfp4_qat config (NVFP4 is LTX2-specific). Pass an
|
||||
# instance — the bare string is not resolved on the from_pretrained path.
|
||||
transformer_quant=get_quantization_config("nvfp4_qat")(),
|
||||
use_fsdp_inference=False, # FSDP shards invalidate the FP4 tensor pointers
|
||||
)
|
||||
gen.generate(request={"prompt": "A raccoon in sunflowers", "output": {"save_video": True}})
|
||||
```
|
||||
|
||||
Or run the example script:
|
||||
|
||||
```bash
|
||||
python examples/inference/optimizations/nvfp4_qat_wan2_1_1_3b.py
|
||||
python examples/inference/optimizations/nvfp4_qat_wan2_1_1_3b.py --bf16 # baseline
|
||||
```
|
||||
|
||||
### Sliding Tile Attention (Archived)
|
||||
|
||||
**`SLIDING_TILE_ATTN`**
|
||||
|
||||
@@ -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,95 @@
|
||||
"""NVFP4 + Attn-QAT (modified SageAttention3) inference on Blackwell.
|
||||
|
||||
Runs Wan2.1-T2V-1.3B fully in 4-bit: NVFP4 linear layers (activations
|
||||
quantized on the fly) together with the modified SageAttention3 FP4 attention
|
||||
backend (``ATTN_QAT_INFER``). This is the inference half of the
|
||||
Quantization-Aware Distillation (QAD) recipe.
|
||||
|
||||
Requirements:
|
||||
- RTX 5090 / consumer Blackwell (sm_120a). The attn_qat_infer kernel hard
|
||||
gates on sm_120; on other GPUs it falls back to Flash Attention.
|
||||
- The attn_qat_infer kernel built into fastvideo-kernel (see #1455) and
|
||||
flashinfer for the NVFP4 linear matmuls.
|
||||
|
||||
Usage:
|
||||
python nvfp4_qat_wan2_1_1_3b.py # NVFP4 linear + Attn-QAT attn
|
||||
python nvfp4_qat_wan2_1_1_3b.py --bf16 # BF16 baseline
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import time
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="NVFP4 + Attn-QAT video generation")
|
||||
parser.add_argument("--bf16", action="store_true",
|
||||
help="BF16 baseline (no NVFP4 linear, default attention)")
|
||||
parser.add_argument("--model", default="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
help="Model path or HuggingFace ID")
|
||||
parser.add_argument("--quant-method", default="nvfp4_qat", choices=["nvfp4_qat", "NVFP4"],
|
||||
help="Linear quantization config. Wan-2.1 uses nvfp4_qat (matches its "
|
||||
"to_q/k/v/out + ffn layers); NVFP4 is LTX2-specific and will NOT "
|
||||
"quantize Wan.")
|
||||
parser.add_argument("--compile", action="store_true", help="Enable torch.compile for the DiT")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--infer_steps", type=int, default=50)
|
||||
args = parser.parse_args()
|
||||
|
||||
# The attention backend is selected via env var before the engine starts.
|
||||
# ATTN_QAT_INFER -> AttnQatInferBackend (modified SageAttention3 FP4).
|
||||
if not args.bf16:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "ATTN_QAT_INFER"
|
||||
|
||||
# Import after the env var so the platform picks up the selection.
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
|
||||
mode = "bf16" if args.bf16 else args.quant_method
|
||||
if args.compile:
|
||||
mode += "_compile"
|
||||
print(f"Mode: {mode.upper()}")
|
||||
|
||||
# transformer_quant needs a QuantizationConfig *instance* — the bare string
|
||||
# is not resolved on the from_pretrained kwarg path.
|
||||
extra = {} if args.bf16 else {"transformer_quant": get_quantization_config(args.quant_method)()}
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
args.model,
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=args.bf16,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
enable_torch_compile=args.compile,
|
||||
**extra,
|
||||
)
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
n_warmup = 2 if args.compile else 1
|
||||
for _ in range(n_warmup):
|
||||
generator.generate(request={"prompt": prompt, "sampling": {"num_inference_steps": 2},
|
||||
"output": {"save_video": False}})
|
||||
|
||||
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
||||
start = time.time()
|
||||
generator.generate(request={
|
||||
"prompt": prompt,
|
||||
"sampling": {"num_inference_steps": args.infer_steps},
|
||||
"output": {"save_video": True, "output_path": os.path.join(OUTPUT_PATH, f"raccoon_{mode}.mp4")},
|
||||
})
|
||||
elapsed = time.time() - start
|
||||
print(f"[{mode.upper()}] {args.infer_steps} steps in {elapsed:.2f}s "
|
||||
f"({args.infer_steps / elapsed:.2f} it/s)")
|
||||
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,119 @@
|
||||
# MixKit training data (QAD 5090 recipe)
|
||||
|
||||
The QAD 5090 models are distilled from Wan2.1-T2V-1.3B on a MixKit subset at
|
||||
**480×832, 77 frames, 16 fps**. FastVideo training consumes **Parquet** shards of
|
||||
precomputed VAE latents + text embeddings (no text encoder / VAE needed at train
|
||||
time).
|
||||
|
||||
## Option A — download the preprocessed data (recommended)
|
||||
|
||||
The encoded dataset is published on the Hugging Face Hub, ready to train:
|
||||
|
||||
```bash
|
||||
# from the repo root
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
|
||||
```
|
||||
|
||||
This pulls [`weizhou03/HD-Mixkit-Finetune-Wan`](https://huggingface.co/datasets/weizhou03/HD-Mixkit-Finetune-Wan)
|
||||
into `data/HD-Mixkit-Finetune-Wan/`:
|
||||
|
||||
```
|
||||
data/HD-Mixkit-Finetune-Wan/
|
||||
├── combined_parquet_dataset/ # training shards -> point --data_path here
|
||||
│ └── worker_0/data_chunk_*.parquet
|
||||
└── validation_parquet_dataset/ # validation shards
|
||||
└── worker_0/data_chunk_0.parquet
|
||||
```
|
||||
|
||||
Each Parquet row holds the VAE latent bytes + text-embedding bytes (plus
|
||||
shape/dtype metadata), matching FastVideo's standard preprocessing output.
|
||||
|
||||
## Option B — build the Parquet from raw videos
|
||||
|
||||
If you want to reproduce the encoding from your own MixKit videos, arrange them as
|
||||
a `merged` dataset (videos + a captions JSON), then run FastVideo's standard
|
||||
preprocessing to VAE-encode and text-embed them into Parquet:
|
||||
|
||||
```bash
|
||||
GPU_NUM=2
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path "data/mixkit_raw/" \
|
||||
--preprocess.dataset_output_dir "data/HD-Mixkit-Finetune-Wan/" \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 77 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.samples_per_file 8
|
||||
```
|
||||
|
||||
The raw videos are full-HD MixKit clips (≈1080p/30fps); preprocessing resizes to
|
||||
480×832, resamples to 16 fps, and extracts 77 frames per clip. See
|
||||
[`docs/training/data_preprocess.md`](../../../../../docs/training/data_preprocess.md)
|
||||
for the full parameter reference.
|
||||
|
||||
## Train (QAT finetune)
|
||||
|
||||
With the data in place, run the quantization-aware finetune. The 4-bit attention
|
||||
path is **config-driven** — selected purely by an env var, no monkey-patching:
|
||||
|
||||
```bash
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
|
||||
# or point at your own parquet dir / GPU count:
|
||||
NUM_GPUS=4 bash .../mixkit/finetune_qat.sh data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/
|
||||
```
|
||||
|
||||
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` routes attention through the
|
||||
fake-quantized Triton kernel (straight-through estimator), so the DiT learns to
|
||||
absorb FP4 attention error. This kernel is Triton, so it runs on both `sm_100`
|
||||
(B200/GB200) and `sm_120` (RTX 5090).
|
||||
|
||||
## Train stage 2 (QAT DMD distillation to 3 steps)
|
||||
|
||||
Distill the QAT-finetuned generator down to **3 sampling steps**. Only the
|
||||
generator is quantized (Attn-QAT); the teacher (`real_score`) and critic
|
||||
(`fake_score`) stay full precision. This is enforced in the loader
|
||||
(`component_loader.py`, via the `_loading_teacher_critic_model` flag), so the
|
||||
same global `ATTN_QAT_TRAIN` env reaches **only** the generator — no per-model
|
||||
flags or monkey-patching.
|
||||
|
||||
```bash
|
||||
# generator init = the stage-1 finetune checkpoint
|
||||
bash examples/training/finetune/wan_t2v_1.3B/mixkit/distill_dmd_qat.sh \
|
||||
data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/ \
|
||||
checkpoints/wan_t2v_qat_finetune/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors
|
||||
```
|
||||
|
||||
DMD runs a double loop (critic every step, generator every
|
||||
`generator_update_interval`), and validation samples the distilled student at
|
||||
3 steps — the final 4-bit-attention model.
|
||||
|
||||
## Inference (NVFP4 4-bit linear)
|
||||
|
||||
For Wan-2.1, enable the FP4 linear layers with the **`nvfp4_qat`** quantization
|
||||
config (it matches Wan's `to_q/k/v/out` + `ffn` layers; the plain `NVFP4` config
|
||||
is LTX2-specific and will not quantize Wan):
|
||||
|
||||
```python
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers", num_gpus=1,
|
||||
transformer_quant=get_quantization_config("nvfp4_qat")(), # a config instance, not the string
|
||||
use_fsdp_inference=False,
|
||||
)
|
||||
gen.generate(request={"prompt": "...", "output": {"save_video": True}})
|
||||
```
|
||||
|
||||
The loader converts the tagged linear weights to FP4 at load time
|
||||
(`_maybe_convert_model_to_nvfp4`). Combine with
|
||||
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_INFER` on an RTX 5090 (`sm_120`) for the
|
||||
full 4-bit path; on other GPUs the attention falls back to Flash while the FP4
|
||||
linear layers still run. `flashinfer` (and a host C++ compiler for its FP4
|
||||
kernel JIT) are required.
|
||||
@@ -0,0 +1,51 @@
|
||||
#!/bin/bash
|
||||
# QAD recipe stage 2 — quantization-aware DMD distillation of Wan2.1-T2V-1.3B
|
||||
# down to 3 sampling steps, with the GENERATOR in fake-quant Attn-QAT and the
|
||||
# teacher (real_score) + critic (fake_score) at full precision.
|
||||
#
|
||||
# Generator-only QAT is config-driven: FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN
|
||||
# is applied to the generator only, because the loader masks it (and the
|
||||
# nvfp4_qat quant) for the teacher/critic via the `_loading_teacher_critic_model`
|
||||
# flag (see fastvideo/models/loader/component_loader.py). No monkey-patching.
|
||||
#
|
||||
# Init the generator from the stage-1 finetune checkpoint (finetune_qat.sh).
|
||||
# Data: run download_mixkit_data.sh first.
|
||||
#
|
||||
# Verified end-to-end on Blackwell (GB200/sm_100): generator loads with
|
||||
# ATTN_QAT_TRAIN while teacher/critic load full-precision; the DMD double loop
|
||||
# runs (generator updates every generator_update_interval steps, critic every
|
||||
# step), 3-step validation generates videos, checkpoint saved.
|
||||
set -euo pipefail
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # generator-only (loader-gated)
|
||||
export WANDB_MODE=${WANDB_MODE:-online}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
BASE="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
|
||||
# Generator init weights = the stage-1 QAT-finetune checkpoint.
|
||||
INIT_WEIGHTS=${2:-"checkpoints/wan_t2v_qat_finetune/checkpoint-2000/transformer/diffusion_pytorch_model.safetensors"}
|
||||
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
--num_gpus "${NUM_GPUS}" --sp_size 1 --tp_size 1 \
|
||||
--hsdp_replicate_dim "${NUM_GPUS}" --hsdp_shard_dim 1 \
|
||||
--model_path "${BASE}" --pretrained_model_name_or_path "${BASE}" \
|
||||
--real_score_model_path "${BASE}" --fake_score_model_path "${BASE}" \
|
||||
--init_weights_from_safetensors "${INIT_WEIGHTS}" \
|
||||
--data_path "${DATA_DIR}" --dataloader_num_workers 4 \
|
||||
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
|
||||
--enable_gradient_checkpointing_type full \
|
||||
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
|
||||
--validation_steps 200 --validation_sampling_steps 3 --validation_guidance_scale 6.0 \
|
||||
--learning_rate 2e-6 --mixed_precision bf16 --weight_decay 0.01 --max_grad_norm 1.0 \
|
||||
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
|
||||
--tracker_project_name wan_t2v_distill_dmd_qat \
|
||||
--output_dir checkpoints/wan_t2v_distill_dmd_qat \
|
||||
--inference_mode False --dit_precision fp32 --ema_start_step 0 --training_cfg_rate 0.0 \
|
||||
--generator_update_interval 5 --real_score_guidance_scale 2.0 \
|
||||
--dmd_denoising_steps '1000,757,522' --min_timestep_ratio 0.02 --max_timestep_ratio 0.98
|
||||
@@ -0,0 +1,22 @@
|
||||
#!/bin/bash
|
||||
# Download the preprocessed MixKit finetune dataset used for the QAD 5090 recipe.
|
||||
#
|
||||
# This is the MixKit subset already VAE-encoded (Wan2.1-T2V-1.3B) and text-embedded
|
||||
# into Parquet shards, so it can be fed straight to training with no further
|
||||
# preprocessing. To build the Parquet from raw videos yourself, see README.md.
|
||||
#
|
||||
# Usage (run from the repo root):
|
||||
# bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh [DATA_ROOT]
|
||||
set -euo pipefail
|
||||
|
||||
DATA_ROOT=${1:-data/HD-Mixkit-Finetune-Wan}
|
||||
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id "weizhou03/HD-Mixkit-Finetune-Wan" \
|
||||
--local_dir "${DATA_ROOT}" \
|
||||
--repo_type "dataset"
|
||||
|
||||
echo "Done."
|
||||
echo " Train data: ${DATA_ROOT}/combined_parquet_dataset"
|
||||
echo " Validation data: ${DATA_ROOT}/validation_parquet_dataset"
|
||||
echo "Point your training script's data path at the combined_parquet_dataset directory."
|
||||
@@ -0,0 +1,44 @@
|
||||
#!/bin/bash
|
||||
# QAD recipe — quantization-aware finetune of Wan2.1-T2V-1.3B with fake-quant
|
||||
# (Attn-QAT) attention.
|
||||
#
|
||||
# The 4-bit attention path is selected purely by env var (config-driven, no
|
||||
# monkey-patching): FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN routes attention
|
||||
# through the fake-quantized Triton kernel (straight-through estimator), so the
|
||||
# DiT learns to absorb FP4 attention error instead of fighting it.
|
||||
#
|
||||
# Data: run download_mixkit_data.sh first (preprocessed Parquet).
|
||||
#
|
||||
# Verified end-to-end on Blackwell (GB200/sm_100): the ATTN_QAT_TRAIN backend is
|
||||
# selected (not a fallback), forward+backward run, loss/grad are healthy, and
|
||||
# validation generates videos. The kernel is Triton so it runs on sm_100 and
|
||||
# sm_120 alike (the FP4 inference kernel, by contrast, is sm_120-only).
|
||||
set -euo pipefail
|
||||
|
||||
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # <-- enables Attn-QAT training
|
||||
export WANDB_MODE=${WANDB_MODE:-online}
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
|
||||
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
|
||||
NUM_GPUS=${NUM_GPUS:-4}
|
||||
|
||||
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
|
||||
fastvideo/training/wan_training_pipeline.py \
|
||||
--num_gpus "${NUM_GPUS}" --sp_size "${NUM_GPUS}" --tp_size 1 \
|
||||
--hsdp_replicate_dim 1 --hsdp_shard_dim "${NUM_GPUS}" \
|
||||
--model_path "${MODEL_PATH}" --pretrained_model_name_or_path "${MODEL_PATH}" \
|
||||
--data_path "${DATA_DIR}" --dataloader_num_workers 1 \
|
||||
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
|
||||
--gradient_accumulation_steps 1 \
|
||||
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
|
||||
--enable_gradient_checkpointing_type full \
|
||||
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
|
||||
--validation_steps 200 --validation_sampling_steps 50 --validation_guidance_scale 3.0 \
|
||||
--learning_rate 5e-5 --mixed_precision bf16 --weight_decay 1e-4 --max_grad_norm 1.0 \
|
||||
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
|
||||
--tracker_project_name wan_t2v_qat_finetune --output_dir checkpoints/wan_t2v_qat_finetune \
|
||||
--inference_mode False --training_cfg_rate 0.1 --not_apply_cfg_solver \
|
||||
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 \
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
@@ -50,6 +50,21 @@ include_directories(
|
||||
set(FASTVIDEO_KERNEL_BUILD_TK "AUTO" CACHE STRING "Build ThunderKittens kernels: AUTO/ON/OFF")
|
||||
set_property(CACHE FASTVIDEO_KERNEL_BUILD_TK PROPERTY STRINGS AUTO ON OFF)
|
||||
|
||||
set(_FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER_DEFAULT "AUTO")
|
||||
if(DEFINED FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3 AND NOT DEFINED CACHE{FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER})
|
||||
set(_FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER_DEFAULT "${FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3}")
|
||||
endif()
|
||||
|
||||
set(FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER "${_FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER_DEFAULT}" CACHE STRING
|
||||
"Build attn_qat_infer Blackwell inference kernels: AUTO/ON/OFF")
|
||||
set_property(CACHE FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER PROPERTY STRINGS AUTO ON OFF)
|
||||
|
||||
if(DEFINED FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3)
|
||||
message(DEPRECATION
|
||||
"FASTVIDEO_KERNEL_BUILD_MODIFIED_SAGE3 is deprecated. "
|
||||
"Use FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER instead.")
|
||||
endif()
|
||||
|
||||
# Prefer environment variable (used by CI) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED TORCH_CUDA_ARCH_LIST AND DEFINED ENV{TORCH_CUDA_ARCH_LIST})
|
||||
set(TORCH_CUDA_ARCH_LIST "$ENV{TORCH_CUDA_ARCH_LIST}")
|
||||
@@ -57,6 +72,7 @@ endif()
|
||||
|
||||
message(STATUS "TORCH_CUDA_ARCH_LIST (cmake/env): ${TORCH_CUDA_ARCH_LIST}")
|
||||
message(STATUS "FASTVIDEO_KERNEL_BUILD_TK: ${FASTVIDEO_KERNEL_BUILD_TK}")
|
||||
message(STATUS "FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER: ${FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER}")
|
||||
|
||||
set(ENABLE_TK_KERNELS OFF)
|
||||
if(FASTVIDEO_KERNEL_BUILD_TK STREQUAL "ON")
|
||||
@@ -91,6 +107,54 @@ else()
|
||||
message(STATUS "ThunderKittens kernels: DISABLED (will use Triton fallbacks at runtime)")
|
||||
endif()
|
||||
|
||||
set(ENABLE_ATTN_QAT_INFER OFF)
|
||||
if(GPU_BACKEND STREQUAL "ROCM")
|
||||
message(STATUS "attn_qat_infer kernels: DISABLED (ROCm build)")
|
||||
else()
|
||||
set(_WANTS_ATTN_QAT_INFER OFF)
|
||||
if(FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER STREQUAL "ON")
|
||||
set(_WANTS_ATTN_QAT_INFER ON)
|
||||
elseif(FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER STREQUAL "AUTO")
|
||||
if(TORCH_CUDA_ARCH_LIST)
|
||||
string(REGEX MATCH
|
||||
"(^|[; ,])((12\\.0a)|(120a)|(sm_120a))([; ,]|$)"
|
||||
_HAS_120A "${TORCH_CUDA_ARCH_LIST}")
|
||||
if(_HAS_120A)
|
||||
set(_WANTS_ATTN_QAT_INFER ON)
|
||||
endif()
|
||||
else()
|
||||
execute_process(
|
||||
COMMAND "${Python_EXECUTABLE}" -c
|
||||
"import torch; print('1' if (torch.cuda.is_available() and torch.version.cuda and torch.cuda.get_device_capability()[0] >= 12) else '0')"
|
||||
OUTPUT_VARIABLE _LOCAL_HAS_BLACKWELL
|
||||
OUTPUT_STRIP_TRAILING_WHITESPACE
|
||||
ERROR_QUIET
|
||||
)
|
||||
if(_LOCAL_HAS_BLACKWELL STREQUAL "1")
|
||||
set(_WANTS_ATTN_QAT_INFER ON)
|
||||
endif()
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if(_WANTS_ATTN_QAT_INFER)
|
||||
if(CUDAToolkit_VERSION VERSION_LESS 12.8)
|
||||
message(WARNING
|
||||
"attn_qat_infer kernels require CUDA Toolkit 12.8+. "
|
||||
"Skipping because CUDAToolkit_VERSION=${CUDAToolkit_VERSION}.")
|
||||
else()
|
||||
set(ENABLE_ATTN_QAT_INFER ON)
|
||||
endif()
|
||||
endif()
|
||||
|
||||
if(ENABLE_ATTN_QAT_INFER)
|
||||
message(STATUS "attn_qat_infer kernels: ENABLED")
|
||||
else()
|
||||
message(STATUS
|
||||
"attn_qat_infer kernels: DISABLED "
|
||||
"(requires CUDA 12.8+ and Blackwell sm_120a)")
|
||||
endif()
|
||||
endif()
|
||||
|
||||
# Always try to build the extension if CUDA is available, but conditionally add sources/flags
|
||||
set(BUILD_CXX_KERNELS ON)
|
||||
|
||||
@@ -183,3 +247,74 @@ if(BUILD_CXX_KERNELS)
|
||||
install(TARGETS fastvideo_kernel_ops LIBRARY DESTINATION fastvideo_kernel/_C)
|
||||
endif()
|
||||
|
||||
if(ENABLE_ATTN_QAT_INFER)
|
||||
set(ATTN_QAT_INFER_DIR ${CMAKE_SOURCE_DIR}/attn_qat_infer)
|
||||
set(ATTN_QAT_INFER_INCLUDE_DIRS
|
||||
${ATTN_QAT_INFER_DIR}
|
||||
${CMAKE_SOURCE_DIR}/include/cutlass/include
|
||||
${CMAKE_SOURCE_DIR}/include/cutlass/tools/util/include
|
||||
${TORCH_INCLUDE_DIRS}
|
||||
)
|
||||
set(ATTN_QAT_INFER_CUDA_FLAGS
|
||||
"-O3"
|
||||
"-std=c++17"
|
||||
"-U__CUDA_NO_HALF_OPERATORS__"
|
||||
"-U__CUDA_NO_HALF_CONVERSIONS__"
|
||||
"-U__CUDA_NO_BFLOAT16_OPERATORS__"
|
||||
"-U__CUDA_NO_BFLOAT16_CONVERSIONS__"
|
||||
"-U__CUDA_NO_BFLOAT162_OPERATORS__"
|
||||
"-U__CUDA_NO_BFLOAT162_CONVERSIONS__"
|
||||
"--expt-relaxed-constexpr"
|
||||
"--expt-extended-lambda"
|
||||
"--use_fast_math"
|
||||
"--ptxas-options=--verbose,--warn-on-local-memory-usage"
|
||||
"-lineinfo"
|
||||
"-DCUTLASS_DEBUG_TRACE_LEVEL=0"
|
||||
"-DNDEBUG"
|
||||
"-DQBLKSIZE=128"
|
||||
"-DKBLKSIZE=128"
|
||||
"-DCTA256"
|
||||
"-DDQINRMEM"
|
||||
)
|
||||
|
||||
Python_add_library(fp4attn_cuda MODULE WITH_SOABI
|
||||
attn_qat_infer/blackwell/api.cu
|
||||
)
|
||||
target_include_directories(fp4attn_cuda PRIVATE ${ATTN_QAT_INFER_INCLUDE_DIRS})
|
||||
target_compile_definitions(fp4attn_cuda PRIVATE TORCH_EXTENSION_NAME=fp4attn_cuda)
|
||||
target_compile_options(fp4attn_cuda PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-O3 -std=c++17>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${ATTN_QAT_INFER_CUDA_FLAGS}>
|
||||
)
|
||||
set_target_properties(fp4attn_cuda PROPERTIES
|
||||
CUDA_ARCHITECTURES "120a"
|
||||
CXX_STANDARD 17
|
||||
CUDA_STANDARD 17
|
||||
)
|
||||
target_link_libraries(fp4attn_cuda PRIVATE ${TORCH_LIBRARIES} CUDA::cudart CUDA::cuda_driver)
|
||||
|
||||
Python_add_library(fp4quant_cuda MODULE WITH_SOABI
|
||||
attn_qat_infer/quantization/fp4_quantization_4d.cu
|
||||
)
|
||||
target_include_directories(fp4quant_cuda PRIVATE ${ATTN_QAT_INFER_INCLUDE_DIRS})
|
||||
target_compile_definitions(fp4quant_cuda PRIVATE TORCH_EXTENSION_NAME=fp4quant_cuda)
|
||||
target_compile_options(fp4quant_cuda PRIVATE
|
||||
$<$<COMPILE_LANGUAGE:CXX>:-O3 -std=c++17>
|
||||
$<$<COMPILE_LANGUAGE:CUDA>:${ATTN_QAT_INFER_CUDA_FLAGS}>
|
||||
)
|
||||
set_target_properties(fp4quant_cuda PROPERTIES
|
||||
CUDA_ARCHITECTURES "120a"
|
||||
CXX_STANDARD 17
|
||||
CUDA_STANDARD 17
|
||||
)
|
||||
target_link_libraries(fp4quant_cuda PRIVATE ${TORCH_LIBRARIES} CUDA::cudart CUDA::cuda_driver)
|
||||
|
||||
if(TORCH_PYTHON_LIBRARY_PATH)
|
||||
target_link_libraries(fp4attn_cuda PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
|
||||
target_link_libraries(fp4quant_cuda PRIVATE "${TORCH_PYTHON_LIBRARY_PATH}")
|
||||
endif()
|
||||
|
||||
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
|
||||
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
|
||||
endif()
|
||||
|
||||
|
||||
@@ -2,5 +2,6 @@ include LICENSE
|
||||
include README.md
|
||||
include pyproject.toml
|
||||
recursive-include python/fastvideo_kernel *.py
|
||||
recursive-include attn_qat_infer *.py *.cu *.cuh *.cpp *.h
|
||||
recursive-include csrc *.cu *.cuh *.cpp *.h
|
||||
recursive-include include/tk *.cu *.cuh *.cpp *.h *.src
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
from .api import sageattn_blackwell
|
||||
@@ -0,0 +1,189 @@
|
||||
# Modified from the original SageATtention3 code
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch.nn.functional as F
|
||||
from typing import Tuple
|
||||
from torch.nn.functional import scaled_dot_product_attention as sdpa
|
||||
import fp4attn_cuda
|
||||
import fp4quant_cuda
|
||||
|
||||
# Centralized block size configuration for sageattn_blackwell kernels
|
||||
# These should match the values in fastvideo/attention/backends/sageattn/blackwell/block_config.h
|
||||
BLOCK_M = 128 # Block size for M dimension (query sequence length)
|
||||
BLOCK_N = 128 # Block size for N dimension (key/value sequence length)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def group_mean_kernel(
|
||||
q_ptr,
|
||||
q_out_ptr,
|
||||
qm_out_ptr,
|
||||
B, H, L, D: tl.constexpr,
|
||||
stride_qb, stride_qh, stride_ql, stride_qd,
|
||||
stride_qmb, stride_qmh, stride_qml, stride_qmd,
|
||||
GROUP_SIZE: tl.constexpr
|
||||
):
|
||||
pid_b = tl.program_id(0)
|
||||
pid_h = tl.program_id(1)
|
||||
pid_group = tl.program_id(2)
|
||||
|
||||
group_start = pid_group * GROUP_SIZE
|
||||
offsets = group_start + tl.arange(0, GROUP_SIZE)
|
||||
|
||||
q_offsets = pid_b * stride_qb + pid_h * stride_qh + offsets[:, None] * stride_ql + tl.arange(0, D)[None, :] * stride_qd
|
||||
q_group = tl.load(q_ptr + q_offsets)
|
||||
|
||||
qm_group = tl.sum(q_group, axis=0) / GROUP_SIZE
|
||||
|
||||
q_group = q_group - qm_group
|
||||
tl.store(q_out_ptr + q_offsets, q_group)
|
||||
|
||||
qm_offset = pid_b * stride_qmb + pid_h * stride_qmh + pid_group * stride_qml + tl.arange(0, D) * stride_qmd
|
||||
tl.store(qm_out_ptr + qm_offset, qm_group)
|
||||
|
||||
|
||||
def triton_group_mean(q: torch.Tensor):
|
||||
B, H, L, D = q.shape
|
||||
GROUP_SIZE = BLOCK_M
|
||||
num_groups = L // GROUP_SIZE
|
||||
|
||||
q_out = torch.empty_like(q) # [B, H, L, D]
|
||||
qm = torch.empty(B, H, num_groups, D, device=q.device, dtype=q.dtype)
|
||||
|
||||
grid = (B, H, num_groups)
|
||||
|
||||
group_mean_kernel[grid](
|
||||
q, q_out, qm,
|
||||
B, H, L, D,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
qm.stride(0), qm.stride(1), qm.stride(2), qm.stride(3),
|
||||
GROUP_SIZE=GROUP_SIZE
|
||||
)
|
||||
return q_out, qm
|
||||
|
||||
|
||||
def preprocess_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, per_block_mean: bool = True, enable_smoothing_q: bool = False, enable_smoothing_k: bool = False):
|
||||
|
||||
def pad_to_block_size(x):
|
||||
L = x.size(2)
|
||||
pad_len = (BLOCK_M - L % BLOCK_M) % BLOCK_M
|
||||
if pad_len == 0:
|
||||
return x.contiguous()
|
||||
return F.pad(x, (0, 0, 0, pad_len), value=0).contiguous()
|
||||
|
||||
if enable_smoothing_k:
|
||||
k -= k.mean(dim=-2, keepdim=True)
|
||||
q, k, v = map(lambda x: pad_to_block_size(x), [q, k, v])
|
||||
if per_block_mean and enable_smoothing_q:
|
||||
q, qm = triton_group_mean(q)
|
||||
elif enable_smoothing_q:
|
||||
qm = q.mean(dim=-2, keepdim=True)
|
||||
q = q - qm
|
||||
if enable_smoothing_q:
|
||||
delta_s = torch.matmul(qm, k.transpose(-2, -1)).to(torch.float32).contiguous()
|
||||
else: # used to disable q smoothing
|
||||
B, H, L, D = q.shape
|
||||
delta_s = torch.zeros((B, H, L // BLOCK_M, k.shape[2]), device=q.device, dtype=torch.float32)
|
||||
|
||||
return q, k, v, delta_s
|
||||
|
||||
def scale_and_quant_fp4(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.ndim == 4
|
||||
B, H, N, D = x.shape
|
||||
packed_fp4 = torch.empty((B, H, N, D // 2), device=x.device, dtype=torch.uint8)
|
||||
fp8_scale = torch.empty((B, H, N, D // 16), device=x.device, dtype=torch.float8_e4m3fn)
|
||||
fp4quant_cuda.scaled_fp4_quant(x, packed_fp4, fp8_scale, 1)
|
||||
return packed_fp4, fp8_scale
|
||||
|
||||
def scale_and_quant_fp4_permute(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.ndim == 4
|
||||
B, H, N, D = x.shape
|
||||
packed_fp4 = torch.empty((B, H, N, D // 2), device=x.device, dtype=torch.uint8)
|
||||
fp8_scale = torch.empty((B, H, N, D // 16), device=x.device, dtype=torch.float8_e4m3fn)
|
||||
fp4quant_cuda.scaled_fp4_quant_permute(x, packed_fp4, fp8_scale, 1)
|
||||
return packed_fp4, fp8_scale
|
||||
|
||||
def scale_and_quant_fp4_transpose(x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
assert x.ndim == 4
|
||||
B, H, N, D = x.shape
|
||||
packed_fp4 = torch.empty((B, H, D, N // 2), device=x.device, dtype=torch.uint8)
|
||||
fp8_scale = torch.empty((B, H, D, N // 16), device=x.device, dtype=torch.float8_e4m3fn)
|
||||
fp4quant_cuda.scaled_fp4_quant_trans(x, packed_fp4, fp8_scale, 1)
|
||||
return packed_fp4, fp8_scale
|
||||
|
||||
def blockscaled_fp4_attn(qlist: Tuple,
|
||||
klist: Tuple,
|
||||
vlist: Tuple,
|
||||
delta_s: torch.Tensor,
|
||||
KL: int,
|
||||
is_causal: bool = False,
|
||||
per_block_mean: bool = True,
|
||||
is_bf16: bool = True,
|
||||
single_level_p_quant: bool = False,
|
||||
sm_scale: float | None = None
|
||||
):
|
||||
softmax_scale = sm_scale if sm_scale is not None else (qlist[0].shape[-1] * 2) ** (-0.5)
|
||||
return fp4attn_cuda.fwd(qlist[0], klist[0], vlist[0], qlist[1], klist[1], vlist[1], delta_s, KL, None, softmax_scale, is_causal, per_block_mean, is_bf16, single_level_p_quant)
|
||||
|
||||
|
||||
def sageattn_blackwell(q, k, v, attn_mask = None, is_causal = False, per_block_mean = True, single_level_p_quant = True, sm_scale: float | None = None, **kwargs):
|
||||
"""
|
||||
SageAttention3 Blackwell kernel for FP4 attention.
|
||||
|
||||
Args:
|
||||
q: Query tensor [B, H, L, D]
|
||||
k: Key tensor [B, H, L, D]
|
||||
v: Value tensor [B, H, L, D]
|
||||
attn_mask: Attention mask (not used)
|
||||
is_causal: Whether to use causal masking
|
||||
per_block_mean: Whether to use per-block mean for Q smoothing
|
||||
single_level_p_quant: If True, use single-level quantization: s_P2, P̂_2 = φ(P̃) directly
|
||||
(standard per-block FP4 quantization like V, no s_P1).
|
||||
If False (default), use two-level quantization:
|
||||
s_P1 = rowmax(P̃)/(448×6), then s_P2, P̂_2 = φ(P̃/s_P1).
|
||||
sm_scale: Softmax scale to pass through to the CUDA kernel. If None,
|
||||
defaults to the kernel's 1/sqrt(D) scale.
|
||||
**kwargs: Additional arguments (ignored)
|
||||
|
||||
Returns:
|
||||
Output tensor [B, H, L, D]
|
||||
"""
|
||||
if q.size(-1) >= 256:
|
||||
print(f"Unsupported Headdim {q.size(-1)}")
|
||||
return sdpa(q, k, v, is_causal = is_causal)
|
||||
QL = q.size(2)
|
||||
KL = k.size(2)
|
||||
is_bf16 = q.dtype == torch.bfloat16
|
||||
q, k, v, delta_s = preprocess_qkv(q, k, v, per_block_mean)
|
||||
qlist_from_cuda = scale_and_quant_fp4(q)
|
||||
klist_from_cuda = scale_and_quant_fp4_permute(k)
|
||||
vlist_from_cuda = scale_and_quant_fp4_transpose(v)
|
||||
o_fp4 = blockscaled_fp4_attn(
|
||||
qlist_from_cuda,
|
||||
klist_from_cuda,
|
||||
vlist_from_cuda,
|
||||
delta_s,
|
||||
KL,
|
||||
is_causal,
|
||||
per_block_mean,
|
||||
is_bf16,
|
||||
single_level_p_quant,
|
||||
sm_scale
|
||||
)[0][:, :, :QL, :].contiguous()
|
||||
return o_fp4
|
||||
@@ -0,0 +1 @@
|
||||
__version__ = "3.0.0.b1"
|
||||
@@ -0,0 +1,347 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
// Include these 2 headers instead of torch/extension.h since we don't need all of the torch headers.
|
||||
#include <torch/python.h>
|
||||
#include <torch/nn/functional.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
#include "params.h"
|
||||
#include "launch.h"
|
||||
#include "static_switch.h"
|
||||
#include "block_config.h"
|
||||
|
||||
#define CHECK_DEVICE(x) TORCH_CHECK(x.is_cuda(), #x " must be on CUDA")
|
||||
#define CHECK_SHAPE(x, ...) TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), #x " must have shape (" #__VA_ARGS__ ")")
|
||||
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
|
||||
|
||||
|
||||
void set_params_fprop(Flash_fwd_params ¶ms,
|
||||
// sizes
|
||||
const size_t b,
|
||||
const size_t seqlen_q,
|
||||
const size_t seqlen_k,
|
||||
const size_t unpadded_seqlen_k,
|
||||
const size_t seqlen_q_rounded,
|
||||
const size_t seqlen_k_rounded,
|
||||
const size_t h,
|
||||
const size_t h_k,
|
||||
const size_t d,
|
||||
const size_t d_rounded,
|
||||
// device pointers
|
||||
const at::Tensor q,
|
||||
const at::Tensor k,
|
||||
const at::Tensor v,
|
||||
const at::Tensor delta_s,
|
||||
at::Tensor out,
|
||||
const at::Tensor sfq,
|
||||
const at::Tensor sfk,
|
||||
const at::Tensor sfv,
|
||||
void *cu_seqlens_q_d,
|
||||
void *cu_seqlens_k_d,
|
||||
void *seqused_k,
|
||||
void *p_d,
|
||||
void *softmax_lse_d,
|
||||
float p_dropout,
|
||||
float softmax_scale,
|
||||
int window_size_left,
|
||||
int window_size_right,
|
||||
bool per_block_mean,
|
||||
bool is_bf16,
|
||||
bool single_level_p_quant=false,
|
||||
bool seqlenq_ngroups_swapped=false) {
|
||||
|
||||
// Reset the parameters
|
||||
params = {};
|
||||
// Set the pointers and strides.
|
||||
params.q_ptr = q.data_ptr();
|
||||
params.k_ptr = k.data_ptr();
|
||||
params.v_ptr = v.data_ptr();
|
||||
params.delta_s_ptr = delta_s.data_ptr();
|
||||
params.sfq_ptr = sfq.data_ptr();
|
||||
params.sfk_ptr = sfk.data_ptr();
|
||||
params.sfv_ptr = sfv.data_ptr();
|
||||
|
||||
// All stride are in elements, not bytes.
|
||||
params.q_row_stride = q.stride(-2) * 2;
|
||||
params.k_row_stride = k.stride(-2) * 2;
|
||||
params.v_row_stride = v.stride(-2) * 2;;
|
||||
params.q_head_stride = q.stride(-3) * 2;
|
||||
params.k_head_stride = k.stride(-3) * 2;
|
||||
params.v_head_stride = v.stride(-3) * 2; // for packed q k v
|
||||
|
||||
params.ds_row_stride = delta_s.stride(-2);
|
||||
params.ds_head_stride = delta_s.stride(-3);
|
||||
|
||||
params.sfq_row_stride = sfq.stride(-2);
|
||||
params.sfk_row_stride = sfk.stride(-2);
|
||||
params.sfv_row_stride = sfv.stride(-2);
|
||||
params.sfq_head_stride = sfq.stride(-3);
|
||||
params.sfk_head_stride = sfk.stride(-3);
|
||||
params.sfv_head_stride = sfv.stride(-3);
|
||||
params.o_ptr = out.data_ptr();
|
||||
params.o_row_stride = out.stride(-2);
|
||||
params.o_head_stride = out.stride(-3);
|
||||
|
||||
if (cu_seqlens_q_d == nullptr) {
|
||||
params.q_batch_stride = q.stride(0) * 2;
|
||||
params.k_batch_stride = k.stride(0) * 2;
|
||||
params.v_batch_stride = v.stride(0) * 2;
|
||||
params.ds_batch_stride = delta_s.stride(0);
|
||||
params.sfq_batch_stride = sfq.stride(0);
|
||||
params.sfk_batch_stride = sfk.stride(0);
|
||||
params.sfv_batch_stride = sfv.stride(0);
|
||||
params.o_batch_stride = out.stride(0);
|
||||
if (seqlenq_ngroups_swapped) {
|
||||
params.q_batch_stride *= seqlen_q;
|
||||
params.o_batch_stride *= seqlen_q;
|
||||
}
|
||||
}
|
||||
|
||||
params.cu_seqlens_q = static_cast<int *>(cu_seqlens_q_d);
|
||||
params.cu_seqlens_k = static_cast<int *>(cu_seqlens_k_d);
|
||||
params.seqused_k = static_cast<int *>(seqused_k);
|
||||
|
||||
// P = softmax(QK^T)
|
||||
params.p_ptr = p_d;
|
||||
|
||||
// Softmax sum
|
||||
params.softmax_lse_ptr = softmax_lse_d;
|
||||
|
||||
// Set the dimensions.
|
||||
params.b = b;
|
||||
params.h = h;
|
||||
params.h_k = h_k;
|
||||
params.h_h_k_ratio = h / h_k;
|
||||
params.seqlen_q = seqlen_q;
|
||||
params.seqlen_k = seqlen_k;
|
||||
params.unpadded_seqlen_k = unpadded_seqlen_k;
|
||||
params.seqlen_q_rounded = seqlen_q_rounded;
|
||||
params.seqlen_k_rounded = seqlen_k_rounded;
|
||||
params.d = d;
|
||||
params.d_rounded = d_rounded;
|
||||
|
||||
params.head_divmod = cutlass::FastDivmod(int(h));
|
||||
|
||||
// Set the different scale values.
|
||||
params.scale_softmax = softmax_scale;
|
||||
params.scale_softmax_log2 = softmax_scale * M_LOG2E;
|
||||
__half scale_softmax_log2_half = __float2half(params.scale_softmax_log2);
|
||||
__half2 scale_softmax_log2_half2 = __half2(scale_softmax_log2_half, scale_softmax_log2_half);
|
||||
params.scale_softmax_log2_half2 = reinterpret_cast<uint32_t&>(scale_softmax_log2_half2);
|
||||
|
||||
// Set this to probability of keeping an element to simplify things.
|
||||
params.p_dropout = 1.f - p_dropout;
|
||||
// Convert p from float to int so we don't have to convert the random uint to float to compare.
|
||||
// [Minor] We want to round down since when we do the comparison we use <= instead of <
|
||||
// params.p_dropout_in_uint = uint32_t(std::floor(params.p_dropout * 4294967295.0));
|
||||
// params.p_dropout_in_uint16_t = uint16_t(std::floor(params.p_dropout * 65535.0));
|
||||
params.p_dropout_in_uint8_t = uint8_t(std::floor(params.p_dropout * 255.0));
|
||||
params.rp_dropout = 1.f / params.p_dropout;
|
||||
params.scale_softmax_rp_dropout = params.rp_dropout * params.scale_softmax;
|
||||
TORCH_CHECK(p_dropout < 1.f);
|
||||
#ifdef FLASHATTENTION_DISABLE_DROPOUT
|
||||
TORCH_CHECK(p_dropout == 0.0f, "This flash attention build does not support dropout.");
|
||||
#endif
|
||||
|
||||
// Causal is the special case where window_size_right == 0 and window_size_left < 0.
|
||||
// Local is the more general case where window_size_right >= 0 or window_size_left >= 0.
|
||||
params.is_causal = window_size_left < 0 && window_size_right == 0;
|
||||
params.per_block_mean = per_block_mean;
|
||||
if (per_block_mean) {
|
||||
params.seqlen_s = seqlen_q;
|
||||
} else {
|
||||
params.seqlen_s = flash::BLOCK_M; // size of BLOCK_M
|
||||
}
|
||||
if (window_size_left < 0 && window_size_right >= 0) { window_size_left = seqlen_k; }
|
||||
if (window_size_left >= 0 && window_size_right < 0) { window_size_right = seqlen_k; }
|
||||
params.window_size_left = window_size_left;
|
||||
params.window_size_right = window_size_right;
|
||||
|
||||
#ifdef FLASHATTENTION_DISABLE_LOCAL
|
||||
TORCH_CHECK(params.is_causal || (window_size_left < 0 && window_size_right < 0),
|
||||
"This flash attention build does not support local attention.");
|
||||
#endif
|
||||
|
||||
params.is_seqlens_k_cumulative = true;
|
||||
params.is_bf16 = is_bf16;
|
||||
params.single_level_p_quant = single_level_p_quant;
|
||||
#ifdef FLASHATTENTION_DISABLE_UNEVEN_K
|
||||
TORCH_CHECK(d == d_rounded, "This flash attention build does not support headdim not being a multiple of 32.");
|
||||
#endif
|
||||
}
|
||||
|
||||
template<bool IsBF16>
|
||||
void run_mha_fwd_dispatch_dtype(Flash_fwd_params ¶ms, cudaStream_t stream) {
|
||||
using OType = std::conditional_t<IsBF16, cutlass::bfloat16_t, cutlass::half_t>;
|
||||
if (params.d == 64) {
|
||||
run_mha_fwd_<cutlass::nv_float4_t<cutlass::float_e2m1_t>, 64, OType>(params, stream);
|
||||
} else if (params.d == 128) {
|
||||
run_mha_fwd_<cutlass::nv_float4_t<cutlass::float_e2m1_t>, 128, OType>(params, stream);
|
||||
}
|
||||
}
|
||||
|
||||
void run_mha_fwd(Flash_fwd_params ¶ms, cudaStream_t stream, bool force_split_kernel = false) {
|
||||
BOOL_SWITCH(params.is_bf16, IsBF16, ([&] {
|
||||
run_mha_fwd_dispatch_dtype<IsBF16>(params, stream);
|
||||
}));
|
||||
}
|
||||
|
||||
std::vector<at::Tensor>
|
||||
mha_fwd(at::Tensor &q, // batch_size x seqlen_q x num_heads x (head_size // 2)
|
||||
const at::Tensor &k, // batch_size x seqlen_k x num_heads_k x (head_size // 2)
|
||||
const at::Tensor &v, // batch_size x seqlen_k x num_heads_k x (head_size // 2)
|
||||
const at::Tensor &sfq,
|
||||
const at::Tensor &sfk,
|
||||
const at::Tensor &sfv,
|
||||
const at::Tensor &delta_s,
|
||||
int unpadded_k,
|
||||
c10::optional<at::Tensor> &out_, // batch_size x seqlen_q x num_heads x head_size
|
||||
const float softmax_scale,
|
||||
bool is_causal,
|
||||
bool per_block_mean,
|
||||
bool is_bf16,
|
||||
bool single_level_p_quant=false // If true, use only per-row scale s_P2 (no per-block s_P1)
|
||||
) {
|
||||
|
||||
auto dprops = at::cuda::getCurrentDeviceProperties();
|
||||
bool is_blackwell_or_newer = dprops->major >= 12;
|
||||
TORCH_CHECK(is_blackwell_or_newer, "only supports Blackwell GPUs or newer.");
|
||||
|
||||
auto q_dtype = q.dtype();
|
||||
auto sfq_dtype = sfq.dtype();
|
||||
TORCH_CHECK(q_dtype == torch::kUInt8, "q dtype must be uint8");
|
||||
TORCH_CHECK(k.dtype() == q_dtype, "query and key must have the same dtype");
|
||||
TORCH_CHECK(v.dtype() == q_dtype, "query and value must have the same dtype");
|
||||
CHECK_DEVICE(q); CHECK_DEVICE(k); CHECK_DEVICE(v);
|
||||
|
||||
TORCH_CHECK(sfq_dtype == torch::kFloat8_e4m3fn, "q dtype must be uint8");
|
||||
TORCH_CHECK(sfk.dtype() == sfq_dtype, "query and key must have the same dtype");
|
||||
TORCH_CHECK(sfv.dtype() == sfq_dtype, "query and value must have the same dtype");
|
||||
CHECK_DEVICE(sfq); CHECK_DEVICE(sfk); CHECK_DEVICE(sfv);
|
||||
|
||||
TORCH_CHECK(q.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
TORCH_CHECK(k.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
TORCH_CHECK(v.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
TORCH_CHECK(delta_s.stride(-1) == 1, "Input tensor must have contiguous last dimension");
|
||||
|
||||
TORCH_CHECK(q.is_contiguous(), "Input tensor must be contiguous");
|
||||
TORCH_CHECK(k.is_contiguous(), "Input tensor must be contiguous");
|
||||
TORCH_CHECK(v.is_contiguous(), "Input tensor must be contiguous");
|
||||
|
||||
const auto sizes = q.sizes();
|
||||
auto opts = q.options();
|
||||
const int batch_size = sizes[0];
|
||||
int seqlen_q = sizes[2];
|
||||
int num_heads = sizes[1];
|
||||
const int head_size_og = sizes[3];
|
||||
const int unpacked_head_size = head_size_og * 2;
|
||||
const int seqlen_k = k.size(2);
|
||||
const int num_heads_k = k.size(1);
|
||||
|
||||
TORCH_CHECK(batch_size > 0, "batch size must be postive");
|
||||
TORCH_CHECK(unpacked_head_size <= 256, "FlashAttention forward only supports head dimension at most 256");
|
||||
TORCH_CHECK(num_heads % num_heads_k == 0, "Number of heads in key/value must divide number of heads in query");
|
||||
TORCH_CHECK(num_heads == num_heads_k, "We do not support MQA/GQA yet");
|
||||
|
||||
TORCH_CHECK(unpacked_head_size == 64 || unpacked_head_size == 128 || unpacked_head_size == 256, "Only support head size 64, 128, and 256 for now");
|
||||
|
||||
CHECK_SHAPE(q, batch_size, num_heads, seqlen_q, head_size_og);
|
||||
CHECK_SHAPE(k, batch_size, num_heads_k, seqlen_k, head_size_og);
|
||||
CHECK_SHAPE(v, batch_size, num_heads_k, unpacked_head_size, seqlen_k/2);
|
||||
// CHECK_SHAPE(delta_s, batch_size, num_heads, seqlen_q / 128, seqlen_k);
|
||||
// CHECK_SHAPE(sfq, batch_size, seqlen_q, num_heads, unpacked_head_size);
|
||||
// CHECK_SHAPE(sfk, batch_size, seqlen_k, num_heads_k, unpacked_head_size);
|
||||
// CHECK_SHAPE(sfv, batch_size, unpacked_head_size, num_heads_k, seqlen_k);
|
||||
TORCH_CHECK(unpacked_head_size % 8 == 0, "head_size must be a multiple of 8");
|
||||
|
||||
auto dtype = is_bf16 ? at::ScalarType::BFloat16 : at::ScalarType::Half;
|
||||
at::Tensor out = torch::empty({batch_size, num_heads, seqlen_q, unpacked_head_size}, opts.dtype(dtype));
|
||||
|
||||
auto round_multiple = [](int x, int m) { return (x + m - 1) / m * m; };
|
||||
// const int head_size = round_multiple(head_size_og, 8);
|
||||
// const int head_size_rounded = round_multiple(head_size, 32);
|
||||
const int seqlen_q_rounded = round_multiple(seqlen_q, flash::BLOCK_M);
|
||||
const int seqlen_k_rounded = round_multiple(seqlen_k, flash::BLOCK_N);
|
||||
|
||||
// Otherwise the kernel will be launched from cuda:0 device
|
||||
// Cast to char to avoid compiler warning about narrowing
|
||||
at::cuda::CUDAGuard device_guard{(char)q.get_device()};
|
||||
|
||||
|
||||
|
||||
auto softmax_lse = torch::empty({batch_size, num_heads, seqlen_q}, opts.dtype(at::kFloat));
|
||||
at::Tensor p;
|
||||
|
||||
Flash_fwd_params params;
|
||||
set_params_fprop(params,
|
||||
batch_size,
|
||||
seqlen_q, seqlen_k, unpadded_k,
|
||||
seqlen_q_rounded, seqlen_k_rounded,
|
||||
num_heads, num_heads_k,
|
||||
unpacked_head_size, unpacked_head_size,
|
||||
q, k, v, delta_s, out,
|
||||
sfq, sfk, sfv,
|
||||
/*cu_seqlens_q_d=*/nullptr,
|
||||
/*cu_seqlens_k_d=*/nullptr,
|
||||
/*seqused_k=*/nullptr,
|
||||
nullptr,
|
||||
softmax_lse.data_ptr(),
|
||||
/*p_dropout=*/0.f,
|
||||
softmax_scale,
|
||||
/*window_size_left=*/-1,
|
||||
/*window_size_right=*/is_causal ? 0 : -1,
|
||||
per_block_mean,
|
||||
is_bf16,
|
||||
single_level_p_quant
|
||||
);
|
||||
// StaticPersistentTileScheduler does not use tile_count_semaphore; avoid a
|
||||
// stack-local tensor whose data pointer would dangle after mha_fwd returns
|
||||
// while the async kernel may still be running.
|
||||
params.tile_count_semaphore = nullptr;
|
||||
|
||||
if (seqlen_k > 0) {
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
run_mha_fwd(params, stream);
|
||||
} else {
|
||||
// If seqlen_k == 0, then we have an empty tensor. We need to set the output to 0.
|
||||
out.zero_();
|
||||
softmax_lse.fill_(std::numeric_limits<float>::infinity());
|
||||
}
|
||||
|
||||
// at::Tensor out_padded = out;
|
||||
// if (head_size_og % 8 != 0) {
|
||||
// out = out.index({"...", torch::indexing::Slice(torch::indexing::None, head_size_og)});
|
||||
// if (out_.has_value()) { out_.value().copy_(out); }
|
||||
// }
|
||||
|
||||
// return {out, q_padded, k_padded, v_padded, out_padded, softmax_lse, p};
|
||||
// cudaDeviceSynchronize();
|
||||
// auto err = cudaGetLastError();
|
||||
// printf("%s\n", cudaGetErrorString(err));
|
||||
return {out, softmax_lse};
|
||||
}
|
||||
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "FlashAttention";
|
||||
m.def("fwd", &mha_fwd, "Forward pass");
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
// Centralized block size configuration for sageattn_blackwell kernels
|
||||
// Block sizes for M and N dimensions
|
||||
namespace flash {
|
||||
// Block size for M dimension (query sequence length)
|
||||
static constexpr int BLOCK_M = 128;
|
||||
|
||||
// Block size for N dimension (key/value sequence length)
|
||||
static constexpr int BLOCK_N = 128;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* This code is based on code from FlashAttention3, https://github.com/Dao-AILab/flash-attention
|
||||
* Copyright (c) 2024, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
namespace flash {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<bool Varlen=true>
|
||||
struct BlockInfo {
|
||||
|
||||
template<typename Params>
|
||||
__device__ BlockInfo(const Params ¶ms, const int bidb)
|
||||
: sum_s_q(!Varlen || params.cu_seqlens_q == nullptr ? -1 : params.cu_seqlens_q[bidb])
|
||||
, sum_s_k(!Varlen || params.cu_seqlens_k == nullptr || !params.is_seqlens_k_cumulative ? -1 : params.cu_seqlens_k[bidb])
|
||||
, actual_seqlen_q(!Varlen || params.cu_seqlens_q == nullptr ? params.seqlen_q : params.cu_seqlens_q[bidb + 1] - sum_s_q)
|
||||
// If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb].
|
||||
// Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K.
|
||||
, seqlen_k_cache(!Varlen || params.cu_seqlens_k == nullptr ? params.seqlen_k : (params.is_seqlens_k_cumulative ? params.cu_seqlens_k[bidb + 1] - sum_s_k : params.cu_seqlens_k[bidb]))
|
||||
, actual_seqlen_k(params.seqused_k ? params.seqused_k[bidb] : seqlen_k_cache + (params.knew_ptr == nullptr ? 0 : params.seqlen_knew))
|
||||
{
|
||||
}
|
||||
|
||||
template <typename index_t>
|
||||
__forceinline__ __device__ index_t q_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const {
|
||||
return sum_s_q == -1 ? bidb * batch_stride : uint32_t(sum_s_q) * row_stride;
|
||||
}
|
||||
|
||||
template <typename index_t>
|
||||
__forceinline__ __device__ index_t k_offset(const index_t batch_stride, const index_t row_stride, const int bidb) const {
|
||||
return sum_s_k == -1 ? bidb * batch_stride : uint32_t(sum_s_k) * row_stride;
|
||||
}
|
||||
|
||||
const int sum_s_q;
|
||||
const int sum_s_k;
|
||||
const int actual_seqlen_q;
|
||||
// We have to have seqlen_k_cache declared before actual_seqlen_k, otherwise actual_seqlen_k is set to 0.
|
||||
const int seqlen_k_cache;
|
||||
const int actual_seqlen_k;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,149 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Blocked Scale configs specific for SM100 BlockScaled MMA
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/layout/matrix.h"
|
||||
|
||||
#include "cute/int_tuple.hpp"
|
||||
#include "cute/atom/mma_traits_sm100.hpp"
|
||||
|
||||
namespace flash {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
using namespace cute;
|
||||
|
||||
template<int SFVecSize, UMMA::Major major = UMMA::Major::K>
|
||||
struct BlockScaledBasicChunk {
|
||||
|
||||
using Blk_MN = _64;
|
||||
using Blk_SF = _4;
|
||||
|
||||
using SfAtom = Layout< Shape< Shape<_16,_4>, Shape<Int<SFVecSize>, _4>>,
|
||||
Stride<Stride<_16,_4>, Stride< _0, _1>>>;
|
||||
};
|
||||
|
||||
template<int SFVecSize_>
|
||||
struct BlockScaledConfig {
|
||||
// We are creating the SFA and SFB tensors' layouts in the collective since they always have the same layout.
|
||||
// k-major order
|
||||
static constexpr int SFVecSize = SFVecSize_;
|
||||
static constexpr int MMA_NSF = 4; // SFVecSize, MMA_NSF
|
||||
using BlkScaledChunk = BlockScaledBasicChunk<SFVecSize>;
|
||||
using Blk_MN = _64;
|
||||
using Blk_SF = _4;
|
||||
using mnBasicBlockShape = Shape<_16,_4>;
|
||||
using mnBasicBlockStride = Stride<_16,_4>;
|
||||
using kBasicBlockShape = Shape<Int<SFVecSize>, Int<MMA_NSF>>; // SFVecSize, MMA_NSF
|
||||
using kBasicBlockStride = Stride<_0, _1>;
|
||||
using SfAtom = Layout< Shape< mnBasicBlockShape, kBasicBlockShape>,
|
||||
Stride<mnBasicBlockStride, kBasicBlockStride>>;
|
||||
|
||||
using LayoutSF = decltype(blocked_product(SfAtom{},
|
||||
make_layout(
|
||||
make_shape(int32_t(0), int32_t(0), int32_t(0), int32_t(0)),
|
||||
make_stride(int32_t(0), _1{}, int32_t(0), int32_t(0)))));
|
||||
// A single indivisible block will hold 4 scale factors of 64 rows/columns (A/B matrix).
|
||||
// 4 is chosen to make consecutive 32bits of data to have scale factors for only a single row (col). 32bits corresponds to the TMEM word size
|
||||
using Blk_Elems = decltype(Blk_MN{} * Blk_SF{});
|
||||
using sSF_strideMN = decltype(prepend(Blk_Elems{}, mnBasicBlockStride{}));
|
||||
|
||||
|
||||
// The following function is provided for user fill dynamic problem size to the layout_SFA.
|
||||
template < class ProblemShape>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
tile_atom_to_shape_SFQKV(ProblemShape problem_shape) {
|
||||
auto [Seqlen, Dim, HeadNum, Batch] = problem_shape;
|
||||
return tile_to_shape(SfAtom{}, make_shape(Seqlen, Dim, HeadNum, Batch), Step<_2,_1,_3,_4>{});
|
||||
}
|
||||
|
||||
// The following function is provided for user fill dynamic problem size to the layout_SFB.
|
||||
template <class ProblemShape>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
tile_atom_to_shape_SFVt(ProblemShape problem_shape) {
|
||||
auto [Dim, Seqlen, HeadNum, Batch] = problem_shape;
|
||||
return tile_to_shape(SfAtom{}, make_shape(Dim, Seqlen, HeadNum, Batch), Step<_2,_1,_3,_4>{});
|
||||
}
|
||||
|
||||
template<class TiledMma, class TileShape_MNK>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
deduce_smem_layoutSFQ(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) {
|
||||
|
||||
using sSFQ_shapeK = decltype(prepend(make_shape(Blk_SF{}/Int<MMA_NSF>{}, size<2>(TileShape_MNK{}) / Int<SFVecSize>{} / Blk_SF{}), kBasicBlockShape{}));
|
||||
using sSFQ_shapeM = decltype(prepend(size<0>(TileShape_MNK{}) / Blk_MN{}, mnBasicBlockShape{}));
|
||||
using sSFQ_strideM = sSF_strideMN;
|
||||
using sSFQ_strideK = decltype(prepend(make_stride(Int<MMA_NSF>{}, size<0>(TileShape_MNK{}) / Blk_MN{} * Blk_Elems{}), kBasicBlockStride{}));
|
||||
using sSFQ_shape = decltype(make_shape(sSFQ_shapeM{}, sSFQ_shapeK{}));
|
||||
using sSFQ_stride = decltype(make_stride(sSFQ_strideM{}, sSFQ_strideK{}));
|
||||
using SmemLayoutAtomSFQ = decltype(make_layout(sSFQ_shape{}, sSFQ_stride{}));
|
||||
return SmemLayoutAtomSFQ{};
|
||||
}
|
||||
|
||||
template<class TiledMma, class TileShape_MNK>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
deduce_smem_layoutSFKV(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) {
|
||||
|
||||
using sSFK_shapeK = decltype(prepend(make_shape(Blk_SF{}/Int<MMA_NSF>{}, size<2>(TileShape_MNK{}) / Int<SFVecSize>{} / Blk_SF{}), kBasicBlockShape{}));
|
||||
using sSFK_shapeN = decltype(prepend(size<1>(TileShape_MNK{}) / Blk_MN{}, mnBasicBlockShape{}));
|
||||
using sSFK_strideN = sSF_strideMN;
|
||||
using sSFK_strideK = decltype(prepend(make_stride(Int<MMA_NSF>{}, size<1>(TileShape_MNK{}) / Blk_MN{} * Blk_Elems{}), kBasicBlockStride{}));
|
||||
using sSFK_shape = decltype(make_shape(sSFK_shapeN{}, sSFK_shapeK{}));
|
||||
using sSFK_stride = decltype(make_stride(sSFK_strideN{}, sSFK_strideK{}));
|
||||
using SmemLayoutAtomSFK = decltype(make_layout(sSFK_shape{}, sSFK_stride{}));
|
||||
return SmemLayoutAtomSFK{};
|
||||
}
|
||||
|
||||
template<class TiledMma, class TileShape_MNK>
|
||||
CUTE_HOST_DEVICE
|
||||
static constexpr auto
|
||||
deduce_smem_layoutSFVt(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) {
|
||||
|
||||
using sSFVt_shapeK = decltype(prepend(make_shape(Blk_SF{}/Int<MMA_NSF>{}, size<2>(TileShape_MNK{}) / Int<SFVecSize>{} / Blk_SF{}), kBasicBlockShape{}));
|
||||
using sSFVt_shapeN = decltype(prepend(size<1>(TileShape_MNK{}) / Blk_MN{}, mnBasicBlockShape{}));
|
||||
using sSFVt_strideN = sSF_strideMN;
|
||||
using sSFVt_strideK = decltype(prepend(make_stride(Int<MMA_NSF>{}, size<1>(TileShape_MNK{}) / Blk_MN{} * Blk_Elems{}), kBasicBlockStride{}));
|
||||
using sSFVt_shape = decltype(make_shape(sSFVt_shapeN{}, sSFVt_shapeK{}));
|
||||
using sSFVt_stride = decltype(make_stride(sSFVt_strideN{}, sSFVt_strideK{}));
|
||||
using SmemLayoutAtomSFVt = decltype(make_layout(sSFVt_shape{}, sSFVt_stride{}));
|
||||
return SmemLayoutAtomSFVt{};
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,327 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#include "cute/arch/mma_sm120.hpp"
|
||||
#include "cute/atom/mma_traits_sm120.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/float8.h"
|
||||
#include "cutlass/float_subbyte.h"
|
||||
|
||||
namespace cute::SM120::BLOCKSCALED {
|
||||
|
||||
using cutlass::float_e2m1_t;
|
||||
using cutlass::float_ue4m3_t;
|
||||
|
||||
// MMA.SF 16x32x64 TN E2M1 x E2M1 with SF E4M3
|
||||
struct SM120_16x32x64_TN_VS_NVFP4 {
|
||||
using DRegisters = float[16];
|
||||
using ARegisters = uint32_t[4];
|
||||
using BRegisters = uint32_t[8];
|
||||
using CRegisters = float[16];
|
||||
|
||||
static constexpr int SFBits = 32;
|
||||
using RegTypeSF = cute::uint_bit_t<SFBits>;
|
||||
|
||||
using SFARegisters = RegTypeSF[1];
|
||||
using SFBRegisters = RegTypeSF[1];
|
||||
|
||||
CUTE_HOST_DEVICE static void
|
||||
fma(float & d0 , float & d1 , float & d2 , float & d3 ,
|
||||
float & d4 , float & d5 , float & d6 , float & d7 ,
|
||||
float & d8 , float & d9 , float & d10, float & d11,
|
||||
float & d12, float & d13, float & d14, float & d15,
|
||||
uint32_t const& a0 , uint32_t const& a1 , uint32_t const& a2 , uint32_t const& a3 ,
|
||||
uint32_t const& b0 , uint32_t const& b1 , uint32_t const& b2 , uint32_t const& b3 ,
|
||||
uint32_t const& b4 , uint32_t const& b5 , uint32_t const& b6 , uint32_t const& b7 ,
|
||||
float const & c0 , float const & c1 , float const & c2 , float const & c3 ,
|
||||
float const & c4 , float const & c5 , float const & c6 , float const & c7 ,
|
||||
float const & c8 , float const & c9 , float const & c10 , float const & c11,
|
||||
float const & c12, float const & c13, float const & c14, float const & c15,
|
||||
RegTypeSF const& sfa0,
|
||||
RegTypeSF const& sfb0)
|
||||
{
|
||||
static constexpr uint16_t tidA = 0;
|
||||
static constexpr uint16_t bidA = 0;
|
||||
static constexpr uint16_t bidB = 0;
|
||||
static constexpr uint16_t tidB0 = 0;
|
||||
static constexpr uint16_t tidB1 = 1;
|
||||
static constexpr uint16_t tidB2 = 2;
|
||||
static constexpr uint16_t tidB3 = 3;
|
||||
|
||||
#if defined(CUTE_ARCH_MXF4NVF4_4X_UE4M3_MMA_ENABLED)
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d0), "=f"(d1), "=f"(d8), "=f"(d9)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b0), "r"(b1),
|
||||
"f"(c0), "f"(c1), "f"(c8), "f"(c9),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB0));
|
||||
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d2), "=f"(d3), "=f"(d10), "=f"(d11)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b2), "r"(b3),
|
||||
"f"(c2), "f"(c3), "f"(c10), "f"(c11),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB1));
|
||||
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d4), "=f"(d5), "=f"(d12), "=f"(d13)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b4), "r"(b5),
|
||||
"f"(c4), "f"(c5), "f"(c12), "f"(c13),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB2));
|
||||
|
||||
asm volatile(
|
||||
"mma.sync.aligned.kind::mxf4nvf4.block_scale.scale_vec::4X.m16n8k64.row.col.f32.e2m1.e2m1.f32.ue4m3 "
|
||||
"{%0, %1, %2, %3},"
|
||||
"{%4, %5, %6, %7},"
|
||||
"{%8, %9},"
|
||||
"{%10, %11, %12, %13},"
|
||||
"{%14},"
|
||||
"{%15, %16},"
|
||||
"{%17},"
|
||||
"{%18, %19};\n"
|
||||
: "=f"(d6), "=f"(d7), "=f"(d14), "=f"(d15)
|
||||
: "r"(a0), "r"(a1), "r"(a2), "r"(a3),
|
||||
"r"(b6), "r"(b7),
|
||||
"f"(c6), "f"(c7), "f"(c14), "f"(c15),
|
||||
"r"(uint32_t(sfa0)) , "h"(bidA), "h"(tidA),
|
||||
"r"(uint32_t(sfb0)) , "h"(bidB), "h"(tidB3));
|
||||
#else
|
||||
CUTE_INVALID_CONTROL_PATH("Attempting to use SM120::BLOCKSCALED::SM120_16x8x64_TN_VS without CUTE_ARCH_MXF4NVF4_4X_UE4M3_MMA_ENABLED");
|
||||
#endif
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace cute::SM120::BLOCKSCALED
|
||||
|
||||
namespace cute {
|
||||
|
||||
// MMA NVFP4 16x32x64 TN
|
||||
template <>
|
||||
struct MMA_Traits<SM120::BLOCKSCALED::SM120_16x32x64_TN_VS_NVFP4>
|
||||
{
|
||||
// The MMA accepts 4-bit inputs regardless of the types for A and B
|
||||
using ValTypeA = uint4_t;
|
||||
using ValTypeB = uint4_t;
|
||||
|
||||
using ValTypeD = float;
|
||||
using ValTypeC = float;
|
||||
|
||||
using ValTypeSF = cutlass::float_ue4m3_t;
|
||||
constexpr static int SFVecSize = 16;
|
||||
|
||||
using Shape_MNK = Shape<_16,_32,_64>;
|
||||
using ThrID = Layout<_32>;
|
||||
|
||||
// (T32,V32) -> (M16,K64)
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _8,_2, _2>>,
|
||||
Stride<Stride<_128,_1>,Stride<_16,_8,_512>>>;
|
||||
// (T32,V64) -> (N32,K64)
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,Shape <_8, _2, _4>>,
|
||||
Stride<Stride<_256,_1>,Stride<_32,_1024, _8>>>;
|
||||
// (T32,V64) -> (M16,K64)
|
||||
using SFALayout = Layout<Shape <Shape <_2,_2,_8>,_64>,
|
||||
Stride<Stride<_8,_0,_1>,_16>>;
|
||||
// (T32,V64) -> (N32,K64)
|
||||
using SFBLayout = Layout<Shape <Shape <_4,_8>,_64>,
|
||||
Stride<Stride<_8,_1>, _32>>;
|
||||
// (T32,V16) -> (M16,N32)
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < Shape<_2, _4>,_2>>,
|
||||
Stride<Stride<_32,_1>,Stride<Stride<_16, _128>,_8>>>;
|
||||
};
|
||||
|
||||
|
||||
template <class SFATensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFA(SFATensor&& sfatensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfatensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFA_TV = typename Atom::Traits::SFALayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<0>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfatensor, t_tile); // (PermM,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<0>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomM,AtomK),(RestM,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFA_TV{},_); // ((ThrV,FrgV),(RestM,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<1>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrM,ThrK)),(FrgV,(RestM,RestK)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFBTensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFB(SFBTensor&& sfbtensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfbtensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFB_TV = typename Atom::Traits::SFBLayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<1>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfbtensor, t_tile); // (PermN,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<1>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomN,AtomK),(RestN,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFB_TV{},_); // ((ThrV,FrgV),(RestN,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<2>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrN,ThrK)),(FrgV,(RestN,RestK)))
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFATensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_SFA(SFATensor&& sfatensor, ThrMma& thread_mma) {
|
||||
auto thr_tensor = make_tensor(static_cast<SFATensor&&>(sfatensor).data(), thrfrg_SFA(sfatensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vmk = make_coord(get<0>(thr_vmnk), make_coord(get<1>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
return thr_tensor(thr_vmk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
}
|
||||
|
||||
template <class SFATensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFA(SFATensor&& sfatensor, ThrMma& thread_mma) {
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
return make_fragment_like<ValTypeSF>(partition_SFA(sfatensor, thread_mma));
|
||||
}
|
||||
|
||||
template <class SFBTensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_SFB(SFBTensor&& sfbtensor, ThrMma& thread_mma) {
|
||||
auto thr_tensor = make_tensor(static_cast<SFBTensor&&>(sfbtensor).data(), thrfrg_SFB(sfbtensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vnk = make_coord(get<0>(thr_vmnk), make_coord(get<2>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
return thr_tensor(thr_vnk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
}
|
||||
|
||||
template <class SFBTensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFB(SFBTensor&& sfbtensor, ThrMma& thread_mma) {
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
return make_fragment_like<ValTypeSF>(partition_SFB(sfbtensor, thread_mma));
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFA_TV(TiledMma& mma)
|
||||
{
|
||||
// (M,K) -> (M,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_A = make_layout(make_shape(size<0>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto atile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<1>{} , Int<0>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFA(ref_A, mma).compose(atile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFB_TV(TiledMma& mma)
|
||||
{
|
||||
// (N,K) -> (N,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_B = make_layout(make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto btile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<0>{} , Int<1>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFB(ref_B, mma).compose(btile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
} // namespace cute
|
||||
@@ -0,0 +1,222 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "named_barrier.h"
|
||||
#include "utils.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename Ktraits>
|
||||
struct CollectiveEpilogueFwd{
|
||||
|
||||
using Element = typename Ktraits::ElementOut;
|
||||
static constexpr int kBlockM = Ktraits::kBlockM;
|
||||
static constexpr int kBlockN = Ktraits::kBlockN;
|
||||
static constexpr int kHeadDim = Ktraits::kHeadDim;
|
||||
using TileShape_MNK = Shape<Int<kBlockM>, Int<kBlockN>, Int<kHeadDim>>;
|
||||
static constexpr int kNWarps = Ktraits::kNWarps;
|
||||
static constexpr int kNThreads = kNWarps * cutlass::NumThreadsPerWarp;
|
||||
static constexpr int NumMmaThreads = kNThreads - cutlass::NumThreadsPerWarpGroup;
|
||||
|
||||
using GmemTiledCopyOTMA = cute::SM90_TMA_STORE;
|
||||
|
||||
// These are for storing the output tensor without TMA (e.g., for setting output to zero)
|
||||
static constexpr int kGmemElemsPerLoad = sizeof(cute::uint128_t) / sizeof(Element);
|
||||
static_assert(kHeadDim % kGmemElemsPerLoad == 0, "kHeadDim must be a multiple of kGmemElemsPerLoad");
|
||||
static constexpr int kGmemThreadsPerRow = kHeadDim / kGmemElemsPerLoad;
|
||||
static_assert(NumMmaThreads % kGmemThreadsPerRow == 0, "NumMmaThreads must be a multiple of kGmemThreadsPerRow");
|
||||
using GmemLayoutAtom = Layout<Shape <Int<NumMmaThreads / kGmemThreadsPerRow>, Int<kGmemThreadsPerRow>>,
|
||||
Stride<Int<kGmemThreadsPerRow>, _1>>;
|
||||
using GmemTiledCopyO = decltype(
|
||||
make_tiled_copy(Copy_Atom<DefaultCopy, Element>{},
|
||||
GmemLayoutAtom{},
|
||||
Layout<Shape<_1, Int<kGmemElemsPerLoad>>>{})); // Val layout, 8 or 16 vals per store
|
||||
|
||||
using SmemLayoutO = typename Ktraits::SmemLayoutO;
|
||||
|
||||
using SmemCopyAtomO = Copy_Atom<SM90_U32x2_STSM_N, Element>;
|
||||
using SharedStorage = cute::array_aligned<Element, cute::cosize_v<SmemLayoutO>>;
|
||||
|
||||
using ShapeO = cute::Shape<int32_t, int32_t, int32_t, int32_t>; // (seqlen_q, d, head, batch)
|
||||
using StrideO = cute::Stride<int64_t, _1, int64_t, int64_t>;
|
||||
using StrideLSE = cute::Stride<_1, int64_t, int64_t>; // (seqlen_q, head, batch)
|
||||
|
||||
using TMA_O = decltype(make_tma_copy(
|
||||
GmemTiledCopyOTMA{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element*>(nullptr)), repeat_like(StrideO{}, int32_t(0)), StrideO{}),
|
||||
SmemLayoutO{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{})); // no mcast for O
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
Element* ptr_O;
|
||||
ShapeO const shape_O;
|
||||
StrideO const stride_O;
|
||||
float* ptr_LSE;
|
||||
StrideLSE const stride_LSE;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
Element* ptr_O;
|
||||
ShapeO const shape_O;
|
||||
StrideO const stride_O;
|
||||
float* ptr_LSE;
|
||||
StrideLSE const stride_LSE;
|
||||
TMA_O tma_store_O;
|
||||
};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
Tensor mO = make_tensor(make_gmem_ptr(args.ptr_O), args.shape_O, args.stride_O);
|
||||
TMA_O tma_store_O = make_tma_copy(
|
||||
GmemTiledCopyOTMA{},
|
||||
mO,
|
||||
SmemLayoutO{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{}); // no mcast for O
|
||||
return {args.ptr_O, args.shape_O, args.stride_O, args.ptr_LSE, args.stride_LSE, tma_store_O};
|
||||
}
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& epilogue_params) {
|
||||
cute::prefetch_tma_descriptor(epilogue_params.tma_store_O.get_tma_descriptor());
|
||||
}
|
||||
|
||||
template <typename SharedStorage, typename FrgTensorO, typename TiledMma>
|
||||
CUTLASS_DEVICE void
|
||||
mma_store(
|
||||
SharedStorage& shared_storage,
|
||||
TiledMma tiled_mma,
|
||||
FrgTensorO const& tOrO,
|
||||
int thread_idx
|
||||
){
|
||||
Tensor sO = cute::as_position_independent_swizzle_tensor(make_tensor(make_smem_ptr(shared_storage.smem_o.begin()), SmemLayoutO{}));
|
||||
auto smem_tiled_copy_O = make_tiled_copy_C(SmemCopyAtomO{}, tiled_mma);
|
||||
auto smem_thr_copy_O = smem_tiled_copy_O.get_thread_slice(thread_idx);
|
||||
constexpr int numel = decltype(size(tOrO))::value;
|
||||
cutlass::NumericArrayConverter<Element, float, numel> convert_op;
|
||||
// HACK: this requires tensor to be "contiguous"
|
||||
auto frag = convert_op(*reinterpret_cast<const cutlass::Array<float, numel> *>(tOrO.data()));
|
||||
auto tOrO_out = make_tensor(make_rmem_ptr<Element>(&frag), tOrO.layout());
|
||||
Tensor taccOrO = smem_thr_copy_O.retile_S(tOrO_out); // ((Atom,AtomNum), MMA_M, MMA_N)
|
||||
Tensor taccOsO = smem_thr_copy_O.partition_D(sO); // ((Atom,AtomNum),PIPE_M,PIPE_N)
|
||||
cute::copy(smem_tiled_copy_O, taccOrO, taccOsO);
|
||||
cutlass::arch::fence_view_async_shared(); // ensure smem writes are visible to TMA
|
||||
}
|
||||
|
||||
template<typename SharedStorage, typename Params, typename WorkTileInfo, typename SchedulerParams>
|
||||
CUTLASS_DEVICE void
|
||||
tma_store(
|
||||
SharedStorage& shared_storage,
|
||||
Params const& epilogue_params,
|
||||
WorkTileInfo work_tile_info,
|
||||
SchedulerParams const& scheduler_params,
|
||||
int thread_idx
|
||||
) {
|
||||
auto [m_block, bidh, bidb] = work_tile_info.get_block_coord(scheduler_params);
|
||||
Tensor sO = cute::as_position_independent_swizzle_tensor(make_tensor(make_smem_ptr(shared_storage.smem_o.begin()), SmemLayoutO{}));
|
||||
Tensor mO = epilogue_params.tma_store_O.get_tma_tensor(epilogue_params.shape_O);
|
||||
Tensor gO = local_tile(mO(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{})); // (M, K)
|
||||
auto block_tma_O = epilogue_params.tma_store_O.get_slice(_0{});
|
||||
Tensor tOgO = block_tma_O.partition_D(gO); // (TMA, TMA_M, TMA_K)
|
||||
Tensor tOsO = block_tma_O.partition_S(sO); // (TMA, TMA_M, TMA_K)
|
||||
|
||||
// auto shape_LSE = select<0, 2, 3>(epilogue_params.shape_O);
|
||||
// Tensor mLSE = make_tensor(make_gmem_ptr(epilogue_params.ptr_LSE), shape_LSE, epilogue_params.stride_LSE);
|
||||
// Tensor gLSE = local_tile(mLSE(_, bidh, bidb), Shape<Int<kBlockM>>{}, make_coord(m_block));
|
||||
|
||||
// Tensor caccO = cute::make_identity_tensor(select<0, 2>(TileShape_MNK{}));
|
||||
// auto thread_mma = tiled_mma.get_thread_slice(thread_idx);
|
||||
// Tensor taccOcO = thread_mma.partition_C(caccO); // (MMA,MMA_M,MMA_K)
|
||||
// static_assert(decltype(size<0, 0>(taccOcO))::value == 2);
|
||||
// static_assert(decltype(size<0, 1>(taccOcO))::value == 2);
|
||||
// // // // taccOcO has shape ((2, 2, V), MMA_M, MMA_K), we only take only the row indices.
|
||||
// Tensor taccOcO_row = taccOcO(make_coord(_0{}, _), _, _0{});
|
||||
// CUTE_STATIC_ASSERT_V(size(lse) == size(taccOcO_row)); // MMA_M
|
||||
// if (get<1>(taccOcO_row(_0{})) == 0) {
|
||||
// #pragma unroll
|
||||
// for (int mi = 0; mi < size(lse); ++mi) {
|
||||
// const int row = get<0>(taccOcO_row(mi));
|
||||
// if (row < get<0>(shape_LSE) - m_block * kBlockM) { gLSE(row) = lse(mi); }
|
||||
// }
|
||||
// }
|
||||
|
||||
// if (cutlass::canonical_warp_idx_sync() == kNWarps - 1) {
|
||||
// cutlass::arch::NamedBarrier::sync(NumMmaThreads + cutlass::NumThreadsPerWarp,
|
||||
// static_cast<uint32_t>(FP4NamedBarriers::EpilogueBarrier));
|
||||
// int const lane_predicate = cute::elect_one_sync();
|
||||
// if (lane_predicate) {
|
||||
// cute::copy(epilogue_params.tma_store_O, tOsO, tOgO);
|
||||
// tma_store_arrive();
|
||||
// }
|
||||
// }
|
||||
cute::copy(epilogue_params.tma_store_O, tOsO, tOgO);
|
||||
tma_store_arrive();
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
store_tail() {
|
||||
tma_store_wait<0>();
|
||||
}
|
||||
|
||||
// Write 0 to output and -inf to LSE
|
||||
CUTLASS_DEVICE void
|
||||
store_zero(
|
||||
Params const& epilogue_params,
|
||||
int thread_idx,
|
||||
cute::tuple<int32_t, int32_t, int32_t> const& block_coord
|
||||
) {
|
||||
auto [m_block, bidh, bidb] = block_coord;
|
||||
Tensor mO = make_tensor(make_gmem_ptr(epilogue_params.ptr_O), epilogue_params.shape_O, epilogue_params.stride_O);
|
||||
Tensor gO = local_tile(mO(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{})); // (M, K)
|
||||
auto shape_LSE = select<0, 2, 3>(epilogue_params.shape_O);
|
||||
Tensor mLSE = make_tensor(make_gmem_ptr(epilogue_params.ptr_LSE), shape_LSE, epilogue_params.stride_LSE);
|
||||
Tensor gLSE = local_tile(mLSE(_, bidh, bidb), Shape<Int<kBlockM>>{}, make_coord(m_block));
|
||||
|
||||
GmemTiledCopyO gmem_tiled_copy_O;
|
||||
auto gmem_thr_copy_O = gmem_tiled_copy_O.get_thread_slice(thread_idx);
|
||||
Tensor tOgO = gmem_thr_copy_O.partition_D(gO);
|
||||
Tensor tOrO = make_fragment_like(tOgO);
|
||||
clear(tOrO);
|
||||
// Construct identity layout for sO
|
||||
Tensor cO = cute::make_identity_tensor(select<0, 2>(TileShape_MNK{})); // (BLK_M,BLK_K) -> (blk_m,blk_k)
|
||||
// Repeat the partitioning with identity layouts
|
||||
Tensor tOcO = gmem_thr_copy_O.partition_D(cO);
|
||||
Tensor tOpO = make_tensor<bool>(make_shape(size<2>(tOgO)));
|
||||
#pragma unroll
|
||||
for (int k = 0; k < size(tOpO); ++k) { tOpO(k) = get<1>(tOcO(_0{}, _0{}, k)) < get<1>(epilogue_params.shape_O); }
|
||||
// Clear_OOB_K must be false since we don't want to write zeros to gmem
|
||||
flash::copy</*Is_even_MN=*/false, /*Is_even_K=*/false, /*Clear_OOB_MN=*/false, /*Clear_OOB_K=*/false>(
|
||||
gmem_tiled_copy_O, tOrO, tOgO, tOcO, tOpO, get<0>(epilogue_params.shape_O) - m_block * kBlockM
|
||||
);
|
||||
static_assert(kBlockM <= NumMmaThreads);
|
||||
if (thread_idx < get<0>(shape_LSE) - m_block * kBlockM) { gLSE(thread_idx) = INFINITY; }
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,202 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cute/algorithm/copy.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/layout/layout.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
|
||||
#include "blockscaled_layout.h"
|
||||
#include "cute_extension.h"
|
||||
#include "named_barrier.h"
|
||||
using namespace cute;
|
||||
|
||||
template <
|
||||
int kStages,
|
||||
int EpiStages,
|
||||
typename Element,
|
||||
typename ElementSF,
|
||||
typename OutputType,
|
||||
typename SmemLayoutQ,
|
||||
typename SmemLayoutK,
|
||||
typename SmemLayoutV,
|
||||
typename SmemLayoutDS,
|
||||
typename SmemLayoutO,
|
||||
typename SmemLayoutSFQ,
|
||||
typename SmemLayoutSFK,
|
||||
typename SmemLayoutSFV
|
||||
>
|
||||
struct SharedStorageQKVOwithSF : cute::aligned_struct<128, _0>{
|
||||
|
||||
alignas(1024) cute::ArrayEngine<Element, cute::cosize_v<SmemLayoutQ>> smem_q;
|
||||
alignas(1024) cute::ArrayEngine<Element, cute::cosize_v<SmemLayoutK>> smem_k;
|
||||
cute::ArrayEngine<ElementSF, cute::cosize_v<SmemLayoutSFQ>> smem_SFQ;
|
||||
cute::ArrayEngine<ElementSF, cute::cosize_v<SmemLayoutSFK>> smem_SFK;
|
||||
cute::ArrayEngine<ElementSF, cute::cosize_v<SmemLayoutSFV>> smem_SFV;
|
||||
alignas(1024) cute::ArrayEngine<float, cute::cosize_v<SmemLayoutDS>> smem_ds;
|
||||
alignas(1024) cute::ArrayEngine<Element, cute::cosize_v<SmemLayoutV>> smem_v;
|
||||
alignas(1024) cute::ArrayEngine<OutputType, cute::cosize_v<SmemLayoutO>> smem_o;
|
||||
|
||||
struct {
|
||||
alignas(16) typename cutlass::PipelineTmaAsync<1>::SharedStorage pipeline_q;
|
||||
alignas(16) typename cutlass::PipelineTmaAsync<kStages>::SharedStorage pipeline_k;
|
||||
alignas(16) typename cutlass::PipelineTmaAsync<kStages>::SharedStorage pipeline_v;
|
||||
alignas(16) typename flash::OrderedSequenceBarrierVarGroupSize<EpiStages, 2>::SharedStorage barrier_o;
|
||||
int tile_count_semaphore;
|
||||
};
|
||||
};
|
||||
|
||||
template <
|
||||
int kHeadDim_,
|
||||
int kBlockM_,
|
||||
int kBlockN_,
|
||||
int kStages_,
|
||||
int kClusterM_,
|
||||
bool BlockMean_,
|
||||
typename ElementPairType_ = cutlass::nv_float4_t<cutlass::float_e2m1_t>,
|
||||
typename ElementOut_ = cutlass::bfloat16_t
|
||||
>
|
||||
struct Flash_fwd_kernel_traits {
|
||||
static constexpr int kBlockM = kBlockM_;
|
||||
static constexpr int kBlockN = kBlockN_;
|
||||
static constexpr int kHeadDim = kHeadDim_;
|
||||
static constexpr bool BlockMean = BlockMean_;
|
||||
static constexpr bool SmoothQ = true;
|
||||
static_assert(kHeadDim % 32 == 0);
|
||||
static_assert(kBlockM == 64 || kBlockM == 128);
|
||||
static constexpr int kNWarps = kBlockM == 128 ? 12 : 8;
|
||||
static constexpr int kNThreads = kNWarps * cutlass::NumThreadsPerWarp;
|
||||
static constexpr int kClusterM = kClusterM_;
|
||||
static constexpr int kStages = kStages_;
|
||||
static constexpr int EpiStages = 1;
|
||||
static constexpr int NumSFQK = kHeadDim / 16;
|
||||
static constexpr int NumSFPV = kBlockN / 16;
|
||||
using ElementSF = cutlass::float_ue4m3_t;
|
||||
using Element = cutlass::float_e2m1_t;
|
||||
using ElementAccum = float;
|
||||
using ElementOut = ElementOut_;
|
||||
using index_t = int64_t;
|
||||
static constexpr auto SFVectorSize = 16;
|
||||
using TileShape_MNK = Shape<Int<kBlockM>, Int<kBlockN>, Int<kHeadDim>>;
|
||||
using ClusterShape_MNK = Shape<_1, _1, _1>;
|
||||
using PermTileM = decltype(cute::min(size<0>(TileShape_MNK{}), _128{}));
|
||||
using PermTileN = _32;
|
||||
using PermTileK = Int<kHeadDim>;
|
||||
|
||||
using ElementQMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element<Element>());
|
||||
using ElementKMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element<Element>());
|
||||
|
||||
using AtomLayoutMNK = std::conditional_t<kBlockM == 128,
|
||||
Layout<Shape<_8, _1, _1>>,
|
||||
Layout<Shape<_4, _1, _1>>
|
||||
>;
|
||||
using TiledMmaQK = decltype(cute::make_tiled_mma(
|
||||
cute::SM120::BLOCKSCALED::SM120_16x32x64_TN_VS_NVFP4{},
|
||||
AtomLayoutMNK{},
|
||||
Tile<PermTileM, PermTileN, PermTileK>{}
|
||||
));
|
||||
|
||||
using TiledMmaPV = decltype(cute::make_tiled_mma(
|
||||
cute::SM120::BLOCKSCALED::SM120_16x32x64_TN_VS_NVFP4{},
|
||||
AtomLayoutMNK{},
|
||||
Tile<PermTileM, _32, PermTileK>{}
|
||||
));
|
||||
|
||||
static constexpr int MMA_NSF = size<2>(typename TiledMmaQK::AtomShape_MNK{}) / SFVectorSize;
|
||||
|
||||
using GmemTiledCopy = SM90_TMA_LOAD;
|
||||
using GmemTiledCopySF = SM90_TMA_LOAD;
|
||||
|
||||
using SmemLayoutAtomQ = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomK = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomV = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomVt = decltype(cutlass::gemm::collective::detail::sm120_rr_smem_selector<Element, decltype(size<1>(TileShape_MNK{}))>());
|
||||
using SmemLayoutQ = decltype(tile_to_shape(SmemLayoutAtomQ{}, select<0, 2>(TileShape_MNK{})));
|
||||
using SmemLayoutK =
|
||||
decltype(tile_to_shape(SmemLayoutAtomK{},
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{}), Int<kStages>{})));
|
||||
using SmemLayoutV =
|
||||
decltype(tile_to_shape(SmemLayoutAtomV{},
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{}), Int<kStages>{})));
|
||||
using SmemLayoutVt =
|
||||
decltype(tile_to_shape(SmemLayoutAtomVt{},
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{}), Int<kStages>{})));
|
||||
using SmemLayoutAtomDS = Layout<Shape<Int<kBlockM>, Int<kBlockN>>, Stride<_0, _1>>;
|
||||
using SmemLayoutDS =
|
||||
decltype(tile_to_shape(SmemLayoutAtomDS{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<1>(TileShape_MNK{}), Int<kStages>{})));
|
||||
|
||||
using SmemCopyAtomQ = Copy_Atom<SM75_U32x4_LDSM_N, Element>;
|
||||
using SmemCopyAtomKV = Copy_Atom<SM75_U32x4_LDSM_N, Element>;
|
||||
using SmemCopyAtomSF = Copy_Atom<UniversalCopy<ElementSF>, ElementSF>;
|
||||
using SmemCopyAtomDS = Copy_Atom<UniversalCopy<float>, float>;
|
||||
|
||||
using BlkScaledConfig = flash::BlockScaledConfig<SFVectorSize>;
|
||||
using LayoutSF = typename BlkScaledConfig::LayoutSF;
|
||||
using SfAtom = typename BlkScaledConfig::SfAtom;
|
||||
using SmemLayoutAtomSFQ = decltype(BlkScaledConfig::deduce_smem_layoutSFQ(TiledMmaQK{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomSFK = decltype(BlkScaledConfig::deduce_smem_layoutSFKV(TiledMmaQK{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomSFV = decltype(BlkScaledConfig::deduce_smem_layoutSFKV(TiledMmaPV{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomSFVt = decltype(BlkScaledConfig::deduce_smem_layoutSFVt(TiledMmaPV{}, Shape<Int<kBlockM>, Int<kHeadDim>, Int<kBlockN>>{}));
|
||||
using LayoutSFP = decltype(
|
||||
make_layout(
|
||||
make_shape(make_shape(_16{}, _4{}), _1{}, Int<kBlockN / 64>{}),
|
||||
make_stride(make_stride(_0{}, _1{}), _0{}, _4{})
|
||||
)
|
||||
);
|
||||
using LayoutP = decltype(
|
||||
make_layout(
|
||||
make_shape(make_shape(_8{}, _2{}, _2{}), _1{}, Int<kBlockN / 64>{}),
|
||||
make_stride(make_stride(_1{}, _8{}, _16{}), _0{}, _32{})
|
||||
)
|
||||
);
|
||||
using SmemLayoutSFQ = decltype(make_layout(
|
||||
shape(SmemLayoutAtomSFQ{}),
|
||||
stride(SmemLayoutAtomSFQ{})
|
||||
));
|
||||
using SmemLayoutSFK = decltype(make_layout(
|
||||
append(shape(SmemLayoutAtomSFK{}), Int<kStages>{}),
|
||||
append(stride(SmemLayoutAtomSFK{}), size(filter_zeros(SmemLayoutAtomSFK{})))
|
||||
));
|
||||
using SmemLayoutSFV = decltype(make_layout(
|
||||
append(shape(SmemLayoutAtomSFV{}), Int<kStages>{}),
|
||||
append(stride(SmemLayoutAtomSFV{}), size(filter_zeros(SmemLayoutAtomSFV{})))
|
||||
));
|
||||
using SmemLayoutSFVt = decltype(make_layout(
|
||||
append(shape(SmemLayoutAtomSFVt{}), Int<kStages>{}),
|
||||
append(stride(SmemLayoutAtomSFVt{}), size(filter_zeros(SmemLayoutAtomSFVt{})))
|
||||
));
|
||||
|
||||
using SmemLayoutAtomO = decltype(cutlass::gemm::collective::detail::ss_smem_selector<GMMA::Major::K, ElementOut,
|
||||
decltype(cute::get<0>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutO = decltype(tile_to_shape(SmemLayoutAtomO{}, select<0, 2>(TileShape_MNK{}), Step<_1, _2>{}));
|
||||
using SharedStorage = SharedStorageQKVOwithSF<kStages, EpiStages, Element, ElementSF, ElementOut,
|
||||
SmemLayoutQ, SmemLayoutK, SmemLayoutV, SmemLayoutDS,
|
||||
SmemLayoutO, SmemLayoutSFQ, SmemLayoutSFK, SmemLayoutSFVt>;
|
||||
using MainloopPipeline = typename cutlass::PipelineTmaAsync<kStages>;
|
||||
using PipelineState = typename cutlass::PipelineState<kStages>;
|
||||
using MainloopPipelineQ = cutlass::PipelineTmaAsync<1>;
|
||||
using PipelineParamsQ = typename MainloopPipelineQ::Params;
|
||||
using PipelineStateQ = typename cutlass::PipelineState<1>;
|
||||
using EpilogueBarrier = typename flash::OrderedSequenceBarrierVarGroupSize<EpiStages, 2>;
|
||||
};
|
||||
|
||||
@@ -0,0 +1,204 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/arch/reg_reconfig.h>
|
||||
#include <cutlass/array.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include <cutlass/numeric_conversion.h>
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
|
||||
#include "params.h"
|
||||
#include "utils.h"
|
||||
#include "tile_scheduler.h"
|
||||
#include "mainloop_tma_ws.h"
|
||||
#include "epilogue_tma_ws.h"
|
||||
#include "named_barrier.h"
|
||||
#include "softmax_fused.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename Ktraits, bool Is_causal, typename TileScheduler>
|
||||
__global__ void __launch_bounds__(Ktraits::kNWarps * cutlass::NumThreadsPerWarp, 1)
|
||||
compute_attn_ws(CUTE_GRID_CONSTANT Flash_fwd_params const params,
|
||||
CUTE_GRID_CONSTANT typename CollectiveMainloopFwd<Ktraits, Is_causal>::Params const mainloop_params,
|
||||
CUTE_GRID_CONSTANT typename CollectiveEpilogueFwd<Ktraits>::Params const epilogue_params,
|
||||
CUTE_GRID_CONSTANT typename TileScheduler::Params const scheduler_params
|
||||
) {
|
||||
|
||||
using Element = typename Ktraits::Element;
|
||||
using ElementAccum = typename Ktraits::ElementAccum;
|
||||
using SoftType = ElementAccum;
|
||||
using TileShape_MNK = typename Ktraits::TileShape_MNK;
|
||||
using ClusterShape = typename Ktraits::ClusterShape_MNK;
|
||||
|
||||
static constexpr int NumMmaThreads = size(typename Ktraits::TiledMmaQK{});
|
||||
static constexpr int NumCopyThreads = cutlass::NumThreadsPerWarpGroup;
|
||||
static constexpr int kBlockM = Ktraits::kBlockM;
|
||||
|
||||
using CollectiveMainloop = CollectiveMainloopFwd<Ktraits, Is_causal>;
|
||||
using CollectiveEpilogue = CollectiveEpilogueFwd<Ktraits>;
|
||||
|
||||
using MainloopPipeline = typename Ktraits::MainloopPipeline;
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename MainloopPipeline::PipelineState;
|
||||
using MainloopPipelineQ = typename Ktraits::MainloopPipelineQ;
|
||||
using PipelineParamsQ = typename Ktraits::PipelineParamsQ;
|
||||
using PipelineStateQ = typename Ktraits::PipelineStateQ;
|
||||
using EpilogueBarrier = typename Ktraits::EpilogueBarrier;
|
||||
|
||||
|
||||
enum class WarpGroupRole {
|
||||
Producer = 0,
|
||||
Consumer0 = 1,
|
||||
Consumer1 = 2
|
||||
};
|
||||
enum class ProducerWarpRole {
|
||||
Mainloop = 0,
|
||||
Epilogue = 1,
|
||||
Warp2 = 2,
|
||||
Warp3 = 3
|
||||
};
|
||||
|
||||
extern __shared__ char shared_memory[];
|
||||
auto &shared_storage = *reinterpret_cast<typename Ktraits::SharedStorage*>(shared_memory);
|
||||
|
||||
int const lane_predicate = cute::elect_one_sync();
|
||||
int const warp_idx = cutlass::canonical_warp_idx_sync();
|
||||
int warp_group_idx = cutlass::canonical_warp_group_idx();
|
||||
int const warp_group_thread_idx = threadIdx.x % cutlass::NumThreadsPerWarpGroup;
|
||||
int warp_idx_in_warp_group = warp_idx % cutlass::NumWarpsPerWarpGroup;
|
||||
auto warp_group_role = WarpGroupRole(warp_group_idx);
|
||||
auto producer_warp_role = ProducerWarpRole(warp_idx_in_warp_group);
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if (warp_idx == 0 && lane_predicate) {
|
||||
CollectiveMainloop::prefetch_tma_descriptors(mainloop_params);
|
||||
CollectiveEpilogue::prefetch_tma_descriptors(epilogue_params);
|
||||
}
|
||||
|
||||
// Obtain warp index
|
||||
|
||||
PipelineParams pipeline_params_v;
|
||||
pipeline_params_v.transaction_bytes = CollectiveMainloop::TmaTransactionBytesV;
|
||||
pipeline_params_v.role = warp_group_role == WarpGroupRole::Producer
|
||||
? MainloopPipeline::ThreadCategory::Producer
|
||||
: MainloopPipeline::ThreadCategory::Consumer;
|
||||
pipeline_params_v.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_v.num_consumers = NumMmaThreads;
|
||||
|
||||
PipelineParams pipeline_params_k;
|
||||
pipeline_params_k.transaction_bytes = CollectiveMainloop::TmaTransactionBytesK;
|
||||
pipeline_params_k.role = warp_group_role == WarpGroupRole::Producer
|
||||
? MainloopPipeline::ThreadCategory::Producer
|
||||
: MainloopPipeline::ThreadCategory::Consumer;
|
||||
pipeline_params_k.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_k.num_consumers = NumMmaThreads;
|
||||
|
||||
PipelineParamsQ pipeline_params_q;
|
||||
pipeline_params_q.transaction_bytes = CollectiveMainloop::TmaTransactionBytesQ;
|
||||
pipeline_params_q.role = warp_group_role == WarpGroupRole::Producer
|
||||
? MainloopPipelineQ::ThreadCategory::Producer
|
||||
: MainloopPipelineQ::ThreadCategory::Consumer;
|
||||
pipeline_params_q.is_leader = warp_group_thread_idx == 0;
|
||||
pipeline_params_q.num_consumers = NumMmaThreads;
|
||||
|
||||
// We're counting on pipeline_k to call cutlass::arch::fence_barrier_init();
|
||||
MainloopPipelineQ pipeline_q(shared_storage.pipeline_q, pipeline_params_q, ClusterShape{});
|
||||
MainloopPipeline pipeline_k(shared_storage.pipeline_k, pipeline_params_k, ClusterShape{});
|
||||
MainloopPipeline pipeline_v(shared_storage.pipeline_v, pipeline_params_v, ClusterShape{});
|
||||
|
||||
uint32_t epilogue_barrier_group_size_list[2] = {cutlass::NumThreadsPerWarp, NumMmaThreads};
|
||||
typename EpilogueBarrier::Params params_epilogue_barrier;
|
||||
params_epilogue_barrier.group_id = (warp_group_role == WarpGroupRole::Producer);
|
||||
params_epilogue_barrier.group_size_list = epilogue_barrier_group_size_list;
|
||||
EpilogueBarrier barrier_o(shared_storage.barrier_o, params_epilogue_barrier);
|
||||
|
||||
CollectiveMainloop collective_mainloop;
|
||||
CollectiveEpilogue collective_epilogue;
|
||||
__syncthreads();
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
cutlass::arch::warpgroup_reg_dealloc<24>();
|
||||
TileScheduler scheduler;
|
||||
|
||||
if (producer_warp_role == ProducerWarpRole::Mainloop) { // Load Q, K, V
|
||||
PipelineStateQ smem_pipe_write_q = cutlass::make_producer_start_state<MainloopPipelineQ>();
|
||||
PipelineState smem_pipe_write_k = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
PipelineState smem_pipe_write_v = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
|
||||
int work_idx = 0;
|
||||
for (auto work_tile_info = scheduler.get_initial_work(); work_tile_info.is_valid(scheduler_params); work_tile_info = scheduler.get_next_work(scheduler_params, work_tile_info)) {
|
||||
int tile_count_semaphore = 0;
|
||||
collective_mainloop.load(mainloop_params, scheduler_params,
|
||||
pipeline_q, pipeline_k, pipeline_v,
|
||||
smem_pipe_write_q, smem_pipe_write_k, smem_pipe_write_v,
|
||||
shared_storage, work_tile_info, work_idx, tile_count_semaphore);
|
||||
}
|
||||
collective_mainloop.load_tail(pipeline_q, pipeline_k, pipeline_v,
|
||||
smem_pipe_write_q, smem_pipe_write_k, smem_pipe_write_v);
|
||||
} else if (producer_warp_role == ProducerWarpRole::Epilogue) {
|
||||
for (auto work_tile_info = scheduler.get_initial_work(); work_tile_info.is_valid(scheduler_params); work_tile_info = scheduler.get_next_work(scheduler_params, work_tile_info)) {
|
||||
barrier_o.wait();
|
||||
collective_epilogue.tma_store(shared_storage, epilogue_params, work_tile_info, scheduler_params, threadIdx.x);
|
||||
collective_epilogue.store_tail();
|
||||
barrier_o.arrive();
|
||||
}
|
||||
|
||||
}
|
||||
} else if (warp_group_role == WarpGroupRole::Consumer0 || warp_group_role == WarpGroupRole::Consumer1) {
|
||||
cutlass::arch::warpgroup_reg_alloc<232>();
|
||||
typename Ktraits::TiledMmaPV tiled_mma_pv;
|
||||
TileScheduler scheduler{};
|
||||
PipelineState smem_pipe_read_k, smem_pipe_read_v;
|
||||
PipelineStateQ smem_pipe_read_q;
|
||||
|
||||
int work_idx = 0;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for (auto work_tile_info = scheduler.get_initial_work(); work_tile_info.is_valid(scheduler_params); work_tile_info = scheduler.get_next_work(scheduler_params, work_tile_info)) {
|
||||
// Attention output (GEMM-II) accumulator.
|
||||
Tensor tOrO = partition_fragment_C(tiled_mma_pv, select<0, 2>(TileShape_MNK{}));
|
||||
// flash::Softmax<2 * (2 * kBlockM / NumMmaThreads)> softmax;
|
||||
// Pass single_level_p_quant flag to control P quantization mode
|
||||
flash::SoftmaxFused<2 * (2 * kBlockM / NumMmaThreads)> softmax_fused(params.single_level_p_quant);
|
||||
auto block_coord = work_tile_info.get_block_coord(scheduler_params);
|
||||
auto [m_block, bidh, bidb] = block_coord;
|
||||
|
||||
int n_block_max = collective_mainloop.get_n_block_max(mainloop_params, m_block);
|
||||
if (Is_causal && n_block_max <= 0) { // We exit early and write 0 to gO and -inf to gLSE.
|
||||
collective_epilogue.store_zero(epilogue_params, threadIdx.x - NumCopyThreads, block_coord);
|
||||
continue;
|
||||
}
|
||||
|
||||
collective_mainloop.mma(mainloop_params, pipeline_q, pipeline_k, pipeline_v, smem_pipe_read_q, smem_pipe_read_k, smem_pipe_read_v,
|
||||
tOrO, softmax_fused, n_block_max, threadIdx.x - NumCopyThreads, work_idx, m_block, shared_storage);
|
||||
barrier_o.wait();
|
||||
collective_epilogue.mma_store(shared_storage, tiled_mma_pv, tOrO, threadIdx.x - NumCopyThreads);
|
||||
barrier_o.arrive();
|
||||
++work_idx;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,114 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cutlass/cluster_launch.hpp"
|
||||
|
||||
#include "static_switch.h"
|
||||
#include "params.h"
|
||||
#include "tile_scheduler.h"
|
||||
#include "kernel_ws.h"
|
||||
#include "kernel_traits.h"
|
||||
#include "block_config.h"
|
||||
|
||||
|
||||
template<typename Kernel_traits, bool Is_causal>
|
||||
void run_flash_fwd(Flash_fwd_params ¶ms, cudaStream_t stream) {
|
||||
using Element = typename Kernel_traits::Element;
|
||||
using ElementSF = typename Kernel_traits::ElementSF;
|
||||
using ElementOut = typename Kernel_traits::ElementOut;
|
||||
using TileShape_MNK = typename Kernel_traits::TileShape_MNK;
|
||||
using ClusterShape = typename Kernel_traits::ClusterShape_MNK;
|
||||
using CollectiveMainloop = flash::CollectiveMainloopFwd<Kernel_traits, Is_causal>;
|
||||
using CollectiveEpilogue = flash::CollectiveEpilogueFwd<Kernel_traits>;
|
||||
// using Scheduler = flash::SingleTileScheduler;
|
||||
using Scheduler = flash::StaticPersistentTileScheduler;
|
||||
typename CollectiveMainloop::Params mainloop_params =
|
||||
CollectiveMainloop::to_underlying_arguments({
|
||||
static_cast<Element const*>(params.q_ptr),
|
||||
{params.seqlen_q, params.d, params.h, params.b}, // shape_Q
|
||||
{params.q_row_stride, _1{}, params.q_head_stride, params.q_batch_stride}, // stride_Q
|
||||
static_cast<Element const*>(params.k_ptr),
|
||||
{params.seqlen_k, params.d, params.h_k, params.b}, // shape_K
|
||||
{params.k_row_stride, _1{}, params.k_head_stride, params.k_batch_stride}, // stride_K
|
||||
{params.unpadded_seqlen_k, params.d, params.h_k, params.b}, // shape_K
|
||||
static_cast<Element const*>(params.v_ptr),
|
||||
{params.d, params.seqlen_k, params.h_k, params.b}, // shape_Vt
|
||||
{params.v_row_stride, _1{}, params.v_head_stride, params.v_batch_stride}, // stride_Vt
|
||||
static_cast<ElementSF const*>(params.sfq_ptr),
|
||||
{params.seqlen_q, params.d, params.h, params.b}, // shape_SFQ
|
||||
static_cast<ElementSF const*>(params.sfk_ptr),
|
||||
{params.seqlen_k, params.d, params.h_k, params.b}, // shape_SFK
|
||||
static_cast<ElementSF const*>(params.sfv_ptr),
|
||||
{params.d, params.seqlen_k, params.h_k, params.b}, // shape_SFVt
|
||||
static_cast<float const*>(params.delta_s_ptr),
|
||||
{params.seqlen_s, params.seqlen_k, params.h_k, params.b},
|
||||
{params.ds_row_stride, _1{}, params.ds_head_stride, params.ds_batch_stride},
|
||||
params.scale_softmax_log2
|
||||
});
|
||||
typename CollectiveEpilogue::Params epilogue_params =
|
||||
CollectiveEpilogue::to_underlying_arguments({
|
||||
static_cast<ElementOut*>(params.o_ptr),
|
||||
{params.seqlen_q, params.d, params.h, params.b}, // shape_O
|
||||
{params.o_row_stride, _1{}, params.o_head_stride, params.o_batch_stride}, // stride_O
|
||||
static_cast<float*>(params.softmax_lse_ptr),
|
||||
{_1{}, params.seqlen_q, params.h * params.seqlen_q}, // stride_LSE
|
||||
});
|
||||
|
||||
int num_blocks_m = cutlass::ceil_div(params.seqlen_q, Kernel_traits::kBlockM);
|
||||
num_blocks_m = cutlass::ceil_div(num_blocks_m, size<0>(ClusterShape{})) * size<0>(ClusterShape{});
|
||||
typename Scheduler::Arguments scheduler_args = {num_blocks_m, params.h, params.b};
|
||||
typename Scheduler::Params scheduler_params = Scheduler::to_underlying_arguments(scheduler_args);
|
||||
// Get the ptr to kernel function.
|
||||
void *kernel;
|
||||
kernel = (void *)flash::compute_attn_ws<Kernel_traits, Is_causal, Scheduler>;
|
||||
int smem_size = sizeof(typename Kernel_traits::SharedStorage);
|
||||
if (smem_size >= 48 * 1024) {
|
||||
C10_CUDA_CHECK(cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size));
|
||||
}
|
||||
static constexpr int ctaSize = Kernel_traits::kNWarps * 32;
|
||||
params.m_block_divmod = cutlass::FastDivmod(num_blocks_m);
|
||||
params.total_blocks = num_blocks_m * params.h * params.b;
|
||||
dim3 grid_dims = Scheduler::get_grid_dim(scheduler_args, 170);
|
||||
dim3 block_dims(ctaSize);
|
||||
dim3 cluster_dims(size<0>(ClusterShape{}), size<1>(ClusterShape{}), size<2>(ClusterShape{}));
|
||||
cutlass::ClusterLaunchParams launch_params{grid_dims, block_dims, cluster_dims, smem_size, stream};
|
||||
cutlass::launch_kernel_on_cluster(launch_params, kernel, params, mainloop_params, epilogue_params, scheduler_params);
|
||||
|
||||
C10_CUDA_KERNEL_LAUNCH_CHECK();
|
||||
}
|
||||
|
||||
|
||||
template<typename T, int Headdim, typename O = cutlass::bfloat16_t>
|
||||
void run_mha_fwd_(Flash_fwd_params ¶ms, cudaStream_t stream) {
|
||||
BOOL_SWITCH(params.is_causal, Is_causal, [&] {
|
||||
BOOL_SWITCH(params.per_block_mean, per_block, [&] {
|
||||
if constexpr (Headdim == 64 || Headdim == 128) {
|
||||
run_flash_fwd<
|
||||
Flash_fwd_kernel_traits<Headdim, flash::BLOCK_M, flash::BLOCK_N, 3, 1, per_block, T, O>,
|
||||
Is_causal
|
||||
>(params, stream);
|
||||
} else {
|
||||
static_assert(Headdim == 64 || Headdim == 128, "Unsupported Headdim");
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
@@ -0,0 +1,920 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/array.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
#include <cutlass/numeric_conversion.h>
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cutlass/gemm/collective/collective_builder.hpp"
|
||||
|
||||
#include "utils.h"
|
||||
#include "named_barrier.h"
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <typename Ktraits, bool Is_causal>
|
||||
struct CollectiveMainloopFwd {
|
||||
|
||||
using Element = typename Ktraits::Element;
|
||||
using ElementSF = typename Ktraits::ElementSF;
|
||||
// using TMAElement = Element;
|
||||
// using TMAElementSF = typename Ktraits::ElementSF;
|
||||
using TileShape_MNK = typename Ktraits::TileShape_MNK;
|
||||
using ClusterShape = typename Ktraits::ClusterShape_MNK;
|
||||
|
||||
static constexpr int kStages = Ktraits::kStages;
|
||||
static constexpr int kHeadDim = Ktraits::kHeadDim;
|
||||
static constexpr int BlockMean = Ktraits::BlockMean;
|
||||
using GmemTiledCopy = typename Ktraits::GmemTiledCopy;
|
||||
using SmemLayoutQ = typename Ktraits::SmemLayoutQ;
|
||||
using SmemLayoutK = typename Ktraits::SmemLayoutK;
|
||||
using SmemLayoutV = typename Ktraits::SmemLayoutV;
|
||||
using SmemLayoutVt = typename Ktraits::SmemLayoutVt;
|
||||
using SmemLayoutDS = typename Ktraits::SmemLayoutDS;
|
||||
using SmemLayoutAtomDS = typename Ktraits::SmemLayoutAtomDS;
|
||||
using LayoutDS = decltype(
|
||||
blocked_product(
|
||||
SmemLayoutAtomDS{},
|
||||
make_layout(
|
||||
make_shape(int32_t(0), int32_t(0), int32_t(0), int32_t(0)),
|
||||
make_stride(int32_t(0), _1{}, int32_t(0), int32_t(0)))
|
||||
)
|
||||
);
|
||||
using ShapeQKV = cute::Shape<int32_t, int32_t, int32_t, int32_t>; // (seqlen, d, head, batch)
|
||||
using StrideQKV = cute::Stride<int64_t, _1, int64_t, int64_t>;
|
||||
using ShapeSF = cute::Shape<int32_t, int32_t, int32_t, int32_t>; // (seqlen, d // 16, head, batch)
|
||||
using LayoutSF = typename Ktraits::LayoutSF;
|
||||
using LayoutP = typename Ktraits::LayoutP;
|
||||
using LayoutSFP = typename Ktraits::LayoutSFP;
|
||||
using SfAtom = typename Ktraits::SfAtom;
|
||||
using TMA_Q = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element const*>(nullptr)), repeat_like(StrideQKV{}, int32_t(0)), StrideQKV{}),
|
||||
SmemLayoutQ{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{}));
|
||||
|
||||
using TMA_KV = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element const*>(nullptr)), repeat_like(StrideQKV{}, int32_t(0)), StrideQKV{}),
|
||||
take<0, 2>(SmemLayoutK{}),
|
||||
select<1, 2>(TileShape_MNK{}),
|
||||
_1{}));
|
||||
|
||||
using TMA_Vt = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<Element const*>(nullptr)), repeat_like(StrideQKV{}, int32_t(0)), StrideQKV{}),
|
||||
take<0, 2>(SmemLayoutVt{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using TMA_DS = decltype(make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
make_tensor(make_gmem_ptr(static_cast<float const*>(nullptr)), LayoutDS{}),
|
||||
take<0, 2>(SmemLayoutDS{}),
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using BlkScaledConfig = typename Ktraits::BlkScaledConfig;
|
||||
using GmemTiledCopySF = typename Ktraits::GmemTiledCopySF;
|
||||
using SmemLayoutSFQ = typename Ktraits::SmemLayoutSFQ;
|
||||
using SmemLayoutSFK = typename Ktraits::SmemLayoutSFK;
|
||||
using SmemLayoutSFV = typename Ktraits::SmemLayoutSFV;
|
||||
using SmemLayoutSFVt = typename Ktraits::SmemLayoutSFVt;
|
||||
|
||||
using TMA_SFQ = decltype(make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
make_tensor(static_cast<ElementSF const*>(nullptr), LayoutSF{}),
|
||||
SmemLayoutSFQ{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{})); // No programmatic multicast
|
||||
|
||||
|
||||
using TMA_SFKV = decltype(make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
make_tensor(static_cast<ElementSF const*>(nullptr), LayoutSF{}),
|
||||
SmemLayoutSFK{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using TMA_SFVt = decltype(make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
make_tensor(static_cast<ElementSF const*>(nullptr), LayoutSF{}),
|
||||
SmemLayoutSFVt{}(_,_,cute::Int<0>{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}));
|
||||
|
||||
using SmemCopyAtomQ = typename Ktraits::SmemCopyAtomQ;
|
||||
using SmemCopyAtomKV = typename Ktraits::SmemCopyAtomKV;
|
||||
using SmemCopyAtomSF = typename Ktraits::SmemCopyAtomSF;
|
||||
using TiledMmaQK = typename Ktraits::TiledMmaQK;
|
||||
using TiledMmaPV = typename Ktraits::TiledMmaPV;
|
||||
static constexpr int NumMmaThreads = size(TiledMmaQK{});
|
||||
using MainloopPipeline = typename Ktraits::MainloopPipeline;
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename MainloopPipeline::PipelineState;
|
||||
using MainloopPipelineQ = typename Ktraits::MainloopPipelineQ;
|
||||
using PipelineParamsQ = typename Ktraits::PipelineParamsQ;
|
||||
using PipelineStateQ = typename Ktraits::PipelineStateQ;
|
||||
using EpilogueBarrier = typename Ktraits::EpilogueBarrier;
|
||||
|
||||
// Set the bytes transferred in this TMA transaction (may involve multiple issues)
|
||||
static constexpr uint32_t TmaTransactionBytesQ = static_cast<uint32_t>(
|
||||
cutlass::bits_to_bytes(cosize((SmemLayoutSFQ{})) * cute::sizeof_bits_v<ElementSF>) +
|
||||
cutlass::bits_to_bytes(size((SmemLayoutQ{})) * sizeof_bits<Element>::value));
|
||||
|
||||
static constexpr uint32_t TmaTransactionBytesK = static_cast<uint32_t>(
|
||||
cutlass::bits_to_bytes(cosize(take<0,2>(SmemLayoutSFK{})) * cute::sizeof_bits_v<ElementSF>) +
|
||||
cutlass::bits_to_bytes(cosize(take<0,2>(SmemLayoutDS{})) * cute::sizeof_bits_v<float>) +
|
||||
cutlass::bits_to_bytes(size(take<0,2>(SmemLayoutK{})) * sizeof_bits<Element>::value));
|
||||
|
||||
static constexpr uint32_t TmaTransactionBytesV = static_cast<uint32_t>(
|
||||
cutlass::bits_to_bytes(cosize(take<0,2>(SmemLayoutSFVt{})) * cute::sizeof_bits_v<ElementSF>) +
|
||||
cutlass::bits_to_bytes(size(take<0,2>(SmemLayoutVt{})) * sizeof_bits<Element>::value));
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
Element const* ptr_Q;
|
||||
ShapeQKV const shape_Q;
|
||||
StrideQKV const stride_Q;
|
||||
Element const* ptr_K;
|
||||
ShapeQKV const shape_K;
|
||||
StrideQKV const stride_K;
|
||||
ShapeQKV const unpadded_shape_K;
|
||||
Element const* ptr_Vt;
|
||||
ShapeQKV const shape_Vt;
|
||||
StrideQKV const stride_Vt;
|
||||
ElementSF const* ptr_SFQ{nullptr};
|
||||
ShapeSF const shape_SFQ{};
|
||||
ElementSF const* ptr_SFK{nullptr};
|
||||
ShapeSF const shape_SFK{};
|
||||
ElementSF const* ptr_SFVt{nullptr};
|
||||
ShapeSF const shape_SFVt{};
|
||||
float const* ptr_ds;
|
||||
ShapeQKV const shape_ds;
|
||||
StrideQKV const stride_ds;
|
||||
float const softmax_scale_log2;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
ShapeQKV const shape_Q;
|
||||
LayoutSF const layout_SFQ;
|
||||
ShapeQKV const shape_K;
|
||||
ShapeQKV const unpadded_shape_K;
|
||||
LayoutSF const layout_SFK;
|
||||
ShapeQKV const shape_Vt;
|
||||
LayoutSF const layout_SFVt;
|
||||
LayoutDS const layout_DS;
|
||||
TMA_Q tma_load_Q;
|
||||
TMA_SFQ tma_load_SFQ;
|
||||
TMA_KV tma_load_K;
|
||||
TMA_SFKV tma_load_SFK;
|
||||
TMA_Vt tma_load_Vt;
|
||||
TMA_SFVt tma_load_SFVt;
|
||||
TMA_DS tma_load_DS;
|
||||
float const softmax_scale_log2;
|
||||
};
|
||||
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
Tensor mQ = make_tensor(make_gmem_ptr(args.ptr_Q), args.shape_Q, args.stride_Q);
|
||||
TMA_Q tma_load_Q = make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
mQ,
|
||||
SmemLayoutQ{},
|
||||
select<0, 2>(TileShape_MNK{}),
|
||||
_1{}); // no mcast for Q
|
||||
Tensor mK = make_tensor(make_gmem_ptr(args.ptr_K), args.shape_K, args.stride_K);
|
||||
TMA_KV tma_load_K = make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
mK,
|
||||
SmemLayoutK{}(_, _, _0{}),
|
||||
select<1, 2>(TileShape_MNK{}),
|
||||
_1{}); // mcast along M mode for this N load, if any
|
||||
Tensor mVt = make_tensor(make_gmem_ptr(args.ptr_Vt), args.shape_Vt, args.stride_Vt);
|
||||
TMA_Vt tma_load_Vt = make_tma_copy(
|
||||
GmemTiledCopy{},
|
||||
mVt,
|
||||
SmemLayoutVt{}(_, _, _0{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{}); // mcast along M mode for this N load, if any
|
||||
auto [Seqlen_Q, Seqlen_K, HeadNum, Batch] = args.shape_ds;
|
||||
LayoutDS layout_ds = tile_to_shape(SmemLayoutAtomDS{}, make_shape(Seqlen_Q, Seqlen_K, HeadNum, Batch), Step<_2,_1,_3,_4>{});
|
||||
Tensor mDS = make_tensor(make_gmem_ptr(args.ptr_ds), layout_ds);
|
||||
TMA_DS tma_load_ds = make_tma_copy (
|
||||
GmemTiledCopy{},
|
||||
mDS,
|
||||
SmemLayoutDS{}(_, _, _0{}),
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{});
|
||||
LayoutSF layout_sfq = BlkScaledConfig::tile_atom_to_shape_SFQKV(args.shape_SFQ);
|
||||
Tensor mSFQ = make_tensor(make_gmem_ptr(args.ptr_SFQ), layout_sfq);
|
||||
TMA_SFQ tma_load_sfq = make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
mSFQ,
|
||||
SmemLayoutSFQ{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{});
|
||||
LayoutSF layout_sfk = BlkScaledConfig::tile_atom_to_shape_SFQKV(args.shape_SFK);
|
||||
Tensor mSFK = make_tensor(make_gmem_ptr(args.ptr_SFK), layout_sfk);
|
||||
TMA_SFKV tma_load_sfk = make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
mSFK,
|
||||
SmemLayoutSFK{}(_, _, _0{}),
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{})),
|
||||
_1{});
|
||||
LayoutSF layout_sfvt = BlkScaledConfig::tile_atom_to_shape_SFVt(args.shape_SFVt);
|
||||
Tensor mSFVt = make_tensor(make_gmem_ptr(args.ptr_SFVt), layout_sfvt);
|
||||
TMA_SFVt tma_load_sfvt = make_tma_copy<uint16_t>(
|
||||
GmemTiledCopySF{},
|
||||
mSFVt,
|
||||
SmemLayoutSFVt{}(_, _, _0{}),
|
||||
make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})),
|
||||
_1{});
|
||||
return {args.shape_Q, layout_sfq,
|
||||
args.shape_K, args.unpadded_shape_K, layout_sfk,
|
||||
args.shape_Vt, layout_sfvt,
|
||||
layout_ds,
|
||||
tma_load_Q, tma_load_sfq,
|
||||
tma_load_K, tma_load_sfk,
|
||||
tma_load_Vt, tma_load_sfvt,
|
||||
tma_load_ds,
|
||||
args.softmax_scale_log2};
|
||||
}
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE
|
||||
static void prefetch_tma_descriptors(Params const& mainloop_params) {
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_Q.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_K.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_Vt.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_SFQ.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_SFK.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_SFVt.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_DS.get_tma_descriptor());
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
int get_n_block_max(Params const& mainloop_params, int m_block) {
|
||||
static constexpr int kBlockM = get<0>(TileShape_MNK{});
|
||||
static constexpr int kBlockN = get<1>(TileShape_MNK{});
|
||||
int const seqlen_q = get<0>(mainloop_params.shape_Q);
|
||||
int const seqlen_k = get<0>(mainloop_params.shape_K);
|
||||
int n_block_max = cute::ceil_div(seqlen_k, kBlockN);
|
||||
if constexpr (Is_causal) {
|
||||
n_block_max = std::min(n_block_max,
|
||||
cute::ceil_div((m_block + 1) * kBlockM + seqlen_k - seqlen_q, kBlockN));
|
||||
}
|
||||
return n_block_max;
|
||||
}
|
||||
|
||||
template <class SFATensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFA(SFATensor&& sfatensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfatensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFA_TV = typename Atom::Traits::SFALayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<0>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfatensor, t_tile); // (PermM,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<0>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomM,AtomK),(RestM,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFA_TV{},_); // ((ThrV,FrgV),(RestM,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<1>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrM,ThrK)),(FrgV,(RestM,RestK)))
|
||||
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFBTensor, class Atom, class TiledThr, class TiledPerm>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
thrfrg_SFB(SFBTensor&& sfbtensor, TiledMMA<Atom, TiledThr, TiledPerm>& mma)
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(sfbtensor) >= Int<2>{});
|
||||
|
||||
using AtomShape_MNK = typename Atom::Shape_MNK;
|
||||
using AtomLayoutSFB_TV = typename Atom::Traits::SFBLayout;
|
||||
|
||||
auto permutation_mnk = TiledPerm{};
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<1>(permutation_mnk),
|
||||
get<2>(permutation_mnk));
|
||||
auto t_tensor = logical_divide(sfbtensor, t_tile); // (PermN,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
auto a_tile = make_tile(make_layout(size<1>(AtomShape_MNK{})),
|
||||
make_layout(size<2>(AtomShape_MNK{})));
|
||||
auto a_tensor = zipped_divide(t_tensor, a_tile); // ((AtomN,AtomK),(RestN,RestK))
|
||||
|
||||
// Transform the Atom mode from (M,K) to (Thr,Val)
|
||||
auto tv_tensor = a_tensor.compose(AtomLayoutSFB_TV{},_); // ((ThrV,FrgV),(RestN,RestK))
|
||||
|
||||
// Tile the tensor for the Thread
|
||||
auto thr_tile = make_tile(_,
|
||||
make_tile(make_layout(size<2>(thr_layout_vmnk)),
|
||||
make_layout(size<3>(thr_layout_vmnk))));
|
||||
auto thr_tensor = zipped_divide(tv_tensor, thr_tile); // ((ThrV,(ThrN,ThrK)),(FrgV,(RestN,RestK)))
|
||||
return thr_tensor;
|
||||
}
|
||||
|
||||
template <class SFATensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFA(SFATensor&& sfatensor, ThrMma& thread_mma)
|
||||
{
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
auto thr_tensor = make_tensor(static_cast<SFATensor&&>(sfatensor).data(), thrfrg_SFA(sfatensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vmk = make_coord(get<0>(thr_vmnk), make_coord(get<1>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
auto partition_SFA = thr_tensor(thr_vmk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
return make_fragment_like<ValTypeSF>(partition_SFA);
|
||||
}
|
||||
|
||||
template <class SFBTensor, class ThrMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
partition_fragment_SFB(SFBTensor&& sfbtensor, ThrMma& thread_mma)
|
||||
{
|
||||
using ValTypeSF = typename ThrMma::Atom::Traits::ValTypeSF;
|
||||
auto thr_tensor = make_tensor(static_cast<SFBTensor&&>(sfbtensor).data(), thrfrg_SFB(sfbtensor.layout(),thread_mma));
|
||||
auto thr_vmnk = thread_mma.thr_vmnk_;
|
||||
auto thr_vnk = make_coord(get<0>(thr_vmnk), make_coord(get<2>(thr_vmnk), get<3>(thr_vmnk)));
|
||||
auto partition_SFB = thr_tensor(thr_vnk, make_coord(_, repeat<rank<1,1>(thr_tensor)>(_)));
|
||||
return make_fragment_like<ValTypeSF>(partition_SFB);
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFA_TV(TiledMma& mma)
|
||||
{
|
||||
// (M,K) -> (M,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_A = make_layout(make_shape(size<0>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto atile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<1>{} , Int<0>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFA(ref_A, mma).compose(atile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
template<class TiledMma>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_layoutSFB_TV(TiledMma& mma)
|
||||
{
|
||||
// (N,K) -> (N,K)
|
||||
auto tile_shape_mnk = tile_shape(mma);
|
||||
auto ref_B = make_layout(make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk)));
|
||||
auto thr_layout_vmnk = mma.get_thr_layout_vmnk();
|
||||
|
||||
// (ThrV,(ThrM,ThrK)) -> (ThrV,(ThrM,ThrN,ThrK))
|
||||
auto btile = make_tile(_,
|
||||
make_tile(make_layout(make_shape (size<1>(thr_layout_vmnk), size<2>(thr_layout_vmnk)),
|
||||
make_stride( Int<0>{} , Int<1>{} )),
|
||||
_));
|
||||
|
||||
// thr_idx -> (ThrV,ThrM,ThrN,ThrK)
|
||||
auto thridx_2_thrid = right_inverse(thr_layout_vmnk);
|
||||
// (thr_idx,val) -> (M,K)
|
||||
return thrfrg_SFB(ref_B, mma).compose(btile, _).compose(thridx_2_thrid, _);
|
||||
}
|
||||
|
||||
template <typename SchedulerParams, typename SharedStorage, typename WorkTileInfo>
|
||||
CUTLASS_DEVICE void
|
||||
load(Params const& mainloop_params,
|
||||
SchedulerParams const& scheduler_params,
|
||||
MainloopPipelineQ pipeline_q,
|
||||
MainloopPipeline pipeline_k,
|
||||
MainloopPipeline pipeline_v,
|
||||
PipelineStateQ& smem_pipe_write_q,
|
||||
PipelineState& smem_pipe_write_k,
|
||||
PipelineState& smem_pipe_write_v,
|
||||
SharedStorage &shared_storage,
|
||||
WorkTileInfo work_tile_info,
|
||||
int& work_idx,
|
||||
int& tile_count_semaphore
|
||||
) {
|
||||
|
||||
static constexpr int kBlockM = get<0>(TileShape_MNK{});
|
||||
static constexpr int kBlockN = get<1>(TileShape_MNK{});
|
||||
|
||||
auto [m_block, bidh, bidb] = work_tile_info.get_block_coord(scheduler_params);
|
||||
|
||||
int n_block_max = get_n_block_max(mainloop_params, m_block);
|
||||
|
||||
Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.begin()), SmemLayoutQ{});
|
||||
Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.begin()), SmemLayoutK{});
|
||||
Tensor sVt = make_tensor(make_smem_ptr(shared_storage.smem_v.begin()), SmemLayoutVt{});
|
||||
Tensor sSFQ = make_tensor(make_smem_ptr(shared_storage.smem_SFQ.begin()), SmemLayoutSFQ{});
|
||||
Tensor sSFK = make_tensor(make_smem_ptr(shared_storage.smem_SFK.begin()), SmemLayoutSFK{});
|
||||
Tensor sSFVt = make_tensor(make_smem_ptr(shared_storage.smem_SFV.begin()), SmemLayoutSFVt{});
|
||||
Tensor sDS = make_tensor(make_smem_ptr(shared_storage.smem_ds.begin()), SmemLayoutDS{});
|
||||
|
||||
Tensor mQ = mainloop_params.tma_load_Q.get_tma_tensor(mainloop_params.shape_Q);
|
||||
Tensor mK = mainloop_params.tma_load_K.get_tma_tensor(mainloop_params.shape_K);
|
||||
Tensor mVt = mainloop_params.tma_load_Vt.get_tma_tensor(mainloop_params.shape_Vt);
|
||||
Tensor mDS = mainloop_params.tma_load_DS.get_tma_tensor(shape(mainloop_params.layout_DS));
|
||||
Tensor mSFQ = mainloop_params.tma_load_SFQ.get_tma_tensor(shape(mainloop_params.layout_SFQ));
|
||||
Tensor mSFK = mainloop_params.tma_load_SFK.get_tma_tensor(shape(mainloop_params.layout_SFK));
|
||||
Tensor mSFVt = mainloop_params.tma_load_SFVt.get_tma_tensor(shape(mainloop_params.layout_SFVt));
|
||||
uint32_t block_rank_in_cluster = cute::block_rank_in_cluster();
|
||||
constexpr uint32_t cluster_shape_x = get<0>(ClusterShape());
|
||||
uint2 cluster_local_block_id = {block_rank_in_cluster % cluster_shape_x, block_rank_in_cluster / cluster_shape_x};
|
||||
Tensor gQ = local_tile(mQ(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{})); // (M, K)
|
||||
Tensor gK = local_tile(mK(_, _, bidh, bidb), select<1, 2>(TileShape_MNK{}), make_coord(_, _0{})); // (N, K, _)
|
||||
Tensor gVt = local_tile(mVt(_, _, bidh, bidb), make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})), make_coord(_0{}, _)); // (N, K, _)
|
||||
Tensor gDS = [&] {
|
||||
if constexpr (BlockMean) {
|
||||
return local_tile(mDS(_, _, bidh, bidb), select<0, 1>(TileShape_MNK{}), make_coord(m_block, _));
|
||||
} else {
|
||||
return local_tile(mDS(_, _, bidh, bidb), select<0, 1>(TileShape_MNK{}), make_coord(_0{}, _));
|
||||
}
|
||||
}();
|
||||
Tensor gSFQ = local_tile(mSFQ(_, _, bidh, bidb), select<0, 2>(TileShape_MNK{}), make_coord(m_block, _0{}));
|
||||
Tensor gSFK = local_tile(mSFK(_, _, bidh, bidb), select<1, 2>(TileShape_MNK{}), make_coord(_, _0{}));
|
||||
Tensor gSFVt = local_tile(mSFVt(_, _, bidh, bidb), make_shape(shape<2>(TileShape_MNK{}), shape<1>(TileShape_MNK{})), make_coord(_0{}, _));
|
||||
auto block_tma_q = mainloop_params.tma_load_Q.get_slice(_0{});
|
||||
Tensor tQgQ = block_tma_q.partition_S(gQ);
|
||||
Tensor tQsQ = block_tma_q.partition_D(sQ);
|
||||
auto block_tma_sfq = mainloop_params.tma_load_SFQ.get_slice(_0{});
|
||||
Tensor tQgSFQ = block_tma_sfq.partition_S(gSFQ);
|
||||
Tensor tQsSFQ = block_tma_sfq.partition_D(sSFQ);
|
||||
auto block_tma_k = mainloop_params.tma_load_K.get_slice(cluster_local_block_id.x);
|
||||
Tensor tKgK = group_modes<0, 3>(block_tma_k.partition_S(gK));
|
||||
Tensor tKsK = group_modes<0, 3>(block_tma_k.partition_D(sK));
|
||||
auto block_tma_sfk = mainloop_params.tma_load_SFK.get_slice(cluster_local_block_id.x);
|
||||
Tensor tKgSFK = group_modes<0, 3>(block_tma_sfk.partition_S(gSFK));
|
||||
Tensor tKsSFK = group_modes<0, 3>(block_tma_sfk.partition_D(sSFK));
|
||||
auto block_tma_vt = mainloop_params.tma_load_Vt.get_slice(cluster_local_block_id.x);
|
||||
Tensor tVgVt = group_modes<0, 3>(block_tma_vt.partition_S(gVt));
|
||||
Tensor tVsVt = group_modes<0, 3>(block_tma_vt.partition_D(sVt));
|
||||
auto block_tma_sfvt = mainloop_params.tma_load_SFVt.get_slice(cluster_local_block_id.x);
|
||||
Tensor tVgSFVt = group_modes<0, 3>(block_tma_sfvt.partition_S(gSFVt));
|
||||
Tensor tVsSFVt = group_modes<0, 3>(block_tma_sfvt.partition_D(sSFVt));
|
||||
auto block_tma_ds = mainloop_params.tma_load_DS.get_slice(cluster_local_block_id.x);
|
||||
Tensor tDSgDS = group_modes<0, 3>(block_tma_ds.partition_S(gDS));
|
||||
Tensor tDSsDS = group_modes<0, 3>(block_tma_ds.partition_D(sDS));
|
||||
uint16_t mcast_mask_kv = 0;
|
||||
|
||||
int n_block = n_block_max - 1;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
if (lane_predicate) {
|
||||
pipeline_q.producer_acquire(smem_pipe_write_q);
|
||||
copy(mainloop_params.tma_load_Q.with(*pipeline_q.producer_get_barrier(smem_pipe_write_q), 0), tQgQ, tQsQ);
|
||||
copy(mainloop_params.tma_load_SFQ.with(*pipeline_q.producer_get_barrier(smem_pipe_write_q), 0), tQgSFQ, tQsSFQ);
|
||||
++smem_pipe_write_q;
|
||||
pipeline_k.producer_acquire(smem_pipe_write_k);
|
||||
copy(mainloop_params.tma_load_K.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgK(_, n_block), tKsK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_SFK.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgSFK(_, n_block), tKsSFK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_DS.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tDSgDS(_, n_block), tDSsDS(_, smem_pipe_write_k.index()));
|
||||
++smem_pipe_write_k;
|
||||
pipeline_v.producer_acquire(smem_pipe_write_v);
|
||||
copy(mainloop_params.tma_load_Vt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgVt(_, n_block), tVsVt(_, smem_pipe_write_v.index()));
|
||||
copy(mainloop_params.tma_load_SFVt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgSFVt(_, n_block), tVsSFVt(_, smem_pipe_write_v.index()));
|
||||
++smem_pipe_write_v;
|
||||
}
|
||||
|
||||
n_block--;
|
||||
if (lane_predicate) {
|
||||
// CUTLASS_PRAGMA_NO_UNROLL
|
||||
#pragma unroll 2
|
||||
for (; n_block >= 0; --n_block) {
|
||||
pipeline_k.producer_acquire(smem_pipe_write_k);
|
||||
copy(mainloop_params.tma_load_K.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgK(_, n_block), tKsK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_SFK.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tKgSFK(_, n_block), tKsSFK(_, smem_pipe_write_k.index()));
|
||||
copy(mainloop_params.tma_load_DS.with(*pipeline_k.producer_get_barrier(smem_pipe_write_k), mcast_mask_kv),
|
||||
tDSgDS(_, n_block), tDSsDS(_, smem_pipe_write_k.index()));
|
||||
++smem_pipe_write_k;
|
||||
pipeline_v.producer_acquire(smem_pipe_write_v);
|
||||
copy(mainloop_params.tma_load_Vt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgVt(_, n_block), tVsVt(_, smem_pipe_write_v.index()));
|
||||
copy(mainloop_params.tma_load_SFVt.with(*pipeline_v.producer_get_barrier(smem_pipe_write_v), mcast_mask_kv),
|
||||
tVgSFVt(_, n_block), tVsSFVt(_, smem_pipe_write_v.index()));
|
||||
++smem_pipe_write_v;
|
||||
}
|
||||
}
|
||||
++work_idx;
|
||||
}
|
||||
|
||||
/// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipelineQ pipeline_q,
|
||||
MainloopPipeline pipeline_k,
|
||||
MainloopPipeline pipeline_v,
|
||||
PipelineStateQ& smem_pipe_write_q,
|
||||
PipelineState& smem_pipe_write_k,
|
||||
PipelineState& smem_pipe_write_v) {
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
// Issue the epilogue waits
|
||||
if (lane_predicate) {
|
||||
pipeline_q.producer_tail(smem_pipe_write_q);
|
||||
pipeline_k.producer_tail(smem_pipe_write_k);
|
||||
pipeline_v.producer_tail(smem_pipe_write_v);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename SharedStorage, typename FrgTensorO, typename SoftmaxFused>
|
||||
CUTLASS_DEVICE void
|
||||
mma(Params const& mainloop_params,
|
||||
MainloopPipelineQ pipeline_q,
|
||||
MainloopPipeline pipeline_k,
|
||||
MainloopPipeline pipeline_v,
|
||||
PipelineStateQ& smem_pipe_read_q,
|
||||
PipelineState& smem_pipe_read_k,
|
||||
PipelineState& smem_pipe_read_v,
|
||||
FrgTensorO& tOrO_store,
|
||||
SoftmaxFused& softmax_fused,
|
||||
int n_block_count,
|
||||
int thread_idx,
|
||||
int work_idx,
|
||||
int m_block,
|
||||
SharedStorage& shared_storage
|
||||
) {
|
||||
|
||||
static_assert(is_rmem<FrgTensorO>::value, "O tensor must be rmem resident.");
|
||||
|
||||
static constexpr int kBlockM = get<0>(TileShape_MNK{});
|
||||
static constexpr int kBlockN = get<1>(TileShape_MNK{});
|
||||
static constexpr int kBlockK = get<2>(TileShape_MNK{});
|
||||
Tensor sQ = make_tensor(make_smem_ptr(shared_storage.smem_q.begin()), SmemLayoutQ{});
|
||||
Tensor sK = make_tensor(make_smem_ptr(shared_storage.smem_k.begin()), SmemLayoutK{});
|
||||
Tensor sVt = make_tensor(make_smem_ptr(shared_storage.smem_v.begin()), SmemLayoutVt{});
|
||||
Tensor sDS = make_tensor(make_smem_ptr(shared_storage.smem_ds.begin()), SmemLayoutDS{});
|
||||
Tensor sSFQ = make_tensor(make_smem_ptr(shared_storage.smem_SFQ.begin()), SmemLayoutSFQ{});
|
||||
Tensor sSFK = make_tensor(make_smem_ptr(shared_storage.smem_SFK.begin()), SmemLayoutSFK{});
|
||||
Tensor sSFVt = make_tensor(make_smem_ptr(shared_storage.smem_SFV.begin()), SmemLayoutSFVt{});
|
||||
|
||||
Tensor cQ = make_identity_tensor(make_shape(size<0>(sQ), size<1>(sQ)));
|
||||
Tensor cKV = make_identity_tensor(make_shape(size<0>(sK), size<1>(sK)));
|
||||
TiledMmaQK tiled_mma_qk;
|
||||
TiledMmaPV tiled_mma_pv;
|
||||
auto thread_mma_qk = tiled_mma_qk.get_thread_slice(thread_idx);
|
||||
auto thread_mma_pv = tiled_mma_pv.get_thread_slice(thread_idx);
|
||||
|
||||
Tensor tSrQ = thread_mma_qk.partition_fragment_A(sQ);
|
||||
Tensor tSrK = thread_mma_qk.partition_fragment_B(sK(_,_,Int<0>{}));
|
||||
Tensor tOrVt = thread_mma_pv.partition_fragment_B(sVt(_,_,Int<0>{}));
|
||||
Tensor tOrP = make_tensor_like<Element>(LayoutP{});
|
||||
Tensor tSrSFQ = partition_fragment_SFA(sSFQ, thread_mma_qk);
|
||||
Tensor tSrSFK = partition_fragment_SFB(sSFK(_,_,Int<0>{}), thread_mma_qk);
|
||||
Tensor tOrSFVt = partition_fragment_SFB(sSFVt(_,_,Int<0>{}), thread_mma_pv);
|
||||
Tensor tOrSFP = make_tensor<ElementSF>(LayoutSFP{});
|
||||
Tensor tOrSFP_flt = filter_zeros(tOrSFP);
|
||||
Tensor tSrDS = make_tensor<float>(make_shape(_8{}, _4{}), make_stride(_1{}, _8{}));
|
||||
// copy qk and sf from smem to rmem
|
||||
auto smem_tiled_copy_Q = make_tiled_copy_A(SmemCopyAtomQ{}, tiled_mma_qk);
|
||||
auto smem_thr_copy_Q = smem_tiled_copy_Q.get_thread_slice(thread_idx);
|
||||
Tensor tSsQ = smem_thr_copy_Q.partition_S(as_position_independent_swizzle_tensor(sQ));
|
||||
Tensor tSrQ_copy_view = smem_thr_copy_Q.retile_D(tSrQ);
|
||||
|
||||
auto smem_tiled_copy_K = make_tiled_copy_B(SmemCopyAtomKV{}, tiled_mma_qk);
|
||||
auto smem_thr_copy_K = smem_tiled_copy_K.get_thread_slice(thread_idx);
|
||||
Tensor tSsK = smem_thr_copy_K.partition_S(as_position_independent_swizzle_tensor(sK));
|
||||
Tensor tSrK_copy_view = smem_thr_copy_K.retile_D(tSrK);
|
||||
|
||||
auto smem_tiled_copy_V = make_tiled_copy_B(SmemCopyAtomKV{}, tiled_mma_pv);
|
||||
auto smem_thr_copy_V = smem_tiled_copy_V.get_thread_slice(thread_idx);
|
||||
Tensor tOsVt = smem_thr_copy_V.partition_S(as_position_independent_swizzle_tensor(sVt));
|
||||
Tensor tOrVt_copy_view = smem_thr_copy_V.retile_D(tOrVt);
|
||||
|
||||
auto tile_shape_mnk = tile_shape(tiled_mma_qk);
|
||||
auto smem_tiled_copy_SFQ = make_tiled_copy_impl(SmemCopyAtomSF{},
|
||||
get_layoutSFA_TV(tiled_mma_qk),
|
||||
make_shape(size<0>(tile_shape_mnk), size<2>(tile_shape_mnk))
|
||||
);
|
||||
auto smem_thr_copy_SFQ = smem_tiled_copy_SFQ.get_thread_slice(thread_idx);
|
||||
Tensor tSsSFQ = smem_thr_copy_SFQ.partition_S(as_position_independent_swizzle_tensor(sSFQ));
|
||||
Tensor tSrSFQ_copy_view = smem_thr_copy_SFQ.retile_D(tSrSFQ);
|
||||
|
||||
auto smem_tiled_copy_SFK = make_tiled_copy_impl(SmemCopyAtomSF{},
|
||||
get_layoutSFB_TV(tiled_mma_qk),
|
||||
make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk))
|
||||
);
|
||||
auto smem_thr_copy_SFK = smem_tiled_copy_SFK.get_thread_slice(thread_idx);
|
||||
Tensor tSsSFK = smem_thr_copy_SFK.partition_S(as_position_independent_swizzle_tensor(sSFK));
|
||||
Tensor tSrSFK_copy_view = smem_thr_copy_SFK.retile_D(tSrSFK);
|
||||
|
||||
auto smem_tiled_copy_SFV = make_tiled_copy_impl(SmemCopyAtomSF{},
|
||||
get_layoutSFB_TV(tiled_mma_pv),
|
||||
make_shape(size<1>(tile_shape_mnk), size<2>(tile_shape_mnk))
|
||||
);
|
||||
auto smem_thr_copy_SFV = smem_tiled_copy_SFV.get_thread_slice(thread_idx);
|
||||
Tensor tOsSFVt = smem_thr_copy_SFV.partition_S(as_position_independent_swizzle_tensor(sSFVt));
|
||||
Tensor tOrSFVt_copy_view = smem_thr_copy_SFV.retile_D(tOrSFVt);
|
||||
|
||||
auto consumer_wait = [](auto& pipeline, auto& smem_pipe_read) {
|
||||
auto barrier_token = pipeline.consumer_try_wait(smem_pipe_read);
|
||||
pipeline.consumer_wait(smem_pipe_read, barrier_token);
|
||||
};
|
||||
|
||||
int const seqlen_q = get<0>(mainloop_params.shape_Q);
|
||||
int const seqlen_k = get<0>(mainloop_params.shape_K);
|
||||
int const unpadded_seqlen_k = get<0>(mainloop_params.unpadded_shape_K);
|
||||
int n_block = n_block_count - 1;
|
||||
|
||||
auto copy_k_block = [&](auto block_id) {
|
||||
auto tSsK_stage = tSsK(_, _, _, smem_pipe_read_k.index());
|
||||
auto tSsSFK_stage = tSsSFK(_, _, _, smem_pipe_read_k.index());
|
||||
copy(smem_tiled_copy_K, tSsK_stage(_, _, block_id), tSrK_copy_view(_, _, block_id));
|
||||
copy(smem_tiled_copy_SFK, tSsSFK_stage(_, _, block_id), tSrSFK_copy_view(_, _, block_id));
|
||||
};
|
||||
|
||||
auto copy_v_block = [&](auto block_id) {
|
||||
auto tOsVt_stage = tOsVt(_, _, _, smem_pipe_read_v.index());
|
||||
auto tOsSFVt_stage = tOsSFVt(_, _, _, smem_pipe_read_v.index());
|
||||
copy(smem_tiled_copy_V, tOsVt_stage(_, _, block_id), tOrVt_copy_view(_, _, block_id));
|
||||
copy(smem_tiled_copy_SFV, tOsSFVt_stage(_, _, block_id), tOrSFVt_copy_view(_, _, block_id));
|
||||
};
|
||||
// auto gemm_qk = [&](auto block_id) {
|
||||
// cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, block_id), tSrSFQ(_, _, block_id)), make_zip_tensor(tSrK(_, _, block_id), tSrSFK(_, _, block_id)), tSrS);
|
||||
// };
|
||||
// auto gemm_pv = [&](auto block_id) {
|
||||
// cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, block_id), tOrSFP(_, _, block_id)), make_zip_tensor(tOrVt(_, _, block_id), tOrSFVt(_, _, block_id)), tOrO);
|
||||
// };
|
||||
auto add_delta_s = [&](auto& acc) {
|
||||
// The MMA atom composites 4 sub-MMA m16n8k64 covering N=0-7, 8-15, 16-23, 24-31.
|
||||
// Each float4 register group spans two sub-MMAs, so N positions are scattered
|
||||
// (e.g., {2t, 2t+1, 8+2t, 9+2t}), not consecutive.
|
||||
float const* ds_ptr = reinterpret_cast<float const*>(
|
||||
&sDS(_0{}, _0{}, smem_pipe_read_k.index()));
|
||||
auto acc_float4 = recast<float4>(acc);
|
||||
int tid = threadIdx.x % 4;
|
||||
for (int i = 0; i < 4; i++) {
|
||||
int base_n = i * 32 + tid * 2;
|
||||
float4 delta_s_0 = make_float4(
|
||||
ds_ptr[base_n], ds_ptr[base_n + 1],
|
||||
ds_ptr[base_n + 8], ds_ptr[base_n + 9]);
|
||||
float4 delta_s_1 = make_float4(
|
||||
ds_ptr[base_n + 16], ds_ptr[base_n + 17],
|
||||
ds_ptr[base_n + 24], ds_ptr[base_n + 25]);
|
||||
acc_float4(make_coord(make_coord(_0{}, _0{}), _0{}), _0{}, i) = delta_s_0;
|
||||
acc_float4(make_coord(make_coord(_0{}, _0{}), _1{}), _0{}, i) = delta_s_0;
|
||||
acc_float4(make_coord(make_coord(_0{}, _1{}), _0{}), _0{}, i) = delta_s_1;
|
||||
acc_float4(make_coord(make_coord(_0{}, _1{}), _1{}), _0{}, i) = delta_s_1;
|
||||
}
|
||||
};
|
||||
consumer_wait(pipeline_q, smem_pipe_read_q);
|
||||
copy(smem_tiled_copy_Q, tSsQ, tSrQ_copy_view);
|
||||
copy(smem_tiled_copy_SFQ, tSsSFQ, tSrSFQ_copy_view);
|
||||
pipeline_q.consumer_release(smem_pipe_read_q);
|
||||
++smem_pipe_read_q;
|
||||
|
||||
Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout()));
|
||||
Tensor AbsMaxP = make_tensor_like<float>(
|
||||
make_layout(shape(group<1, 4>(flatten(tSrS_converion_view.layout()(make_coord(_0{}, _), _, _)))))
|
||||
);
|
||||
consumer_wait(pipeline_k, smem_pipe_read_k);
|
||||
copy_k_block(_0{});
|
||||
add_delta_s(tSrS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tSrQ); ++k_block) {
|
||||
cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, k_block), tSrSFQ(_, _, k_block)),
|
||||
make_zip_tensor(tSrK(_, _, k_block), tSrSFK(_, _, k_block)), tSrS);
|
||||
if (k_block < size<2>(tSrQ) - 1) {
|
||||
copy_k_block(k_block + 1);
|
||||
} else {
|
||||
pipeline_k.consumer_release(smem_pipe_read_k);
|
||||
++smem_pipe_read_k;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
auto col_limit_causal = [&](int row, int n_block) {
|
||||
return row + 1 + seqlen_k - n_block * kBlockN - seqlen_q + m_block * kBlockM;
|
||||
};
|
||||
{
|
||||
Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tScS = thread_mma_qk.partition_C(cS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(tSrS); ++i) {
|
||||
if constexpr (!Is_causal) { // Just masking based on col
|
||||
if (int(get<1>(tScS(i))) >= int(unpadded_seqlen_k - n_block * kBlockN)) { tSrS(i) = -INFINITY; }
|
||||
} else {
|
||||
if (int(get<1>(tScS(i))) >= std::min(seqlen_k - n_block * kBlockN,
|
||||
col_limit_causal(int(get<0>(tScS(i))), n_block))) {
|
||||
tSrS(i) = -INFINITY;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
auto quantize = [&](auto mma_k, auto acc_conversion_view) {
|
||||
Tensor AbsMaxP_stagek = AbsMaxP(_, make_coord(_, _, mma_k));
|
||||
Tensor acc_conversion_stagek = acc_conversion_view(_, _, mma_k);
|
||||
Tensor SFP = make_tensor_like<cutlass::float_ue4m3_t>(AbsMaxP_stagek.layout());
|
||||
Tensor SFP_uint32_view = recast<uint32_t>(SFP);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(AbsMaxP_stagek); i += 4) {
|
||||
uint32_t& tmp = SFP_uint32_view(i / 4);
|
||||
flash::packed_float_to_ue4m3(
|
||||
AbsMaxP_stagek(i),
|
||||
AbsMaxP_stagek(i + 1),
|
||||
AbsMaxP_stagek(i + 2),
|
||||
AbsMaxP_stagek(i + 3),
|
||||
tmp
|
||||
);
|
||||
}
|
||||
int const quad_id = threadIdx.x & 3;
|
||||
uint32_t MASK = (0xFF00FF) << ((quad_id & 1) * 8);
|
||||
Tensor tOrSFP_uint32_view = recast<uint32_t>(tOrSFP(_, _, mma_k));
|
||||
Tensor tOrP_uint32_view = recast<uint32_t>(tOrP(_, _, mma_k));
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mma_m = 0; mma_m < size<1>(tOrP); ++mma_m) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
flash::packed_float_to_e2m1(
|
||||
acc_conversion_stagek(make_coord(_0{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_1{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_2{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_3{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_4{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_5{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_6{}, i), mma_m),
|
||||
acc_conversion_stagek(make_coord(_7{}, i), mma_m),
|
||||
tOrP_uint32_view(i, mma_m)
|
||||
);
|
||||
}
|
||||
uint32_t local_sfp = SFP_uint32_view(_0{}, _0{}, mma_m);
|
||||
uint32_t peer_sfp = __shfl_xor_sync(int32_t(-1), local_sfp, 2);
|
||||
if ((quad_id & 1) == 0) {
|
||||
uint32_t sfp = (local_sfp & MASK) | ((peer_sfp & MASK) << 8);
|
||||
tOrSFP_uint32_view(_0{}, mma_m) = sfp;
|
||||
} else {
|
||||
uint32_t sfp = (peer_sfp & MASK) | ((local_sfp & MASK) >> 8);
|
||||
tOrSFP_uint32_view(_0{}, mma_m) = sfp;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
softmax_fused.template online_softmax_with_quant</*Is_first=*/true>(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2);
|
||||
|
||||
consumer_wait(pipeline_v, smem_pipe_read_v);
|
||||
copy_v_block(_0{});
|
||||
quantize(_0{}, tSrS_converion_view);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) {
|
||||
cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)),
|
||||
make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO_store);
|
||||
if (v_block < size<2>(tOrP) - 1) {
|
||||
copy_v_block(v_block + 1);
|
||||
quantize(v_block + 1, tSrS_converion_view);
|
||||
} else {
|
||||
pipeline_v.consumer_release(smem_pipe_read_v);
|
||||
++smem_pipe_read_v;
|
||||
}
|
||||
}
|
||||
|
||||
n_block--;
|
||||
constexpr int n_masking_steps = !Is_causal ? 1 : cute::ceil_div(kBlockM, kBlockN) + 1;
|
||||
// // Only go through these if Is_causal, since n_masking_steps = 1 when !Is_causal
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int masking_step = 0; masking_step < n_masking_steps - 1 && n_block >= 0; ++masking_step, --n_block) {
|
||||
Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout()));
|
||||
consumer_wait(pipeline_k, smem_pipe_read_k);
|
||||
copy_k_block(_0{});
|
||||
add_delta_s(tSrS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tSrQ); ++k_block) {
|
||||
cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, k_block), tSrSFQ(_, _, k_block)),
|
||||
make_zip_tensor(tSrK(_, _, k_block), tSrSFK(_, _, k_block)), tSrS);
|
||||
if (k_block < size<2>(tSrQ) - 1) {
|
||||
copy_k_block(k_block + 1);
|
||||
}
|
||||
}
|
||||
pipeline_k.consumer_release(smem_pipe_read_k); // release K
|
||||
++smem_pipe_read_k;
|
||||
Tensor cS = cute::make_identity_tensor(select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tScS = thread_mma_qk.partition_C(cS);
|
||||
#pragma unroll
|
||||
for (int i = 0; i < size(tSrS); ++i) {
|
||||
if (int(get<1>(tScS(i))) >= col_limit_causal(int(get<0>(tScS(i))), n_block)) {
|
||||
tSrS(i) = -INFINITY;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
softmax_fused.template online_softmax_with_quant</*Is_first=*/false>(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2);
|
||||
Tensor tOrO = make_fragment_like(tOrO_store);
|
||||
consumer_wait(pipeline_v, smem_pipe_read_v);
|
||||
copy_v_block(_0{});
|
||||
quantize(_0{}, tSrS_converion_view);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) {
|
||||
cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)),
|
||||
make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO);
|
||||
if (v_block < size<2>(tOrP) - 1) {
|
||||
copy_v_block(v_block + 1);
|
||||
quantize(v_block + 1, tSrS_converion_view);
|
||||
}
|
||||
}
|
||||
pipeline_v.consumer_release(smem_pipe_read_v);
|
||||
++smem_pipe_read_v;
|
||||
if (masking_step > 0) { softmax_fused.rescale_o(tOrO_store, tOrO); }
|
||||
}
|
||||
|
||||
#pragma unroll 1
|
||||
for (; n_block >= 0; --n_block) {
|
||||
Tensor tSrS = partition_fragment_C(tiled_mma_qk, select<0, 1>(TileShape_MNK{}));
|
||||
Tensor tSrS_converion_view = make_tensor(tSrS.data(), flash::convert_to_conversion_layout(tSrS.layout()));
|
||||
consumer_wait(pipeline_k, smem_pipe_read_k);
|
||||
copy_k_block(_0{});
|
||||
add_delta_s(tSrS);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tSrQ); ++k_block) {
|
||||
cute::gemm(tiled_mma_qk, make_zip_tensor(tSrQ(_, _, k_block), tSrSFQ(_, _, k_block)),
|
||||
make_zip_tensor(tSrK(_, _, k_block), tSrSFK(_, _, k_block)), tSrS);
|
||||
if (k_block < size<2>(tSrQ) - 1) {
|
||||
copy_k_block(k_block + 1);
|
||||
} else {
|
||||
pipeline_k.consumer_release(smem_pipe_read_k);
|
||||
++smem_pipe_read_k;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
softmax_fused.template online_softmax_with_quant</*Is_first=*/false>(tSrS, AbsMaxP, mainloop_params.softmax_scale_log2);
|
||||
Tensor tOrO = make_fragment_like(tOrO_store);
|
||||
consumer_wait(pipeline_v, smem_pipe_read_v);
|
||||
copy_v_block(_0{});
|
||||
quantize(_0{}, tSrS_converion_view);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int v_block = 0; v_block < size<2>(tOrP); ++v_block) {
|
||||
cute::gemm(tiled_mma_pv, make_zip_tensor(tOrP(_, _, v_block), tOrSFP(_, _, v_block)),
|
||||
make_zip_tensor(tOrVt(_, _, v_block), tOrSFVt(_, _, v_block)), tOrO);
|
||||
if (v_block < size<2>(tOrP) - 1) {
|
||||
copy_v_block(v_block + 1);
|
||||
quantize(v_block + 1, tSrS_converion_view);
|
||||
} else {
|
||||
pipeline_v.consumer_release(smem_pipe_read_v);
|
||||
++smem_pipe_read_v;
|
||||
}
|
||||
}
|
||||
softmax_fused.rescale_o(tOrO_store, tOrO);
|
||||
}
|
||||
softmax_fused.finalize(tOrO_store);
|
||||
return;
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace flash
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/arch/barrier.h"
|
||||
#include "cutlass/pipeline/sm90_pipeline.hpp"
|
||||
|
||||
namespace flash {
|
||||
|
||||
enum class FP4NamedBarriers {
|
||||
QueryEmpty = 1,
|
||||
WarpSpecializedConsumer = 2,
|
||||
WarpSpecializedPingPongConsumer1 = 3,
|
||||
WarpSpecializedPingPongConsumer2 = 4,
|
||||
ProducerEnd = 5,
|
||||
ConsumerEnd = 6,
|
||||
EpilogueBarrier = 7
|
||||
};
|
||||
|
||||
template<int SequenceDepth, int SequenceLength>
|
||||
struct OrderedSequenceBarrierVarGroupSizeSharedStorage {
|
||||
using Barrier = cutlass::arch::ClusterBarrier;
|
||||
Barrier barrier_[SequenceDepth][SequenceLength];
|
||||
};
|
||||
|
||||
template<int SequenceDepth_, int SequenceLength_>
|
||||
class OrderedSequenceBarrierVarGroupSize {
|
||||
public:
|
||||
static constexpr int SequenceDepth = SequenceDepth_;
|
||||
static constexpr int SequenceLength = SequenceLength_;
|
||||
using Barrier = cutlass::arch::ClusterBarrier;
|
||||
using SharedStorage = flash::OrderedSequenceBarrierVarGroupSizeSharedStorage<SequenceDepth, SequenceLength>;
|
||||
|
||||
|
||||
struct Params {
|
||||
uint32_t group_id;
|
||||
uint32_t* group_size_list;
|
||||
};
|
||||
|
||||
private :
|
||||
// In future this Params object can be replaced easily with a CG object
|
||||
Params params_;
|
||||
Barrier *barrier_ptr_;
|
||||
cutlass::PipelineState<SequenceDepth> stage_;
|
||||
|
||||
static constexpr int Depth = SequenceDepth;
|
||||
static constexpr int Length = SequenceLength;
|
||||
|
||||
public:
|
||||
OrderedSequenceBarrierVarGroupSize() = delete;
|
||||
OrderedSequenceBarrierVarGroupSize(const OrderedSequenceBarrierVarGroupSize&) = delete;
|
||||
OrderedSequenceBarrierVarGroupSize(OrderedSequenceBarrierVarGroupSize&&) = delete;
|
||||
OrderedSequenceBarrierVarGroupSize& operator=(const OrderedSequenceBarrierVarGroupSize&) = delete;
|
||||
OrderedSequenceBarrierVarGroupSize& operator=(OrderedSequenceBarrierVarGroupSize&&) = delete;
|
||||
~OrderedSequenceBarrierVarGroupSize() = default;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
OrderedSequenceBarrierVarGroupSize(SharedStorage& storage, Params const& params) :
|
||||
params_(params),
|
||||
barrier_ptr_(&storage.barrier_[0][0]),
|
||||
// Group 0 - starts with an opposite phase
|
||||
stage_({0, params.group_id == 0, 0}) {
|
||||
int warp_idx = cutlass::canonical_warp_idx_sync();
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Barrier FULL, EMPTY init
|
||||
// Init is done only by the one elected thread of the block
|
||||
if (warp_idx == 0 && lane_predicate) {
|
||||
for (int d = 0; d < Depth; ++d) {
|
||||
for (int l = 0; l < Length; ++l) {
|
||||
barrier_ptr_[d * Length + l].init(*(params.group_size_list + l));
|
||||
}
|
||||
}
|
||||
}
|
||||
cutlass::arch::fence_barrier_init();
|
||||
}
|
||||
|
||||
// Wait on a stage to be unlocked
|
||||
CUTLASS_DEVICE
|
||||
void wait() {
|
||||
get_barrier_for_current_stage(params_.group_id).wait(stage_.phase());
|
||||
}
|
||||
|
||||
// Signal completion of Stage and move to the next stage
|
||||
// (group_id) signals to (group_id+1)
|
||||
CUTLASS_DEVICE
|
||||
void arrive() {
|
||||
int signalling_id = (params_.group_id + 1) % Length;
|
||||
get_barrier_for_current_stage(signalling_id).arrive();
|
||||
++stage_;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void advance() {
|
||||
++stage_;
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
CUTLASS_DEVICE
|
||||
Barrier& get_barrier_for_current_stage(int group_id) {
|
||||
return barrier_ptr_[stage_.index() * Length + group_id];
|
||||
}
|
||||
};
|
||||
|
||||
} // flash
|
||||
@@ -0,0 +1,180 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cuda.h>
|
||||
#include <vector>
|
||||
|
||||
#ifdef OLD_GENERATOR_PATH
|
||||
#include <ATen/CUDAGeneratorImpl.h>
|
||||
#else
|
||||
#include <ATen/cuda/CUDAGeneratorImpl.h>
|
||||
#endif
|
||||
|
||||
#include <ATen/cuda/CUDAGraphsUtils.cuh> // For at::cuda::philox::unpack
|
||||
|
||||
#include "cutlass/fast_math.h" // For cutlass::FastDivmod
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct Qkv_params {
|
||||
using index_t = int64_t;
|
||||
// The QKV matrices.
|
||||
void *__restrict__ q_ptr;
|
||||
void *__restrict__ k_ptr;
|
||||
void *__restrict__ v_ptr;
|
||||
void *__restrict__ delta_s_ptr;
|
||||
// The QKV scale factor matrices.
|
||||
void *__restrict__ sfq_ptr;
|
||||
void *__restrict__ sfk_ptr;
|
||||
void *__restrict__ sfv_ptr;
|
||||
// The stride between rows of the Q, K and V matrices.
|
||||
index_t q_batch_stride;
|
||||
index_t k_batch_stride;
|
||||
index_t v_batch_stride;
|
||||
index_t q_row_stride;
|
||||
index_t k_row_stride;
|
||||
index_t v_row_stride;
|
||||
index_t q_head_stride;
|
||||
index_t k_head_stride;
|
||||
index_t v_head_stride;
|
||||
index_t ds_batch_stride;
|
||||
index_t ds_row_stride;
|
||||
index_t ds_head_stride;
|
||||
// The stride of the Q, K and V scale factor matrices.
|
||||
index_t sfq_batch_stride;
|
||||
index_t sfk_batch_stride;
|
||||
index_t sfv_batch_stride;
|
||||
index_t sfq_row_stride;
|
||||
index_t sfk_row_stride;
|
||||
index_t sfv_row_stride;
|
||||
index_t sfq_head_stride;
|
||||
index_t sfk_head_stride;
|
||||
index_t sfv_head_stride;
|
||||
|
||||
// The number of heads.
|
||||
int h, h_k;
|
||||
// In the case of multi-query and grouped-query attention (MQA/GQA), nheads_k could be
|
||||
// different from nheads (query).
|
||||
int h_h_k_ratio; // precompute h / h_k,
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
struct Flash_fwd_params : public Qkv_params {
|
||||
|
||||
// The O matrix (output).
|
||||
void * __restrict__ o_ptr;
|
||||
void * __restrict__ oaccum_ptr;
|
||||
void * __restrict__ s_ptr;
|
||||
|
||||
// The stride between rows of O.
|
||||
index_t o_batch_stride;
|
||||
index_t o_row_stride;
|
||||
index_t o_head_stride;
|
||||
|
||||
// The pointer to the P matrix.
|
||||
void * __restrict__ p_ptr;
|
||||
|
||||
// The pointer to the softmax sum.
|
||||
void * __restrict__ softmax_lse_ptr;
|
||||
void * __restrict__ softmax_lseaccum_ptr;
|
||||
|
||||
// The dimensions.
|
||||
int b, seqlen_q, seqlen_k, seqlen_knew, d, seqlen_q_rounded, seqlen_k_rounded, d_rounded, rotary_dim, unpadded_seqlen_k;
|
||||
cutlass::FastDivmod head_divmod, m_block_divmod;
|
||||
int total_blocks;
|
||||
int seqlen_s;
|
||||
|
||||
// The scaling factors for the kernel.
|
||||
float scale_softmax;
|
||||
float scale_softmax_log2;
|
||||
uint32_t scale_softmax_log2_half2;
|
||||
|
||||
// array of length b+1 holding starting offset of each sequence.
|
||||
int * __restrict__ cu_seqlens_q;
|
||||
int * __restrict__ cu_seqlens_k;
|
||||
|
||||
// If provided, the actual length of each k sequence.
|
||||
int * __restrict__ seqused_k;
|
||||
|
||||
int *__restrict__ blockmask;
|
||||
|
||||
// The K_new and V_new matrices.
|
||||
void * __restrict__ knew_ptr;
|
||||
void * __restrict__ vnew_ptr;
|
||||
|
||||
// The stride between rows of the Q, K and V matrices.
|
||||
index_t knew_batch_stride;
|
||||
index_t vnew_batch_stride;
|
||||
index_t knew_row_stride;
|
||||
index_t vnew_row_stride;
|
||||
index_t knew_head_stride;
|
||||
index_t vnew_head_stride;
|
||||
|
||||
// The cos and sin matrices for rotary embedding.
|
||||
void * __restrict__ rotary_cos_ptr;
|
||||
void * __restrict__ rotary_sin_ptr;
|
||||
|
||||
// The indices to index into the KV cache.
|
||||
int * __restrict__ cache_batch_idx;
|
||||
|
||||
// Paged KV cache
|
||||
int * __restrict__ block_table;
|
||||
index_t block_table_batch_stride;
|
||||
int page_block_size;
|
||||
|
||||
// The dropout probability (probability of keeping an activation).
|
||||
float p_dropout;
|
||||
// uint32_t p_dropout_in_uint;
|
||||
// uint16_t p_dropout_in_uint16_t;
|
||||
uint8_t p_dropout_in_uint8_t;
|
||||
|
||||
// Scale factor of 1 / (1 - p_dropout).
|
||||
float rp_dropout;
|
||||
float scale_softmax_rp_dropout;
|
||||
|
||||
// Local window size
|
||||
int window_size_left, window_size_right;
|
||||
|
||||
// Random state.
|
||||
at::PhiloxCudaState philox_args;
|
||||
|
||||
// Pointer to the RNG seed (idx 0) and offset (idx 1).
|
||||
uint64_t * rng_state;
|
||||
|
||||
bool is_bf16;
|
||||
bool is_e4m3;
|
||||
bool is_causal;
|
||||
bool per_block_mean;
|
||||
bool single_level_p_quant; // If true, use single-level 1x16 block scale quantization for P (like V), instead of two-level quantization
|
||||
// If is_seqlens_k_cumulative, then seqlen_k is cu_seqlens_k[bidb + 1] - cu_seqlens_k[bidb].
|
||||
// Otherwise it's cu_seqlens_k[bidb], i.e., we use cu_seqlens_k to store the sequence lengths of K.
|
||||
bool is_seqlens_k_cumulative;
|
||||
|
||||
bool is_rotary_interleaved;
|
||||
|
||||
int num_splits; // For split-KV version
|
||||
|
||||
void * __restrict__ alibi_slopes_ptr;
|
||||
index_t alibi_slopes_batch_stride;
|
||||
|
||||
int * __restrict__ tile_count_semaphore;
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,190 @@
|
||||
// Modified from the original SageAttention3 code
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include <cmath>
|
||||
#include "cute/tensor.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "utils.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
template <int Rows>
|
||||
struct SoftmaxFused{
|
||||
|
||||
using TensorT = decltype(make_fragment_like<float>(Shape<Int<Rows>>{}));
|
||||
TensorT row_sum, row_max, scores_scale;
|
||||
static constexpr float fp8_scalexfp4_scale = 1.f / (448 * 6);
|
||||
static constexpr float fp8_scalexfp4_scale_log2 = -11.392317422778762f; //log2f(fp8_scalexfp4_scale)
|
||||
static constexpr float fp4_scale_log2 = -2.584962500721156f; // log2f(fp4_scale)
|
||||
static constexpr int RowReductionThr = 4;
|
||||
|
||||
// If true, use single-level quantization: s_P2, P̂_2 = φ(P̃) directly (standard per-block FP4 quantization like V)
|
||||
// If false (default), use two-level quantization: s_P1 = rowmax(P̃)/(448×6), then s_P2, P̂_2 = φ(P̃/s_P1)
|
||||
bool single_level_p_quant;
|
||||
|
||||
CUTLASS_DEVICE SoftmaxFused(bool single_level = false) : single_level_p_quant(single_level) {};
|
||||
|
||||
template<bool FirstTile, bool InfCheck = false, typename TensorAcc, typename TensorMax>
|
||||
CUTLASS_DEVICE auto online_softmax_with_quant(
|
||||
TensorAcc& acc,
|
||||
TensorMax& AbsMaxP,
|
||||
const float softmax_scale_log2
|
||||
) {
|
||||
Tensor acc_reduction_view = make_tensor(acc.data(), flash::convert_to_reduction_layout(acc.layout()));
|
||||
Tensor acc_conversion_view = make_tensor(acc.data(), flash::convert_to_conversion_layout(acc.layout()));
|
||||
Tensor acc_conversion_flatten = group_modes<1, 5>(group_modes<0, 2>(flatten(acc_conversion_view)));
|
||||
|
||||
if constexpr (FirstTile) {
|
||||
fill(row_max, -INFINITY);
|
||||
clear(row_sum);
|
||||
fill(scores_scale, 1.f);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size<0>(acc_reduction_view); mi++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1, 1>(acc_reduction_view); ni++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ei = 0; ei < size<1, 0>(acc_reduction_view); ei++) {
|
||||
AbsMaxP(mi, ni) = fmaxf(AbsMaxP(mi, ni), acc_reduction_view(mi, make_coord(ei, ni)));
|
||||
}
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), AbsMaxP(mi, ni), 1); // exchange max with neighbour thread of 8 elements
|
||||
AbsMaxP(mi, ni) = fmaxf(AbsMaxP(mi, ni), max_recv);
|
||||
row_max(mi) = fmaxf(row_max(mi), AbsMaxP(mi, ni));
|
||||
}
|
||||
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), row_max(mi), 2); // exchange max in a quad in a row
|
||||
row_max(mi) = fmaxf(row_max(mi), max_recv);
|
||||
|
||||
// Two-level P quantization (default): s_P1 = rowmax(P̃)/(448×6), then s_P2,P̂_2 = φ(P̃/s_P1)
|
||||
// - Pre-scales P to [0, 448×6] range before φ, output scaled by s_P1
|
||||
// Single-level P quantization: s_P2, P̂_2 = φ(P̃) directly (like V quantization)
|
||||
// - No s_P1, just standard per-block FP4 quantization φ
|
||||
const float s_P1_offset = single_level_p_quant ? 0.f : fp8_scalexfp4_scale_log2;
|
||||
const float max_scaled = InfCheck
|
||||
? (row_max(mi) == -INFINITY ? 0.f : (row_max(mi) * softmax_scale_log2 + s_P1_offset))
|
||||
: (row_max(mi) * softmax_scale_log2 + s_P1_offset);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(acc_reduction_view); ni++) {
|
||||
acc_reduction_view(mi, ni) = flash::ptx_exp2(acc_reduction_view(mi, ni) * softmax_scale_log2 - max_scaled);
|
||||
}
|
||||
// s_P2 = max(P_block)/6 — per-block scale factor from φ function (same formula for both modes)
|
||||
// The difference is in max_scaled: two-level includes 448×6 pre-scaling, single-level doesn't
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int sfi = 0; sfi < size<1>(AbsMaxP); sfi++) {
|
||||
AbsMaxP(mi, sfi) = flash::ptx_exp2(AbsMaxP(mi, sfi) * softmax_scale_log2 - max_scaled + fp4_scale_log2);
|
||||
}
|
||||
}
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size<0>(acc_reduction_view); mi++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(acc_reduction_view); ni++) {
|
||||
row_sum(mi) += acc_reduction_view(mi, ni);
|
||||
}
|
||||
}
|
||||
}
|
||||
else {
|
||||
Tensor scores_max_prev = make_fragment_like(row_max);
|
||||
cute::copy(row_max, scores_max_prev);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size<0>(acc_reduction_view); mi++) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1, 1>(acc_reduction_view); ni++) {
|
||||
float local_max = -INFINITY;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ei = 0; ei < size<1, 0>(acc_reduction_view); ei++) {
|
||||
local_max = fmaxf(local_max, acc_reduction_view(mi, make_coord(ei, ni)));
|
||||
}
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), local_max, 1); // exchange max with neighbour thread of 8 elements
|
||||
AbsMaxP(mi, ni) = fmaxf(local_max, max_recv);
|
||||
row_max(mi) = fmaxf(row_max(mi), AbsMaxP(mi, ni));
|
||||
}
|
||||
|
||||
float max_recv = __shfl_xor_sync(int32_t(-1), row_max(mi), 2); // exchange max in a quad in a row
|
||||
row_max(mi) = fmaxf(row_max(mi), max_recv);
|
||||
|
||||
float scores_max_cur = !InfCheck
|
||||
? row_max(mi)
|
||||
: (row_max(mi) == -INFINITY ? 0.0f : row_max(mi));
|
||||
scores_scale(mi) = flash::ptx_exp2((scores_max_prev(mi) - scores_max_cur) * softmax_scale_log2);
|
||||
|
||||
// Two-level P quantization (default): s_P1 = rowmax(P̃)/(448×6), then s_P2,P̂_2 = φ(P̃/s_P1)
|
||||
// Single-level P quantization: s_P2, P̂_2 = φ(P̃) directly (like V quantization)
|
||||
const float s_P1_offset = single_level_p_quant ? 0.f : fp8_scalexfp4_scale_log2;
|
||||
const float max_scaled = InfCheck
|
||||
? (row_max(mi) == -INFINITY ? 0.f : (row_max(mi) * softmax_scale_log2 + s_P1_offset))
|
||||
: (row_max(mi) * softmax_scale_log2 + s_P1_offset);
|
||||
row_sum(mi) = row_sum(mi) * scores_scale(mi);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(acc_reduction_view); ni++) {
|
||||
acc_reduction_view(mi, ni) = flash::ptx_exp2(acc_reduction_view(mi, ni) * softmax_scale_log2 - max_scaled);
|
||||
row_sum(mi) += acc_reduction_view(mi, ni);
|
||||
}
|
||||
// s_P2 = max(P_block)/6 — per-block scale factor from φ function
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int sfi = 0; sfi < size<1>(AbsMaxP); sfi++) {
|
||||
AbsMaxP(mi, sfi) = flash::ptx_exp2(AbsMaxP(mi, sfi) * softmax_scale_log2 - max_scaled + fp4_scale_log2);
|
||||
}
|
||||
// scores_scale(mi) = max_scaled;
|
||||
}
|
||||
}
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 0; i < size(AbsMaxP); ++i) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int j = 0; j < size<0>(acc_conversion_flatten); ++j)
|
||||
acc_conversion_flatten(j, i) /= AbsMaxP(i);
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TensorAcc>
|
||||
CUTLASS_DEVICE void finalize(TensorAcc& o_store) {
|
||||
Tensor o_store_reduction_view = make_tensor(o_store.data(), flash::convert_to_reduction_layout(o_store.layout()));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size(row_max); ++mi) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int i = 1; i < RowReductionThr; i <<= 1) {
|
||||
float sum_recv = __shfl_xor_sync(int32_t(-1), row_sum(mi), i);
|
||||
row_sum(mi) += sum_recv;
|
||||
}
|
||||
float sum = row_sum(mi);
|
||||
float inv_sum = (sum == 0.f || sum != sum) ? 0.f : 1 / sum;
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(o_store_reduction_view); ++ni) {
|
||||
o_store_reduction_view(mi, ni) *= inv_sum;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<typename TensorAcc>
|
||||
CUTLASS_DEVICE void rescale_o(TensorAcc& o_store, TensorAcc const& o_tmp) {
|
||||
Tensor o_store_reduction_view = make_tensor(o_store.data(), flash::convert_to_reduction_layout(o_store.layout()));
|
||||
Tensor o_tmp_reduction_view = make_tensor(o_tmp.data(), flash::convert_to_reduction_layout(o_tmp.layout()));
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int mi = 0; mi < size(row_max); ++mi) {
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int ni = 0; ni < size<1>(o_store_reduction_view); ++ni) {
|
||||
o_store_reduction_view(mi, ni) = o_store_reduction_view(mi, ni) * scores_scale(mi) + o_tmp_reduction_view(mi, ni);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
|
||||
};
|
||||
} // namespace flash
|
||||
@@ -0,0 +1,83 @@
|
||||
// Inspired by
|
||||
// https://github.com/NVIDIA/DALI/blob/main/include/dali/core/static_switch.h
|
||||
// and https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/Dispatch.h
|
||||
|
||||
#pragma once
|
||||
|
||||
/// @param COND - a boolean expression to switch by
|
||||
/// @param CONST_NAME - a name given for the constexpr bool variable.
|
||||
/// @param ... - code to execute for true and false
|
||||
///
|
||||
/// Usage:
|
||||
/// ```
|
||||
/// BOOL_SWITCH(flag, BoolConst, [&] {
|
||||
/// some_function<BoolConst>(...);
|
||||
/// });
|
||||
/// ```
|
||||
//
|
||||
|
||||
#define BOOL_SWITCH(COND, CONST_NAME, ...) \
|
||||
[&] { \
|
||||
if (COND) { \
|
||||
constexpr static bool CONST_NAME = true; \
|
||||
return __VA_ARGS__(); \
|
||||
} else { \
|
||||
constexpr static bool CONST_NAME = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define PREC_SWITCH(PRECTYPE, ...) \
|
||||
[&] { \
|
||||
if (PRECTYPE == 1) { \
|
||||
using kPrecType = cutlass::half_t; \
|
||||
constexpr static bool kSoftFp16 = false; \
|
||||
constexpr static bool kHybrid = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (PRECTYPE == 2) { \
|
||||
using kPrecType = cutlass::float_e4m3_t; \
|
||||
constexpr static bool kSoftFp16 = false; \
|
||||
constexpr static bool kHybrid = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (PRECTYPE == 3) { \
|
||||
using kPrecType = cutlass::float_e4m3_t; \
|
||||
constexpr static bool kSoftFp16 = false; \
|
||||
constexpr static bool kHybrid = true; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (PRECTYPE == 4) { \
|
||||
using kPrecType = cutlass::float_e4m3_t; \
|
||||
constexpr static bool kSoftFp16 = true; \
|
||||
constexpr static bool kHybrid = false; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define HEADDIM_SWITCH(HEADDIM, ...) \
|
||||
[&] { \
|
||||
if (HEADDIM == 64) { \
|
||||
constexpr static int kHeadSize = 64; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (HEADDIM == 128) { \
|
||||
constexpr static int kHeadSize = 128; \
|
||||
return __VA_ARGS__(); \
|
||||
} else if (HEADDIM == 256) { \
|
||||
constexpr static int kHeadSize = 256; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
|
||||
#define SEQLEN_SWITCH(USE_VAR_SEQ_LEN, SEQ_LEN_OUT_OF_BOUND_CHECK, ...) \
|
||||
[&] { \
|
||||
if (!USE_VAR_SEQ_LEN) { \
|
||||
if (SEQ_LEN_OUT_OF_BOUND_CHECK) { \
|
||||
using kSeqLenTraitsType = FixedSeqLenTraits<true>; \
|
||||
return __VA_ARGS__(); \
|
||||
} else { \
|
||||
using kSeqLenTraitsType = FixedSeqLenTraits<false>; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
} else { \
|
||||
using kSeqLenTraitsType = VarSeqLenTraits; \
|
||||
return __VA_ARGS__(); \
|
||||
} \
|
||||
}()
|
||||
@@ -0,0 +1,304 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* This code is based on code from FlashAttention3, https://github.com/Dao-AILab/flash-attention
|
||||
* Copyright (c) 2024, Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao.
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/fast_math.h"
|
||||
|
||||
namespace flash {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class StaticPersistentTileSchedulerOld {
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
private:
|
||||
int current_work_linear_idx_;
|
||||
cutlass::FastDivmod const &m_block_divmod, &head_divmod;
|
||||
int const total_blocks;
|
||||
|
||||
public:
|
||||
struct WorkTileInfo {
|
||||
int M_idx = 0;
|
||||
int H_idx = 0;
|
||||
int B_idx = 0;
|
||||
bool is_valid_tile = false;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_valid() const {
|
||||
return is_valid_tile;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static WorkTileInfo
|
||||
invalid_work_tile() {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE explicit StaticPersistentTileSchedulerOld(cutlass::FastDivmod const &m_block_divmod_,
|
||||
cutlass::FastDivmod const &head_divmod_,
|
||||
int const total_blocks_) :
|
||||
m_block_divmod(m_block_divmod_), head_divmod(head_divmod_), total_blocks(total_blocks_) {
|
||||
|
||||
// MSVC requires protecting use of CUDA-specific nonstandard syntax,
|
||||
// like blockIdx and gridDim, with __CUDA_ARCH__.
|
||||
#if defined(__CUDA_ARCH__)
|
||||
// current_work_linear_idx_ = blockIdx.x + blockIdx.y * gridDim.x + blockIdx.z * gridDim.x * gridDim.y;
|
||||
current_work_linear_idx_ = blockIdx.x;
|
||||
#else
|
||||
CUTLASS_ASSERT(false && "This line should never be reached");
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work() const {
|
||||
return get_current_work_for_linear_idx(current_work_linear_idx_);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work_for_linear_idx(int linear_idx) const {
|
||||
if (linear_idx >= total_blocks) {
|
||||
return WorkTileInfo::invalid_work_tile();
|
||||
}
|
||||
|
||||
// Map worker's linear index into the CTA tiled problem shape to the corresponding MHB indices
|
||||
int M_idx, H_idx, B_idx;
|
||||
int quotient = m_block_divmod.divmod(M_idx, linear_idx);
|
||||
B_idx = head_divmod.divmod(H_idx, quotient);
|
||||
return {M_idx, H_idx, B_idx, true};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
// advance_to_next_work(int advance_count = 1) {
|
||||
advance_to_next_work() {
|
||||
// current_work_linear_idx_ += int(gridDim.x * gridDim.y * gridDim.z);
|
||||
current_work_linear_idx_ += int(gridDim.x);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
fetch_next_work() {
|
||||
WorkTileInfo new_work_tile_info;
|
||||
advance_to_next_work();
|
||||
new_work_tile_info = get_current_work();
|
||||
return new_work_tile_info;
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class SingleTileScheduler {
|
||||
|
||||
public:
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
int const num_blocks_m, num_head, num_batch;
|
||||
int const* tile_count_semaphore = nullptr;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
return {};
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_dim(Arguments const& args, int num_sm) {
|
||||
return {uint32_t(args.num_blocks_m), uint32_t(args.num_head), uint32_t(args.num_batch)};
|
||||
}
|
||||
|
||||
struct WorkTileInfo {
|
||||
int M_idx = 0;
|
||||
int H_idx = 0;
|
||||
int B_idx = 0;
|
||||
bool is_valid_tile = false;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_valid(Params const& params) const {
|
||||
return is_valid_tile;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
cute::tuple<int32_t, int32_t, int32_t>
|
||||
get_block_coord(Params const& params) const {
|
||||
return {M_idx, H_idx, B_idx};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params) const {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_initial_work() const {
|
||||
return {int(blockIdx.x), int(blockIdx.y), int(blockIdx.z), true};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params, WorkTileInfo const& current_work) const {
|
||||
return {-1, -1, -1, false};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
class StaticPersistentTileScheduler {
|
||||
|
||||
public:
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
int const num_blocks_m, num_head, num_batch;
|
||||
int const* tile_count_semaphore = nullptr;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
int total_blocks;
|
||||
cutlass::FastDivmod m_block_divmod, head_divmod;
|
||||
};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
return {args.num_blocks_m * args.num_head * args.num_batch,
|
||||
cutlass::FastDivmod(args.num_blocks_m), cutlass::FastDivmod(args.num_head)};
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_dim(Arguments const& args, int num_sm) {
|
||||
return {uint32_t(num_sm)};
|
||||
}
|
||||
|
||||
struct WorkTileInfo {
|
||||
int tile_idx;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
is_valid(Params const& params) const {
|
||||
return tile_idx < params.total_blocks;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
cute::tuple<int32_t, int32_t, int32_t>
|
||||
get_block_coord(Params const& params) const {
|
||||
int m_block, bidh, bidb;
|
||||
bidb = params.head_divmod.divmod(bidh, params.m_block_divmod.divmod(m_block, tile_idx));
|
||||
return {m_block, bidh, bidb};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_initial_work() const {
|
||||
return {int(blockIdx.x)};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params, WorkTileInfo const& current_work) const {
|
||||
return {current_work.tile_idx + int(gridDim.x)};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
class DynamicPersistentTileScheduler {
|
||||
|
||||
public:
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
int const num_blocks_m, num_head, num_batch;
|
||||
int const* tile_count_semaphore;
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
int const total_blocks;
|
||||
cutlass::FastDivmod const m_block_divmod, head_divmod;
|
||||
int const* tile_count_semaphore;
|
||||
};
|
||||
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args) {
|
||||
return {args.num_blocks_m * args.num_head * args.num_batch,
|
||||
cutlass::FastDivmod(args.num_blocks_m), cutlass::FastDivmod(args.num_head),
|
||||
args.tile_count_semaphore};
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_dim(Arguments const& args, int num_sm) {
|
||||
return {uint32_t(num_sm)};
|
||||
}
|
||||
|
||||
using WorkTileInfo = StaticPersistentTileScheduler::WorkTileInfo;
|
||||
// struct WorkTileInfo {
|
||||
// int tile_idx;
|
||||
|
||||
// CUTLASS_DEVICE
|
||||
// bool
|
||||
// is_valid(Params const& params) const {
|
||||
// return tile_idx < params.total_blocks;
|
||||
// }
|
||||
|
||||
// CUTLASS_DEVICE
|
||||
// cute::tuple<int32_t, int32_t, int32_t>
|
||||
// get_block_coord(Params const& params) const {
|
||||
// int m_block, bidh, bidb;
|
||||
// bidb = params.head_divmod.divmod(bidh, params.m_block_divmod.divmod(m_block, tile_idx));
|
||||
// return {m_block, bidh, bidb};
|
||||
// }
|
||||
|
||||
// };
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_initial_work() const {
|
||||
return {int(blockIdx.x)};
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_next_work(Params const& params, WorkTileInfo const& current_work) const {
|
||||
return {current_work.tile_idx + int(gridDim.x)};
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // flash
|
||||
@@ -0,0 +1,408 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <assert.h>
|
||||
#include <stdint.h>
|
||||
#include <stdlib.h>
|
||||
|
||||
#include <cuda_fp16.h>
|
||||
|
||||
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 800
|
||||
#include <cuda_bf16.h>
|
||||
#endif
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
#include <cutlass/array.h>
|
||||
#include <cutlass/cutlass.h>
|
||||
#include <cutlass/numeric_conversion.h>
|
||||
#include <cutlass/numeric_types.h>
|
||||
|
||||
namespace flash {
|
||||
|
||||
using namespace cute;
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename T>
|
||||
struct MaxOp {
|
||||
__device__ __forceinline__ T operator()(T const & x, T const & y) { return x > y ? x : y; }
|
||||
};
|
||||
|
||||
template <>
|
||||
struct MaxOp<float> {
|
||||
// This is slightly faster
|
||||
__device__ __forceinline__ float operator()(float const &x, float const &y) { return max(x, y); }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<typename T>
|
||||
struct SumOp {
|
||||
__device__ __forceinline__ T operator()(T const & x, T const & y) { return x + y; }
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<int THREADS>
|
||||
struct Allreduce {
|
||||
static_assert(THREADS == 32 || THREADS == 16 || THREADS == 8 || THREADS == 4);
|
||||
template<typename T, typename Operator>
|
||||
static __device__ __forceinline__ T run(T x, Operator &op) {
|
||||
constexpr int OFFSET = THREADS / 2;
|
||||
x = op(x, __shfl_xor_sync(uint32_t(-1), x, OFFSET));
|
||||
return Allreduce<OFFSET>::run(x, op);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<>
|
||||
struct Allreduce<2> {
|
||||
template<typename T, typename Operator>
|
||||
static __device__ __forceinline__ T run(T x, Operator &op) {
|
||||
x = op(x, __shfl_xor_sync(uint32_t(-1), x, 1));
|
||||
return x;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template<bool zero_init=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1, typename Operator>
|
||||
__device__ __forceinline__ void thread_reduce_(Tensor<Engine0, Layout0> const &tensor, Tensor<Engine1, Layout1> &summary, Operator &op) {
|
||||
static_assert(Layout0::rank == 2, "Only support 2D Tensor");
|
||||
static_assert(Layout1::rank == 1, "Only support 1D Tensor");
|
||||
CUTE_STATIC_ASSERT_V(size<0>(summary) == size<0>(tensor));
|
||||
#pragma unroll
|
||||
for (int mi = 0; mi < size<0>(tensor); mi++) {
|
||||
summary(mi) = zero_init ? tensor(mi, 0) : op(summary(mi), tensor(mi, 0));
|
||||
#pragma unroll
|
||||
for (int ni = 1; ni < size<1>(tensor); ni++) {
|
||||
summary(mi) = op(summary(mi), tensor(mi, ni));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template<typename Engine0, typename Layout0, typename Engine1, typename Layout1, typename Operator>
|
||||
__device__ __forceinline__ void quad_allreduce_(Tensor<Engine0, Layout0> &dst, Tensor<Engine1, Layout1> &src, Operator &op) {
|
||||
CUTE_STATIC_ASSERT_V(size(dst) == size(src));
|
||||
#pragma unroll
|
||||
for (int i = 0; i < size(dst); i++){
|
||||
dst(i) = Allreduce<4>::run(src(i), op);
|
||||
}
|
||||
}
|
||||
|
||||
template<bool zero_init=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1, typename Operator>
|
||||
__device__ __forceinline__ void reduce_(Tensor<Engine0, Layout0> const& tensor, Tensor<Engine1, Layout1> &summary, Operator &op) {
|
||||
thread_reduce_<zero_init>(tensor, summary, op);
|
||||
quad_allreduce_(summary, summary, op);
|
||||
}
|
||||
|
||||
template<bool zero_init=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__device__ __forceinline__ void reduce_max(Tensor<Engine0, Layout0> const& tensor, Tensor<Engine1, Layout1> &max){
|
||||
MaxOp<float> max_op;
|
||||
reduce_<zero_init>(tensor, max, max_op);
|
||||
}
|
||||
|
||||
template<bool zero_init=true, bool warp_reduce=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__device__ __forceinline__ void reduce_sum(Tensor<Engine0, Layout0> const& tensor, Tensor<Engine1, Layout1> &sum){
|
||||
SumOp<float> sum_op;
|
||||
thread_reduce_<zero_init>(tensor, sum, sum_op);
|
||||
if constexpr (warp_reduce) { quad_allreduce_(sum, sum, sum_op); }
|
||||
}
|
||||
|
||||
__forceinline__ __device__ __half2 half_exp(__half2 x) {
|
||||
uint32_t tmp_out, tmp_in;
|
||||
tmp_in = reinterpret_cast<uint32_t&>(x);
|
||||
asm ("ex2.approx.f16x2 %0, %1;\n"
|
||||
: "=r"(tmp_out)
|
||||
: "r"(tmp_in));
|
||||
__half2 out = reinterpret_cast<__half2&>(tmp_out);
|
||||
return out;
|
||||
}
|
||||
|
||||
// Apply the exp to all the elements.
|
||||
template <bool zero_init=false, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__forceinline__ __device__ void max_scale_exp2_sum(Tensor<Engine0, Layout0> &tensor, Tensor<Engine1, Layout1> &max, Tensor<Engine1, Layout1> &sum, const float scale) {
|
||||
static_assert(Layout0::rank == 2, "Only support 2D Tensor"); static_assert(Layout1::rank == 1, "Only support 1D Tensor"); CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor));
|
||||
#pragma unroll
|
||||
for (int mi = 0; mi < size<0>(tensor); ++mi) {
|
||||
MaxOp<float> max_op;
|
||||
max(mi) = zero_init ? tensor(mi, 0) : max_op(max(mi), tensor(mi, 0));
|
||||
#pragma unroll
|
||||
for (int ni = 1; ni < size<1>(tensor); ni++) {
|
||||
max(mi) = max_op(max(mi), tensor(mi, ni));
|
||||
}
|
||||
max(mi) = Allreduce<4>::run(max(mi), max_op);
|
||||
// If max is -inf, then all elements must have been -inf (possibly due to masking).
|
||||
// We don't want (-inf - (-inf)) since that would give NaN.
|
||||
const float max_scaled = max(mi) == -INFINITY ? 0.f : max(mi) * scale;
|
||||
sum(mi) = 0;
|
||||
#pragma unroll
|
||||
for (int ni = 0; ni < size<1>(tensor); ++ni) {
|
||||
// Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
// max * log_2(e)) This allows the compiler to use the ffma
|
||||
// instruction instead of fadd and fmul separately.
|
||||
tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled);
|
||||
sum(mi) += tensor(mi, ni);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Apply the exp to all the elements.
|
||||
template <bool Scale_max=true, bool Check_inf=true, typename Engine0, typename Layout0, typename Engine1, typename Layout1>
|
||||
__forceinline__ __device__ void scale_apply_exp2(Tensor<Engine0, Layout0> &tensor, Tensor<Engine1, Layout1> const &max, const float scale) {
|
||||
static_assert(Layout0::rank == 2, "Only support 2D Tensor");
|
||||
static_assert(Layout1::rank == 1, "Only support 1D Tensor");
|
||||
CUTE_STATIC_ASSERT_V(size<0>(max) == size<0>(tensor));
|
||||
#pragma unroll
|
||||
for (int mi = 0; mi < size<0>(tensor); ++mi) {
|
||||
// If max is -inf, then all elements must have been -inf (possibly due to masking).
|
||||
// We don't want (-inf - (-inf)) since that would give NaN.
|
||||
// If we don't have float around M_LOG2E the multiplication is done in fp64.
|
||||
const float max_scaled = Check_inf
|
||||
? (max(mi) == -INFINITY ? 0.f : (max(mi) * (Scale_max ? scale : float(M_LOG2E))))
|
||||
: (max(mi) * (Scale_max ? scale : float(M_LOG2E)));
|
||||
#pragma unroll
|
||||
for (int ni = 0; ni < size<1>(tensor); ++ni) {
|
||||
// Instead of computing exp(x - max), we compute exp2(x * log_2(e) -
|
||||
// max * log_2(e)) This allows the compiler to use the ffma
|
||||
// instruction instead of fadd and fmul separately.
|
||||
tensor(mi, ni) = exp2f(tensor(mi, ni) * scale - max_scaled);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
__forceinline__ __device__ float ptx_exp2(float x) {
|
||||
float y;
|
||||
asm volatile("ex2.approx.ftz.f32 %0, %1;" : "=f"(y) : "f"(x));
|
||||
return y;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
packed_float_to_ue4m3(
|
||||
float const &f0, float const &f1, float const &f2, float const &f3,
|
||||
uint32_t &out
|
||||
) {
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
".reg .b16 lo;\n" \
|
||||
".reg .b16 hi;\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f32 lo, %2, %1;\n" \
|
||||
"cvt.rn.satfinite.e4m3x2.f32 hi, %4, %3;\n" \
|
||||
"mov.b32 %0, {lo, hi};\n" \
|
||||
"}" \
|
||||
: "=r"(out) : "f"(f0), "f"(f1), "f"(f2), "f"(f3));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
packed_float_to_e2m1(
|
||||
float const &f0, float const &f1, float const &f2, float const& f3,
|
||||
float const &f4, float const &f5, float const &f6, float const& f7,
|
||||
uint32_t &out
|
||||
) {
|
||||
|
||||
asm volatile( \
|
||||
"{\n" \
|
||||
".reg .b8 byte0;\n" \
|
||||
".reg .b8 byte1;\n" \
|
||||
".reg .b8 byte2;\n" \
|
||||
".reg .b8 byte3;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte0, %2, %1;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte1, %4, %3;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte2, %6, %5;\n" \
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte3, %8, %7;\n" \
|
||||
"mov.b32 %0, {byte0, byte1, byte2, byte3};\n" \
|
||||
"}" \
|
||||
: "=r"(out) : "f"(f0), "f"(f1), "f"(f2), "f"(f3),
|
||||
"f"(f4), "f"(f5), "f"(f6), "f"(f7));
|
||||
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
add(float2 & c,
|
||||
float2 const& a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("add.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(c))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
add_inplace(float2 &a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("add.f32x2 %0, %0, %1;\n"
|
||||
: "+l"(reinterpret_cast<uint64_t &>(a)) // a: input/output
|
||||
: "l"(reinterpret_cast<uint64_t const&>(b)) // b: input
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
sub(float2 & c,
|
||||
float2 const& a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("sub.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(c))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
sub_inplace(float2 &a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("sub.f32x2 %0, %0, %1;\n"
|
||||
: "+l"(reinterpret_cast<uint64_t &>(a)) // a: input/output
|
||||
: "l"(reinterpret_cast<uint64_t const&>(b)) // b: input
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
mul(float2 & c,
|
||||
float2 const& a,
|
||||
float2 const& b)
|
||||
{
|
||||
asm volatile("mul.f32x2 %0, %1, %2;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(c))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
fma(float2 & d,
|
||||
float2 const& a,
|
||||
float2 const& b,
|
||||
float2 const& c)
|
||||
{
|
||||
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;\n"
|
||||
: "=l"(reinterpret_cast<uint64_t &>(d))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(a)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(b)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(c)));
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE void
|
||||
fma_inplace(float2 &a,
|
||||
float2 const& b,
|
||||
float2 const& c)
|
||||
{
|
||||
asm volatile("fma.rn.f32x2 %0, %0, %1, %2;\n"
|
||||
: "+l"(reinterpret_cast<uint64_t &>(a))
|
||||
: "l"(reinterpret_cast<uint64_t const&>(b)),
|
||||
"l"(reinterpret_cast<uint64_t const&>(c)));
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class Layout
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_reduction_layout(Layout mma_layout) {
|
||||
static_assert(rank(mma_layout) == 3, "Mma Layout should be (MmaAtom, MmaM, MmaN)");
|
||||
static_assert(rank(get<0>(shape(mma_layout))) == 2, "MmaAtom should be (AtomN, AtomM)");
|
||||
|
||||
return make_layout(
|
||||
make_layout(get<0,1>(mma_layout), get<1>(mma_layout)),
|
||||
make_layout(get<0,0>(mma_layout), get<2>(mma_layout))
|
||||
);
|
||||
}
|
||||
|
||||
template <
|
||||
class Tensor
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_reduction_tensor(Tensor mma_tensor) {
|
||||
return make_tensor(mma_tensor.data(), convert_to_reduction_layout(mma_tensor.layout()));
|
||||
}
|
||||
|
||||
|
||||
template <
|
||||
class Layout
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_conversion_layout(Layout mma_layout) {
|
||||
static_assert(rank(mma_layout) == 3, "Mma Layout should be (MmaAtom, MmaM, MmaN)");
|
||||
static_assert(rank(get<0>(shape(mma_layout))) == 2, "MmaAtom should be (AtomN, AtomM)");
|
||||
|
||||
constexpr int MmaAtomN = size<0, 0>(mma_layout);
|
||||
constexpr int MmaAtomM = size<0, 1>(mma_layout);
|
||||
constexpr int MmaM = size<1>(mma_layout);
|
||||
constexpr int MmaN = size<2>(mma_layout);
|
||||
|
||||
static_assert(MmaAtomN == 8, "MmaAtomN should be 8.");
|
||||
static_assert(MmaAtomM == 2, "MmaAtomM should be 2.");
|
||||
static_assert(MmaN % 2 == 0, "MmaN should be multiple of 2.");
|
||||
|
||||
auto mma_n_division = zipped_divide(
|
||||
layout<2>(mma_layout), make_tile(_2{})
|
||||
);
|
||||
return make_layout(
|
||||
make_layout(layout<0,0>(mma_layout), make_layout(layout<0,1>(mma_layout), layout<0>(mma_n_division))),
|
||||
layout<1>(mma_layout), layout<1>(mma_n_division)
|
||||
);
|
||||
}
|
||||
|
||||
template <
|
||||
class Tensor
|
||||
>
|
||||
CUTLASS_DEVICE constexpr
|
||||
auto convert_to_conversion_tensor(Tensor mma_tensor) {
|
||||
return make_tensor(mma_tensor.data(), convert_to_conversion_layout(mma_tensor.layout()));
|
||||
}
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <bool Is_even_MN=true, bool Is_even_K=true, bool Clear_OOB_MN=false, bool Clear_OOB_K=true,
|
||||
typename TiledCopy, typename Engine0, typename Layout0, typename Engine1, typename Layout1,
|
||||
typename Engine2, typename Layout2, typename Engine3, typename Layout3>
|
||||
CUTLASS_DEVICE void copy(TiledCopy tiled_copy, Tensor<Engine0, Layout0> const &S,
|
||||
Tensor<Engine1, Layout1> &D, Tensor<Engine2, Layout2> const &identity_MN,
|
||||
Tensor<Engine3, Layout3> const &predicate_K, const int max_MN=0) {
|
||||
CUTE_STATIC_ASSERT_V(rank(S) == Int<3>{});
|
||||
CUTE_STATIC_ASSERT_V(rank(D) == Int<3>{});
|
||||
CUTE_STATIC_ASSERT_V(size<0>(S) == size<0>(D)); // MMA
|
||||
CUTE_STATIC_ASSERT_V(size<1>(S) == size<1>(D)); // MMA_M
|
||||
CUTE_STATIC_ASSERT_V(size<2>(S) == size<2>(D)); // MMA_K
|
||||
// There's no case where !Clear_OOB_K && Clear_OOB_MN
|
||||
static_assert(!(Clear_OOB_MN && !Clear_OOB_K));
|
||||
#pragma unroll
|
||||
for (int m = 0; m < size<1>(S); ++m) {
|
||||
if (Is_even_MN || get<0>(identity_MN(0, m, 0)) < max_MN) {
|
||||
#pragma unroll
|
||||
for (int k = 0; k < size<2>(S); ++k) {
|
||||
if (Is_even_K || predicate_K(k)) {
|
||||
cute::copy(tiled_copy, S(_, m, k), D(_, m, k));
|
||||
} else if (Clear_OOB_K) {
|
||||
cute::clear(D(_, m, k));
|
||||
}
|
||||
}
|
||||
} else if (Clear_OOB_MN) {
|
||||
cute::clear(D(_, m, _));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace flash
|
||||
@@ -0,0 +1 @@
|
||||
__version__ = "3.0.0.b1"
|
||||
@@ -0,0 +1,90 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import fp4quant
|
||||
from triton.tools.mxfp import MXFP4Tensor
|
||||
|
||||
from bench_utils import bench_kineto
|
||||
b = 1
|
||||
h = 32
|
||||
n = 16384
|
||||
d = 128
|
||||
|
||||
def test():
|
||||
q = torch.randn((b, h, n, d), device="cuda", dtype=torch.float16)
|
||||
o = torch.empty((b, h, n, d // 2), device="cuda", dtype=torch.uint8)
|
||||
o_s = torch.empty((b, h, n, d // 16), device="cuda", dtype=torch.float8_e4m3fn)
|
||||
fp4quant.scaled_fp4_quant_permute(q, o, o_s, 1)
|
||||
|
||||
test()
|
||||
|
||||
t = bench_kineto(test, "scaled_fp4_quant_kernel", suppress_kineto_output=True)
|
||||
|
||||
IO = b * h * n * d * 2 + b * h * n * d * 0.5 + b * h * n * d // 16 * 1
|
||||
throughput = IO / t * 1e-9
|
||||
|
||||
print(f"Throughput: {throughput:.2f} GB/s")
|
||||
|
||||
def scale_and_fp4_tensor(x: torch.Tensor, packed_dim: int = 3, all_ones: bool = False, permuted: bool = False):
|
||||
assert x.is_contiguous() and x.ndim == 4 and x.shape[-1] % 16 == 0
|
||||
B, H, M, N = x.shape
|
||||
x = x.view(B, H, M, N // 16, 16)
|
||||
scales = (x.abs().amax(dim=-1, keepdim=True) / 6).to(torch.float32)
|
||||
if all_ones:
|
||||
scales = torch.ones_like(scales)
|
||||
x_scaled = x / scales
|
||||
packed_fp4 = MXFP4Tensor(x_scaled.flatten(start_dim=-2)).to_packed_tensor(dim=packed_dim)
|
||||
dequant_x = (MXFP4Tensor(x_scaled).to(torch.float32) * scales.to(torch.float8_e4m3fn).to(torch.float32)).flatten(start_dim=-2)
|
||||
fp8_scale = scales.flatten(start_dim=-2).to(torch.float8_e4m3fn)
|
||||
permuted_fp8_scale = None
|
||||
if permuted:
|
||||
scales = scales.view(B, H // 64, 4, 16, M, N // 16).permute(0, 1, 3, 2, 4, 5).reshape(B, H, M, N // 16)
|
||||
permuted_fp8_scale = scales.view(B, H // 64, 64, M, N // 64, 4).permute(0, 1, 4, 3, 2, 5).reshape(B, H, M, N // 16).to(torch.float8_e4m3fn)
|
||||
return fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale
|
||||
|
||||
b = 2
|
||||
h = 4
|
||||
n = 251
|
||||
n_padded = (n + 127) // 128 * 128
|
||||
d = 128
|
||||
|
||||
q = torch.randn(b, h, n, d, dtype=torch.float16, device='cuda')
|
||||
o = torch.empty((b, h, n, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s = torch.empty((b, h, n, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
|
||||
fp4quant.scaled_fp4_quant(q, o, o_s, 1)
|
||||
|
||||
k_permute = [0, 1, 8, 9, 16, 17, 24, 25, 2, 3, 10, 11, 18, 19, 26, 27, 4, 5, 12, 13, 20, 21, 28, 29, 6, 7, 14, 15, 22, 23, 30, 31]
|
||||
o_permuted = torch.empty((b, h, n_padded, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s_permuted = torch.empty((b, h, n_padded, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
fp4quant.scaled_fp4_quant_permute(q, o_permuted, o_s_permuted, 1)
|
||||
|
||||
# padding
|
||||
if n % 128 != 0:
|
||||
o_permuted_gt = torch.cat([o, torch.zeros((b, h, n_padded - n, d // 2), dtype=torch.uint8, device='cuda')], dim=2)
|
||||
o_s_permuted_gt = torch.cat([o_s, torch.zeros((b, h, n_padded - n, d // 16), dtype=torch.float8_e4m3fn, device='cuda')], dim=2)
|
||||
else:
|
||||
o_permuted_gt = o
|
||||
o_s_permuted_gt = o_s
|
||||
|
||||
# use scale_and_fp4_tensor + torch permutation to get the ground truth
|
||||
o_permuted_gt = o_permuted_gt.reshape(b, h, n_padded // 32, 32, d // 2)[:, :, :, k_permute, :].reshape(b, h, n_padded, d // 2)
|
||||
o_s_permuted_gt = o_s_permuted_gt.reshape(b, h, n_padded // 32, 32, d // 16)[:, :, :, k_permute, :].reshape(b, h, n_padded, d // 16)
|
||||
|
||||
assert((o_permuted - o_permuted_gt).abs().max() == 0)
|
||||
assert((o_s_permuted.float() - o_s_permuted_gt.float()).abs().max() == 0)
|
||||
|
||||
print("All tests passed!")
|
||||
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import fp4quant
|
||||
from triton.tools.mxfp import MXFP4Tensor
|
||||
|
||||
from bench_utils import bench_kineto
|
||||
b = 1
|
||||
h = 32
|
||||
n = 16384
|
||||
d = 128
|
||||
|
||||
def test():
|
||||
q = torch.randn((b, h, n, d), device="cuda", dtype=torch.float16)
|
||||
o = torch.empty((b, h, n, d // 2), device="cuda", dtype=torch.uint8)
|
||||
o_s = torch.empty((b, h, n, d // 16), device="cuda", dtype=torch.float8_e4m3fn)
|
||||
fp4quant.scaled_fp4_quant(q, o, o_s, 1)
|
||||
|
||||
test()
|
||||
|
||||
t = bench_kineto(test, "scaled_fp4_quant_kernel", suppress_kineto_output=True)
|
||||
|
||||
IO = b * h * n * d * 2 + b * h * n * d * 0.5 + b * h * n * d // 16 * 1
|
||||
throughput = IO / t * 1e-9
|
||||
|
||||
print(f"Throughput: {throughput:.2f} GB/s")
|
||||
|
||||
def scale_and_fp4_tensor(x: torch.Tensor, packed_dim: int = 3, all_ones: bool = False, permuted: bool = False):
|
||||
assert x.is_contiguous() and x.ndim == 4 and x.shape[-1] % 16 == 0
|
||||
B, H, M, N = x.shape
|
||||
x = x.view(B, H, M, N // 16, 16)
|
||||
scales = (x.abs().amax(dim=-1, keepdim=True) / 6).to(torch.float32)
|
||||
if all_ones:
|
||||
scales = torch.ones_like(scales)
|
||||
x_scaled = x / scales
|
||||
packed_fp4 = MXFP4Tensor(x_scaled.flatten(start_dim=-2)).to_packed_tensor(dim=packed_dim)
|
||||
dequant_x = (MXFP4Tensor(x_scaled).to(torch.float32) * scales.to(torch.float8_e4m3fn).to(torch.float32)).flatten(start_dim=-2)
|
||||
fp8_scale = scales.flatten(start_dim=-2).to(torch.float8_e4m3fn)
|
||||
permuted_fp8_scale = None
|
||||
if permuted:
|
||||
scales = scales.view(B, H // 64, 4, 16, M, N // 16).permute(0, 1, 3, 2, 4, 5).reshape(B, H, M, N // 16)
|
||||
permuted_fp8_scale = scales.view(B, H // 64, 64, M, N // 64, 4).permute(0, 1, 4, 3, 2, 5).reshape(B, H, M, N // 16).to(torch.float8_e4m3fn)
|
||||
return fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale
|
||||
|
||||
b = 2
|
||||
h = 4
|
||||
n = 251
|
||||
d = 128
|
||||
|
||||
q = torch.randn(b, h, n, d, dtype=torch.float16, device='cuda')
|
||||
o = torch.empty((b, h, n, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s = torch.empty((b, h, n, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
|
||||
fp4quant.scaled_fp4_quant(q, o, o_s, 1)
|
||||
|
||||
fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale = scale_and_fp4_tensor(q, packed_dim=3)
|
||||
|
||||
assert((fp8_scale.float() - o_s.float()).abs().max() == 0)
|
||||
|
||||
o_binary = [
|
||||
(int(bin_str[:4], 2), int(bin_str[4:], 2))
|
||||
for bin_str in [format(x.item(), '08b') for x in o.view(-1)]
|
||||
]
|
||||
o_binary_gt = [
|
||||
(int(bin_str[:4], 2), int(bin_str[4:], 2))
|
||||
for bin_str in [format(x.item(), '08b') for x in packed_fp4.view(-1)]
|
||||
]
|
||||
for i in range(len(o_binary)):
|
||||
# check contiguous 4 bits. Difference should be at most one
|
||||
assert(abs(o_binary[i][0] - o_binary_gt[i][0]) <= 1)
|
||||
assert(abs(o_binary[i][1] - o_binary_gt[i][1]) <= 1)
|
||||
|
||||
print("All tests passed!")
|
||||
@@ -0,0 +1,86 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import torch
|
||||
import fp4quant
|
||||
from triton.tools.mxfp import MXFP4Tensor
|
||||
|
||||
from bench_utils import bench_kineto
|
||||
b = 1
|
||||
h = 32
|
||||
n = 16384
|
||||
d = 128
|
||||
|
||||
def test():
|
||||
q = torch.randn((b, h, n, d), device="cuda", dtype=torch.float16)
|
||||
o = torch.empty((b, h, d, n // 2), device="cuda", dtype=torch.uint8)
|
||||
o_s = torch.empty((b, h, d, n // 16), device="cuda", dtype=torch.float8_e4m3fn)
|
||||
fp4quant.scaled_fp4_quant_trans(q, o, o_s, 1)
|
||||
|
||||
test()
|
||||
|
||||
t = bench_kineto(test, "scaled_fp4_quant_trans_kernel", suppress_kineto_output=True)
|
||||
|
||||
IO = b * h * n * d * 2 + b * h * n * d * 0.5 + b * h * n * d // 16 * 1
|
||||
throughput = IO / t * 1e-9
|
||||
|
||||
print(f"Throughput: {throughput:.2f} GB/s")
|
||||
|
||||
def scale_and_fp4_tensor(x: torch.Tensor, packed_dim: int = 3, all_ones: bool = False, permuted: bool = False):
|
||||
assert x.is_contiguous() and x.ndim == 4 and x.shape[-1] % 16 == 0
|
||||
B, H, M, N = x.shape
|
||||
x = x.view(B, H, M, N // 16, 16)
|
||||
scales = (x.abs().amax(dim=-1, keepdim=True) / 6).to(torch.float32)
|
||||
if all_ones:
|
||||
scales = torch.ones_like(scales)
|
||||
x_scaled = x / scales
|
||||
packed_fp4 = MXFP4Tensor(x_scaled.flatten(start_dim=-2)).to_packed_tensor(dim=packed_dim)
|
||||
dequant_x = (MXFP4Tensor(x_scaled).to(torch.float32) * scales.to(torch.float8_e4m3fn).to(torch.float32)).flatten(start_dim=-2)
|
||||
fp8_scale = scales.flatten(start_dim=-2).to(torch.float8_e4m3fn)
|
||||
permuted_fp8_scale = None
|
||||
if permuted:
|
||||
scales = scales.view(B, H // 64, 4, 16, M, N // 16).permute(0, 1, 3, 2, 4, 5).reshape(B, H, M, N // 16)
|
||||
permuted_fp8_scale = scales.view(B, H // 64, 64, M, N // 64, 4).permute(0, 1, 4, 3, 2, 5).reshape(B, H, M, N // 16).to(torch.float8_e4m3fn)
|
||||
return fp8_scale, packed_fp4, dequant_x, permuted_fp8_scale
|
||||
|
||||
b = 2
|
||||
h = 4
|
||||
n = 491
|
||||
n_padded = (n + 127) // 128 * 128
|
||||
d = 128
|
||||
|
||||
q = torch.randn(b, h, n, d, dtype=torch.float16, device='cuda')
|
||||
o = torch.empty((b, h, d, n_padded // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s = torch.empty((b, h, d, n_padded // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
|
||||
fp4quant.scaled_fp4_quant_trans(q, o, o_s, 1)
|
||||
|
||||
if n % 128 != 0:
|
||||
q_padded = torch.cat([q, torch.zeros((b, h, n_padded - n, d), dtype=torch.float16, device='cuda')], dim=2)
|
||||
else:
|
||||
q_padded = q
|
||||
|
||||
# use torch transpose + scaled_fp4_quant to get the ground truth
|
||||
q_padded = q_padded.transpose(2, 3).reshape(b, h, n_padded, d).contiguous()
|
||||
o_gt = torch.empty((b, h, n_padded, d // 2), dtype=torch.uint8, device='cuda')
|
||||
o_s_gt = torch.empty((b, h, n_padded, d // 16), dtype=torch.float8_e4m3fn, device='cuda')
|
||||
fp4quant.scaled_fp4_quant(q_padded, o_gt, o_s_gt, 1)
|
||||
o_gt = o_gt.reshape(b, h, d, n_padded // 2).contiguous()
|
||||
o_s_gt = o_s_gt.reshape(b, h, d, n_padded // 16).contiguous()
|
||||
|
||||
assert((o_s_gt.float() - o_s.float()).abs().max() == 0)
|
||||
assert((o_gt - o).abs().max() == 0)
|
||||
|
||||
print("All tests passed!")
|
||||
@@ -0,0 +1,169 @@
|
||||
"""
|
||||
Copyright (c) 2025 by SageAttention team.
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
|
||||
def bench(fn, num_warmups: int = 5, num_tests: int = 10,
|
||||
high_precision: bool = False):
|
||||
# Flush L2 cache with 256 MB data
|
||||
torch.cuda.synchronize()
|
||||
cache = torch.empty(int(256e6 // 4), dtype=torch.int, device='cuda')
|
||||
cache.zero_()
|
||||
|
||||
# Warmup
|
||||
for _ in range(num_warmups):
|
||||
fn()
|
||||
|
||||
# Add a large kernel to eliminate the CPU launch overhead
|
||||
if high_precision:
|
||||
x = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
y = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
x @ y
|
||||
|
||||
# Testing
|
||||
start_event = torch.cuda.Event(enable_timing=True)
|
||||
end_event = torch.cuda.Event(enable_timing=True)
|
||||
start_event.record()
|
||||
for i in range(num_tests):
|
||||
fn()
|
||||
end_event.record()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
return start_event.elapsed_time(end_event) / num_tests
|
||||
|
||||
|
||||
class empty_suppress:
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
pass
|
||||
|
||||
|
||||
class suppress_stdout_stderr:
|
||||
def __enter__(self):
|
||||
self.outnull_file = open(os.devnull, 'w')
|
||||
self.errnull_file = open(os.devnull, 'w')
|
||||
|
||||
self.old_stdout_fileno_undup = sys.stdout.fileno()
|
||||
self.old_stderr_fileno_undup = sys.stderr.fileno()
|
||||
|
||||
self.old_stdout_fileno = os.dup(sys.stdout.fileno())
|
||||
self.old_stderr_fileno = os.dup(sys.stderr.fileno())
|
||||
|
||||
self.old_stdout = sys.stdout
|
||||
self.old_stderr = sys.stderr
|
||||
|
||||
os.dup2(self.outnull_file.fileno(), self.old_stdout_fileno_undup)
|
||||
os.dup2(self.errnull_file.fileno(), self.old_stderr_fileno_undup)
|
||||
|
||||
sys.stdout = self.outnull_file
|
||||
sys.stderr = self.errnull_file
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
sys.stdout = self.old_stdout
|
||||
sys.stderr = self.old_stderr
|
||||
|
||||
os.dup2(self.old_stdout_fileno, self.old_stdout_fileno_undup)
|
||||
os.dup2(self.old_stderr_fileno, self.old_stderr_fileno_undup)
|
||||
|
||||
os.close(self.old_stdout_fileno)
|
||||
os.close(self.old_stderr_fileno)
|
||||
|
||||
self.outnull_file.close()
|
||||
self.errnull_file.close()
|
||||
|
||||
|
||||
def bench_kineto(fn, kernel_names, num_tests: int = 30, suppress_kineto_output: bool = False,
|
||||
trace_path: str = None, barrier_comm_profiling: bool = False, flush_l2: bool = False):
|
||||
# Conflict with Nsight Systems
|
||||
using_nsys = os.environ.get('DG_NSYS_PROFILING', False)
|
||||
|
||||
# For some auto-tuning kernels with prints
|
||||
fn()
|
||||
|
||||
# Profile
|
||||
suppress = suppress_stdout_stderr if suppress_kineto_output and not using_nsys else empty_suppress
|
||||
with suppress():
|
||||
schedule = torch.profiler.schedule(wait=0, warmup=1, active=1, repeat=1) if not using_nsys else None
|
||||
profiler = torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CUDA], schedule=schedule) if not using_nsys else empty_suppress()
|
||||
with profiler:
|
||||
for i in range(2):
|
||||
# NOTES: use a large kernel and a barrier to eliminate the unbalanced CPU launch overhead
|
||||
if barrier_comm_profiling:
|
||||
lhs = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
rhs = torch.randn((8192, 8192), dtype=torch.float, device='cuda')
|
||||
lhs @ rhs
|
||||
dist.all_reduce(torch.ones(1, dtype=torch.float, device='cuda'))
|
||||
for _ in range(num_tests):
|
||||
if flush_l2:
|
||||
torch.empty(int(256e6 // 4), dtype=torch.int, device='cuda').zero_()
|
||||
fn()
|
||||
|
||||
if not using_nsys:
|
||||
profiler.step()
|
||||
|
||||
# Return 1 if using Nsight Systems
|
||||
if using_nsys:
|
||||
return 1
|
||||
|
||||
# Parse the profiling table
|
||||
assert isinstance(kernel_names, str) or isinstance(kernel_names, tuple)
|
||||
is_tupled = isinstance(kernel_names, tuple)
|
||||
prof_lines = profiler.key_averages().table(sort_by='cuda_time_total', max_name_column_width=100).split('\n')
|
||||
kernel_names = (kernel_names, ) if isinstance(kernel_names, str) else kernel_names
|
||||
assert all([isinstance(name, str) for name in kernel_names])
|
||||
for name in kernel_names:
|
||||
assert sum([name in line for line in prof_lines]) == 1, f'Errors of the kernel {name} in the profiling table'
|
||||
|
||||
# Save chrome traces
|
||||
if trace_path is not None:
|
||||
profiler.export_chrome_trace(trace_path)
|
||||
|
||||
# Return average kernel times
|
||||
units = {'ms': 1e3, 'us': 1e6}
|
||||
kernel_times = []
|
||||
for name in kernel_names:
|
||||
for line in prof_lines:
|
||||
if name in line:
|
||||
time_str = line.split()[-2]
|
||||
for unit, scale in units.items():
|
||||
if unit in time_str:
|
||||
kernel_times.append(float(time_str.replace(unit, '')) / scale)
|
||||
break
|
||||
break
|
||||
return tuple(kernel_times) if is_tupled else kernel_times[0]
|
||||
|
||||
|
||||
def calc_diff(x, y):
|
||||
x, y = x.double(), y.double()
|
||||
denominator = (x * x + y * y).sum()
|
||||
sim = 2 * (x * y).sum() / denominator
|
||||
return 1 - sim
|
||||
|
||||
|
||||
def count_bytes(tensors):
|
||||
total = 0
|
||||
for t in tensors:
|
||||
if isinstance(t, tuple):
|
||||
total += count_bytes(t)
|
||||
else:
|
||||
total += t.numel() * t.element_size()
|
||||
return total
|
||||
@@ -0,0 +1,52 @@
|
||||
#pragma once
|
||||
|
||||
#include <stdio.h>
|
||||
|
||||
#if defined(__HIPCC__)
|
||||
#define HOST_DEVICE_INLINE __host__ __device__
|
||||
#define DEVICE_INLINE __device__
|
||||
#define HOST_INLINE __host__
|
||||
#elif defined(__CUDACC__) || defined(_NVHPC_CUDA)
|
||||
#define HOST_DEVICE_INLINE __host__ __device__ __forceinline__
|
||||
#define DEVICE_INLINE __device__ __forceinline__
|
||||
#define HOST_INLINE __host__ __forceinline__
|
||||
#else
|
||||
#define HOST_DEVICE_INLINE inline
|
||||
#define DEVICE_INLINE inline
|
||||
#define HOST_INLINE inline
|
||||
#endif
|
||||
|
||||
#define CUDA_CHECK(cmd) \
|
||||
do { \
|
||||
cudaError_t e = cmd; \
|
||||
if (e != cudaSuccess) { \
|
||||
printf("Failed: Cuda error %s:%d '%s'\n", __FILE__, __LINE__, \
|
||||
cudaGetErrorString(e)); \
|
||||
exit(EXIT_FAILURE); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
int64_t get_device_attribute(int64_t attribute, int64_t device_id) {
|
||||
static int value = [=]() {
|
||||
int device = static_cast<int>(device_id);
|
||||
if (device < 0) {
|
||||
CUDA_CHECK(cudaGetDevice(&device));
|
||||
}
|
||||
int value;
|
||||
CUDA_CHECK(cudaDeviceGetAttribute(
|
||||
&value, static_cast<cudaDeviceAttr>(attribute), device));
|
||||
return static_cast<int>(value);
|
||||
}();
|
||||
|
||||
return value;
|
||||
}
|
||||
|
||||
namespace cuda_utils {
|
||||
|
||||
template <typename T>
|
||||
HOST_DEVICE_INLINE constexpr std::enable_if_t<std::is_integral_v<T>, T>
|
||||
ceil_div(T a, T b) {
|
||||
return (a + b - 1) / b;
|
||||
}
|
||||
|
||||
}; // namespace cuda_utils
|
||||
@@ -0,0 +1,644 @@
|
||||
/*
|
||||
* Copyright (c) 2025 by SageAttention team.
|
||||
*
|
||||
* Licensed under the Apache License, Version 2.0 (the "License");
|
||||
* you may not use this file except in compliance with the License.
|
||||
* You may obtain a copy of the License at
|
||||
*
|
||||
* http://www.apache.org/licenses/LICENSE-2.0
|
||||
*
|
||||
* Unless required by applicable law or agreed to in writing, software
|
||||
* distributed under the License is distributed on an "AS IS" BASIS,
|
||||
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
* See the License for the specific language governing permissions and
|
||||
* limitations under the License.
|
||||
*/
|
||||
#include <torch/all.h>
|
||||
#include <torch/python.h>
|
||||
#include <torch/nn/functional.h>
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
#include <cuda_runtime_api.h>
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
#include <cuda_fp8.h>
|
||||
|
||||
#include "cuda_utils.h"
|
||||
#include "../blackwell/block_config.h"
|
||||
|
||||
#define DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(pytorch_dtype, c_type, ...) \
|
||||
if (pytorch_dtype == at::ScalarType::Half) { \
|
||||
using c_type = half; \
|
||||
__VA_ARGS__ \
|
||||
} else if (pytorch_dtype == at::ScalarType::BFloat16) { \
|
||||
using c_type = nv_bfloat16; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
std::ostringstream oss; \
|
||||
oss << __PRETTY_FUNCTION__ << " failed to dispatch data type " << pytorch_dtype; \
|
||||
TORCH_CHECK(false, oss.str()); \
|
||||
}
|
||||
|
||||
#define DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, ...) \
|
||||
if (head_dim == 64) { \
|
||||
constexpr int HEAD_DIM = 64; \
|
||||
__VA_ARGS__ \
|
||||
} else if (head_dim == 128) { \
|
||||
constexpr int HEAD_DIM = 128; \
|
||||
__VA_ARGS__ \
|
||||
} else { \
|
||||
std::ostringstream err_msg; \
|
||||
err_msg << "Unsupported head dim: " << int(head_dim); \
|
||||
throw std::invalid_argument(err_msg.str()); \
|
||||
}
|
||||
|
||||
#define CHECK_CUDA(x) \
|
||||
TORCH_CHECK(x.is_cuda(), "Tensor " #x " must be on CUDA")
|
||||
#define CHECK_DTYPE(x, true_dtype) \
|
||||
TORCH_CHECK(x.dtype() == true_dtype, \
|
||||
"Tensor " #x " must have dtype (" #true_dtype ")")
|
||||
#define CHECK_DIMS(x, true_dim) \
|
||||
TORCH_CHECK(x.dim() == true_dim, \
|
||||
"Tensor " #x " must have dimension number (" #true_dim ")")
|
||||
#define CHECK_SHAPE(x, ...) \
|
||||
TORCH_CHECK(x.sizes() == torch::IntArrayRef({__VA_ARGS__}), \
|
||||
"Tensor " #x " must have shape (" #__VA_ARGS__ ")")
|
||||
#define CHECK_CONTIGUOUS(x) \
|
||||
TORCH_CHECK(x.is_contiguous(), "Tensor " #x " must be contiguous")
|
||||
#define CHECK_LASTDIM_CONTIGUOUS(x) \
|
||||
TORCH_CHECK(x.stride(-1) == 1, \
|
||||
"Tensor " #x " must be contiguous at the last dimension")
|
||||
|
||||
constexpr int CVT_FP4_ELTS_PER_THREAD = 16;
|
||||
|
||||
// Convert 4 float2 values into 8 e2m1 values (represented as one uint32_t).
|
||||
inline __device__ uint32_t fp32_vec_to_e2m1(float2 *array) {
|
||||
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 1000)
|
||||
uint32_t val;
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .b8 byte0;\n"
|
||||
".reg .b8 byte1;\n"
|
||||
".reg .b8 byte2;\n"
|
||||
".reg .b8 byte3;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte0, %2, %1;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte1, %4, %3;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte2, %6, %5;\n"
|
||||
"cvt.rn.satfinite.e2m1x2.f32 byte3, %8, %7;\n"
|
||||
"mov.b32 %0, {byte0, byte1, byte2, byte3};\n"
|
||||
"}"
|
||||
: "=r"(val)
|
||||
: "f"(array[0].x), "f"(array[0].y), "f"(array[1].x), "f"(array[1].y),
|
||||
"f"(array[2].x), "f"(array[2].y), "f"(array[3].x), "f"(array[3].y));
|
||||
return val;
|
||||
#else
|
||||
return 0;
|
||||
#endif
|
||||
}
|
||||
|
||||
// Get type2 from type or vice versa (applied to half and bfloat16)
|
||||
template <typename T>
|
||||
struct TypeConverter {
|
||||
using Type = half2;
|
||||
}; // keep for generality
|
||||
|
||||
template <>
|
||||
struct TypeConverter<half2> {
|
||||
using Type = half;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct TypeConverter<half> {
|
||||
using Type = half2;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct TypeConverter<__nv_bfloat162> {
|
||||
using Type = __nv_bfloat16;
|
||||
};
|
||||
|
||||
template <>
|
||||
struct TypeConverter<__nv_bfloat16> {
|
||||
using Type = __nv_bfloat162;
|
||||
};
|
||||
|
||||
// Define a 32 bytes packed data type.
|
||||
template <class Type>
|
||||
struct PackedVec {
|
||||
typename TypeConverter<Type>::Type elts[8];
|
||||
};
|
||||
|
||||
template <uint32_t head_dim, uint32_t BLOCK_SIZE, bool permute, typename T>
|
||||
__global__ void scaled_fp4_quant_kernel(
|
||||
const T* input, uint8_t* output, uint8_t* output_sf,
|
||||
int batch_size, int num_heads, int num_tokens,
|
||||
int stride_bz_input, int stride_h_input, int stride_seq_input,
|
||||
int stride_bz_output, int stride_h_output, int stride_seq_output,
|
||||
int stride_bz_output_sf, int stride_h_output_sf, int stride_seq_output_sf) {
|
||||
static_assert(std::is_same<T, half>::value || std::is_same<T, nv_bfloat16>::value, "Only half and bfloat16 input are supported");
|
||||
using PackedVec = PackedVec<T>;
|
||||
|
||||
const int batch_id = blockIdx.y;
|
||||
const int head_id = blockIdx.z;
|
||||
const int token_block_id = blockIdx.x;
|
||||
|
||||
static_assert(CVT_FP4_ELTS_PER_THREAD == 8 || CVT_FP4_ELTS_PER_THREAD == 16,
|
||||
"CVT_FP4_ELTS_PER_THREAD must be 8 or 16");
|
||||
static_assert(sizeof(PackedVec) == sizeof(T) * CVT_FP4_ELTS_PER_THREAD,
|
||||
"Vec size is not matched.");
|
||||
|
||||
constexpr uint32_t NUM_THREADS_PER_TOKEN = head_dim / CVT_FP4_ELTS_PER_THREAD;
|
||||
|
||||
// load input
|
||||
const int token_id = token_block_id * BLOCK_SIZE + threadIdx.x / NUM_THREADS_PER_TOKEN;
|
||||
|
||||
int load_token_id;
|
||||
if constexpr (!permute) {
|
||||
load_token_id = token_id;
|
||||
} else {
|
||||
int local_token_id = threadIdx.x / NUM_THREADS_PER_TOKEN;
|
||||
int local_token_id_residue = local_token_id % 32;
|
||||
// [0, 1, 8, 9, 16, 17, 24, 25, 2, 3, 10, 11, 18, 19, 26, 27, 4, 5, 12, 13, 20, 21, 28, 29, 6, 7, 14, 15, 22, 23, 30, 31]
|
||||
load_token_id = token_block_id * BLOCK_SIZE + (local_token_id / 32) * 32 +
|
||||
(local_token_id_residue / 8) * 2 +
|
||||
((local_token_id_residue % 8) / 2) * 8 +
|
||||
(local_token_id_residue % 8) % 2;
|
||||
}
|
||||
|
||||
PackedVec in_vec;
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
reinterpret_cast<uint32_t&>(in_vec.elts[i]) = 0;
|
||||
}
|
||||
|
||||
if (load_token_id < num_tokens) {
|
||||
in_vec = reinterpret_cast<PackedVec const*>(input +
|
||||
batch_id * stride_bz_input + // batch dim
|
||||
head_id * stride_h_input + // head dim
|
||||
load_token_id * stride_seq_input + // seq dim
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD)[0]; // feature dim
|
||||
}
|
||||
|
||||
// calculate max of every consecutive 16 elements
|
||||
auto localMax = __habs2(in_vec.elts[0]);
|
||||
#pragma unroll
|
||||
for (int i = 1; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) { // local max
|
||||
localMax = __hmax2(localMax, __habs2(in_vec.elts[i]));
|
||||
}
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) { // shuffle across two threads
|
||||
localMax = __hmax2(__shfl_xor_sync(0xffffffff, localMax, 1, 32), localMax);
|
||||
}
|
||||
|
||||
float vecMax = float(__hmax(localMax.x, localMax.y));
|
||||
|
||||
// scaling factor
|
||||
float SFValue = vecMax / 6.0f;
|
||||
uint8_t SFValueFP8;
|
||||
reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8) = __nv_fp8_e4m3(SFValue);
|
||||
SFValue = float(reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8));
|
||||
|
||||
float SFValueInv = (SFValue == 0.0f) ? 0.0f : 1.0f / SFValue;
|
||||
|
||||
// convert input to float2 and apply scale
|
||||
float2 fp2Vals[CVT_FP4_ELTS_PER_THREAD / 2];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
if constexpr (std::is_same<T, half>::value) {
|
||||
fp2Vals[i] = __half22float2(in_vec.elts[i]);
|
||||
} else {
|
||||
fp2Vals[i] = __bfloat1622float2(in_vec.elts[i]);
|
||||
}
|
||||
fp2Vals[i].x = fp2Vals[i].x * SFValueInv;
|
||||
fp2Vals[i].y = fp2Vals[i].y * SFValueInv;
|
||||
}
|
||||
|
||||
// convert to e2m1
|
||||
uint32_t e2m1Vals[CVT_FP4_ELTS_PER_THREAD / 8];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 8; i++) {
|
||||
e2m1Vals[i] = fp32_vec_to_e2m1(fp2Vals + i * 4);
|
||||
}
|
||||
|
||||
// Skip out-of-range tokens: never write past the sequence when num_tokens
|
||||
// is not a multiple of BLOCK_SIZE (out-of-bounds write fix).
|
||||
if (token_id >= num_tokens) return;
|
||||
|
||||
// save
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) {
|
||||
reinterpret_cast<uint32_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
token_id * stride_seq_output +
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD / 2)[0] = e2m1Vals[0];
|
||||
} else {
|
||||
reinterpret_cast<uint64_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
token_id * stride_seq_output +
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD / 2)[0] = reinterpret_cast<uint64_t*>(e2m1Vals)[0];
|
||||
}
|
||||
|
||||
uint8_t* output_sf_save_base = output_sf + batch_id * stride_bz_output_sf + head_id * stride_h_output_sf + (token_id / 64) * 64 * stride_seq_output_sf;
|
||||
uint32_t token_id_local = token_id % 64;
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 16) {
|
||||
uint32_t col_id_local = threadIdx.x % NUM_THREADS_PER_TOKEN;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(token_id_local / 16) * 4 + (token_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
} else {
|
||||
if (threadIdx.x % 2 == 0) {
|
||||
uint32_t col_id_local = (threadIdx.x % NUM_THREADS_PER_TOKEN) / 2;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(token_id_local / 16) * 4 + (token_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <uint32_t head_dim, uint32_t BLOCK_SIZE, typename T>
|
||||
__global__ void scaled_fp4_quant_trans_kernel(
|
||||
const T* input, uint8_t* output, uint8_t* output_sf,
|
||||
int batch_size, int num_heads, int num_tokens,
|
||||
int stride_bz_input, int stride_h_input, int stride_seq_input,
|
||||
int stride_bz_output, int stride_h_output, int stride_d_output,
|
||||
int stride_bz_output_sf, int stride_h_output_sf, int stride_d_output_sf) {
|
||||
static_assert(std::is_same<T, half>::value || std::is_same<T, nv_bfloat16>::value, "Only half and bfloat16 input are supported");
|
||||
using PackedVec = PackedVec<T>;
|
||||
|
||||
const int batch_id = blockIdx.y;
|
||||
const int head_id = blockIdx.z;
|
||||
const int token_block_id = blockIdx.x;
|
||||
|
||||
static_assert(CVT_FP4_ELTS_PER_THREAD == 8 || CVT_FP4_ELTS_PER_THREAD == 16,
|
||||
"CVT_FP4_ELTS_PER_THREAD must be 8 or 16");
|
||||
static_assert(sizeof(PackedVec) == sizeof(T) * CVT_FP4_ELTS_PER_THREAD,
|
||||
"Vec size is not matched.");
|
||||
|
||||
constexpr uint32_t NUM_THREADS_PER_TOKEN = head_dim / CVT_FP4_ELTS_PER_THREAD;
|
||||
constexpr uint32_t NUM_THREADS_PER_SEQ = BLOCK_SIZE / CVT_FP4_ELTS_PER_THREAD;
|
||||
|
||||
// load input
|
||||
const int token_id = token_block_id * BLOCK_SIZE + threadIdx.x / NUM_THREADS_PER_TOKEN;
|
||||
// Permute V rows within each 32-element block so the PV MMA K-indexed
|
||||
// access reads the correct CLayout N-indexed values (Edenzzzz causal fix).
|
||||
const int k_intra = token_id & 31;
|
||||
const int load_token_id = (token_id & ~31)
|
||||
| ((k_intra & 6) << 2) | ((k_intra & 24) >> 2) | (k_intra & 1);
|
||||
|
||||
PackedVec in_vec;
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
reinterpret_cast<uint32_t&>(in_vec.elts[i]) = 0;
|
||||
}
|
||||
|
||||
if (load_token_id < num_tokens) {
|
||||
in_vec = reinterpret_cast<PackedVec const*>(input +
|
||||
batch_id * stride_bz_input + // batch dim
|
||||
head_id * stride_h_input + // head dim
|
||||
load_token_id * stride_seq_input + // seq dim (permuted)
|
||||
(threadIdx.x % NUM_THREADS_PER_TOKEN) * CVT_FP4_ELTS_PER_THREAD)[0]; // feature dim
|
||||
}
|
||||
|
||||
// transpose
|
||||
__shared__ T shared_input[BLOCK_SIZE * head_dim];
|
||||
reinterpret_cast<PackedVec*>(shared_input)[threadIdx.x] = in_vec;
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
in_vec.elts[i].x = shared_input[(threadIdx.x / NUM_THREADS_PER_SEQ) + ((threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD + 2 * i) * head_dim];
|
||||
in_vec.elts[i].y = shared_input[(threadIdx.x / NUM_THREADS_PER_SEQ) + ((threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD + 2 * i + 1) * head_dim];
|
||||
}
|
||||
|
||||
// calculate max of every consecutive 16 elements
|
||||
auto localMax = __habs2(in_vec.elts[0]);
|
||||
#pragma unroll
|
||||
for (int i = 1; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) { // local max
|
||||
localMax = __hmax2(localMax, __habs2(in_vec.elts[i]));
|
||||
}
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) { // shuffle across two threads
|
||||
localMax = __hmax2(__shfl_xor_sync(0xffffffff, localMax, 1, 32), localMax);
|
||||
}
|
||||
|
||||
float vecMax = float(__hmax(localMax.x, localMax.y));
|
||||
|
||||
// scaling factor
|
||||
float SFValue = vecMax / 6.0f;
|
||||
uint8_t SFValueFP8;
|
||||
reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8) = __nv_fp8_e4m3(SFValue);
|
||||
SFValue = float(reinterpret_cast<__nv_fp8_e4m3&>(SFValueFP8));
|
||||
|
||||
float SFValueInv = (SFValue == 0.0f) ? 0.0f : 1.0f / SFValue;
|
||||
|
||||
// convert input to float2 and apply scale
|
||||
float2 fp2Vals[CVT_FP4_ELTS_PER_THREAD / 2];
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 2; i++) {
|
||||
if constexpr (std::is_same<T, half>::value) {
|
||||
fp2Vals[i] = __half22float2(in_vec.elts[i]);
|
||||
} else {
|
||||
fp2Vals[i] = __bfloat1622float2(in_vec.elts[i]);
|
||||
}
|
||||
fp2Vals[i].x = fp2Vals[i].x * SFValueInv;
|
||||
fp2Vals[i].y = fp2Vals[i].y * SFValueInv;
|
||||
}
|
||||
|
||||
// convert to e2m1
|
||||
uint32_t e2m1Vals[CVT_FP4_ELTS_PER_THREAD / 8];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < CVT_FP4_ELTS_PER_THREAD / 8; i++) {
|
||||
e2m1Vals[i] = fp32_vec_to_e2m1(fp2Vals + i * 4);
|
||||
}
|
||||
|
||||
// Skip out-of-range tokens: never write past the sequence when num_tokens
|
||||
// is not a multiple of BLOCK_SIZE (out-of-bounds write fix).
|
||||
const int write_token_id = token_block_id * BLOCK_SIZE +
|
||||
(threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD;
|
||||
if (write_token_id >= num_tokens) return;
|
||||
|
||||
// save
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 8) {
|
||||
reinterpret_cast<uint32_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
(threadIdx.x / NUM_THREADS_PER_SEQ) * stride_d_output +
|
||||
(token_block_id * BLOCK_SIZE + (threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD) / 2)[0] = e2m1Vals[0];
|
||||
} else {
|
||||
reinterpret_cast<uint64_t*>(output +
|
||||
batch_id * stride_bz_output +
|
||||
head_id * stride_h_output +
|
||||
(threadIdx.x / NUM_THREADS_PER_SEQ) * stride_d_output +
|
||||
(token_block_id * BLOCK_SIZE + (threadIdx.x % NUM_THREADS_PER_SEQ) * CVT_FP4_ELTS_PER_THREAD) / 2)[0] = reinterpret_cast<uint64_t*>(e2m1Vals)[0];
|
||||
}
|
||||
|
||||
uint8_t *output_sf_save_base = output_sf +
|
||||
batch_id * stride_bz_output_sf +
|
||||
head_id * stride_h_output_sf +
|
||||
(threadIdx.x / NUM_THREADS_PER_SEQ / 64) * 64 * stride_d_output_sf;
|
||||
uint32_t row_id_local = (threadIdx.x / NUM_THREADS_PER_SEQ) % 64;
|
||||
|
||||
if constexpr (CVT_FP4_ELTS_PER_THREAD == 16) {
|
||||
uint32_t col_id_local = token_block_id * BLOCK_SIZE / CVT_FP4_ELTS_PER_THREAD + threadIdx.x % NUM_THREADS_PER_SEQ;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(row_id_local / 16) * 4 + (row_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
} else {
|
||||
if (threadIdx.x % 2 == 0) {
|
||||
uint32_t col_id_local = token_block_id * BLOCK_SIZE / CVT_FP4_ELTS_PER_THREAD + (threadIdx.x % NUM_THREADS_PER_SEQ) / 2;
|
||||
uint32_t offset_local = (col_id_local / 4) * 256 + (col_id_local % 4) +
|
||||
(row_id_local / 16) * 4 + (row_id_local % 16) * 16;
|
||||
reinterpret_cast<uint8_t*>(output_sf_save_base + offset_local)[0] = SFValueFP8;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void scaled_fp4_quant(torch::Tensor const& input,
|
||||
torch::Tensor const& output,
|
||||
torch::Tensor const& output_sf,
|
||||
int tensor_layout) {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
|
||||
CHECK_CUDA(input);
|
||||
CHECK_CUDA(output);
|
||||
CHECK_CUDA(output_sf);
|
||||
|
||||
CHECK_LASTDIM_CONTIGUOUS(input);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output_sf);
|
||||
|
||||
CHECK_DTYPE(output, at::ScalarType::Byte);
|
||||
CHECK_DTYPE(output_sf, at::ScalarType::Float8_e4m3fn);
|
||||
|
||||
CHECK_DIMS(input, 4);
|
||||
CHECK_DIMS(output, 4);
|
||||
CHECK_DIMS(output_sf, 4);
|
||||
|
||||
const int batch_size = input.size(0);
|
||||
const int head_dim = input.size(3);
|
||||
|
||||
const int stride_bz_input = input.stride(0);
|
||||
const int stride_bz_output = output.stride(0);
|
||||
const int stride_bz_output_sf = output_sf.stride(0);
|
||||
|
||||
int num_tokens, num_heads;
|
||||
int stride_seq_input, stride_seq_output, stride_seq_output_sf;
|
||||
int stride_h_input, stride_h_output, stride_h_output_sf;
|
||||
if (tensor_layout == 0) {
|
||||
num_tokens = input.size(1);
|
||||
num_heads = input.size(2);
|
||||
stride_seq_input = input.stride(1);
|
||||
stride_seq_output = output.stride(1);
|
||||
stride_seq_output_sf = output_sf.stride(1);
|
||||
stride_h_input = input.stride(2);
|
||||
stride_h_output = output.stride(2);
|
||||
stride_h_output_sf = output_sf.stride(2);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_tokens, num_heads, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_tokens, num_heads, head_dim / 16);
|
||||
} else {
|
||||
num_tokens = input.size(2);
|
||||
num_heads = input.size(1);
|
||||
stride_seq_input = input.stride(2);
|
||||
stride_seq_output = output.stride(2);
|
||||
stride_seq_output_sf = output_sf.stride(2);
|
||||
stride_h_input = input.stride(1);
|
||||
stride_h_output = output.stride(1);
|
||||
stride_h_output_sf = output_sf.stride(1);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_heads, num_tokens, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_heads, num_tokens, head_dim / 16);
|
||||
}
|
||||
|
||||
auto input_dtype = input.scalar_type();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input_dtype, c_type, {
|
||||
DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, {
|
||||
dim3 block(BLOCK_SIZE * HEAD_DIM / CVT_FP4_ELTS_PER_THREAD, 1, 1);
|
||||
dim3 grid((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE, batch_size, num_heads);
|
||||
|
||||
scaled_fp4_quant_kernel<HEAD_DIM, BLOCK_SIZE, false, c_type>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<c_type*>(input.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output_sf.data_ptr()),
|
||||
batch_size, num_heads, num_tokens,
|
||||
stride_bz_input, stride_h_input, stride_seq_input,
|
||||
stride_bz_output, stride_h_output, stride_seq_output,
|
||||
stride_bz_output_sf, stride_h_output_sf, stride_seq_output_sf);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
void scaled_fp4_quant_permute(torch::Tensor const& input,
|
||||
torch::Tensor const& output,
|
||||
torch::Tensor const& output_sf,
|
||||
int tensor_layout) {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
|
||||
CHECK_CUDA(input);
|
||||
CHECK_CUDA(output);
|
||||
CHECK_CUDA(output_sf);
|
||||
|
||||
CHECK_LASTDIM_CONTIGUOUS(input);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output_sf);
|
||||
|
||||
CHECK_DTYPE(output, at::ScalarType::Byte);
|
||||
CHECK_DTYPE(output_sf, at::ScalarType::Float8_e4m3fn);
|
||||
|
||||
CHECK_DIMS(input, 4);
|
||||
CHECK_DIMS(output, 4);
|
||||
CHECK_DIMS(output_sf, 4);
|
||||
|
||||
const int batch_size = input.size(0);
|
||||
const int head_dim = input.size(3);
|
||||
|
||||
const int stride_bz_input = input.stride(0);
|
||||
const int stride_bz_output = output.stride(0);
|
||||
const int stride_bz_output_sf = output_sf.stride(0);
|
||||
|
||||
int num_tokens, num_heads;
|
||||
int stride_seq_input, stride_seq_output, stride_seq_output_sf;
|
||||
int stride_h_input, stride_h_output, stride_h_output_sf;
|
||||
if (tensor_layout == 0) {
|
||||
num_tokens = input.size(1);
|
||||
num_heads = input.size(2);
|
||||
stride_seq_input = input.stride(1);
|
||||
stride_seq_output = output.stride(1);
|
||||
stride_seq_output_sf = output_sf.stride(1);
|
||||
stride_h_input = input.stride(2);
|
||||
stride_h_output = output.stride(2);
|
||||
stride_h_output_sf = output_sf.stride(2);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, num_heads, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, num_heads, head_dim / 16);
|
||||
} else {
|
||||
num_tokens = input.size(2);
|
||||
num_heads = input.size(1);
|
||||
stride_seq_input = input.stride(2);
|
||||
stride_seq_output = output.stride(2);
|
||||
stride_seq_output_sf = output_sf.stride(2);
|
||||
stride_h_input = input.stride(1);
|
||||
stride_h_output = output.stride(1);
|
||||
stride_h_output_sf = output_sf.stride(1);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, head_dim / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE, head_dim / 16);
|
||||
}
|
||||
|
||||
auto input_dtype = input.scalar_type();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input_dtype, c_type, {
|
||||
DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
dim3 block(BLOCK_SIZE * HEAD_DIM / CVT_FP4_ELTS_PER_THREAD, 1, 1);
|
||||
dim3 grid((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE, batch_size, num_heads);
|
||||
|
||||
scaled_fp4_quant_kernel<HEAD_DIM, BLOCK_SIZE, true, c_type>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<c_type*>(input.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output_sf.data_ptr()),
|
||||
batch_size, num_heads, num_tokens,
|
||||
stride_bz_input, stride_h_input, stride_seq_input,
|
||||
stride_bz_output, stride_h_output, stride_seq_output,
|
||||
stride_bz_output_sf, stride_h_output_sf, stride_seq_output_sf);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
void scaled_fp4_quant_trans(torch::Tensor const& input,
|
||||
torch::Tensor const& output,
|
||||
torch::Tensor const& output_sf,
|
||||
int tensor_layout) {
|
||||
constexpr int BLOCK_SIZE = flash::BLOCK_M;
|
||||
|
||||
CHECK_CUDA(input);
|
||||
CHECK_CUDA(output);
|
||||
CHECK_CUDA(output_sf);
|
||||
|
||||
CHECK_LASTDIM_CONTIGUOUS(input);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output);
|
||||
CHECK_LASTDIM_CONTIGUOUS(output_sf);
|
||||
|
||||
CHECK_DTYPE(output, at::ScalarType::Byte);
|
||||
CHECK_DTYPE(output_sf, at::ScalarType::Float8_e4m3fn);
|
||||
|
||||
CHECK_DIMS(input, 4);
|
||||
CHECK_DIMS(output, 4);
|
||||
CHECK_DIMS(output_sf, 4);
|
||||
|
||||
const int batch_size = input.size(0);
|
||||
const int head_dim = input.size(3);
|
||||
|
||||
const int stride_bz_input = input.stride(0);
|
||||
const int stride_bz_output = output.stride(0);
|
||||
const int stride_bz_output_sf = output_sf.stride(0);
|
||||
|
||||
int num_tokens, num_heads;
|
||||
int stride_seq_input;
|
||||
int stride_d_output, stride_d_output_sf;
|
||||
int stride_h_input, stride_h_output, stride_h_output_sf;
|
||||
if (tensor_layout == 0) {
|
||||
num_tokens = input.size(1);
|
||||
num_heads = input.size(2);
|
||||
stride_seq_input = input.stride(1);
|
||||
stride_d_output = output.stride(1);
|
||||
stride_d_output_sf = output_sf.stride(1);
|
||||
stride_h_input = input.stride(2);
|
||||
stride_h_output = output.stride(2);
|
||||
stride_h_output_sf = output_sf.stride(2);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, head_dim, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, head_dim, num_heads, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 16);
|
||||
} else {
|
||||
num_tokens = input.size(2);
|
||||
num_heads = input.size(1);
|
||||
stride_seq_input = input.stride(2);
|
||||
stride_d_output = output.stride(2);
|
||||
stride_d_output_sf = output_sf.stride(2);
|
||||
stride_h_input = input.stride(1);
|
||||
stride_h_output = output.stride(1);
|
||||
stride_h_output_sf = output_sf.stride(1);
|
||||
|
||||
CHECK_SHAPE(output, batch_size, num_heads, head_dim, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 2);
|
||||
CHECK_SHAPE(output_sf, batch_size, num_heads, head_dim, ((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE) * BLOCK_SIZE / 16);
|
||||
}
|
||||
|
||||
auto input_dtype = input.scalar_type();
|
||||
auto stream = at::cuda::getCurrentCUDAStream(input.get_device());
|
||||
|
||||
DISPATCH_PYTORCH_DTYPE_TO_CTYPE_FP16(input_dtype, c_type, {
|
||||
DISPATCH_HEAD_DIM(head_dim, HEAD_DIM, {
|
||||
dim3 block(BLOCK_SIZE * HEAD_DIM / CVT_FP4_ELTS_PER_THREAD, 1, 1);
|
||||
dim3 grid((num_tokens + BLOCK_SIZE - 1) / BLOCK_SIZE, batch_size, num_heads);
|
||||
|
||||
scaled_fp4_quant_trans_kernel<HEAD_DIM, BLOCK_SIZE, c_type>
|
||||
<<<grid, block, 0, stream>>>(
|
||||
reinterpret_cast<c_type*>(input.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output.data_ptr()),
|
||||
reinterpret_cast<uint8_t*>(output_sf.data_ptr()),
|
||||
batch_size, num_heads, num_tokens,
|
||||
stride_bz_input, stride_h_input, stride_seq_input,
|
||||
stride_bz_output, stride_h_output, stride_d_output,
|
||||
stride_bz_output_sf, stride_h_output_sf, stride_d_output_sf);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("scaled_fp4_quant", &scaled_fp4_quant);
|
||||
m.def("scaled_fp4_quant_permute", &scaled_fp4_quant_permute);
|
||||
m.def("scaled_fp4_quant_trans", &scaled_fp4_quant_trans);
|
||||
}
|
||||
@@ -32,4 +32,4 @@ dependencies = [
|
||||
[tool.scikit-build]
|
||||
cmake.build-type = "Release"
|
||||
minimum-version = "build-system.requires"
|
||||
wheel.packages = ["python/fastvideo_kernel"]
|
||||
wheel.packages = ["python/fastvideo_kernel", "attn_qat_infer"]
|
||||
|
||||
@@ -25,12 +25,17 @@ from fastvideo_kernel.turbodiffusion_ops import (
|
||||
int8_quant,
|
||||
)
|
||||
|
||||
from fastvideo_kernel.block_sparse_attn_varlen import (
|
||||
block_sparse_attn_varlen,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"sliding_tile_attention",
|
||||
"video_sparse_attn",
|
||||
"video_sparse_attn_bshd",
|
||||
"block_sparse_attn",
|
||||
"block_sparse_attn_from_indices",
|
||||
"block_sparse_attn_varlen",
|
||||
"moba_attn_varlen",
|
||||
"process_moba_input",
|
||||
"process_moba_output",
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
"""Variable-length block-sparse attention via sequence packing.
|
||||
|
||||
Packs multiple variable-length sequences into a single [1, H, T_total, D]
|
||||
tensor and delegates to the existing block_sparse_attn_from_indices kernel
|
||||
in a single launch. No kernel modifications required.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Sequence
|
||||
|
||||
import torch
|
||||
|
||||
from .block_sparse_attn import block_sparse_attn_from_indices
|
||||
|
||||
BLOCK_SIZE = 64
|
||||
|
||||
|
||||
def _scatter_to_padded(
|
||||
src: torch.Tensor,
|
||||
block_sizes: torch.Tensor,
|
||||
block_size: int,
|
||||
dst: torch.Tensor,
|
||||
dst_offset: int,
|
||||
src_start: int,
|
||||
src_end: int,
|
||||
) -> None:
|
||||
"""Copy tokens from a flat source into block-aligned positions in dst.
|
||||
|
||||
Each block occupies exactly `block_size` slots in dst. The first
|
||||
`block_sizes[b]` slots of block *b* receive real tokens; the remainder
|
||||
stays zero (padding the kernel expects).
|
||||
|
||||
src: [total_tokens, H, D]
|
||||
dst: [1, H, total_padded, D]
|
||||
block_sizes: [num_blocks] int32, actual token count per block.
|
||||
"""
|
||||
src_pos = src_start
|
||||
dst_pos = dst_offset
|
||||
sizes = block_sizes.cpu().tolist()
|
||||
for actual in sizes:
|
||||
actual = min(actual, src_end - src_pos)
|
||||
if actual > 0:
|
||||
dst[:, :, dst_pos:dst_pos + actual, :] = (
|
||||
src[src_pos:src_pos + actual].transpose(0, 1).unsqueeze(0)
|
||||
)
|
||||
src_pos += actual
|
||||
dst_pos += block_size
|
||||
|
||||
|
||||
def _gather_from_padded(
|
||||
src: torch.Tensor,
|
||||
block_sizes: torch.Tensor,
|
||||
block_size: int,
|
||||
dst: torch.Tensor,
|
||||
src_offset: int,
|
||||
dst_start: int,
|
||||
dst_end: int,
|
||||
) -> None:
|
||||
"""Inverse of _scatter_to_padded: extract real tokens from padded blocks.
|
||||
|
||||
src: [1, H, total_padded, D]
|
||||
dst: [total_tokens, H, D]
|
||||
"""
|
||||
src_pos = src_offset
|
||||
dst_pos = dst_start
|
||||
sizes = block_sizes.cpu().tolist()
|
||||
for actual in sizes:
|
||||
actual = min(actual, dst_end - dst_pos)
|
||||
if actual > 0:
|
||||
dst[dst_pos:dst_pos + actual] = (
|
||||
src[0, :, src_pos:src_pos + actual, :].transpose(0, 1)
|
||||
)
|
||||
dst_pos += actual
|
||||
src_pos += block_size
|
||||
|
||||
|
||||
def block_sparse_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens_q: torch.Tensor,
|
||||
cu_seqlens_kv: torch.Tensor,
|
||||
q2k_idx_list: Sequence[torch.Tensor],
|
||||
q2k_num_list: Sequence[torch.Tensor],
|
||||
variable_block_sizes_list: Sequence[torch.Tensor],
|
||||
q_variable_block_sizes_list: Sequence[torch.Tensor] | None = None,
|
||||
block_size: int = BLOCK_SIZE,
|
||||
) -> torch.Tensor:
|
||||
"""Block-sparse attention over packed variable-length sequences.
|
||||
|
||||
Args:
|
||||
q: [total_q_tokens, H, D] packed query tensor.
|
||||
k: [total_kv_tokens, H, D] packed key tensor.
|
||||
v: [total_kv_tokens, H, D] packed value tensor.
|
||||
cu_seqlens_q: [N+1] int32, cumulative Q token offsets.
|
||||
cu_seqlens_kv: [N+1] int32, cumulative KV token offsets.
|
||||
q2k_idx_list: Per-sequence q2k_idx tensors, each [1, H, Nq_i, Mk].
|
||||
q2k_num_list: Per-sequence q2k_num tensors, each [1, H, Nq_i].
|
||||
variable_block_sizes_list: Per-sequence KV block sizes, each [Nkv_i].
|
||||
q_variable_block_sizes_list: Per-sequence Q block sizes, each [Nq_i].
|
||||
If None, each Q block is assumed to be exactly `block_size` tokens.
|
||||
block_size: Attention block size (default 64).
|
||||
|
||||
Returns:
|
||||
out: [total_q_tokens, H, D] packed output tensor.
|
||||
"""
|
||||
device = q.device
|
||||
dtype = q.dtype
|
||||
num_heads = q.shape[1]
|
||||
head_dim = q.shape[2]
|
||||
num_seqs = cu_seqlens_q.shape[0] - 1
|
||||
|
||||
cu_q = cu_seqlens_q.cpu().tolist()
|
||||
cu_kv = cu_seqlens_kv.cpu().tolist()
|
||||
|
||||
padded_q_lens = []
|
||||
padded_kv_lens = []
|
||||
q_block_offsets = [0]
|
||||
kv_block_offsets = [0]
|
||||
q_vbs_resolved = []
|
||||
|
||||
for i in range(num_seqs):
|
||||
n_q_blocks = q2k_num_list[i].shape[-1]
|
||||
n_kv_blocks = variable_block_sizes_list[i].numel()
|
||||
padded_q_lens.append(n_q_blocks * block_size)
|
||||
padded_kv_lens.append(n_kv_blocks * block_size)
|
||||
q_block_offsets.append(q_block_offsets[-1] + n_q_blocks)
|
||||
kv_block_offsets.append(kv_block_offsets[-1] + n_kv_blocks)
|
||||
|
||||
if q_variable_block_sizes_list is not None:
|
||||
q_vbs_resolved.append(q_variable_block_sizes_list[i])
|
||||
else:
|
||||
q_vbs_resolved.append(
|
||||
torch.full((n_q_blocks,), block_size, dtype=torch.int32)
|
||||
)
|
||||
|
||||
total_padded_q = sum(padded_q_lens)
|
||||
total_padded_kv = sum(padded_kv_lens)
|
||||
|
||||
q_packed = torch.zeros(1, num_heads, total_padded_q, head_dim, device=device, dtype=dtype)
|
||||
k_packed = torch.zeros(1, num_heads, total_padded_kv, head_dim, device=device, dtype=dtype)
|
||||
v_packed = torch.zeros(1, num_heads, total_padded_kv, head_dim, device=device, dtype=dtype)
|
||||
|
||||
q_offset = 0
|
||||
kv_offset = 0
|
||||
for i in range(num_seqs):
|
||||
_scatter_to_padded(
|
||||
q, q_vbs_resolved[i], block_size,
|
||||
q_packed, q_offset, cu_q[i], cu_q[i + 1],
|
||||
)
|
||||
_scatter_to_padded(
|
||||
k, variable_block_sizes_list[i], block_size,
|
||||
k_packed, kv_offset, cu_kv[i], cu_kv[i + 1],
|
||||
)
|
||||
_scatter_to_padded(
|
||||
v, variable_block_sizes_list[i], block_size,
|
||||
v_packed, kv_offset, cu_kv[i], cu_kv[i + 1],
|
||||
)
|
||||
q_offset += padded_q_lens[i]
|
||||
kv_offset += padded_kv_lens[i]
|
||||
|
||||
total_q_blocks = q_block_offsets[-1]
|
||||
max_kv_per_q = max(t.shape[-1] for t in q2k_idx_list)
|
||||
|
||||
global_q2k_idx = torch.zeros(
|
||||
1, num_heads, total_q_blocks, max_kv_per_q,
|
||||
dtype=torch.int32, device=device,
|
||||
)
|
||||
global_q2k_num = torch.zeros(
|
||||
1, num_heads, total_q_blocks,
|
||||
dtype=torch.int32, device=device,
|
||||
)
|
||||
global_vbs_parts = []
|
||||
|
||||
for i in range(num_seqs):
|
||||
qb_start = q_block_offsets[i]
|
||||
qb_end = q_block_offsets[i + 1]
|
||||
n_q_blocks = qb_end - qb_start
|
||||
kv_offset_blocks = kv_block_offsets[i]
|
||||
|
||||
idx = q2k_idx_list[i]
|
||||
num = q2k_num_list[i]
|
||||
vbs = variable_block_sizes_list[i]
|
||||
|
||||
mk = idx.shape[-1]
|
||||
global_q2k_idx[:, :, qb_start:qb_end, :mk] = idx[:, :, :n_q_blocks, :] + kv_offset_blocks
|
||||
global_q2k_num[:, :, qb_start:qb_end] = num[:, :, :n_q_blocks]
|
||||
global_vbs_parts.append(vbs)
|
||||
|
||||
global_vbs = torch.cat(global_vbs_parts, dim=0).to(torch.int32).contiguous()
|
||||
|
||||
out_packed, _ = block_sparse_attn_from_indices(
|
||||
q_packed, k_packed, v_packed,
|
||||
global_q2k_idx, global_q2k_num, global_vbs,
|
||||
)
|
||||
|
||||
out = torch.zeros(cu_q[-1], num_heads, head_dim, device=device, dtype=dtype)
|
||||
q_offset = 0
|
||||
for i in range(num_seqs):
|
||||
_gather_from_padded(
|
||||
out_packed, q_vbs_resolved[i], block_size,
|
||||
out, q_offset, cu_q[i], cu_q[i + 1],
|
||||
)
|
||||
q_offset += padded_q_lens[i]
|
||||
|
||||
return out
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,55 @@
|
||||
"""Compatibility shim for the legacy non-QAT Triton attention import path.
|
||||
|
||||
Historically callers imported
|
||||
``fastvideo_kernel.triton_kernels.fused_attention`` directly. The shared
|
||||
implementation now lives in ``attn_qat_train.py`` and is parameterized by the
|
||||
``IS_QAT`` flag. This module preserves the original public API for tests and
|
||||
downstream users while always dispatching to the non-QAT configuration.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from .attn_qat_train import attention as _attention
|
||||
|
||||
|
||||
def attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
causal: bool,
|
||||
sm_scale: float,
|
||||
warp_specialize: bool = True,
|
||||
) -> torch.Tensor:
|
||||
"""Run the shared Triton attention kernel in non-QAT mode."""
|
||||
use_qat_qkv_backward = True
|
||||
smooth_k = False
|
||||
is_qat = False
|
||||
two_level_quant_p = False
|
||||
fake_quant_p = False
|
||||
use_high_prec_o = False
|
||||
smooth_q = False
|
||||
use_global_sf_p = False
|
||||
use_global_sf_qkv = False
|
||||
|
||||
return _attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
causal,
|
||||
sm_scale,
|
||||
use_qat_qkv_backward,
|
||||
smooth_k,
|
||||
warp_specialize,
|
||||
is_qat,
|
||||
two_level_quant_p,
|
||||
fake_quant_p,
|
||||
use_high_prec_o,
|
||||
smooth_q,
|
||||
use_global_sf_p,
|
||||
use_global_sf_qkv,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["attention"]
|
||||
@@ -0,0 +1,237 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/triton-lang/triton/blob/main/python/triton_kernels/triton_kernels/numerics_details/mxfp_details/_upcast_from_mxfp.py
|
||||
# and https://github.com/triton-lang/triton/blob/main/python/triton_kernels/triton_kernels/numerics_details/mxfp_details/_downcast_to_mxfp.py
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from triton.language.target_info import cuda_capability_geq
|
||||
|
||||
MXFP_BLOCK_SIZE = tl.constexpr(16)
|
||||
|
||||
@triton.jit
|
||||
def _compute_quant_and_scale(
|
||||
src_tensor,
|
||||
valid_src_mask,
|
||||
mx_tensor_dtype: tl.constexpr = tl.uint8,
|
||||
use_global_sf=True,
|
||||
two_level_quant_P=False,
|
||||
):
|
||||
BLOCK_SIZE_OUT_DIM: tl.constexpr = src_tensor.shape[0]
|
||||
BLOCK_SIZE_QUANT_DIM: tl.constexpr = src_tensor.shape[1]
|
||||
BLOCK_SIZE_QUANT_MX_SCALE: tl.constexpr = src_tensor.shape[1] // MXFP_BLOCK_SIZE
|
||||
is_fp4: tl.constexpr = mx_tensor_dtype == tl.uint8
|
||||
|
||||
tl.static_assert(
|
||||
is_fp4
|
||||
or mx_tensor_dtype == tl.float8e4nv
|
||||
or mx_tensor_dtype == tl.float8e5,
|
||||
"mx_tensor_dtype must be uint8, float8e4nv, or float8e5",
|
||||
)
|
||||
|
||||
# Explicit cast to fp32 since most ops are not supported on bfloat16. We avoid needless conversions to and from bf16
|
||||
f32_tensor = src_tensor.to(tl.float32)
|
||||
abs_tensor = tl.abs(f32_tensor)
|
||||
abs_tensor = tl.where(valid_src_mask, abs_tensor, -1.0) # Don't consider padding tensors in scale computation
|
||||
|
||||
if two_level_quant_P:
|
||||
# row max from SageAttn3 paper
|
||||
global_max_val = tl.max(f32_tensor, axis=1, keep_dims=True) # (BLOCK_SIZE_OUT_DIM, 1)
|
||||
global_max_val = tl.maximum(global_max_val, 1e-8)
|
||||
s_enc = ((6 * 448) / global_max_val).reshape([BLOCK_SIZE_OUT_DIM, 1, 1])
|
||||
s_dec = (1 / s_enc)
|
||||
|
||||
abs_tensor = tl.reshape(abs_tensor, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, MXFP_BLOCK_SIZE])
|
||||
|
||||
if use_global_sf and not two_level_quant_P:
|
||||
global_max_val = tl.max(abs_tensor)
|
||||
# Avoid division by zero: if all values are padding (max is 0), use a default scale
|
||||
global_max_val = tl.maximum(global_max_val, 1e-8)
|
||||
s_enc = (6 * 448) / global_max_val
|
||||
s_dec = (1 / s_enc)
|
||||
elif not two_level_quant_P and not use_global_sf:
|
||||
s_dec = 1.0
|
||||
s_enc = 1.0
|
||||
|
||||
max_val = tl.max(abs_tensor, axis=2, keep_dims=True) # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1) # per block maxima
|
||||
s_dec_b = max_val / 6 # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1)
|
||||
s_dec_b_e4m3 = (s_dec_b * s_enc).to(tl.float8e4nv) # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1)
|
||||
s_enc_b = 1 / (s_dec_b_e4m3.to(tl.float32) * s_dec) # (BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1)
|
||||
|
||||
f32_tensor = tl.reshape(f32_tensor, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, MXFP_BLOCK_SIZE])
|
||||
quant_tensor = f32_tensor * s_enc_b
|
||||
|
||||
# Reshape the tensors after scaling
|
||||
quant_tensor = quant_tensor.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM])
|
||||
# Set the invalid portions of the tensor to 0. This will ensure that any padding tensors are 0 in the mx format.
|
||||
quant_tensor = tl.where(valid_src_mask, quant_tensor, 0.0)
|
||||
dequant_scale = s_dec_b_e4m3.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE])
|
||||
|
||||
if is_fp4 and cuda_capability_geq(10, 0):
|
||||
# Convert scaled values to two f32 lanes and use PTX cvt to e2m1x2 with two f32 operands.
|
||||
pairs = tl.reshape(quant_tensor, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM // 2, 2])
|
||||
lo_f, hi_f = tl.split(pairs)
|
||||
lo_f32 = lo_f.to(tl.float32)
|
||||
hi_f32 = hi_f.to(tl.float32)
|
||||
|
||||
# Inline PTX: cvt.rn.satfinite.e2m1x2.f32 takes two f32 sources and produces one .b8 packed e2m1x2.
|
||||
out_tensor = tl.inline_asm_elementwise(
|
||||
"""
|
||||
{
|
||||
.reg .b8 r;
|
||||
cvt.rn.satfinite.e2m1x2.f32 r, $1, $2;
|
||||
mov.b32 $0, {r, r, r, r};
|
||||
}
|
||||
""",
|
||||
constraints="=r,f,f",
|
||||
args=[hi_f32, lo_f32],
|
||||
dtype=tl.uint8,
|
||||
is_pure=True,
|
||||
pack=1,
|
||||
)
|
||||
elif is_fp4:
|
||||
quant_tensor = quant_tensor.to(tl.uint32, bitcast=True)
|
||||
signs = quant_tensor & 0x80000000
|
||||
exponents = (quant_tensor >> 23) & 0xFF
|
||||
mantissas_orig = (quant_tensor & 0x7FFFFF)
|
||||
|
||||
# For RTNE: 0.25 < x < 0.75 maps to 0.5 (denormal); exactly 0.25 maps to 0.0
|
||||
E8_BIAS = 127
|
||||
E2_BIAS = 1
|
||||
# Move implicit bit 1 at the beginning to mantissa for denormals
|
||||
is_subnormal = exponents < E8_BIAS
|
||||
adjusted_exponents = tl.core.sub(E8_BIAS, exponents + 1, sanitize_overflow=False)
|
||||
mantissas_pre = (0x400000 | (mantissas_orig >> 1))
|
||||
mantissas = tl.where(is_subnormal, mantissas_pre >> adjusted_exponents, mantissas_orig)
|
||||
|
||||
# For normal numbers, we change the bias from 127 to 1, and for subnormals, we keep exponent as 0.
|
||||
exponents = tl.maximum(exponents, E8_BIAS - E2_BIAS) - (E8_BIAS - E2_BIAS)
|
||||
|
||||
# Combine sign, exponent, and mantissa, while saturating
|
||||
# Round to nearest, ties to even (RTNE): use guard/sticky and LSB to decide increment
|
||||
m2bits = mantissas >> 21
|
||||
lsb_keep = (m2bits >> 1) & 0x1
|
||||
guard = m2bits & 0x1
|
||||
IS_SRC_FP32: tl.constexpr = src_tensor.dtype == tl.float32
|
||||
if IS_SRC_FP32:
|
||||
bit0_dropped = (mantissas_orig & 0x1) != 0
|
||||
mask = (1 << tl.minimum(adjusted_exponents, 31)) - 1
|
||||
dropped_post = (mantissas_pre & mask) != 0
|
||||
sticky = is_subnormal & (bit0_dropped | dropped_post)
|
||||
sticky |= ((mantissas & 0x1FFFFF) != 0).to(tl.uint32)
|
||||
else:
|
||||
sticky = ((mantissas & 0x1FFFFF) != 0).to(tl.uint32)
|
||||
round_inc = guard & (sticky | lsb_keep)
|
||||
e2m1_tmp = tl.minimum((((exponents << 2) | m2bits) + round_inc) >> 1, 0x7)
|
||||
e2m1_value = ((signs >> 28) | e2m1_tmp).to(tl.uint8)
|
||||
|
||||
e2m1_value = tl.reshape(e2m1_value, [BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM // 2, 2])
|
||||
evens, odds = tl.split(e2m1_value)
|
||||
out_tensor = evens | (odds << 4)
|
||||
else:
|
||||
out_tensor = quant_tensor.to(mx_tensor_dtype)
|
||||
|
||||
return out_tensor, dequant_scale, s_dec
|
||||
|
||||
@triton.jit
|
||||
def _compute_dequant(
|
||||
mx_tensor,
|
||||
scale,
|
||||
s_dec,
|
||||
BLOCK_SIZE_OUT_DIM: tl.constexpr,
|
||||
BLOCK_SIZE_QUANT_DIM: tl.constexpr,
|
||||
dst_dtype: tl.constexpr,
|
||||
):
|
||||
tl.static_assert(BLOCK_SIZE_QUANT_DIM % MXFP_BLOCK_SIZE == 0, f"Block size along quantization block must be a multiple of {MXFP_BLOCK_SIZE=}")
|
||||
# uint8 signifies two fp4 e2m1 values packed into a single byte
|
||||
mx_tensor_dtype: tl.constexpr = mx_tensor.dtype
|
||||
tl.static_assert(dst_dtype == tl.float16 or dst_dtype == tl.bfloat16 or dst_dtype == tl.float32)
|
||||
tl.static_assert(
|
||||
mx_tensor_dtype == tl.uint8
|
||||
or ((mx_tensor_dtype == tl.float8e4nv or mx_tensor_dtype == tl.float8e5) or mx_tensor_dtype == dst_dtype),
|
||||
"mx_tensor_ptr must be uint8 or float8 or dst_dtype")
|
||||
tl.static_assert(scale.dtype == tl.float8e4nv, "scale must be float8e4nv")
|
||||
|
||||
# Determine if we are dealing with fp8 types.
|
||||
is_fp4: tl.constexpr = mx_tensor_dtype == tl.uint8
|
||||
BLOCK_SIZE_QUANT_MX_SCALE: tl.constexpr = BLOCK_SIZE_QUANT_DIM // MXFP_BLOCK_SIZE
|
||||
|
||||
# Upcast the scale to the destination type.
|
||||
if dst_dtype == tl.bfloat16:
|
||||
dst_scale = scale.to(tl.bfloat16)
|
||||
else:
|
||||
dst_scale = scale.to(tl.float32)
|
||||
if dst_dtype == tl.float16:
|
||||
dst_scale = dst_scale.to(tl.float16)
|
||||
|
||||
# Now upcast the tensor.
|
||||
intermediate_dtype: tl.constexpr = tl.bfloat16 if dst_dtype == tl.float32 else dst_dtype
|
||||
if cuda_capability_geq(10, 0):
|
||||
assert is_fp4
|
||||
packed_u32 = tl.inline_asm_elementwise(
|
||||
asm="""
|
||||
{
|
||||
.reg .b8 in_8;
|
||||
.reg .f16x2 out;
|
||||
cvt.u8.u32 in_8, $1;
|
||||
cvt.rn.f16x2.e2m1x2 out, in_8;
|
||||
mov.b32 $0, out;
|
||||
}
|
||||
""",
|
||||
constraints="=r,r",
|
||||
args=[mx_tensor], # tl.uint8 passed in as a 32-bit reg with value in low 8 bits
|
||||
dtype=tl.uint32,
|
||||
is_pure=True,
|
||||
pack=1,
|
||||
)
|
||||
lo_u16 = (packed_u32 & 0xFFFF).to(tl.uint16)
|
||||
hi_u16 = (packed_u32 >> 16).to(tl.uint16)
|
||||
lo_f16 = lo_u16.to(tl.float16, bitcast=True)
|
||||
hi_f16 = hi_u16.to(tl.float16, bitcast=True)
|
||||
|
||||
if intermediate_dtype == tl.float16:
|
||||
x0, x1 = lo_f16, hi_f16
|
||||
else:
|
||||
x0 = lo_f16.to(intermediate_dtype)
|
||||
x1 = hi_f16.to(intermediate_dtype)
|
||||
|
||||
dst_tensor = tl.interleave(x0, x1)
|
||||
|
||||
else:
|
||||
assert is_fp4
|
||||
dst_bias: tl.constexpr = 127 if intermediate_dtype == tl.bfloat16 else 15 # exponent bias
|
||||
dst_0p5: tl.constexpr = 16128 if intermediate_dtype == tl.bfloat16 else 0x3800
|
||||
dst_m_bits: tl.constexpr = 7 if intermediate_dtype == tl.bfloat16 else 10 # mantissa bits
|
||||
# e2m1
|
||||
em0 = mx_tensor & 0x07
|
||||
em1 = mx_tensor & 0x70
|
||||
x0 = (em0.to(tl.uint16) << (dst_m_bits - 1)) | ((mx_tensor & 0x08).to(tl.uint16) << 12)
|
||||
x1 = (em1.to(tl.uint16) << (dst_m_bits - 5)) | ((mx_tensor & 0x80).to(tl.uint16) << 8)
|
||||
# Three cases:
|
||||
# 1) x is normal and non-zero: Correct bias
|
||||
x0 = tl.where((em0 & 0x06) != 0, x0 + ((dst_bias - 1) << dst_m_bits), x0)
|
||||
x1 = tl.where((em1 & 0x60) != 0, x1 + ((dst_bias - 1) << dst_m_bits), x1)
|
||||
# 2) x is subnormal (x == 0bs001 where s is the sign): Map to +-0.5 in the dst type
|
||||
x0 = tl.where(em0 == 0x01, dst_0p5 | (x0 & 0x8000), x0)
|
||||
x1 = tl.where(em1 == 0x10, dst_0p5 | (x1 & 0x8000), x1)
|
||||
# 3) x is zero, do nothing
|
||||
dst_tensor = tl.interleave(x0, x1).to(intermediate_dtype, bitcast=True)
|
||||
|
||||
dst_tensor = dst_tensor.to(dst_dtype)
|
||||
|
||||
# Reshape for proper broadcasting: the scale was stored with a 16‐sized “inner” grouping.
|
||||
dst_tensor = dst_tensor.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, MXFP_BLOCK_SIZE])
|
||||
dst_scale = dst_scale.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_MX_SCALE, 1])
|
||||
scale = scale.reshape(dst_scale.shape)
|
||||
|
||||
out_tensor = dst_tensor * dst_scale * s_dec # NVFP4 has the additional global scale factor
|
||||
if dst_dtype == tl.float32:
|
||||
max_fin = 3.4028234663852886e+38
|
||||
elif dst_dtype == tl.bfloat16:
|
||||
max_fin = 3.3895313892515355e+38
|
||||
else:
|
||||
tl.static_assert(dst_dtype == tl.float16)
|
||||
max_fin = 65504
|
||||
out_tensor = tl.clamp(out_tensor, min=-max_fin, max=max_fin)
|
||||
out_tensor = out_tensor.reshape([BLOCK_SIZE_OUT_DIM, BLOCK_SIZE_QUANT_DIM])
|
||||
out_tensor = out_tensor.to(dst_dtype)
|
||||
return out_tensor
|
||||
@@ -0,0 +1,80 @@
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from .nvfp4_utils import _compute_quant_and_scale, _compute_dequant
|
||||
|
||||
@triton.jit
|
||||
def fake_quantize(src_tensor, valid_src_mask, BLOCK_SIZE_OUT_DIM: tl.constexpr,
|
||||
BLOCK_SIZE_QUANT_DIM: tl.constexpr,
|
||||
dst_dtype: tl.constexpr,
|
||||
mx_tensor_dtype: tl.constexpr = tl.uint8,
|
||||
use_global_sf: tl.constexpr = True,
|
||||
two_level_quant_P: tl.constexpr = False):
|
||||
high_prec_src_tensor = src_tensor
|
||||
src_tensor, src_scale, src_s_dec = _compute_quant_and_scale(src_tensor=src_tensor,
|
||||
valid_src_mask=valid_src_mask,
|
||||
mx_tensor_dtype=mx_tensor_dtype,
|
||||
use_global_sf=use_global_sf,
|
||||
two_level_quant_P=two_level_quant_P)
|
||||
src_tensor = _compute_dequant(mx_tensor=src_tensor,
|
||||
scale=src_scale,
|
||||
s_dec=src_s_dec,
|
||||
BLOCK_SIZE_OUT_DIM=BLOCK_SIZE_OUT_DIM,
|
||||
BLOCK_SIZE_QUANT_DIM=BLOCK_SIZE_QUANT_DIM,
|
||||
dst_dtype=dst_dtype)
|
||||
return src_tensor, high_prec_src_tensor.to(src_tensor.dtype)
|
||||
|
||||
@triton.jit
|
||||
def fake_quantize_q(Q, fake_Q, stride_z_q, stride_h_q,
|
||||
stride_tok_q, stride_d_q,
|
||||
fake_stride_z_q, fake_stride_h_q,
|
||||
fake_stride_tok_q, fake_stride_d_q,
|
||||
H, N_CTX_Q,
|
||||
BLOCK_M: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
use_global_sf: tl.constexpr = True):
|
||||
bhid = tl.program_id(1)
|
||||
adj_q = (stride_h_q * (bhid % H) + stride_z_q * (bhid // H))
|
||||
fake_adj_q = (fake_stride_h_q * (bhid % H) + fake_stride_z_q * (bhid // H))
|
||||
Q += adj_q
|
||||
fake_Q += fake_adj_q
|
||||
|
||||
pid = tl.program_id(0)
|
||||
start_m = pid * BLOCK_M
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
q_valid = offs_m < N_CTX_Q
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok_q + offs_k[None, :] * stride_d_q, mask=q_valid[:, None], other=0.0)
|
||||
q, _ = fake_quantize(src_tensor=q, valid_src_mask=q_valid[:, None], BLOCK_SIZE_OUT_DIM=BLOCK_M, BLOCK_SIZE_QUANT_DIM=HEAD_DIM, dst_dtype=q.dtype, use_global_sf=use_global_sf)
|
||||
tl.store(fake_Q + offs_m[:, None] * fake_stride_tok_q + offs_k[None, :] * fake_stride_d_q, q, mask=q_valid[:, None])
|
||||
|
||||
@triton.jit
|
||||
def fake_quantize_kv(K, V, fake_K, fake_V, stride_z_kv, stride_h_kv,
|
||||
stride_tok_kv, stride_d_kv,
|
||||
fake_stride_z_kv, fake_stride_h_kv,
|
||||
fake_stride_tok_kv, fake_stride_d_kv,
|
||||
H, N_CTX_KV,
|
||||
BLOCK_N: tl.constexpr,
|
||||
HEAD_DIM: tl.constexpr,
|
||||
use_global_sf: tl.constexpr = True):
|
||||
bhid = tl.program_id(1)
|
||||
adj_kv = (stride_h_kv * (bhid % H) + stride_z_kv * (bhid // H))
|
||||
fake_adj_kv = (fake_stride_h_kv * (bhid % H) + fake_stride_z_kv * (bhid // H))
|
||||
K += adj_kv
|
||||
V += adj_kv
|
||||
fake_K += fake_adj_kv
|
||||
fake_V += fake_adj_kv
|
||||
|
||||
pid = tl.program_id(0)
|
||||
start_n = pid * BLOCK_N
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
kv_valid = offs_n < N_CTX_KV
|
||||
k_block = tl.load(K + offs_n[:, None] * stride_tok_kv + offs_k[None, :] * stride_d_kv, mask=kv_valid[:, None], other=0.0)
|
||||
v_block = tl.load(V + offs_n[:, None] * stride_tok_kv + offs_k[None, :] * stride_d_kv, mask=kv_valid[:, None], other=0.0)
|
||||
k, _ = fake_quantize(src_tensor=k_block, valid_src_mask=kv_valid[:, None], BLOCK_SIZE_OUT_DIM=BLOCK_N, BLOCK_SIZE_QUANT_DIM=HEAD_DIM, dst_dtype=k_block.dtype, use_global_sf=use_global_sf)
|
||||
v, _ = fake_quantize(src_tensor=v_block, valid_src_mask=kv_valid[:, None], BLOCK_SIZE_OUT_DIM=BLOCK_N, BLOCK_SIZE_QUANT_DIM=HEAD_DIM, dst_dtype=v_block.dtype, use_global_sf=use_global_sf)
|
||||
tl.store(fake_K + offs_n[:, None] * fake_stride_tok_kv + offs_k[None, :] * fake_stride_d_kv, k, mask=kv_valid[:, None])
|
||||
tl.store(fake_V + offs_n[:, None] * fake_stride_tok_kv + offs_k[None, :] * fake_stride_d_kv, v, mask=kv_valid[:, None])
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,434 @@
|
||||
"""Correctness tests for variable-length block-sparse attention.
|
||||
|
||||
Reference: per-sequence calls to block_sparse_attn_from_indices.
|
||||
Test: single-launch via block_sparse_attn_varlen.
|
||||
Tests cover both forward and backward (gradient) correctness.
|
||||
"""
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
|
||||
from .test_vsa import (
|
||||
BLOCK_M,
|
||||
generate_variable_block_sizes,
|
||||
get_non_pad_index,
|
||||
vsa_pad,
|
||||
generate_tensor,
|
||||
)
|
||||
from .utils import generate_block_sparse_mask_for_function
|
||||
from fastvideo_kernel.block_sparse_attn import (
|
||||
block_sparse_attn_from_indices,
|
||||
_map_to_index,
|
||||
)
|
||||
from fastvideo_kernel.block_sparse_attn_varlen import block_sparse_attn_varlen
|
||||
|
||||
|
||||
def _reference_per_sequence(
|
||||
q_list, k_list, v_list,
|
||||
block_masks, vbs_list,
|
||||
non_pad_q_list, non_pad_kv_list,
|
||||
q_nblocks_list, kv_nblocks_list,
|
||||
):
|
||||
"""Run per-sequence block_sparse_attn and concat outputs."""
|
||||
outs = []
|
||||
for i in range(len(q_list)):
|
||||
q_pad = vsa_pad(q_list[i], non_pad_q_list[i], q_nblocks_list[i], BLOCK_M)
|
||||
k_pad = vsa_pad(k_list[i], non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
v_pad = vsa_pad(v_list[i], non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0))
|
||||
o_pad, _ = block_sparse_attn_from_indices(
|
||||
q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i],
|
||||
)
|
||||
o = o_pad[:, :, non_pad_q_list[i], :]
|
||||
outs.append(o.squeeze(0).transpose(0, 1))
|
||||
return torch.cat(outs, dim=0)
|
||||
|
||||
|
||||
def _run_varlen_test(
|
||||
seq_configs: list,
|
||||
h: int = 8,
|
||||
d: int = 64,
|
||||
topk: int = 2,
|
||||
atol: float = 0.05,
|
||||
rtol: float = 0.02,
|
||||
):
|
||||
"""Core test: compare varlen vs per-sequence reference.
|
||||
|
||||
seq_configs: list of (num_q_blocks, num_kv_blocks) per sequence.
|
||||
"""
|
||||
device = "cuda"
|
||||
num_seqs = len(seq_configs)
|
||||
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
block_masks = []
|
||||
vbs_list = []
|
||||
q_vbs_list = []
|
||||
non_pad_q_list = []
|
||||
non_pad_kv_list = []
|
||||
q_nblocks_list = []
|
||||
kv_nblocks_list = []
|
||||
q2k_idx_list = []
|
||||
q2k_num_list = []
|
||||
q_vbs_for_varlen = []
|
||||
|
||||
cu_q = [0]
|
||||
cu_kv = [0]
|
||||
|
||||
for nq, nkv in seq_configs:
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
vbs_q = generate_variable_block_sizes(nq, device=device)
|
||||
sq = int(vbs_q.sum().item())
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
|
||||
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
npq = get_non_pad_index(vbs_q, nq, BLOCK_M)
|
||||
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
q_list.append(q)
|
||||
k_list.append(k)
|
||||
v_list.append(v)
|
||||
block_masks.append(mask)
|
||||
vbs_list.append(vbs_kv)
|
||||
q_vbs_list.append(vbs_q)
|
||||
non_pad_q_list.append(npq)
|
||||
non_pad_kv_list.append(npkv)
|
||||
q_nblocks_list.append(nq)
|
||||
kv_nblocks_list.append(nkv)
|
||||
q2k_idx_list.append(q2k_idx)
|
||||
q2k_num_list.append(q2k_num)
|
||||
q_vbs_for_varlen.append(vbs_q)
|
||||
|
||||
cu_q.append(cu_q[-1] + sq)
|
||||
cu_kv.append(cu_kv[-1] + skv)
|
||||
|
||||
ref_out = _reference_per_sequence(
|
||||
q_list, k_list, v_list,
|
||||
block_masks, vbs_list,
|
||||
non_pad_q_list, non_pad_kv_list,
|
||||
q_nblocks_list, kv_nblocks_list,
|
||||
)
|
||||
|
||||
q_packed = torch.cat(
|
||||
[qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0,
|
||||
)
|
||||
k_packed = torch.cat(
|
||||
[ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0,
|
||||
)
|
||||
v_packed = torch.cat(
|
||||
[vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0,
|
||||
)
|
||||
|
||||
cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device)
|
||||
|
||||
varlen_out = block_sparse_attn_varlen(
|
||||
q_packed, k_packed, v_packed,
|
||||
cu_seqlens_q, cu_seqlens_kv,
|
||||
q2k_idx_list, q2k_num_list,
|
||||
vbs_list,
|
||||
q_variable_block_sizes_list=q_vbs_for_varlen,
|
||||
)
|
||||
|
||||
max_abs = (ref_out - varlen_out).abs().max().item()
|
||||
mean_abs = ref_out.abs().mean().item()
|
||||
max_rel = max_abs / (mean_abs + 1e-8)
|
||||
|
||||
print(f" seqs={[c for c in seq_configs]}, max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
|
||||
assert max_rel < rtol, f"max relative error {max_rel:.4e} exceeds threshold {rtol}"
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
class TestVSAVarlen:
|
||||
|
||||
def test_equal_length(self):
|
||||
"""Two sequences with same number of blocks."""
|
||||
_run_varlen_test([(4, 4), (4, 4)], h=8, d=64)
|
||||
|
||||
def test_different_lengths(self):
|
||||
"""Three sequences with different block counts."""
|
||||
_run_varlen_test([(2, 3), (5, 4), (3, 6)], h=8, d=64)
|
||||
|
||||
def test_single_sequence(self):
|
||||
"""Degenerate case: single sequence should match non-varlen path."""
|
||||
_run_varlen_test([(8, 8)], h=8, d=64)
|
||||
|
||||
def test_many_heads(self):
|
||||
"""More heads to stress the packing logic."""
|
||||
_run_varlen_test([(3, 4), (5, 3)], h=16, d=128)
|
||||
|
||||
def test_many_sequences(self):
|
||||
"""Stress test: 8 sequences with varying block counts."""
|
||||
configs = [(i + 2, i + 3) for i in range(8)]
|
||||
_run_varlen_test(configs, h=8, d=64)
|
||||
|
||||
def test_topk_equals_num_blocks(self):
|
||||
"""Edge: topk covers all KV blocks (dense attention)."""
|
||||
_run_varlen_test([(3, 3), (4, 4)], h=8, d=64, topk=8)
|
||||
|
||||
def test_single_block_per_sequence(self):
|
||||
"""Minimal: each sequence has exactly 1 Q block and 1 KV block."""
|
||||
_run_varlen_test([(1, 1), (1, 1), (1, 1)], h=8, d=64, topk=1)
|
||||
|
||||
def test_asymmetric_q_kv(self):
|
||||
"""Q and KV have very different block counts."""
|
||||
_run_varlen_test([(1, 8), (8, 1)], h=8, d=64, topk=1)
|
||||
|
||||
def test_without_q_vbs(self):
|
||||
"""Test the default path where q_variable_block_sizes_list is None.
|
||||
|
||||
Uses full block_size=64 for Q blocks so the None path is valid.
|
||||
"""
|
||||
device = "cuda"
|
||||
h, d, topk = 4, 64, 2
|
||||
nq, nkv = 3, 4
|
||||
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
sq = nq * BLOCK_M
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
|
||||
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
k_pad = vsa_pad(k, npkv, nkv, BLOCK_M)
|
||||
v_pad = vsa_pad(v, npkv, nkv, BLOCK_M)
|
||||
ref_out, _ = block_sparse_attn_from_indices(q, k_pad, v_pad, q2k_idx, q2k_num, vbs_kv)
|
||||
ref_flat = ref_out.squeeze(0).transpose(0, 1)
|
||||
|
||||
q_flat = q.squeeze(0).transpose(0, 1)
|
||||
k_flat = k.squeeze(0).transpose(0, 1)
|
||||
v_flat = v.squeeze(0).transpose(0, 1)
|
||||
cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device)
|
||||
cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device)
|
||||
|
||||
varlen_out = block_sparse_attn_varlen(
|
||||
q_flat, k_flat, v_flat,
|
||||
cu_q, cu_kv,
|
||||
[q2k_idx], [q2k_num], [vbs_kv],
|
||||
)
|
||||
|
||||
max_abs = (ref_flat - varlen_out).abs().max().item()
|
||||
mean_abs = ref_flat.abs().mean().item()
|
||||
max_rel = max_abs / (mean_abs + 1e-8)
|
||||
print(f" without_q_vbs: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
|
||||
assert max_rel < 0.02, f"max relative error {max_rel:.4e} exceeds threshold"
|
||||
|
||||
|
||||
def _run_varlen_backward_test(
|
||||
seq_configs: list,
|
||||
h: int = 8,
|
||||
d: int = 64,
|
||||
topk: int = 2,
|
||||
grad_rtol: float = 0.05,
|
||||
):
|
||||
"""Backward correctness: compare dQ/dK/dV from varlen vs per-sequence reference.
|
||||
|
||||
Both paths use the same underlying block_sparse_attn_from_indices kernel
|
||||
(which has registered autograd). The varlen wrapper's scatter/gather must
|
||||
correctly propagate gradients through PyTorch's in-place slice assignment.
|
||||
"""
|
||||
device = "cuda"
|
||||
num_seqs = len(seq_configs)
|
||||
|
||||
q_list = []
|
||||
k_list = []
|
||||
v_list = []
|
||||
block_masks = []
|
||||
vbs_list = []
|
||||
q_vbs_list = []
|
||||
non_pad_q_list = []
|
||||
non_pad_kv_list = []
|
||||
q_nblocks_list = []
|
||||
kv_nblocks_list = []
|
||||
q2k_idx_list = []
|
||||
q2k_num_list = []
|
||||
|
||||
cu_q = [0]
|
||||
cu_kv = [0]
|
||||
|
||||
for nq, nkv in seq_configs:
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
vbs_q = generate_variable_block_sizes(nq, device=device)
|
||||
sq = int(vbs_q.sum().item())
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = generate_tensor((1, h, sq, d), torch.bfloat16, device)
|
||||
k = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
v = generate_tensor((1, h, skv, d), torch.bfloat16, device)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
npq = get_non_pad_index(vbs_q, nq, BLOCK_M)
|
||||
npkv = get_non_pad_index(vbs_kv, nkv, BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
q_list.append(q)
|
||||
k_list.append(k)
|
||||
v_list.append(v)
|
||||
block_masks.append(mask)
|
||||
vbs_list.append(vbs_kv)
|
||||
q_vbs_list.append(vbs_q)
|
||||
non_pad_q_list.append(npq)
|
||||
non_pad_kv_list.append(npkv)
|
||||
q_nblocks_list.append(nq)
|
||||
kv_nblocks_list.append(nkv)
|
||||
q2k_idx_list.append(q2k_idx)
|
||||
q2k_num_list.append(q2k_num)
|
||||
|
||||
cu_q.append(cu_q[-1] + sq)
|
||||
cu_kv.append(cu_kv[-1] + skv)
|
||||
|
||||
# --- Reference: per-sequence backward ---
|
||||
ref_q_grads = []
|
||||
ref_k_grads = []
|
||||
ref_v_grads = []
|
||||
ref_outs = []
|
||||
for i in range(num_seqs):
|
||||
qi = q_list[i].detach().requires_grad_(True)
|
||||
ki = k_list[i].detach().requires_grad_(True)
|
||||
vi = v_list[i].detach().requires_grad_(True)
|
||||
|
||||
q_pad = vsa_pad(qi, non_pad_q_list[i], q_nblocks_list[i], BLOCK_M)
|
||||
k_pad = vsa_pad(ki, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
v_pad = vsa_pad(vi, non_pad_kv_list[i], kv_nblocks_list[i], BLOCK_M)
|
||||
|
||||
q2k_idx, q2k_num = _map_to_index(block_masks[i].unsqueeze(0))
|
||||
o_pad, _ = block_sparse_attn_from_indices(
|
||||
q_pad, k_pad, v_pad, q2k_idx, q2k_num, vbs_list[i],
|
||||
)
|
||||
o = o_pad[:, :, non_pad_q_list[i], :]
|
||||
o_flat = o.squeeze(0).transpose(0, 1)
|
||||
ref_outs.append(o_flat)
|
||||
|
||||
dO = torch.ones_like(o_flat)
|
||||
o_flat.backward(dO)
|
||||
|
||||
ref_q_grads.append(qi.grad.squeeze(0).transpose(0, 1))
|
||||
ref_k_grads.append(ki.grad.squeeze(0).transpose(0, 1))
|
||||
ref_v_grads.append(vi.grad.squeeze(0).transpose(0, 1))
|
||||
|
||||
ref_dq = torch.cat(ref_q_grads, dim=0)
|
||||
ref_dk = torch.cat(ref_k_grads, dim=0)
|
||||
ref_dv = torch.cat(ref_v_grads, dim=0)
|
||||
|
||||
# --- Varlen backward ---
|
||||
q_packed = torch.cat(
|
||||
[qi.squeeze(0).transpose(0, 1) for qi in q_list], dim=0,
|
||||
).detach().requires_grad_(True)
|
||||
k_packed = torch.cat(
|
||||
[ki.squeeze(0).transpose(0, 1) for ki in k_list], dim=0,
|
||||
).detach().requires_grad_(True)
|
||||
v_packed = torch.cat(
|
||||
[vi.squeeze(0).transpose(0, 1) for vi in v_list], dim=0,
|
||||
).detach().requires_grad_(True)
|
||||
|
||||
cu_seqlens_q = torch.tensor(cu_q, dtype=torch.int32, device=device)
|
||||
cu_seqlens_kv = torch.tensor(cu_kv, dtype=torch.int32, device=device)
|
||||
|
||||
varlen_out = block_sparse_attn_varlen(
|
||||
q_packed, k_packed, v_packed,
|
||||
cu_seqlens_q, cu_seqlens_kv,
|
||||
q2k_idx_list, q2k_num_list,
|
||||
vbs_list,
|
||||
q_variable_block_sizes_list=q_vbs_list,
|
||||
)
|
||||
|
||||
dO = torch.ones_like(varlen_out)
|
||||
varlen_out.backward(dO)
|
||||
|
||||
varlen_dq = q_packed.grad
|
||||
varlen_dk = k_packed.grad
|
||||
varlen_dv = v_packed.grad
|
||||
|
||||
for name, ref, actual in [
|
||||
("dQ", ref_dq, varlen_dq),
|
||||
("dK", ref_dk, varlen_dk),
|
||||
("dV", ref_dv, varlen_dv),
|
||||
]:
|
||||
assert actual is not None, f"{name}: gradient is None (autograd chain broken)"
|
||||
max_abs = (ref - actual).abs().max().item()
|
||||
mean_abs = ref.abs().mean().item()
|
||||
max_rel = max_abs / (mean_abs + 1e-8)
|
||||
print(f" {name}: max_abs={max_abs:.4e}, max_rel={max_rel:.4e}")
|
||||
assert max_rel < grad_rtol, (
|
||||
f"{name}: max relative error {max_rel:.4e} exceeds threshold {grad_rtol}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
|
||||
class TestVSAVarlenBackward:
|
||||
|
||||
def test_backward_equal_length(self):
|
||||
"""Backward: two sequences with same number of blocks."""
|
||||
_run_varlen_backward_test([(4, 4), (4, 4)], h=8, d=64)
|
||||
|
||||
def test_backward_different_lengths(self):
|
||||
"""Backward: three sequences with different block counts."""
|
||||
_run_varlen_backward_test([(2, 3), (5, 4), (3, 6)], h=8, d=64)
|
||||
|
||||
def test_backward_single_sequence(self):
|
||||
"""Backward: single sequence should match non-varlen gradient path."""
|
||||
_run_varlen_backward_test([(8, 8)], h=8, d=64)
|
||||
|
||||
def test_backward_many_heads(self):
|
||||
"""Backward: more heads to stress gradient routing."""
|
||||
_run_varlen_backward_test([(3, 4), (5, 3)], h=16, d=128)
|
||||
|
||||
def test_backward_asymmetric_q_kv(self):
|
||||
"""Backward: Q and KV have very different block counts."""
|
||||
_run_varlen_backward_test([(1, 8), (8, 1)], h=8, d=64, topk=1)
|
||||
|
||||
def test_backward_grad_nonzero(self):
|
||||
"""Smoke test: gradients are non-zero (autograd chain is connected)."""
|
||||
device = "cuda"
|
||||
h, d, topk = 4, 64, 2
|
||||
nq, nkv = 3, 4
|
||||
|
||||
vbs_kv = generate_variable_block_sizes(nkv, device=device)
|
||||
vbs_q = generate_variable_block_sizes(nq, device=device)
|
||||
sq = int(vbs_q.sum().item())
|
||||
skv = int(vbs_kv.sum().item())
|
||||
|
||||
q = torch.randn(sq, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
k = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
v = torch.randn(skv, h, d, device=device, dtype=torch.bfloat16, requires_grad=True)
|
||||
|
||||
mask = generate_block_sparse_mask_for_function(h, nq, nkv, topk, device)
|
||||
q2k_idx, q2k_num = _map_to_index(mask.unsqueeze(0))
|
||||
|
||||
cu_q = torch.tensor([0, sq], dtype=torch.int32, device=device)
|
||||
cu_kv = torch.tensor([0, skv], dtype=torch.int32, device=device)
|
||||
|
||||
out = block_sparse_attn_varlen(
|
||||
q, k, v,
|
||||
cu_q, cu_kv,
|
||||
[q2k_idx], [q2k_num], [vbs_kv],
|
||||
q_variable_block_sizes_list=[vbs_q],
|
||||
)
|
||||
|
||||
loss = out.sum()
|
||||
loss.backward()
|
||||
|
||||
assert q.grad is not None, "q.grad is None"
|
||||
assert k.grad is not None, "k.grad is None"
|
||||
assert v.grad is not None, "v.grad is None"
|
||||
assert q.grad.abs().sum().item() > 0, "q.grad is all zeros"
|
||||
assert k.grad.abs().sum().item() > 0, "k.grad is all zeros"
|
||||
assert v.grad.abs().sum().item() > 0, "v.grad is all zeros"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__, "-v", "-s"])
|
||||
@@ -49,6 +49,10 @@ def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
|
||||
return _attn_qat_train_attention
|
||||
|
||||
|
||||
def is_attn_qat_train_available() -> bool:
|
||||
return _get_attn_qat_train_attention() is not None
|
||||
|
||||
|
||||
def attn_qat_train(q_BLHD: torch.Tensor,
|
||||
k_BLHD: torch.Tensor,
|
||||
v_BLHD: torch.Tensor,
|
||||
|
||||
@@ -24,7 +24,9 @@ class DiTArchConfig(ArchConfig):
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.SAGE_ATTN_THREE,
|
||||
AttentionBackendEnum.SLA_ATTN, AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
AttentionBackendEnum.ATTN_QAT_INFER,
|
||||
AttentionBackendEnum.ATTN_QAT_TRAIN, AttentionBackendEnum.SLA_ATTN,
|
||||
AttentionBackendEnum.SAGE_SLA_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -942,6 +942,20 @@ class TransformerLoader(ComponentLoader):
|
||||
dit_config = deepcopy(fastvideo_args.pipeline_config.dit_config)
|
||||
dit_config.update_model_arch(config)
|
||||
|
||||
# Generator-only QAT for DMD distillation: the teacher (real_score) and
|
||||
# critic (fake_score) transformers load with this flag set and must stay
|
||||
# full precision. Drop the nvfp4_qat quant from their copied config, and
|
||||
# mask the global ATTN_QAT_TRAIN env so their attention falls back to dense
|
||||
# (the backend is read globally at build time). The generator loads without
|
||||
# the flag and keeps both.
|
||||
_qat_generator_only = hasattr(fastvideo_args, "_loading_teacher_critic_model")
|
||||
_qat_prev_attn_env = None
|
||||
if _qat_generator_only:
|
||||
dit_config.quant_config = None
|
||||
from fastvideo.attention.selector import _cached_get_attn_backend
|
||||
_qat_prev_attn_env = os.environ.pop("FASTVIDEO_ATTENTION_BACKEND", None)
|
||||
_cached_get_attn_backend.cache_clear()
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
# Find all safetensors files
|
||||
@@ -1028,6 +1042,12 @@ class TransformerLoader(ComponentLoader):
|
||||
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
|
||||
)
|
||||
|
||||
if _qat_generator_only:
|
||||
from fastvideo.attention.selector import _cached_get_attn_backend
|
||||
if _qat_prev_attn_env is not None:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = _qat_prev_attn_env
|
||||
_cached_get_attn_backend.cache_clear()
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
|
||||
@@ -50,13 +50,20 @@ def _maybe_convert_model_to_nvfp4(model: nn.Module) -> None:
|
||||
from fastvideo.layers.quantization.nvfp4_config import (
|
||||
NVFP4QuantizeMethod, convert_model_to_nvfp4,
|
||||
)
|
||||
from fastvideo.layers.quantization.nvfp4_qat_config import (
|
||||
NVFP4QATQuantizeMethod, convert_model_to_fp4,
|
||||
)
|
||||
|
||||
for mod in model.modules():
|
||||
if isinstance(getattr(mod, "quant_method", None),
|
||||
NVFP4QuantizeMethod):
|
||||
qm = getattr(mod, "quant_method", None)
|
||||
if isinstance(qm, NVFP4QuantizeMethod):
|
||||
logger.info("Converting loaded model weights for NVFP4 linear layers")
|
||||
convert_model_to_nvfp4(model)
|
||||
return
|
||||
if isinstance(qm, NVFP4QATQuantizeMethod):
|
||||
logger.info("Converting loaded model weights for NVFP4-QAT linear layers")
|
||||
convert_model_to_fp4(model)
|
||||
return
|
||||
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Local dashboard helpers for FastVideo performance tracking."""
|
||||
@@ -0,0 +1,26 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Run the local performance dashboard server."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
import uvicorn
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description="Run the FastVideo performance dashboard")
|
||||
parser.add_argument("--host", default="127.0.0.1")
|
||||
parser.add_argument("--port", type=int, default=8000)
|
||||
parser.add_argument("--reload", action="store_true")
|
||||
args = parser.parse_args()
|
||||
uvicorn.run(
|
||||
"fastvideo.performance_dashboard.api:app",
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
reload=args.reload,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,191 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastAPI app for the local performance dashboard."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import threading
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, Query
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import FileResponse
|
||||
from fastapi.staticfiles import StaticFiles
|
||||
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
from .service import build_latest_summary, build_trends, filter_records
|
||||
|
||||
DEFAULT_TRACKING_ROOT = "/tmp/fastvideo-perf-dashboard"
|
||||
DEFAULT_DAYS = 90
|
||||
FRONTEND_DIST = os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "..", "..", "performance_dashboard", "frontend", "dist"))
|
||||
|
||||
|
||||
class PerformanceDataStore:
|
||||
|
||||
def __init__(self, tracking_root: str | None = None) -> None:
|
||||
self.tracking_root = tracking_root or os.environ.get("PERFORMANCE_TRACKING_ROOT", DEFAULT_TRACKING_ROOT)
|
||||
self._lock = threading.RLock()
|
||||
self.last_sync_at: str | None = None
|
||||
self.last_sync_error: str | None = None
|
||||
|
||||
@property
|
||||
def repo_id(self) -> str:
|
||||
return hf_store.HF_REPO_ID
|
||||
|
||||
def sync(self) -> dict[str, Any]:
|
||||
with self._lock:
|
||||
try:
|
||||
local_dir = hf_store.sync_from_hf(self.tracking_root, reuse_existing=False)
|
||||
self.last_sync_at = datetime.now(timezone.utc).isoformat()
|
||||
self.last_sync_error = None
|
||||
return {
|
||||
"ok": True,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": local_dir,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": None,
|
||||
}
|
||||
except Exception as exc:
|
||||
self.last_sync_error = str(exc)
|
||||
return {
|
||||
"ok": False,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": self.tracking_root,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": self.last_sync_error,
|
||||
}
|
||||
|
||||
def ensure_synced(self) -> None:
|
||||
if self.last_sync_at is not None:
|
||||
return
|
||||
with self._lock:
|
||||
if self.last_sync_at is None:
|
||||
hf_store.sync_from_hf(self.tracking_root, reuse_existing=True)
|
||||
self.last_sync_at = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
def load_records(self, *, days: int | None = None, successful_only: bool = False) -> list[dict[str, Any]]:
|
||||
self.ensure_synced()
|
||||
return hf_store.load_records(self.tracking_root, days=days, successful_only=successful_only)
|
||||
|
||||
def health(self) -> dict[str, Any]:
|
||||
return {
|
||||
"ok": self.last_sync_error is None,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": self.tracking_root,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": self.last_sync_error,
|
||||
}
|
||||
|
||||
|
||||
def create_app(store: PerformanceDataStore | None = None) -> FastAPI:
|
||||
data_store = store or PerformanceDataStore()
|
||||
app = FastAPI(title="FastVideo Performance Dashboard")
|
||||
app.state.performance_store = data_store
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=[
|
||||
"http://localhost:5173",
|
||||
"http://127.0.0.1:5173",
|
||||
"http://localhost:3000",
|
||||
"http://127.0.0.1:3000",
|
||||
],
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
@app.get("/api/performance/health")
|
||||
def health() -> dict[str, Any]:
|
||||
return data_store.health()
|
||||
|
||||
@app.post("/api/performance/refresh")
|
||||
def refresh() -> dict[str, Any]:
|
||||
return data_store.sync()
|
||||
|
||||
@app.get("/api/performance/records")
|
||||
def records(
|
||||
days: int = Query(DEFAULT_DAYS, ge=1, le=3650),
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
success: bool | None = None,
|
||||
) -> dict[str, Any]:
|
||||
loaded = data_store.load_records(days=days)
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type, success=success)
|
||||
return {
|
||||
"records": filtered,
|
||||
"count": len(filtered),
|
||||
"filters": {
|
||||
"days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"success": success,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
@app.get("/api/performance/summary")
|
||||
def summary(
|
||||
days: int = Query(DEFAULT_DAYS, ge=1, le=3650),
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
# Latest status should be stable when users change the trend window.
|
||||
# Use all cached records for latest/baseline computation; the ``days``
|
||||
# query is kept only so the frontend can share filter state across
|
||||
# endpoints without affecting the summary semantics.
|
||||
loaded = data_store.load_records(days=None)
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
|
||||
rows = build_latest_summary(filtered, max_regression=float(os.environ.get("PERF_MAX_REGRESSION", "0.05")))
|
||||
return {
|
||||
"rows": rows,
|
||||
"count": len(rows),
|
||||
"status_counts": {
|
||||
"pass": sum(1 for row in rows if row["status"] == "pass"),
|
||||
"fail": sum(1 for row in rows if row["status"] == "fail"),
|
||||
},
|
||||
"filters": {
|
||||
"days": None,
|
||||
"trend_window_days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
@app.get("/api/performance/trends")
|
||||
def trends(
|
||||
days: int = Query(DEFAULT_DAYS, ge=1, le=3650),
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
loaded = data_store.load_records(days=days)
|
||||
filtered = filter_records(loaded, model_id=model_id, gpu_type=gpu_type)
|
||||
groups = build_trends(filtered)
|
||||
return {
|
||||
"groups": groups,
|
||||
"count": len(groups),
|
||||
"filters": {
|
||||
"days": days,
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
},
|
||||
"sync": data_store.health(),
|
||||
}
|
||||
|
||||
assets_dir = os.path.join(FRONTEND_DIST, "assets")
|
||||
index_file = os.path.join(FRONTEND_DIST, "index.html")
|
||||
if os.path.isdir(assets_dir) and os.path.isfile(index_file):
|
||||
app.mount("/assets", StaticFiles(directory=assets_dir), name="performance-dashboard-assets")
|
||||
|
||||
@app.get("/{full_path:path}", include_in_schema=False)
|
||||
def frontend(full_path: str) -> FileResponse:
|
||||
return FileResponse(index_file)
|
||||
|
||||
return app
|
||||
|
||||
|
||||
app = create_app()
|
||||
@@ -0,0 +1,26 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Metric definitions shared by the performance dashboard backend."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MetricDefinition:
|
||||
key: str
|
||||
label: str
|
||||
precision: int
|
||||
lower_is_better: bool
|
||||
|
||||
|
||||
METRICS: tuple[MetricDefinition, ...] = (
|
||||
MetricDefinition("latency", "Latency", 3, True),
|
||||
MetricDefinition("throughput", "Throughput", 3, False),
|
||||
MetricDefinition("memory", "Memory", 1, True),
|
||||
MetricDefinition("text_encoder_time_s", "Text Encoder", 3, True),
|
||||
MetricDefinition("dit_time_s", "DiT", 3, True),
|
||||
MetricDefinition("vae_decode_time_s", "VAE Decode", 3, True),
|
||||
)
|
||||
|
||||
METRIC_BY_KEY = {metric.key: metric for metric in METRICS}
|
||||
@@ -0,0 +1,165 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Pure data transforms for the local performance dashboard.
|
||||
|
||||
The functions in this module operate on normalized records from
|
||||
``fastvideo/tests/performance/compare_baseline.py``. They intentionally avoid
|
||||
network and FastAPI concerns so they can be tested with in-memory fixtures.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import statistics
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.tests.performance.hf_store import safe_float
|
||||
|
||||
from .metrics import METRICS
|
||||
|
||||
Record = dict[str, Any]
|
||||
|
||||
|
||||
def parse_timestamp(value: Any) -> datetime | None:
|
||||
if not value:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
ts = value
|
||||
else:
|
||||
try:
|
||||
ts = datetime.fromisoformat(str(value))
|
||||
except ValueError:
|
||||
return None
|
||||
if ts.tzinfo is None:
|
||||
return ts.replace(tzinfo=timezone.utc)
|
||||
return ts.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def record_sort_key(record: Record) -> tuple[datetime, str]:
|
||||
ts = parse_timestamp(record.get("timestamp"))
|
||||
return (ts or datetime.min.replace(tzinfo=timezone.utc), str(record.get("commit_sha") or ""))
|
||||
|
||||
|
||||
def filter_records(
|
||||
records: list[Record],
|
||||
*,
|
||||
model_id: str | None = None,
|
||||
gpu_type: str | None = None,
|
||||
success: bool | None = None,
|
||||
) -> list[Record]:
|
||||
filtered = records
|
||||
if model_id:
|
||||
filtered = [record for record in filtered if record.get("model_id") == model_id]
|
||||
if gpu_type:
|
||||
filtered = [record for record in filtered if record.get("gpu_type") == gpu_type]
|
||||
if success is not None:
|
||||
filtered = [record for record in filtered if bool(record.get("success", True)) == success]
|
||||
return sorted(filtered, key=record_sort_key)
|
||||
|
||||
|
||||
def group_by_model_gpu(records: list[Record]) -> dict[tuple[str, str], list[Record]]:
|
||||
groups: dict[tuple[str, str], list[Record]] = defaultdict(list)
|
||||
for record in records:
|
||||
model_id = str(record.get("model_id") or "unknown")
|
||||
gpu_type = str(record.get("gpu_type") or "unknown")
|
||||
groups[(model_id, gpu_type)].append(record)
|
||||
return {key: sorted(value, key=record_sort_key) for key, value in groups.items()}
|
||||
|
||||
|
||||
def baseline_value(records: list[Record], metric_key: str) -> float | None:
|
||||
values = [safe_float(record.get(metric_key)) for record in records]
|
||||
values = [value for value in values if value is not None]
|
||||
if not values:
|
||||
return None
|
||||
return float(statistics.median(values))
|
||||
|
||||
|
||||
def regression_percent(metric_key: str, current: float | None, baseline: float | None) -> float | None:
|
||||
if current is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
metric = next(metric for metric in METRICS if metric.key == metric_key)
|
||||
if metric.lower_is_better:
|
||||
return (current - baseline) / baseline * 100.0
|
||||
return (baseline - current) / baseline * 100.0
|
||||
|
||||
|
||||
def build_latest_summary(records: list[Record],
|
||||
*,
|
||||
baseline_window: int = 5,
|
||||
max_regression: float = 0.05) -> list[Record]:
|
||||
rows: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
latest = group[-1]
|
||||
earlier_successes = [record for record in group[:-1] if record.get("success", True)]
|
||||
baseline_records = earlier_successes[-baseline_window:]
|
||||
|
||||
metrics: dict[str, Record] = {}
|
||||
regressions: list[float] = []
|
||||
for metric in METRICS:
|
||||
current = safe_float(latest.get(metric.key))
|
||||
baseline = baseline_value(baseline_records, metric.key)
|
||||
regression = regression_percent(metric.key, current, baseline)
|
||||
metrics[metric.key] = {
|
||||
"current": current,
|
||||
"baseline": baseline,
|
||||
"regression_pct": regression,
|
||||
"label": metric.label,
|
||||
"lower_is_better": metric.lower_is_better,
|
||||
"precision": metric.precision,
|
||||
}
|
||||
if regression is not None:
|
||||
regressions.append(regression)
|
||||
|
||||
worst_regression = max(regressions) if regressions else None
|
||||
success = bool(latest.get("success", True))
|
||||
status = "pass" if success else "fail"
|
||||
|
||||
rows.append({
|
||||
"model_id":
|
||||
model_id,
|
||||
"gpu_type":
|
||||
gpu_type,
|
||||
"timestamp":
|
||||
latest.get("timestamp"),
|
||||
"commit_sha":
|
||||
latest.get("commit_sha"),
|
||||
"success":
|
||||
success,
|
||||
"baseline_n":
|
||||
len(baseline_records),
|
||||
"worst_regression_pct":
|
||||
worst_regression,
|
||||
"regression_threshold_pct":
|
||||
max_regression * 100.0,
|
||||
"computed_regression_status":
|
||||
"fail" if worst_regression is not None and worst_regression > max_regression * 100.0 else "pass",
|
||||
"status":
|
||||
status,
|
||||
"metrics":
|
||||
metrics,
|
||||
})
|
||||
|
||||
return sorted(rows, key=lambda row: (row["status"] != "fail", row["model_id"], row["gpu_type"]))
|
||||
|
||||
|
||||
def build_trends(records: list[Record]) -> list[Record]:
|
||||
trends: list[Record] = []
|
||||
for (model_id, gpu_type), group in group_by_model_gpu(records).items():
|
||||
points = []
|
||||
for record in group:
|
||||
point = {
|
||||
"timestamp": record.get("timestamp"),
|
||||
"commit_sha": record.get("commit_sha"),
|
||||
"success": bool(record.get("success", True)),
|
||||
"metrics": {
|
||||
metric.key: safe_float(record.get(metric.key))
|
||||
for metric in METRICS
|
||||
},
|
||||
}
|
||||
points.append(point)
|
||||
trends.append({
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"points": points,
|
||||
})
|
||||
return sorted(trends, key=lambda trend: (trend["model_id"], trend["gpu_type"]))
|
||||
@@ -140,6 +140,23 @@ class CudaPlatformBase(Platform):
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info("Sage Attention 3 backend is not installed. Fall back to Flash Attention.")
|
||||
elif selected_backend == AttentionBackendEnum.ATTN_QAT_INFER:
|
||||
from fastvideo.attention.backends.attn_qat_infer import ( # noqa: F401
|
||||
AttnQatInferBackend, is_attn_qat_infer_available)
|
||||
if is_attn_qat_infer_available():
|
||||
logger.info("Using Attn-QAT inference (modified SageAttention3 FP4) backend.")
|
||||
return "fastvideo.attention.backends.attn_qat_infer.AttnQatInferBackend"
|
||||
logger.info("Attn-QAT inference kernel is not built. Fall back to Flash Attention.")
|
||||
elif selected_backend == AttentionBackendEnum.ATTN_QAT_TRAIN:
|
||||
from fastvideo.attention.backends.attn_qat_train import ( # noqa: F401
|
||||
AttnQatTrainBackend, is_attn_qat_train_available)
|
||||
if is_attn_qat_train_available():
|
||||
logger.info("Using Attn-QAT training (fake-quantized attention) backend.")
|
||||
return "fastvideo.attention.backends.attn_qat_train.AttnQatTrainBackend"
|
||||
raise ImportError(
|
||||
"ATTN_QAT_TRAIN selected but fastvideo_kernel.triton_kernels.attn_qat_train is not built. "
|
||||
"Silent fallback would produce a non-QAT training run; refusing to proceed. "
|
||||
"Install the training kernel or pick a different FASTVIDEO_ATTENTION_BACKEND.")
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from fastvideo_kernel import video_sparse_attn # noqa: F401
|
||||
|
||||
@@ -15,6 +15,8 @@ class AttentionBackendEnum(enum.Enum):
|
||||
TORCH_SDPA = enum.auto()
|
||||
SAGE_ATTN = enum.auto()
|
||||
SAGE_ATTN_THREE = enum.auto()
|
||||
ATTN_QAT_INFER = enum.auto()
|
||||
ATTN_QAT_TRAIN = enum.auto()
|
||||
VIDEO_SPARSE_ATTN = enum.auto()
|
||||
BSA_ATTN = enum.auto()
|
||||
VMOBA_ATTN = enum.auto()
|
||||
|
||||
@@ -24,7 +24,7 @@ from huggingface_hub import HfApi, snapshot_download
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
HF_REPO_ID: str = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
|
||||
HF_TOKEN: str | None = os.environ.get("HF_API_KEY")
|
||||
HF_TOKEN_ENV_VARS = ("HF_API_KEY", "HUGGINGFACE_HUB_TOKEN", "HF_TOKEN")
|
||||
SYNC_MARKER = ".hf_sync_complete"
|
||||
SYNC_REUSE_TTL_SECONDS = int(os.environ.get("PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS", "3600"))
|
||||
|
||||
@@ -48,6 +48,15 @@ def safe_float(value: Any) -> float | None:
|
||||
return None
|
||||
|
||||
|
||||
def resolve_hf_token() -> str | None:
|
||||
"""Return the first configured Hugging Face token env var."""
|
||||
for env_var in HF_TOKEN_ENV_VARS:
|
||||
token = os.environ.get(env_var)
|
||||
if token:
|
||||
return token
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# HF I/O
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -120,7 +129,7 @@ def sync_from_hf(
|
||||
repo_id=HF_REPO_ID,
|
||||
repo_type="dataset",
|
||||
local_dir=local_dir,
|
||||
token=HF_TOKEN,
|
||||
token=resolve_hf_token(),
|
||||
allow_patterns="*.json",
|
||||
)
|
||||
os.makedirs(local_dir, exist_ok=True)
|
||||
@@ -150,8 +159,9 @@ def upload_record(
|
||||
that must not silently lose records — otherwise the rolling baseline can
|
||||
stop advancing without any signal in the build log.
|
||||
"""
|
||||
if not HF_TOKEN:
|
||||
msg = "hf_store: HF_API_KEY not set"
|
||||
token = resolve_hf_token()
|
||||
if not token:
|
||||
msg = f"hf_store: none of {', '.join(HF_TOKEN_ENV_VARS)} set"
|
||||
if strict:
|
||||
raise RuntimeError(f"{msg}; cannot upload.")
|
||||
print(f"{msg}, skipping upload.")
|
||||
@@ -161,7 +171,7 @@ def upload_record(
|
||||
path_in_repo = f"{sanitize(model_id)}/{os.path.basename(local_path)}"
|
||||
commit_sha = (record.get("commit_sha") or "unknown")[:7]
|
||||
|
||||
api = HfApi(token=HF_TOKEN)
|
||||
api = HfApi(token=token)
|
||||
try:
|
||||
api.upload_file(
|
||||
path_or_fileobj=local_path,
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastapi.testclient import TestClient
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from fastvideo.performance_dashboard.api import PerformanceDataStore, create_app
|
||||
|
||||
|
||||
class FakeStore(PerformanceDataStore):
|
||||
def __init__(self, records):
|
||||
super().__init__(tracking_root="/tmp/fake-fastvideo-perf-dashboard")
|
||||
self._records = records
|
||||
self.last_sync_at = "2026-01-03T00:00:00+00:00"
|
||||
|
||||
@property
|
||||
def repo_id(self):
|
||||
return "FastVideo/performance-tracking"
|
||||
|
||||
def sync(self):
|
||||
return {
|
||||
"ok": True,
|
||||
"repo_id": self.repo_id,
|
||||
"tracking_root": self.tracking_root,
|
||||
"last_sync_at": self.last_sync_at,
|
||||
"last_sync_error": None,
|
||||
}
|
||||
|
||||
def load_records(self, *, days=None, successful_only=False):
|
||||
records = list(self._records)
|
||||
if days is not None:
|
||||
latest_ts = max(datetime.fromisoformat(record["timestamp"]) for record in records) if records else None
|
||||
if latest_ts is not None:
|
||||
cutoff = latest_ts - timedelta(days=days)
|
||||
records = [record for record in records if datetime.fromisoformat(record["timestamp"]) >= cutoff]
|
||||
if successful_only:
|
||||
return [record for record in records if record.get("success", True)]
|
||||
return records
|
||||
|
||||
|
||||
def _record(model_id, gpu_type, ts, commit, latency, throughput, success=True):
|
||||
return {
|
||||
"model_id": model_id,
|
||||
"gpu_type": gpu_type,
|
||||
"timestamp": ts,
|
||||
"commit_sha": commit,
|
||||
"latency": latency,
|
||||
"throughput": throughput,
|
||||
"memory": 10000.0,
|
||||
"text_encoder_time_s": 2.0,
|
||||
"dit_time_s": 8.0,
|
||||
"vae_decode_time_s": 3.0,
|
||||
"success": success,
|
||||
}
|
||||
|
||||
|
||||
def test_summary_endpoint_returns_latest_group_status():
|
||||
app = create_app(FakeStore([
|
||||
_record("wan", "NVIDIA L40S", "2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("wan", "NVIDIA L40S", "2026-01-02T00:00:00+00:00", "b" * 40, 11.0, 9.0),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.get("/api/performance/summary")
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["count"] == 1
|
||||
assert body["status_counts"] == {"pass": 1, "fail": 0}
|
||||
assert body["rows"][0]["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert body["rows"][0]["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_summary_status_is_independent_of_days_window():
|
||||
app = create_app(FakeStore([
|
||||
_record("wan", "NVIDIA L40S", "2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("wan", "NVIDIA L40S", "2026-02-15T00:00:00+00:00", "b" * 40, 11.0, 9.0),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
narrow = client.get("/api/performance/summary", params={"days": 1}).json()
|
||||
wide = client.get("/api/performance/summary", params={"days": 365}).json()
|
||||
trends = client.get("/api/performance/trends", params={"days": 1}).json()
|
||||
|
||||
assert narrow["rows"][0]["status"] == wide["rows"][0]["status"]
|
||||
assert narrow["rows"][0]["baseline_n"] == wide["rows"][0]["baseline_n"] == 1
|
||||
assert narrow["filters"]["days"] is None
|
||||
assert narrow["filters"]["trend_window_days"] == 1
|
||||
assert len(trends["groups"][0]["points"]) == 1
|
||||
|
||||
|
||||
def test_records_and_trends_endpoints_filter_by_model_and_gpu():
|
||||
app = create_app(FakeStore([
|
||||
_record("wan", "NVIDIA L40S", "2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("ltx", "NVIDIA A100", "2026-01-01T00:00:00+00:00", "b" * 40, 20.0, 5.0),
|
||||
]))
|
||||
client = TestClient(app)
|
||||
|
||||
records = client.get("/api/performance/records", params={"model_id": "wan"}).json()
|
||||
trends = client.get("/api/performance/trends", params={"gpu_type": "NVIDIA L40S"}).json()
|
||||
|
||||
assert records["count"] == 1
|
||||
assert records["records"][0]["model_id"] == "wan"
|
||||
assert trends["count"] == 1
|
||||
assert trends["groups"][0]["gpu_type"] == "NVIDIA L40S"
|
||||
|
||||
|
||||
def test_refresh_endpoint_reports_sync_metadata():
|
||||
app = create_app(FakeStore([]))
|
||||
client = TestClient(app)
|
||||
|
||||
response = client.post("/api/performance/refresh")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["ok"] is True
|
||||
@@ -0,0 +1,73 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from fastvideo.performance_dashboard.service import build_latest_summary, build_trends, filter_records
|
||||
from fastvideo.tests.performance import hf_store
|
||||
|
||||
|
||||
def _record(ts, commit, latency, throughput, success=True):
|
||||
return {
|
||||
"model_id": "wan-t2v-1.3b-2gpu",
|
||||
"gpu_type": "NVIDIA L40S",
|
||||
"timestamp": ts,
|
||||
"commit_sha": commit,
|
||||
"latency": latency,
|
||||
"throughput": throughput,
|
||||
"memory": 10000.0,
|
||||
"text_encoder_time_s": None,
|
||||
"dit_time_s": 8.0,
|
||||
"vae_decode_time_s": 3.0,
|
||||
"success": success,
|
||||
}
|
||||
|
||||
|
||||
def test_build_latest_summary_uses_previous_successful_records_for_baseline():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 12.0, 8.0, success=False),
|
||||
_record("2026-01-03T00:00:00+00:00", "c" * 40, 11.0, 9.0),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records, max_regression=0.05)
|
||||
|
||||
assert len(rows) == 1
|
||||
row = rows[0]
|
||||
assert row["baseline_n"] == 1
|
||||
assert row["metrics"]["latency"]["baseline"] == 10.0
|
||||
assert row["metrics"]["latency"]["regression_pct"] == 10.0
|
||||
assert row["metrics"]["throughput"]["regression_pct"] == 10.0
|
||||
assert row["status"] == "pass"
|
||||
assert row["computed_regression_status"] == "fail"
|
||||
|
||||
|
||||
def test_build_latest_summary_status_uses_latest_record_success_field():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 10.0, 10.0, success=False),
|
||||
]
|
||||
|
||||
rows = build_latest_summary(records)
|
||||
|
||||
assert rows[0]["status"] == "fail"
|
||||
assert rows[0]["success"] is False
|
||||
|
||||
|
||||
def test_filter_records_and_trends_preserve_metric_points():
|
||||
records = [
|
||||
_record("2026-01-01T00:00:00+00:00", "a" * 40, 10.0, 10.0),
|
||||
_record("2026-01-02T00:00:00+00:00", "b" * 40, 12.0, 8.0, success=False),
|
||||
]
|
||||
|
||||
failed = filter_records(records, success=False)
|
||||
trends = build_trends(records)
|
||||
|
||||
assert [record["commit_sha"] for record in failed] == ["b" * 40]
|
||||
assert len(trends) == 1
|
||||
assert trends[0]["points"][1]["metrics"]["latency"] == 12.0
|
||||
|
||||
|
||||
def test_hf_token_resolution_accepts_standard_env_names(monkeypatch):
|
||||
for env_var in hf_store.HF_TOKEN_ENV_VARS:
|
||||
monkeypatch.delenv(env_var, raising=False)
|
||||
|
||||
monkeypatch.setenv("HF_TOKEN", "hf_local")
|
||||
|
||||
assert hf_store.resolve_hf_token() == "hf_local"
|
||||
@@ -0,0 +1,80 @@
|
||||
# FastVideo Performance Dashboard
|
||||
|
||||
Local FastAPI + React dashboard for records stored in the Hugging Face
|
||||
performance tracking dataset.
|
||||
|
||||
## Data Source
|
||||
|
||||
The dashboard reads the same normalized JSON records used by
|
||||
`fastvideo/tests/performance/compare_baseline.py`.
|
||||
|
||||
Defaults:
|
||||
|
||||
- `HF_REPO_ID=FastVideo/performance-tracking`
|
||||
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
|
||||
- `PERF_MAX_REGRESSION=0.05`
|
||||
|
||||
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
|
||||
configured dataset repo requires authenticated access:
|
||||
|
||||
```bash
|
||||
export HF_TOKEN=hf_...
|
||||
```
|
||||
|
||||
If Hugging Face returns `401 Unauthorized`, confirm that `HF_REPO_ID` points to
|
||||
the dataset repo you expect and that your token has access to it.
|
||||
|
||||
## Development
|
||||
|
||||
Run the API:
|
||||
|
||||
```bash
|
||||
python -m fastvideo.performance_dashboard --host 127.0.0.1 --port 8000 --reload
|
||||
```
|
||||
|
||||
Run the React dev server:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
npm install
|
||||
npm run dev
|
||||
```
|
||||
|
||||
Open `http://127.0.0.1:5173`. Vite proxies `/api/*` to the FastAPI server on
|
||||
port 8000.
|
||||
|
||||
## Single-Port Mode For ngrok
|
||||
|
||||
Build the frontend:
|
||||
|
||||
```bash
|
||||
cd performance_dashboard/frontend
|
||||
npm install
|
||||
npm run build
|
||||
```
|
||||
|
||||
Serve API and built frontend from one FastAPI process:
|
||||
|
||||
```bash
|
||||
python -m fastvideo.performance_dashboard --host 0.0.0.0 --port 8000
|
||||
```
|
||||
|
||||
Expose it:
|
||||
|
||||
```bash
|
||||
ngrok http 8000
|
||||
```
|
||||
|
||||
The ngrok URL will serve the dashboard UI and all `/api/performance/*`
|
||||
endpoints from the same local port.
|
||||
|
||||
## API
|
||||
|
||||
- `GET /api/performance/health`
|
||||
- `POST /api/performance/refresh`
|
||||
- `GET /api/performance/summary?days=90`
|
||||
- `GET /api/performance/trends?days=90`
|
||||
- `GET /api/performance/records?days=90`
|
||||
|
||||
The current v1 grouping key is `(model_id, gpu_type)`. Baselines are computed
|
||||
from the latest five previous successful records in each group.
|
||||
@@ -0,0 +1,3 @@
|
||||
node_modules/
|
||||
dist/
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>FastVideo Performance Dashboard</title>
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
<script type="module" src="/src/main.tsx"></script>
|
||||
</body>
|
||||
</html>
|
||||
|
||||
+1824
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,23 @@
|
||||
{
|
||||
"name": "fastvideo-performance-dashboard",
|
||||
"version": "0.1.0",
|
||||
"private": true,
|
||||
"type": "module",
|
||||
"scripts": {
|
||||
"dev": "vite dev --host 0.0.0.0",
|
||||
"build": "tsc && node scripts/build.mjs",
|
||||
"preview": "vite preview --host 0.0.0.0"
|
||||
},
|
||||
"dependencies": {
|
||||
"react": "^19.2.3",
|
||||
"react-dom": "^19.2.3"
|
||||
},
|
||||
"devDependencies": {
|
||||
"@types/react": "^19.2.7",
|
||||
"@types/react-dom": "^19.2.3",
|
||||
"@vitejs/plugin-react": "^5.1.1",
|
||||
"esbuild": "^0.25.12",
|
||||
"typescript": "^5.9.3",
|
||||
"vite": "^6.4.2"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
import * as esbuild from "esbuild";
|
||||
import { mkdir, writeFile } from "node:fs/promises";
|
||||
import { resolve } from "node:path";
|
||||
|
||||
const root = resolve(import.meta.dirname, "..");
|
||||
const dist = resolve(root, "dist");
|
||||
const assets = resolve(dist, "assets");
|
||||
|
||||
await mkdir(assets, { recursive: true });
|
||||
|
||||
await esbuild.build({
|
||||
entryPoints: [resolve(root, "src/main.tsx")],
|
||||
bundle: true,
|
||||
format: "esm",
|
||||
minify: true,
|
||||
sourcemap: true,
|
||||
target: ["es2020"],
|
||||
outfile: resolve(assets, "dashboard.js"),
|
||||
loader: {
|
||||
".svg": "file"
|
||||
}
|
||||
});
|
||||
|
||||
await writeFile(
|
||||
resolve(dist, "index.html"),
|
||||
`<!doctype html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="UTF-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
|
||||
<title>FastVideo Performance Dashboard</title>
|
||||
<link rel="stylesheet" href="/assets/dashboard.css" />
|
||||
</head>
|
||||
<body>
|
||||
<div id="root"></div>
|
||||
<script type="module" src="/assets/dashboard.js"></script>
|
||||
</body>
|
||||
</html>
|
||||
`
|
||||
);
|
||||
@@ -0,0 +1,291 @@
|
||||
import { useEffect, useMemo, useState } from "react";
|
||||
|
||||
import { fetchSummary, fetchTrends, refreshData, SummaryResponse, TrendGroup } from "./api";
|
||||
|
||||
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
|
||||
|
||||
function formatNumber(value: number | null | undefined, precision = 2) {
|
||||
if (value === null || value === undefined || Number.isNaN(value)) {
|
||||
return "n/a";
|
||||
}
|
||||
return value.toFixed(precision);
|
||||
}
|
||||
|
||||
function shortSha(value: string | null | undefined) {
|
||||
return value ? value.slice(0, 7) : "unknown";
|
||||
}
|
||||
|
||||
function formatTime(value: string | null | undefined) {
|
||||
if (!value) {
|
||||
return "never";
|
||||
}
|
||||
const date = new Date(value);
|
||||
if (Number.isNaN(date.getTime())) {
|
||||
return value;
|
||||
}
|
||||
return date.toLocaleString();
|
||||
}
|
||||
|
||||
function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string }) {
|
||||
const points = group.points
|
||||
.map((point, index) => ({
|
||||
index,
|
||||
value: point.metrics[metricKey],
|
||||
success: point.success
|
||||
}))
|
||||
.filter((point) => point.value !== null && point.value !== undefined) as Array<{
|
||||
index: number;
|
||||
value: number;
|
||||
success: boolean;
|
||||
}>;
|
||||
|
||||
if (points.length === 0) {
|
||||
return <div className="empty-chart">No data</div>;
|
||||
}
|
||||
|
||||
const width = 280;
|
||||
const height = 96;
|
||||
const pad = 12;
|
||||
const min = Math.min(...points.map((point) => point.value));
|
||||
const max = Math.max(...points.map((point) => point.value));
|
||||
const span = max - min || 1;
|
||||
const maxIndex = Math.max(...points.map((point) => point.index)) || 1;
|
||||
const xy = (point: { index: number; value: number }) => {
|
||||
const x = pad + (point.index / maxIndex) * (width - pad * 2);
|
||||
const y = height - pad - ((point.value - min) / span) * (height - pad * 2);
|
||||
return `${x},${y}`;
|
||||
};
|
||||
|
||||
return (
|
||||
<svg className="trend-chart" viewBox={`0 0 ${width} ${height}`} role="img">
|
||||
<polyline points={points.map(xy).join(" ")} fill="none" stroke="currentColor" strokeWidth="2.2" />
|
||||
{points.map((point) => {
|
||||
const [cx, cy] = xy(point).split(",");
|
||||
return (
|
||||
<circle
|
||||
key={`${point.index}-${point.value}`}
|
||||
cx={cx}
|
||||
cy={cy}
|
||||
r="3"
|
||||
className={point.success ? "point-pass" : "point-fail"}
|
||||
/>
|
||||
);
|
||||
})}
|
||||
</svg>
|
||||
);
|
||||
}
|
||||
|
||||
export default function App() {
|
||||
const [days, setDays] = useState(90);
|
||||
const [modelFilter, setModelFilter] = useState("");
|
||||
const [gpuFilter, setGpuFilter] = useState("");
|
||||
const [summary, setSummary] = useState<SummaryResponse | null>(null);
|
||||
const [trends, setTrends] = useState<TrendGroup[]>([]);
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [refreshing, setRefreshing] = useState(false);
|
||||
const [error, setError] = useState<string | null>(null);
|
||||
|
||||
async function load() {
|
||||
setLoading(true);
|
||||
setError(null);
|
||||
try {
|
||||
const [summaryData, trendData] = await Promise.all([
|
||||
fetchSummary(days, modelFilter || undefined, gpuFilter || undefined),
|
||||
fetchTrends(days, modelFilter || undefined, gpuFilter || undefined)
|
||||
]);
|
||||
setSummary(summaryData);
|
||||
setTrends(trendData.groups);
|
||||
} catch (err) {
|
||||
setError(err instanceof Error ? err.message : String(err));
|
||||
} finally {
|
||||
setLoading(false);
|
||||
}
|
||||
}
|
||||
|
||||
async function refresh() {
|
||||
setRefreshing(true);
|
||||
setError(null);
|
||||
try {
|
||||
await refreshData();
|
||||
await load();
|
||||
} catch (err) {
|
||||
setError(err instanceof Error ? err.message : String(err));
|
||||
} finally {
|
||||
setRefreshing(false);
|
||||
}
|
||||
}
|
||||
|
||||
useEffect(() => {
|
||||
load();
|
||||
const interval = window.setInterval(load, 5 * 60 * 1000);
|
||||
return () => window.clearInterval(interval);
|
||||
}, [days, modelFilter, gpuFilter]);
|
||||
|
||||
const models = useMemo(() => {
|
||||
const values = new Set(summary?.rows.map((row) => row.model_id) ?? []);
|
||||
trends.forEach((trend) => values.add(trend.model_id));
|
||||
return [...values].sort();
|
||||
}, [summary, trends]);
|
||||
|
||||
const gpus = useMemo(() => {
|
||||
const values = new Set(summary?.rows.map((row) => row.gpu_type) ?? []);
|
||||
trends.forEach((trend) => values.add(trend.gpu_type));
|
||||
return [...values].sort();
|
||||
}, [summary, trends]);
|
||||
|
||||
const latestRows = summary?.rows ?? [];
|
||||
const totalRuns = trends.reduce((total, group) => total + group.points.length, 0);
|
||||
const sync = summary?.sync;
|
||||
|
||||
return (
|
||||
<main className="dashboard">
|
||||
<header className="topbar">
|
||||
<div>
|
||||
<p className="eyebrow">FastVideo CI</p>
|
||||
<h1>Performance Dashboard</h1>
|
||||
</div>
|
||||
<button className="refresh-button" onClick={refresh} disabled={refreshing || loading}>
|
||||
{refreshing ? "Refreshing" : "Refresh"}
|
||||
</button>
|
||||
</header>
|
||||
|
||||
<section className="filters" aria-label="Filters">
|
||||
<label>
|
||||
Days
|
||||
<input
|
||||
type="number"
|
||||
min="1"
|
||||
max="3650"
|
||||
value={days}
|
||||
onChange={(event) => setDays(Number(event.target.value) || 90)}
|
||||
/>
|
||||
</label>
|
||||
<label>
|
||||
Model
|
||||
<select value={modelFilter} onChange={(event) => setModelFilter(event.target.value)}>
|
||||
<option value="">All models</option>
|
||||
{models.map((model) => (
|
||||
<option key={model} value={model}>
|
||||
{model}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</label>
|
||||
<label>
|
||||
GPU
|
||||
<select value={gpuFilter} onChange={(event) => setGpuFilter(event.target.value)}>
|
||||
<option value="">All GPUs</option>
|
||||
{gpus.map((gpu) => (
|
||||
<option key={gpu} value={gpu}>
|
||||
{gpu}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</label>
|
||||
</section>
|
||||
|
||||
{error && <div className="notice error">Failed to load dashboard data: {error}</div>}
|
||||
{loading && <div className="notice">Loading performance data</div>}
|
||||
|
||||
<section className="cards" aria-label="Overview">
|
||||
<div className="stat">
|
||||
<span>Groups</span>
|
||||
<strong>{summary?.count ?? 0}</strong>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<span>Failing</span>
|
||||
<strong>{summary?.status_counts.fail ?? 0}</strong>
|
||||
</div>
|
||||
<div className="stat">
|
||||
<span>Runs</span>
|
||||
<strong>{totalRuns}</strong>
|
||||
</div>
|
||||
<div className="stat wide">
|
||||
<span>Last sync</span>
|
||||
<strong>{formatTime(sync?.last_sync_at)}</strong>
|
||||
<small>{sync?.repo_id ?? "FastVideo/performance-tracking"}</small>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Latest Status</h2>
|
||||
<span>{latestRows.length} model/GPU groups</span>
|
||||
</div>
|
||||
{latestRows.length === 0 ? (
|
||||
<div className="empty">No records match the selected filters.</div>
|
||||
) : (
|
||||
<div className="table-wrap">
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Stored Status</th>
|
||||
<th>Recomputed</th>
|
||||
<th>Model</th>
|
||||
<th>GPU</th>
|
||||
<th>Commit</th>
|
||||
<th>Baseline N</th>
|
||||
<th>Latency</th>
|
||||
<th>Throughput</th>
|
||||
<th>Memory</th>
|
||||
<th>Worst</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{latestRows.map((row) => (
|
||||
<tr key={`${row.model_id}-${row.gpu_type}`}>
|
||||
<td>
|
||||
<span className={`badge ${row.status}`}>{row.status}</span>
|
||||
</td>
|
||||
<td>
|
||||
<span className={`badge muted ${row.computed_regression_status}`}>
|
||||
{row.computed_regression_status}
|
||||
</span>
|
||||
</td>
|
||||
<td>{row.model_id}</td>
|
||||
<td>{row.gpu_type}</td>
|
||||
<td>{shortSha(row.commit_sha)}</td>
|
||||
<td>{row.baseline_n}</td>
|
||||
<td>{formatNumber(row.metrics.latency?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
|
||||
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
|
||||
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
|
||||
</tr>
|
||||
))}
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
)}
|
||||
</section>
|
||||
|
||||
<section className="panel">
|
||||
<div className="panel-header">
|
||||
<h2>Trends</h2>
|
||||
<span>{days} day window</span>
|
||||
</div>
|
||||
<div className="trend-grid">
|
||||
{trends.length === 0 ? (
|
||||
<div className="empty full-width">
|
||||
No trend records found in the selected time window. Increase the day range or refresh after new CI
|
||||
performance records are uploaded.
|
||||
</div>
|
||||
) : (
|
||||
trends.map((group) =>
|
||||
METRIC_KEYS.map((metricKey) => (
|
||||
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
|
||||
<div>
|
||||
<h3>{summary?.rows[0]?.metrics[metricKey]?.label ?? metricKey}</h3>
|
||||
<p>
|
||||
{group.model_id} | {group.gpu_type}
|
||||
</p>
|
||||
</div>
|
||||
<TrendChart group={group} metricKey={metricKey} />
|
||||
</article>
|
||||
))
|
||||
)
|
||||
)}
|
||||
</div>
|
||||
</section>
|
||||
</main>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
export type MetricValue = {
|
||||
current: number | null;
|
||||
baseline: number | null;
|
||||
regression_pct: number | null;
|
||||
label: string;
|
||||
lower_is_better: boolean;
|
||||
precision: number;
|
||||
};
|
||||
|
||||
export type SummaryRow = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
timestamp: string | null;
|
||||
commit_sha: string | null;
|
||||
success: boolean;
|
||||
baseline_n: number;
|
||||
worst_regression_pct: number | null;
|
||||
regression_threshold_pct: number;
|
||||
computed_regression_status: "pass" | "fail";
|
||||
status: "pass" | "fail";
|
||||
metrics: Record<string, MetricValue>;
|
||||
};
|
||||
|
||||
export type SummaryResponse = {
|
||||
rows: SummaryRow[];
|
||||
count: number;
|
||||
status_counts: {
|
||||
pass: number;
|
||||
fail: number;
|
||||
};
|
||||
filters: {
|
||||
days: number;
|
||||
model_id: string | null;
|
||||
gpu_type: string | null;
|
||||
};
|
||||
sync: SyncState;
|
||||
};
|
||||
|
||||
export type TrendPoint = {
|
||||
timestamp: string | null;
|
||||
commit_sha: string | null;
|
||||
success: boolean;
|
||||
metrics: Record<string, number | null>;
|
||||
};
|
||||
|
||||
export type TrendGroup = {
|
||||
model_id: string;
|
||||
gpu_type: string;
|
||||
points: TrendPoint[];
|
||||
};
|
||||
|
||||
export type TrendsResponse = {
|
||||
groups: TrendGroup[];
|
||||
count: number;
|
||||
sync: SyncState;
|
||||
};
|
||||
|
||||
export type SyncState = {
|
||||
ok: boolean;
|
||||
repo_id: string;
|
||||
tracking_root: string;
|
||||
last_sync_at: string | null;
|
||||
last_sync_error: string | null;
|
||||
};
|
||||
|
||||
const jsonHeaders = {
|
||||
Accept: "application/json"
|
||||
};
|
||||
|
||||
function params(values: Record<string, string | number | null | undefined>) {
|
||||
const out = new URLSearchParams();
|
||||
for (const [key, value] of Object.entries(values)) {
|
||||
if (value !== null && value !== undefined && value !== "") {
|
||||
out.set(key, String(value));
|
||||
}
|
||||
}
|
||||
return out.toString();
|
||||
}
|
||||
|
||||
async function getJson<T>(path: string): Promise<T> {
|
||||
const response = await fetch(path, { headers: jsonHeaders });
|
||||
if (!response.ok) {
|
||||
throw new Error(`${response.status} ${response.statusText}`);
|
||||
}
|
||||
return response.json() as Promise<T>;
|
||||
}
|
||||
|
||||
export async function fetchSummary(days = 90, modelId?: string, gpuType?: string) {
|
||||
return getJson<SummaryResponse>(
|
||||
`/api/performance/summary?${params({ days, model_id: modelId, gpu_type: gpuType })}`
|
||||
);
|
||||
}
|
||||
|
||||
export async function fetchTrends(days = 90, modelId?: string, gpuType?: string) {
|
||||
return getJson<TrendsResponse>(
|
||||
`/api/performance/trends?${params({ days, model_id: modelId, gpu_type: gpuType })}`
|
||||
);
|
||||
}
|
||||
|
||||
export async function refreshData() {
|
||||
const response = await fetch("/api/performance/refresh", { method: "POST", headers: jsonHeaders });
|
||||
if (!response.ok) {
|
||||
throw new Error(`${response.status} ${response.statusText}`);
|
||||
}
|
||||
return response.json() as Promise<SyncState>;
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
import React from "react";
|
||||
import { createRoot } from "react-dom/client";
|
||||
|
||||
import App from "./App";
|
||||
import "./styles.css";
|
||||
|
||||
createRoot(document.getElementById("root") as HTMLElement).render(
|
||||
<React.StrictMode>
|
||||
<App />
|
||||
</React.StrictMode>
|
||||
);
|
||||
|
||||
@@ -0,0 +1,302 @@
|
||||
:root {
|
||||
color: #1f2933;
|
||||
background: #eef2f5;
|
||||
font-family:
|
||||
Inter, ui-sans-serif, system-ui, -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
|
||||
line-height: 1.4;
|
||||
}
|
||||
|
||||
* {
|
||||
box-sizing: border-box;
|
||||
}
|
||||
|
||||
body {
|
||||
margin: 0;
|
||||
min-width: 320px;
|
||||
}
|
||||
|
||||
button,
|
||||
input,
|
||||
select {
|
||||
font: inherit;
|
||||
}
|
||||
|
||||
.dashboard {
|
||||
width: min(1440px, 100%);
|
||||
margin: 0 auto;
|
||||
padding: 28px;
|
||||
}
|
||||
|
||||
.topbar {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
gap: 20px;
|
||||
margin-bottom: 22px;
|
||||
}
|
||||
|
||||
.eyebrow {
|
||||
margin: 0 0 4px;
|
||||
color: #607080;
|
||||
font-size: 0.78rem;
|
||||
font-weight: 700;
|
||||
letter-spacing: 0;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
|
||||
h1,
|
||||
h2,
|
||||
h3,
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
|
||||
h1 {
|
||||
color: #12202f;
|
||||
font-size: clamp(2rem, 4vw, 3.5rem);
|
||||
letter-spacing: 0;
|
||||
}
|
||||
|
||||
h2 {
|
||||
color: #182736;
|
||||
font-size: 1.08rem;
|
||||
}
|
||||
|
||||
h3 {
|
||||
color: #233242;
|
||||
font-size: 0.94rem;
|
||||
}
|
||||
|
||||
.refresh-button {
|
||||
min-height: 42px;
|
||||
border: 1px solid #0f6b8f;
|
||||
border-radius: 6px;
|
||||
padding: 0 18px;
|
||||
color: #ffffff;
|
||||
background: #0f6b8f;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.refresh-button:disabled {
|
||||
cursor: default;
|
||||
opacity: 0.6;
|
||||
}
|
||||
|
||||
.filters {
|
||||
display: grid;
|
||||
grid-template-columns: 120px minmax(220px, 1fr) minmax(220px, 1fr);
|
||||
gap: 14px;
|
||||
margin-bottom: 18px;
|
||||
}
|
||||
|
||||
.filters label {
|
||||
display: grid;
|
||||
gap: 6px;
|
||||
color: #425466;
|
||||
font-size: 0.82rem;
|
||||
font-weight: 700;
|
||||
}
|
||||
|
||||
.filters input,
|
||||
.filters select {
|
||||
width: 100%;
|
||||
min-height: 40px;
|
||||
border: 1px solid #c8d2dc;
|
||||
border-radius: 6px;
|
||||
padding: 0 10px;
|
||||
color: #17212b;
|
||||
background: #ffffff;
|
||||
}
|
||||
|
||||
.notice {
|
||||
border: 1px solid #c8d2dc;
|
||||
border-radius: 6px;
|
||||
margin-bottom: 16px;
|
||||
padding: 12px 14px;
|
||||
background: #ffffff;
|
||||
}
|
||||
|
||||
.notice.error {
|
||||
border-color: #d94f4f;
|
||||
color: #8a1f1f;
|
||||
background: #fff4f4;
|
||||
}
|
||||
|
||||
.cards {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, minmax(120px, 1fr)) minmax(260px, 1.8fr);
|
||||
gap: 12px;
|
||||
margin-bottom: 18px;
|
||||
}
|
||||
|
||||
.stat,
|
||||
.panel,
|
||||
.trend-card {
|
||||
border: 1px solid #d5dde5;
|
||||
border-radius: 8px;
|
||||
background: #ffffff;
|
||||
}
|
||||
|
||||
.stat {
|
||||
min-height: 96px;
|
||||
padding: 16px;
|
||||
}
|
||||
|
||||
.stat span,
|
||||
.panel-header span,
|
||||
.trend-card p {
|
||||
color: #607080;
|
||||
font-size: 0.82rem;
|
||||
}
|
||||
|
||||
.stat strong {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
color: #132232;
|
||||
font-size: 1.9rem;
|
||||
}
|
||||
|
||||
.stat.wide strong {
|
||||
font-size: 1rem;
|
||||
}
|
||||
|
||||
.stat small {
|
||||
display: block;
|
||||
margin-top: 8px;
|
||||
color: #607080;
|
||||
}
|
||||
|
||||
.panel {
|
||||
margin-bottom: 18px;
|
||||
overflow: hidden;
|
||||
}
|
||||
|
||||
.panel-header {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
gap: 12px;
|
||||
border-bottom: 1px solid #d5dde5;
|
||||
padding: 14px 16px;
|
||||
}
|
||||
|
||||
.table-wrap {
|
||||
overflow-x: auto;
|
||||
}
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
min-width: 900px;
|
||||
border-collapse: collapse;
|
||||
}
|
||||
|
||||
th,
|
||||
td {
|
||||
border-bottom: 1px solid #e6ebf0;
|
||||
padding: 11px 12px;
|
||||
text-align: left;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
th {
|
||||
color: #536171;
|
||||
font-size: 0.78rem;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
|
||||
td {
|
||||
color: #1b2836;
|
||||
font-size: 0.9rem;
|
||||
}
|
||||
|
||||
.badge {
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
min-width: 52px;
|
||||
border-radius: 999px;
|
||||
padding: 4px 9px;
|
||||
font-size: 0.75rem;
|
||||
font-weight: 800;
|
||||
text-transform: uppercase;
|
||||
}
|
||||
|
||||
.badge.pass {
|
||||
color: #166534;
|
||||
background: #dcfce7;
|
||||
}
|
||||
|
||||
.badge.fail {
|
||||
color: #991b1b;
|
||||
background: #fee2e2;
|
||||
}
|
||||
|
||||
.badge.muted {
|
||||
opacity: 0.78;
|
||||
}
|
||||
|
||||
.empty {
|
||||
padding: 28px 16px;
|
||||
color: #607080;
|
||||
}
|
||||
|
||||
.empty.full-width {
|
||||
grid-column: 1 / -1;
|
||||
}
|
||||
|
||||
.trend-grid {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(auto-fill, minmax(280px, 1fr));
|
||||
gap: 12px;
|
||||
padding: 14px;
|
||||
}
|
||||
|
||||
.trend-card {
|
||||
display: grid;
|
||||
gap: 10px;
|
||||
min-width: 0;
|
||||
padding: 14px;
|
||||
}
|
||||
|
||||
.trend-chart {
|
||||
width: 100%;
|
||||
min-height: 96px;
|
||||
color: #0f6b8f;
|
||||
}
|
||||
|
||||
.point-pass {
|
||||
fill: #0f6b8f;
|
||||
}
|
||||
|
||||
.point-fail {
|
||||
fill: #d94f4f;
|
||||
}
|
||||
|
||||
.empty-chart {
|
||||
display: grid;
|
||||
min-height: 96px;
|
||||
place-items: center;
|
||||
color: #72808f;
|
||||
background: #f5f7f9;
|
||||
}
|
||||
|
||||
@media (max-width: 760px) {
|
||||
.dashboard {
|
||||
padding: 18px;
|
||||
}
|
||||
|
||||
.topbar,
|
||||
.panel-header {
|
||||
align-items: flex-start;
|
||||
flex-direction: column;
|
||||
}
|
||||
|
||||
.refresh-button {
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
.filters,
|
||||
.cards {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
{
|
||||
"compilerOptions": {
|
||||
"target": "ES2022",
|
||||
"useDefineForClassFields": true,
|
||||
"lib": ["DOM", "DOM.Iterable", "ES2022"],
|
||||
"allowJs": false,
|
||||
"skipLibCheck": true,
|
||||
"esModuleInterop": true,
|
||||
"allowSyntheticDefaultImports": true,
|
||||
"strict": true,
|
||||
"forceConsistentCasingInFileNames": true,
|
||||
"module": "ESNext",
|
||||
"moduleResolution": "Node",
|
||||
"resolveJsonModule": true,
|
||||
"isolatedModules": true,
|
||||
"noEmit": true,
|
||||
"jsx": "react-jsx"
|
||||
},
|
||||
"include": ["src"],
|
||||
"references": []
|
||||
}
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
import react from "@vitejs/plugin-react";
|
||||
import { defineConfig } from "vite";
|
||||
|
||||
export default defineConfig({
|
||||
plugins: [react()],
|
||||
server: {
|
||||
port: 5173,
|
||||
proxy: {
|
||||
"/api": "http://127.0.0.1:8000"
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
+4
-3
@@ -220,11 +220,12 @@ follow_imports = "silent"
|
||||
# ``*/_vendored/*`` matches upstream-provenance files vendored under any
|
||||
# ``_vendored/`` subdir (project-wide convention; mirrors the
|
||||
# ``_``-prefixed auto-discovery skip).
|
||||
skip = "./data,./wandb,ui/package-lock.json,*/_vendored/*"
|
||||
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,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"
|
||||
]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user