Compare commits

...
Author SHA1 Message Date
Peiyuan Zhang 67f5b53595 [fix]: keep v2 examples runnable 2026-07-04 20:22:18 +00:00
Peiyuan Zhang aa6bf5079b [docs]: scope v2 docs to inference runtime 2026-07-04 19:33:05 +00:00
Peiyuan Zhang 4ca56d8951 [fix]: route v2 serving tasks by capabilities 2026-07-04 19:32:46 +00:00
Peiyuan Zhang 1e19650eb3 [misc]: simplify v2 program execution model 2026-07-04 19:32:30 +00:00
Peiyuan Zhang de354f804c [feat]: add FastWan QAD FP8 card 2026-07-04 19:32:18 +00:00
Peiyuan Zhang 6feee4e4fa [feat]: add FP8 vendor support for Wan weights 2026-07-04 19:32:05 +00:00
Peiyuan Zhang 82e00615e5 [feat]: wire v2 VideoGenerator to real torch backend 2026-07-04 19:31:55 +00:00
Peiyuan Zhang 30cfbfc4f6 [refactor]: remove training semantics from v2 inference contracts 2026-07-04 19:31:35 +00:00
Peiyuan Zhang a4d8978c37 [refactor]: remove v2 training package 2026-07-04 19:31:12 +00:00
Will Lin 803a6c99ae [refactor] v2: drop unwired ConditioningInjector policy
ConditioningInjector (ABC) + PassthroughConditioning were defined and exported
but never instantiated or called — zero call sites anywhere. Unlike the policies
the §5 thesis actually names (CFG / flow-shift / precision / expert-routing),
which every recipe wires (cfg=ClassicCFG(), expert=NoRouting(), ...), conditioning
is injected INLINE by the loops: `st.cond["prompt_embeds"] = ctx.slots.get(...)`
(wan21/loop.py, ltx2/loop.py); the qwen_omni cascade conditions via loop wiring.
PassthroughConditioning even described a dataflow (state.scratch["cond"]) that is
not how conditioning actually flows (loops read ctx.slots). So it was a designed-
but-bypassed seam, not forward-design — and conditioning isn't in the README §5
policy list. Removed the ABC + impl + exports; kept the live `cond` field (the
loops fill it directly) and fixed its comment.

Tests: 143 passed, 2 skipped.
2026-07-04 17:15:35 +00:00
Will Lin 8a637c3215 [refactor] v2: remove unused DataRef provenance spec
DataRef (dataset_id/revision/description — "what a recipe trained on") was only
the type of RecipeSpec.data_contract, which no recipe authored and no code read
(0/0). It is absent from the README RecipeSpec contract (§2.1: method, parents,
assumes_loop, assumes_precision, consistency_required) and not among the §18
"wire the inert metadata" roadmap items — i.e. unwired governance metadata, not
deliberate forward-design. RecipeSpec now matches §2.1 exactly.

Kept CheckpointManifest: unlike DataRef it IS wired (ModelCard.checkpoint), the
declarative "explicit components + key maps, no name-detector guessing" load
contract — declared intent, not dead.

Tests: 143 passed, 2 skipped.
2026-07-04 17:15:35 +00:00
Will Lin 51bae33c7b [refactor] v2: drop unused typed-schema slots from Component/LoopSpec
state_schema / step_schema / result_schema (LoopSpec) and config_schema /
io_schema (ComponentSpec), plus valid_parallel_plans and parallel_constraints,
were authored by no recipe and read by no executor (audited: 0 reads / 0 sets
across native + tests, no dynamic dataclasses.fields/asdict/__dict__ access, no
consumer in scripts/_vendor/examples). They duplicated mechanisms that already
exist: the concrete LoopState/WorkPlan/StepResult classes the driver uses
directly, and the card-level ParallelismContract. Removing them slims the core
spec surface with zero behavior change.

Thesis untouched: cards still own components/loops/recipe/parity; kept the
parity, precision/placement, behavior-capture (behavior_schema), wired
extension_schema, and roadmap required_for/optional_for/resident_for fields.
Also drop the stale README step_cost_model mention (that field went with the
cost mechanism in 352c1b28).

Tests: 143 passed, 2 skipped (toy backend).
2026-07-04 17:15:34 +00:00
Will Lin 54b7fe3a55 [refactor] v2: drop dead toy components + unused Karras schedule
ToyLoRA / ToyControlNet / ToyTargetModel / ToyDraftModel / ToyRewardModel (+ the
_spec_target_next helper) in the toy backend, and build_karras_sigmas in the
sampler, were defined but referenced nowhere — no recipe, card, loop, test, or
example used them. They were toy stand-ins for capabilities (adapter plane,
speculative decode, served reward) and an EDM/Karras noise schedule that were
written ahead of being wired. Remove them (-134 LOC); a toy can be re-added when
the capability is actually wired. No thesis impact, no behavior change.

Suite green: 143 passed, 2 skipped (toy backend); import v2 stays torch-free.
2026-07-04 17:15:34 +00:00
Will Lin 1fe50c0092 [refactor] v2: group flat top-level into planes; isolate vendored under _vendor/
The v2 top level had grown to ~27 dirs + 13 loose files — half of them
single-concept abstraction shells (memory/ 86 LOC, transport/ 222, parity/ 178,
extend/ 226) and half vendored fastvideo code sitting as peers to the actual v2
design. Regroup to mirror README section 3 "Planes & dependency order":

  core/     enums+types, card, loop, program, parity, request, parallel
            (the model-native contracts; no kernels)
  runtime/  + folded-in substrate: cache, memory, transport, extend
            (the import graph shows only runtime consumes them; compile/cudagraph
            already lived here)
  serving/  + deploy/ (products / fleet)
  _vendor/  all copied fastvideo: models, layers, attention, configs, distributed,
            platforms, api, hooks, logging_utils, third_party + fastvideo_args/
            utils/logger/envs/forward_context — internal layout unchanged, still
            mirrors upstream for diffing

28 dirs + 13 root files -> 8 dirs + 6 root files. Pure mechanical move: 947
absolute-import paths rewritten (v2.X -> v2.{core,runtime,serving,_vendor}.X),
boundary-anchored so platform/ (native dispatch) and platforms/ (vendored CUDA
detect) no longer collide and neither does hooks/ vs extend/. Deleted the empty
v2/loader/. README section 3 + 16 updated; stale test count corrected
(34 files/216 tests -> 22 files/143 tests).

Validated: `import v2` stays torch-free; `v2/run_tests.py` and `pytest v2/tests/`
-> 143 passed, 2 skipped on the numpy toy backend. (On a GPU box force the toy
backend with CUDA_VISIBLE_DEVICES="" or detect() picks cuda.) No external importer
changed — examples use the re-exported `from v2 import VideoGenerator`.
2026-07-04 17:15:34 +00:00
SolitaryThinker 290795daf8 [refactor] v2: remove interleave; pooled run-to-completion serving (P2)
Second step of the runtime simplification (after cost removal). Drop the
coordinated step-interleave scheduler + the interleave-parity gate; serving is
now pooled run-to-completion.

- Engine: remove run_interleaved + the WorkUnit/BatchScheduler imports; run /
  run_serial drive each request to completion (tick/run_to_completion kept as the
  per-request stepper).
- scheduler.py: remove BatchScheduler + WorkUnit + the batches metric; the
  AdmissionController is now a pure refundable memory/OOM guard.
- AsyncEngine: bound concurrency with a serving pool (asyncio.Semaphore,
  max_concurrent) — each request waits for a slot, then runs to completion.
- parity: remove assert_interleave_parity (the run_serial==run_interleaved gate);
  rename interleave_gate.py -> compare.py (compare_outputs stays — bit-parity
  between execution paths, e.g. disaggregated==inline).
- card specs: drop ParitySpec.interleave_required + LoopSpec.allows_interleaving;
  stripped interleave_required from all cards.
- tests: delete the interleave-gate/parity tests; refocus the ones that exercised
  real behavior (residual-skip, compare_outputs symmetric-empty).
- README: removed the design-doc references at the top; simplified the thesis /
  scheduler (§6) / parity (§9) / package-layout / comparison sections to pooled
  run-to-completion (no cost model, no interleave gate).

CPU mini: 143 passed / 2 skipped. Native omni port (P3) still to come. (pyproject
kernel hack excluded.)
2026-07-04 17:15:34 +00:00
SolitaryThinker 0047d54a1b [refactor] v2: remove the cost mechanism (P1 of runtime simplification)
First step toward lean pooled run-to-completion serving: rip out the GPU-time
cost/budget machinery entirely (it priced nothing useful for the target design).

Removed: CostModel + LoopSpec.step_cost_model; ResourceRequest.compute_seconds;
StepResult.actual_seconds; AdmissionController's compute budget + SchedulerMetrics
.gpu_seconds (the memory/OOM reservation guard stays); the Profiler observer
(cost calibration); per-step timing in RuntimeLoopContext; cost-based fleet/Dynamo
routing (now a coarse step-count load proxy); DeploymentCard.cost_model. Stripped
cost from all 9 recipe cards + their loops.

Tests: dropped the 3 cost-specific tests (cost routing, cost_model aliasing,
compute-budget gate); refocused 2 (loop cache validation, NaNWatch-clean).

CPU mini: 151 passed / 2 skipped. Interleave removal + pooled serving (P2) and the
native omni port (P3) follow. (pyproject kernel hack excluded as always.)
2026-07-04 17:15:34 +00:00
SolitaryThinker b2d55a7ba0 [refactor] v2: full vendor cutover — copy fastvideo modeling + layer code into v2 (zero fastvideo imports)
Replace the re-export stubs with real vendored copies of the fastvideo modeling
+ layer + supporting infra, for the kept diffusion models (wan21, wan_causal,
ltx2, flux2, matrixgame2). v2 now imports ZERO `fastvideo.*` — it is
self-contained. (bagel/qwen_omni load from vllm_omni, an external pkg, not
fastvideo; cosmos3's load_id was already dangling — both out of scope here.)

Vendored (cp + `sed fastvideo. -> v2.`):
- models/  the 5 models' nn.Module dits/vaes/encoders/audio/upsamplers + the
           component loader/ + the lazy class registry (other families' rows are
           dormant/lazy — only the 5 resolve).
- layers/ attention/ platforms/ distributed/ configs/ logging_utils/ hooks/
  third_party/pynvml + top-level forward_context/fastvideo_args/envs/logger/
  utils/version — copied verbatim (layers et al. 'as is').
- api/  slimmed to schema + results (the VideoGenerator's config dataclasses);
  the fastvideo parser/presets/overrides (which pull the pipeline runtime) are
  intentionally NOT vendored — v2 has its own runtime/loop.

Decoupling surgery (cut the loader's coupling to the fastvideo runtime):
- configs/pipeline_registry.py (vendored from fastvideo/registry.py, renamed to
  avoid colliding with v2/registry.py): dropped the _register_presets() auto-call
  and matrixgame3 (removed model); config-class resolution preserved.
- configs/pipelines/__init__.py: dropped the registry back-edge (fixes an import
  cycle) — base.py imports the registry lazily where used.
- torch_backend: load_component now uses v2.models.loader.

The only remaining external 'fastvideo*' refs are `fastvideo_kernel` (the
separate optional CUDA-kernel pkg for sparse/MoBA attention) — guarded; the dense
TORCH_SDPA path v2 uses never imports it.

Vendored subtrees added to the pre-commit exclude (faithful copies, mirroring the
existing fastvideo/models exclusion — not re-linted, to stay re-syncable).

Verified: grep finds zero fastvideo-package imports in v2/; Wan2.1 T2V on H100 is
BIT-IDENTICAL to the fastvideo-backed path (same .npy SHA256, byte-for-byte); CPU
mini 156 (154 passed + 2 env-skipped on x86, torch present). Backup: branch v2_backup.
2026-07-04 17:15:34 +00:00
SolitaryThinker 24cfe281b3 [refactor] v2: prune recipes to 8 models (+ omni shared infra)
Keep: bagel, cosmos3, flux2, matrixgame2, ltx2, qwen_omni, wan21, wan_causal
(plus the shared omni/ package that bagel/cosmos3/qwen_omni depend on). Remove
the other 26 recipe packages.

- Delete 26 recipe dirs (adapters, adaptive, cosmos2, cosmos25, fastwan, gen3c,
  hunyuangamecraft, hunyuan_video(15), hyworld, image_video, kandinsky5,
  lingbotworld, longcat, lucy_edit, matrixgame3, multi_expert, reward, sd35,
  sfwan22, speculative, stable_audio, tiled, turbowan, unified, wan_fun_control).
- recipes/__init__.py: keep-closure builders + build_default_engine /
  build_omni_engine (dropped the workflow/tiled/unified/image_video engine helpers).
- registry.py: _BUCKET_C pruned to flux2 + matrixgame2; removed the cosmos2
  ModelEntry + the CosmosTransformer3DModel arch branch + the TurboWan-14B entry.
- Delete 11 tests for removed recipes/features; patch test_bucket_c_ports to drop
  the cosmos2 reference (it now auto-derives from the pruned _BUCKET_C).
- README: correct the recipes/ roster to the kept families.

Backup of the full pre-prune tree is on branch v2_backup. CPU mini 156 passed / 0
failed (the prior 5 torch-absent bucket_c failures are gone with sd35/stable_audio);
no dangling references to any removed recipe; all kept model ids still resolve.
2026-07-04 17:15:34 +00:00
SolitaryThinker dc086b207f [perf] v2: on-device denoise loop — kill the per-step numpy<->torch round-trip (Wan2.1)
The torch adapter boundary marshalled the latent host<->device on EVERY denoise
step: _t uploaded the latent (and re-uploaded the text embeds) and _n downloaded
the velocity with a forced CUDA sync — 2*N PCIe copies + N syncs per generation,
buying nothing, since the latent could stay resident on the GPU the whole loop.

Root cause was a numpy loop surface. But the loop MATH is already array-agnostic
(CFG combine + flow-match Euler are pure arithmetic; the solver kernel already
passes torch through). So introduce a per-platform array namespace (v2/platform/
array_ns): numpy on CPU (torch-free — the parity mini is unchanged), torch-on-
device on cuda. The latent is seeded with numpy and uploaded ONCE; it then stays
resident through forward -> CFG combine -> solver -> next step; a single host
marshal happens at the request/output boundary (engine._to_artifact).

Opt-in per recipe via ModelCard.device_io (set on the Wan cards). When set on a
GPU box, build_component flips the components' TorchComponent.device_io so _out
keeps tensors on-device (in fp32, matching the old _to_numpy cast so the combine
dtype is unchanged). Un-migrated families and the CPU toy keep numpy in/out.
Also: PrecisionPolicy.cast is array-preserving; _t accepts resident tensors.

Verified BIT-IDENTICAL on real Wan2.1-1.3B / H100: the on-device latent equals
the pre-change numpy-path latent exactly (max_abs_diff 0.0, np.array_equal True).
CPU mini holds 237 passed / 5 pre-existing; pre-commit clean.

Other WanDenoiseLoop families can flip device_io next (per-family GPU re-verify);
non-Wan loops migrate to the xp namespace later.
2026-07-04 17:15:34 +00:00
SolitaryThinker b9db151658 [refactor] v2: co-locate per-model torch adapters into their recipe packages
platform/backends/ had become a flat dump of 15 per-model torch_<model>.py
adapters next to the genuinely-shared infra. Each adapter is referenced from
its card by a plain 'module:Class' string loaded via importlib, so there was
no real coupling forcing it into platform/ — the Cosmos/Flux/etc adapter
belongs WITH its recipe (card/loop/program).

Move each torch_<model>.py -> v2/recipes/<model>/adapter.py and flip the card
strings to v2.recipes.<model>.adapter:<Class>. backends/ now holds only the
shared substrate (torch_backend base, torch_cuda registration, torch_kernels,
toy/cpu/accel). Each recipe is now a self-contained package.

Cross-refs updated: gen3c/adapter imports CosmosT5Encoder from cosmos2/adapter;
sd35/program + stable_audio/card import from their own package. Recipes still
import torch-free (adapters pulled only via the string on a GPU box) — CPU mini
holds 237 passed / 5 pre-existing (bucket_c torch-absent). pre-commit clean.
2026-07-04 17:15:34 +00:00
SolitaryThinker 9124963238 [misc] v2: full pre-commit clean (ruff UP038/SIM/UP031 + mypy annotations)
Sweep all of v2/ through pre-commit (was previously only run on changed
files). Fixes surfaced across untouched modules:

- ruff: isinstance-tuple -> X | Y (UP038), try/except/pass ->
  contextlib.suppress (SIM105), negated-return (SIM103), %-format ->
  f-string (UP031).
- mypy: add annotations for no-untyped-call + var-annotated across recipes,
  training methods, torch/toy backends, and serving.
- yapf reflow of the SF-Wan KV-cache call sites (semantics unchanged).

yapf/ruff/codespell/mypy all pass; v2 tests 237 passed / 5 pre-existing
(bucket_c torch-absent on CPU venv).
2026-07-04 17:15:34 +00:00
SolitaryThinker f90e8f3e76 [bugfix] v2: SF-Wan cross-chunk KV cache — condition each chunk on prior clean chunks
The causal adapter never passed a kv_cache, so CausalWanTransformer3DModel.forward routed to
_forward_train (no cross-chunk KV) on every chunk instead of _forward_inference (the CausVid
Alg-2 KV-cache path). Each chunk denoised blind to the previous ones; the loop's cross-chunk
"context" was a toy mean(prior_latents) the adapter ignored. Result: hard discontinuities at
every chunk boundary (frame-to-frame absdiff spikes 43-55 every ~12 frames).

Fix (cuda path only; toy/CPU path and the 237-test suite untouched):
- WanDiT.alloc_causal_caches(): allocate the persistent per-block KV + cross-attn caches sized
  from the model config (mirrors CausalDenoisingStage._initialize_kv_cache).
- WanDiT.__call__: thread kv_cache/crossattn_cache/current_start/cache_start/start_frame/
  frame_seqlen so the model runs _forward_inference.
- wan_causal/loop.py: own the caches in LoopState (per-request -> interleave-safe); pass
  current_start = chunk_idx*chunk_size*frame_seqlen per chunk; do the clean-KV write
  (timestep ~0) after each chunk so the next attends to it.

Verified on H100: frame-to-frame absdiff mean 12.8->4.4, max 55.4->8.3; chunk-boundary spikes
eliminated; coherent across all 7 chunks. CPU causal toy tests unchanged (24 passed).
2026-07-04 17:15:34 +00:00
SolitaryThinker 4ff33a28b7 [misc] v2: simplify docstrings/comments + drop deleted-design-doc citations
Sweep all 296 v2 modules: simplify verbose docstrings/comments and remove 376 dangling
"(design_vN §X)" citations to the now-deleted design docs (v2/README.md is the source of
truth). Comment/docstring-only — AST-verified code-identical; the CPU suite holds at 237
passed / 5 pre-existing. Also applies yapf + ruff --fix auto-fixes (import ordering,
forward-ref annotation de-quoting under `from __future__ import annotations`; behavior-
neutral, suite-confirmed) and adds the legitimate domain terms mot/clen/te to the codespell
ignore-list. Remaining ruff (24) + mypy (68 no-untyped-call) findings are pre-existing v2
debt, untouched here.
2026-07-04 17:15:34 +00:00
SolitaryThinker 940f94f435 [docs] v2: make v2/README.md the design source of truth + one-page philosophy + M* roadmap
Unify the four design docs (design.md, designv2.md, design_v3.md, designv4.md) into a single
authoritative v2/README.md: the (recipe, runtime) thesis, driven loops, planes, one-WorkUnit
scheduler, the parity ladder + interleave gate, training-on-shared-loops, the weight-sharing
topology catalog, the current GPU status (20+ models + the BAGEL/Qwen-Omni/Cosmos3 trio
verified), and a prioritized roadmap. Recast design_summary.md as a one-page design philosophy
pointing to it. Add .agents/exploration/mstar-v2-roadmap.md (the adversarially-verified M*
Walk-Graph gap analysis driving the roadmap). Delete the four superseded design docs.
2026-07-04 17:15:34 +00:00
SolitaryThinker ac9dcb63ea [bugfix] v2: MatrixGame2/3 causal-loop progress counter + MG3 patch alignment
MatrixGame2 (causal DMD loop): bump st.step_idx on every executed work unit (each
DMD step and each clean-context pass). The loop drives its own control flow off
block_idx/dmd_idx/phase, but the runtime's no-progress watchdog keys on
st.step_idx, so a multi-block causal rollout was seen as stalled. Mirrors what
every other recipe loop does.

MatrixGame3 (5B WanModel): patch-align the latent H/W (patch_size (1,2,2)) before
denoise. The model folds (H/2, W/2) tokens, so an odd latent dim made the
unpatchified velocity come back one row/col short of the noise latent. Crop to
(latent // patch) * patch, faithful to MatrixGame3DenoisingStage.

v2 mini: 240 passed.
2026-07-04 17:15:34 +00:00
SolitaryThinker 861f87e843 [bugfix] v2: correct SF-Wan + LTX2 2-stage SR sampling defaults (GPU frame-verified)
Two distilled few-step video models rendered incorrectly on GPU; root-caused via
dense frame sampling (contact sheets) and fixed in the recipe cards.

SF-Wan2.1 (self-forcing causal, wan_causal card): was oversaturated/overcooked.
The distilled student is CFG-FREE (guidance 1.0, single forward/step) and denoises
with the 4-step DMD schedule [1000,750,500,250] (warped by FlowShiftPolicy(5.0)),
at a native causal block of 3 latent frames. Defaults were ClassicCFG@6.0 + 2 steps
+ block 2 -> overcooked AND under-denoised. Fixed: num_chunks=7, chunk_size=3,
steps_per_chunk=4; SamplingDefaults num_steps=4, guidance_scale=1.0 (7x3=21 latent
-> 81 frames). Renders a clean raccoon-in-sunflowers across all 81 frames.

LTX2-Distilled 2-stage SR (ltx2 card + LTX2VAE): was temporally blocky. Root cause
was an OOM-forced 57-frame reduction (only 8 latent temporal frames); the model is
designed for 121 (16 latent frames). Enable VAE tiling in LTX2VAE so the 121-frame
full-res decode fits the 80 GiB GPU; keep base cfg_scale=3.0 (drives brightness;
cfg=1 washed out) with stg_scale=0.0 (v2's drop-text perturbation is not real
skip-layer STG). Renders the on-prompt backyard shot, bright + temporally coherent.

torch_backend.py: enable LTX2VAE tiling; clear pre-existing mypy no-untyped-call /
yapf debt on the file (surfaced once per-file linting bypassed the duplicate-module
flakiness) by annotating the helper/constructor/maker signatures.

Tests: update the 3 affected CPU defaults/chunk-count tests. v2 mini: 240 passed.
2026-07-04 17:15:34 +00:00
SolitaryThinker 7ef2d9083e [docs] v2: GPU bring-up results — 20 models verified on H100
V2_PORTING_STATUS.md now records the GPU bring-up outcome: 20 models generate
real video/audio on H100 (the 7 prior + 13 newly-ported), with the remaining
split into fastvideo/env-blocked (SLA/VSA kernels, transformers incompat, fastvideo
registry/flash_attn gaps) and HF-access-blocked (gated cosmos2/flux2/sd35) — none
a v2 recipe bug.
2026-07-04 17:15:34 +00:00
SolitaryThinker 0ac5367b54 [feat] v2 GPU bring-up: 5 more models verified on H100 (huge MoE/world + DMD)
Second GPU pass (distilled + huge dense, 2-wide). 5 verified end-to-end with real
weights; CPU toy path kept green (240 passed, 2 skipped).

VERIFIED:
  * matrixgame3   — mp4 (9,256,256,3), zero fixes. 6.47B, standard Wan attn (NOT
                    sparse-attn-blocked, like its mg2 sibling); degenerate single-clip.
  * fastwan       — FastWan2.2-TI2V-5B-FullAttn DMD 3-step, mp4 (17,256,256,3), zero
                    fixes. The FULL-ATTENTION variant has no VSA params -> the generic
                    Wan loader maps it cleanly (reuses WanDiT via load_id, no adapter).
  * longcat       — LongCat-Video-T2V 13.58B, mp4 (17,256,256,3), zero fixes, CPU
                    offload (~40GB peak).
  * sfwan22       — Self-Forcing Wan2.2-A14B causal+MoE (2x14B), mp4 (29,288,288,3),
                    CPU expert offload (~80GB peak).
  * lingbotworld  — Wan2.2-class 2x14B camera world model, mp4 (9,256,256,3), offload
                    (~98GB peak transient).

BLOCKED: turbowan-i2v-a14b — SLA sparse-attn params (attn1.attn_impl.proj_l on all
40 layers of both experts) cannot load into the dense Wan build + needs the
fastvideo-kernel SLA Triton kernels (no nvcc here). Confirmed via the safetensors
header (no 60GB download). Same SLA family as turbowan-1.3b.

Fixes (own-port only): sfwan22/loop.py, lingbotworld/{card,program}.py +
torch_lingbotworld.py. pre-commit clean per-package.
2026-07-04 17:15:34 +00:00
SolitaryThinker 56f84018df [bugfix] v2 VideoGenerator: modality-aware result path (audio/image, not only video)
VideoGenerator._result hardcoded out.artifacts['video'].frames, so an audio-only
(Stable Audio) or image-only (SD3.5 / FLUX.2 T2I) generation crashed with
KeyError 'video' even though the engine had correctly produced the AudioArtifact /
image TensorArtifact (surfaced during stable_audio GPU bring-up). _result now
guards on the artifact present: video -> mp4 (unchanged), else image -> png
([C,H,W]/[B,..] normalized to [H,W,C]), plus the existing audio -> sibling .wav;
image_path recorded in result.extra. The video path is byte-for-byte unchanged.

v2 mini 240 passed, 2 skipped; pre-commit clean.
2026-07-04 17:15:34 +00:00
SolitaryThinker d04127fe6f [feat] v2 GPU bring-up: 8 ported models verified end-to-end on H100 (+CPU-safe fixes)
Ran each dense public port through the real VideoGenerator on H100 (2-wide across
both GPUs). 8 produce real finite output end-to-end; per-model adapter/loop fixes
landed in each port's OWN files (no shared/fastvideo edits). CPU toy path kept
green (240 passed, 2 skipped) — GPU-only conditioning gated to the cuda backend.

VERIFIED (real GPU output):
  * stable_audio    — stereo audio (2, 441000) @44.1kHz. Fixes: dedicated
                      'conditioner' component kind (SA owns its T5, empty
                      text_encoder_configs) + ConditionerLoader from conditioner/
                      + VDenoiser c_noise = atan(sigma)/(pi/2).
  * matrixgame2     — mp4 (9,256,256,3). Loads CLEAN (not sparse-attn-blocked).
                      Fixes: 20ch cond_concat (4ch mask + 16ch img), mandatory i2v
                      (synth blank first frame), pre-sized kv_cache/crossattn_cache
                      (SDPA inference path, avoids flex_attention compile), bf16
                      autocast, per-request reset_caches.
  * gen3c           — video (9,256,256,3), zero fixes (worked first try).
  * wan_fun_control — video (9,256,256,3).
  * lucy_edit       — video (17,256,256,3).
  * hunyuangamecraft— video (9,256,256,3).
  * hunyuan_video   — video (3,9,256,256), dual LLaMA+CLIP + Hunyuan VAE.
  * hunyuan_video15 — video (9,256,256,3), dual Qwen+ByT5 (gated to cuda; CPU passes
                      single embed).

BLOCKED (not v2 recipe bugs — load/run reached, then a fastvideo/env wall):
  * cosmos25 — DiT + VAE ran finite on GPU; the Qwen2.5-VL Reason1 encoder hits a
               transformers 5.12.1 incompat in fastvideo shared code
               (Qwen2_5_VLConfig.pad_token_id).
  * kandinsky5 — fastvideo registry.py registers it with a bare PipelineConfig (no
                 Kandinsky5 config) -> load fails fastvideo-side. (latent z=16 +
                 visual_cond adapter corrected; toy decoupled to its own channels.)
  * hyworld — fastvideo's hyworld DiT hardcodes flash_attn (no SDPA fallback);
              flash_attn kernel not built here.

CPU-safety fixes (mine): kandinsky5 toy ToyDiT/ToyVAE use the toy LATENT_CHANNELS
(not the real z=16); hunyuan_video15 dual-encoder packing gated to cuda.
Sparse-attn distilled (turbowan/SLA, fastwan/VSA) + gated (cosmos2/flux2/sd35) +
huge (>80GB) handled separately. pre-commit clean per-package.
2026-07-04 17:15:34 +00:00
SolitaryThinker ae0cf3e2d1 [docs] v2: porting status — ALL fastvideo models ported (63/64 by-id; VSA env-blocked)
Rewrites V2_PORTING_STATUS.md to reflect completion: the scope is now ALL
fastvideo models (not Wan+LTX-2 only). Documents the self-contained recipe-package
porting mechanism (ComponentSpec.adapter), the 15 net-new architectures + 5
Wan-family variants newly ported (CPU-verified end-to-end; GPU=BRINGUP), the 7
GPU-verified models, and the single env-blocked id (VSA-14B, needs nvcc).
2026-07-04 17:15:34 +00:00
SolitaryThinker f4af3cf886 [feat] v2 registry: LTX-2/2.3 repo aliases -> by-id resolution 63/64
Adds explicit ModelEntry aliases for the LTX-2 (FastVideo/LTX2-Diffusers,
LTX2-base, Lightricks/LTX-2 -> single-stage base) and LTX-2.3 (LTX2.3-Diffusers,
LTX2.3-Distilled-Diffusers, LTX2.3-base, Lightricks/LTX-2.3, lightricks/ltx-2.3
-> the distilled joint-A/V card) naming variants of already-ported LTX
checkpoints (the arch fallback also resolves LTX2Transformer3DModel from a root).

v2 now resolves 63/64 fastvideo registry ids by exact id; the only remaining id,
FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers, is ENV-BLOCKED (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).

v2 mini 240 passed, 2 skipped; pre-commit clean.
2026-07-04 17:15:34 +00:00
SolitaryThinker 39580ad73a [feat] v2: port the residual Wan-family variants (rCM/DMD/v2v/control/causal-MoE)
Closes the bucket-B sampler/conditioning gap — each reuses the Wan/Causal ARCH
(no new torch adapter) with a new in-package loop/sampler/conditioning, declared
in _BUCKET_C as explicit-HF-id-only (transformer_cls="" so the generic Wan/Causal
arch fallback is NOT hijacked — only the exact id distinguishes the capability
variant from a base Wan of the same class).

  * turbowan      — TurboWan rCM (Reparameterized Consistency Model) few-step: a
                    faithful in-package RCMScheduler port (TrigFlow->RectifiedFlow
                    schedule + stochastic consistency SDE step), 1.3B/14B T2V +
                    TurboWan2.2-I2V-A14B (MoE i2v, boundary 0.9 in raw-sigma space)
  * lucy_edit     — Lucy-Edit v2v editor: a video_vae_encode node (the input video
                    -> 48ch cond latent) threaded via the shared i2v_cond hook ->
                    96ch Lucy DiT input (faithful to denoising.py is_lucy_edit)
  * wan_fun_control — Wan2.1-Fun-Control: control-video conditioning (reuses the
                    i2v [mask|cond] concat pattern)
  * sfwan22       — Self-Forcing Wan2.2-A14B: causal chunk_rollout + Wan2.2 MoE
                    boundary routing, i2v (boundary 0.9) + t2v (boundary 0.875)
  * fastwan       — FastWan DMD 3-step: TI2V-5B-FullAttn loadable; the VSA-trained
                    variants + non-strict to_gate_compress load are BRINGUP

All 12 residual ids resolve+build; base Wan/Causal resolution unchanged (arch
fallback not hijacked). The _BUCKET_C regression test auto-extended -> v2 mini
240 passed, 2 skipped; pre-commit clean (per-package + registry).
2026-07-04 17:15:34 +00:00
SolitaryThinker 521c2845e6 [test] v2: end-to-end CPU regression guard for the bucket-C ports
Data-driven from registry._BUCKET_C (+ cosmos2): each net-new ported arch
resolves through the registry (exact id + arch fallback) AND runs end-to-end on
the CPU toy backend via the public Engine path (resolve -> build card+program ->
load_card -> Engine.run), emitting exactly one modality-correct artifact
(video / image / audio) + latents. Auto-covers future _BUCKET_C rows.

21 tests pass; full v2 mini 232 passed, 2 skipped.
2026-07-04 17:15:34 +00:00
SolitaryThinker 0467edbd07 [feat] v2: port the 14 remaining bucket-C archs as self-contained recipe packages
Completes the bucket-C porting backlog. Each arch is a self-contained recipe
package (card-declared torch adapter via ComponentSpec.adapter + a new/forked
loop + program, NO edit to the shared torch_backend dispatch), following the
cosmos2 reference pattern. One _BUCKET_C table in v2/registry.py drives both the
explicit HF-id registry (PRIMARY) and the select_by_architecture fallback.

Ported (CPU-verified: import + card/program build + registry resolve + denoise
loop runs end-to-end on the CPU toy backend; GPU load/run is BRINGUP):
  * cosmos25       — Cosmos-Predict2.5 (flow-match, per-frame plain-sigma timestep;
                     reuse FLOW_MATCH_STEP; Reason1/Qwen2.5-VL encoder adapter)
  * hunyuan_video  — HunyuanVideo (reuses WanDenoiseLoop; dual LLaMA+CLIP encoders;
                     Hunyuan VAE scaling_factor) + FastHunyuan variant
  * hunyuan_video15 — HunyuanVideo 1.5 (480p/720p cards)
  * longcat        — LongCat-Video T2V/I2V/VC
  * sd35           — SD3.5 MMDiT (flow-match, image; triple-encoder joint embed +
                     pooled_projections)
  * gen3c          — GEN3C (EDM; 82ch pose-buffer DiT; camera/depth -> BRINGUP)
  * kandinsky5     — Kandinsky 5.0 T2V Lite
  * flux2          — FLUX.2 dev/klein (MMDiT, image; gated weights -> BRINGUP)
  * stable_audio   — Stable Audio Open (audio modality)
  * hunyuangamecraft, hyworld, lingbotworld, matrixgame2, matrixgame3 — interactive
                     world models; t2v/degenerate path CPU-verified, action/camera/
                     memory conditioning is BRINGUP (needs request-API extension)

Adapters declared via ComponentSpec.adapter (the ac29750b enabler) so each port
adds only NEW files (recipe package + per-arch torch_<arch>.py + optional facade
stub) — zero shared-file edits. Registry resolves all 31 bucket-C HF ids by exact
id + 14 architecture fallbacks; no regression (cosmos2/wan/ltx2 unchanged).
v2 mini green (211 passed, 2 skipped); pre-commit clean (per-file/registry).
2026-07-04 17:15:34 +00:00
SolitaryThinker bfbd90ea3a [feat] v2: Cosmos-Predict2-2B-Video2World port (EDM-Karras) — reference bucket-C recipe
First net-new architecture ported via the self-contained recipe-package pattern
(card-declared adapter + new loop, no shared-dispatch edit):

* CosmosDenoiseLoop (v2/recipes/cosmos2/loop.py): EDM preconditioning folded into
  a flow-match Euler integrator. Faithful port of CosmosDenoisingStage — Karras
  sigma schedule (rho=7, sigma_max=80 -> sigma_min=0.002, terminal clamp), latent
  init randn*sigma_max, per-step c_in/c_skip/c_out (sigma_data=1) -> x0, CFG in x0
  space, x0 -> velocity (x-x0)/sigma, FLOW_MATCH_STEP. video2world frame-replace
  conditioning threaded but inert for the t2v preset.
* build_karras_sigmas helper added to v2/loop/sampler.py.
* CosmosDiT + CosmosT5Encoder adapters (v2/platform/backends/torch_cosmos.py),
  declared on the card via ComponentSpec.adapter (the ac29750b enabler) — DiT
  returns raw EDM output + builds the mandatory zero condition/padding masks + fps;
  T5 uses the raw last_hidden_state (no Wan zero-pad). Reuses the WanVAE adapter.
* card/program/registry (HF id nvidia/Cosmos-Predict2-2B-Video2World + arch
  fallback on CosmosTransformer3DModel) + COSMOS_NEG prompt + SamplingDefaults
  (35 steps, gs 7, 704x1280, 93f, 16fps).

CPU-verified: Karras schedule, card/program build, registry resolve (id + arch),
EDM loop runs end-to-end on the CPU toy backend. GPU load/run is BRINGUP.
v2 mini green (211 passed, 2 skipped); pre-commit clean.
2026-07-04 17:15:34 +00:00
SolitaryThinker 5fd6e23e30 [feat] v2 torch backend: ComponentSpec.adapter — card-declared per-arch TorchComponent
A new architecture can declare its own torch adapter on the card
(ComponentSpec.adapter="module:Class") instead of editing the shared _make_dit/
_make_vae/_make_text_encoder dispatch. _explicit_adapter() constructs it as
cls(module, *extra, device=, dtype=) and short-circuits the built-in Wan/LTX2
class-name dispatch when set. This makes each bucket-C port a self-contained
recipe package (card + adapter module + loop + program) with no shared-file edit
-> conflict-free parallel porting. Unset -> unchanged built-in dispatch.

CPU mini green (211 passed, 2 skipped).
2026-07-04 17:15:34 +00:00
Will Lin fc3332550c [feat] v2: Wan2.2-I2V-A14B (MoE i2v) — combine boundary-routed experts + i2v conditioning
Reuses everything: 2 WanTransformer3DModel experts + BoundaryTimestepRouting (from the A14B MoE pattern),
the CLIP image encoder + first-frame [mask|cond] conditioning + the i2v program (from the Fun-InP i2v
port), and the shared WanDenoiseLoop (i2v hooks + the boundary expert). No new adapter. CPU-verified: the
toy MoE i2v runs end-to-end (2 experts + boundary + conditioning -> finite video); resolves with i2v caps.
Structural (GPU-pending: 2x14B, like the A14B T2V). Wan family now largely covered (T2V 1.3B/14B/TI2V-5B/
A14B, causal SF, i2v 1.3B/14B/A14B). CPU mini 211/2.
2026-07-04 17:15:34 +00:00
Will Lin 7464ef8308 [feat] v2: register Wan2.1-I2V-14B 480P/720P (reuse the GPU-verified i2v card)
The 14B i2v variants reuse the Wan2.1 i2v card/path proven on Fun-1.3B-InP — just per-variant params
(480P flow_shift 3.0 / 480x832, 720P flow_shift 5.0 / 720x1280). Registry resolution + caps verified;
specific 14B weights GPU-pending (same generic Wan i2v loader path that Fun-InP validated). i2v cluster
now supported; roadmap updated (11 models ported).
2026-07-04 17:15:34 +00:00
Will Lin 4a56274d10 [feat] v2: Wan2.1 i2v port (Wan2.1-Fun-1.3B-InP) — CLIP encoder + first-frame conditioning, GPU-verified
Real image-to-video, unlocking the i2v cluster. v2/recipes/wan21/i2v.py: CLIP image-encode -> the DiT's
encoder_hidden_states_image; first-frame VAE conditioning + a 4-channel mask -> the 20ch [mask|cond] that
the Wan adapter concatenates with the 16ch noise -> the 36ch i2v DiT input (mirrors fastvideo's
ImageEncodingStage + ImageVAEEncodingStage; v2's WanVAE.encode already applies the matching (z-mean)/std).
Reuses the shared WanDenoiseLoop (its None-default i2v hooks) + the Wan torch adapter unchanged. Adds
ToyImageEncoder + the image_encoder checkpoint subfolder stamp. Registered Wan2.1-Fun-1.3B-InP.

GPU-verified: loads via the generic Wan loader (1.56B, no param-mapping issue), runs the full i2v
conditioning, produces real video (3,9,256,384, std 0.44, finite, motion 0.041). CPU mini 211/2 (T2V
unregressed). BRINGUP: visual confirmation that the output follows the conditioning image is
human-in-the-loop.
2026-07-04 17:15:34 +00:00
Will Lin 9ede9af123 [feat] v2 Wan loop+adapter: optional i2v conditioning hooks (cond concat + CLIP context); T2V unchanged
Threads i2v conditioning through the SHARED WanDenoiseLoop with zero T2V risk: init() reads optional
slots i2v_cond (the [mask|cond] latent) + i2v_img_embeds (CLIP) into scratch; _velocity passes them to
the dit (context=, cond=); WanDiT concats cond (16->36ch) and uses the embeds as
encoder_hidden_states_image; capture is disabled only when i2v conditioning is present. For T2V both are
None -> the dit call, CFG, and cudagraph capture are byte-identical (CPU mini 211 pass, 2 skip — no
regression). ToyDiT accepts+ignores cond (image-conditioning is a GPU-path concern). Completes the i2v
backend seam; the program's mask+cond construction + the Wan i2v card/registry + GPU verify follow.
2026-07-04 17:15:34 +00:00
Will Lin 3d405f6dfa [feat] v2 torch backend: CLIP image-encoder adapter + image_encoder component kind (i2v groundwork)
Adds CLIPImageEncoder (encode_image -> the DiT's encoder_hidden_states_image) + the generic builder's
image_encoder maker (ImageEncoderLoader + ImageProcessorLoader), registered as the cuda 'image_encoder'
kind. Mirrors fastvideo's ImageEncodingStage. CPU-verified (component-kinds + lazy invariant, 211/2);
GPU path marked BRINGUP/written-not-run (processor subfolder + dtype to confirm on a real i2v checkpoint),
matching how the rest of the torch backend was originally landed. Reusable by the Wan i2v cluster + many
bucket-C models (Hunyuan/Cosmos i2v). Next i2v increments: the mask+cond latent construction + the
concat-into-DiT-input loop, then the card + registry + GPU-verify with Wan2.1-Fun-1.3B-InP.
2026-07-04 17:15:34 +00:00
Will Lin d891771ba3 [revert] v2: drop FastWan/VSA registry entries — generic Wan loader can't map their gated-attn params
GPU verification (Wan2.1-Fun... no: FastWan2.1-T2V-1.3B) failed at load: 'Parameter blocks.0.to_gate_compress.bias
not found in custom model state dict' — the FastVideo/* DMD-distilled checkpoints carry gated-attention
params (to_gate_compress) that the generic WanTransformer3DModel loader can't map. v2/registry.py's
select_by_architecture ALREADY rejects WanDMDPipeline for exactly this reason; my explicit ModelEntry
wrongly bypassed it. Reverted FastWan (1.3B + 14B-480P) and the unverified VSA-14B alias (same FastVideo/*
risk). Kept the official Wan2.1-T2V-14B (standard weights, same loader path as the GPU-verified 1.3B).

Lesson recorded in V2_PORTING_STATUS.md: FastWan/Turbo/VSA need a param-mapping fix (like LTX-2.3 did),
not just a schedule — bucket-C-effort. GPU-verify every port before claiming support. CPU mini 211/2.
2026-07-04 17:15:34 +00:00
Will Lin 67dea39052 [feat] v2: alias FastVideo/Wan2.1-VSA-T2V-14B-720P to the Wan-14B card (bucket B)
Same WanTransformer3DModel arch (VSA is an attention-backend choice, not a weight/arch difference); v2
runs dense TORCH_SDPA, so it resolves to build_wan_t2v_14b_card. Registry resolution verified; the GPU
forward path is the 1.3B-proven Wan adapter (specific 14B/VSA weights not separately GPU-run).
2026-07-04 17:15:34 +00:00
Will Lin 5f1d2ef7d2 [feat] v2: port Wan2.1-T2V-14B (bucket B) + document the all-models backlog
First bucket-B port toward 'support every fastvideo model': Wan2.1-T2V-14B reuses the Wan recipe +
torch adapter unchanged (same WanTransformer3DModel/AutoencoderKLWan/UMT5) — only a registry entry +
build_wan_t2v_14b_card (720p, flow_shift 5.0) + SamplingDefaults differ. Without the entry the arch
fallback would give it the 1.3B 480p defaults; the explicit entry gives 50 steps / 720x1280.

Also recorded the full backlog in V2_PORTING_STATUS.md: 63 fastvideo models = 8 ported / 21 bucket-B
(reuse Wan/Causal/LTX2 arch — registry+recipe+defaults, no new adapter) / 34 bucket-C (13 new
architectures needing a TorchComponent adapter). Updated the stale 'how to add a model' steps to the
post-redesign structure (v2/recipes/, torch_backend.py, SamplingDefaults). CPU mini 211 pass, 2 skip.
2026-07-04 17:15:34 +00:00
Will Lin 440b99523e [refactor] v2 torch backend: TorchComponent base + one generic builder + v2.* facade (Phase 1b/1c)
Addresses the adapter-setup pains: collapses the torch_cuda(trampolines)/torch_adapters/torch_ltx2 split
+ 11 near-identical adapter classes + 6 build_torch_* builders into:
- v2/platform/backends/torch_backend.py: a TorchComponent base centralizing .to/.eval, the numpy<->torch
  marshalling (ONE place), the set_forward_context wrap, and the weight surface; thin per-model subclasses
  (WanDiT/LTX2DiT/WanVAE/LTX2VAE/T5Encoder/Gemma/LTX2Upsampler/LTX2AudioVAE/LTX2Vocoder) carrying only
  forward semantics; and ONE build_component(spec) dispatching by spec.kind via _MAKERS.
- torch_cuda.py: registers that single generic builder for all 6 cuda kinds (no per-kind trampolines).
- v2 owns its namespace via re-export STUBS (facade, marked '# STUB'): v2/forward_context, v2/fastvideo_args,
  v2/distributed, v2/loader (the load_component seam), v2/api, v2/models/{dits,audio,upsamplers}/*. All v2
  code imports v2.*; 'from fastvideo' now lives ONLY in those 8 stub files -> a future per-module vendored
  cutover swaps a stub body, no caller changes. No divergence (stubs run fastvideo's live code).
- Deleted torch_adapters.py + torch_ltx2.py.

Verified: CPU mini 210 pass/2 skip; lazy invariant (platform load imports no torch); GPU bit-parity LTX-2.3
T2VS (audio std 0.04304, identical to pre-redesign) + Wan2.1 (video std 31.94, motion 5.997).
2026-07-04 17:15:34 +00:00
Will Lin 1cb2b4e84c [feat] v2: per-model sampling defaults on ModelCard (Phase 1a)
v2 had no per-model defaults — generate_video hardcoded 30 steps/25 frames/480x832/cfg5/16fps for
every model, badly wrong for e.g. LTX-2 distilled (wants 8 steps @1024x1536) or Wan2.2-TI2V (704x1280@24fps).

- New SamplingDefaults dataclass + ModelCard.sampling_defaults (v2/card/specs.py), exported from v2.card.
- Populated all 7 supported cards from fastvideo's InferencePreset defaults (steps/guidance/HxW/frames/fps
  + per-modality guidance for LTX-2.3 A/V). Negative prompts copied verbatim into v2/recipes/_prompts.py
  (recipe DATA, not model code -> v2-owned, no fastvideo import).
- VideoGenerator stores the resolved card; generate_video applies card defaults with precedence
  kwargs > SamplingParam > card > generic fallback (pure _resolve_default helper, unit-tested).
- test_sampling_defaults.py: per-card values + precedence (incl. empty-neg-prompt edge). CPU mini 210 pass, 2 skip.
2026-07-04 17:15:34 +00:00
Will Lin 51898f48e9 [refactor] v2: rename models/ (recipe layer) -> recipes/; move toy backend -> platform/backends/toy.py
Frees v2/models/ to become the vendored-architecture namespace that mirrors fastvideo/models
(part of making v2 self-contained / able to replace fastvideo). The v2 recipe layer (per-family
card.py/loop.py/program.py + common.py + the build_*_engine re-exports) is the recipe, not the
architectures, so it moves to v2/recipes/. The pure-numpy toy/parity implementations (ToyDiT etc.)
move from v2/models/backend.py to v2/platform/backends/toy.py (alongside cpu.py/accel.py/torch_*).

Mechanical: all imports are absolute, so v2.models.<x> -> v2.recipes.<x> and v2.models.backend ->
v2.platform.backends.toy across v2/ + examples/ (89 files, 178 refs). No behavior change.
CPU mini green (202 passed, 2 skipped); all v2 files compile.
2026-07-04 17:15:34 +00:00
SolitaryThinker 4e331eb7ff [refactor] v2: use absolute imports (v2.*) everywhere instead of relative
Mechanical conversion of every relative import under v2/ to an absolute v2.* path
(from .x / ..x / ...x -> from v2.<pkg>.x) so imports are unambiguous, grep-able, and
stable when code is copied/moved between entrypoints (VideoGenerator, CLI, server).

Surgical prefix-only rewrite: only the 'from <dots><module>' prefix changed — import
names, parentheses, multi-line formatting, comments, and ordering are byte-for-byte
preserved (no collapsing, no reorder, no unrelated reformatting).

- 471 imports across 125 files; v2/tests/ was already absolute (untouched).
- Validated: all 125 files compile, every 'from v2.* import' target resolves to a real
  module/package, zero relative imports remain (full sweep), CPU mini suite green
  (202 passed, 2 skipped).
2026-07-04 17:15:34 +00:00
SolitaryThinker c6d2976fc2 [feat] v2 VideoGenerator: A/V convenience path (generate_video -> T2VS -> mp4 + 24kHz wav)
Makes the 'Full A/V' LTX-2.3 deliverable reachable from the user-facing entrypoint, not just the
engine. A model advertising TEXT_TO_VIDEO_SOUND (LTX-2.3) now auto-issues a T2VS request, so generate()
/ generate_video() return BOTH modalities in one joint pass:
- VideoGenerator stores the resident instance + a supports_av flag (from card.capabilities); generate()
  gains want_audio (None=auto-by-capability, True/False to force) and routes T2V vs T2VS+{video,audio}.
- _result saves the stereo waveform as a sibling .wav at the vocoder's REAL rate (24000) — read off the
  built audio_vae adapter (TorchLTX2AudioVAE.sample_rate = Vocoder.output_sample_rate), since the
  AudioArtifact default rate is a placeholder. Populates GenerationResult.audio/.audio_sample_rate and
  extra['audio_path']. scipy IEEE-float WAV; [channels,samples] auto-transposed.
- Rewrote v2_basic_ltx2_3_distilled.py: registry routes to build_ltx2_3_card (its own joint T2VS A/V
  card, not the LTX-2 base/2-stage card); the example prints both the mp4 and the wav.
- GPU-verified via the convenience API: ev.mp4 + ev.wav (24000 Hz, stereo 61920x2, nonzero, std 0.043).
  CPU mini green (202 passed, 2 skipped); engine/program/toy paths untouched (test_ltx2_av pins 44100).
2026-07-04 17:15:34 +00:00
SolitaryThinker 6096b00aeb [feat] v2 LTX-2.3 T2VS GPU-verified: audio VAE/vocoder wiring + dual-connector audio fix
The full joint text->video+audio LTX-2.3 path now generates on the real 18.99B model:
- GPU audio components: build_torch_audio_vae (AudioDecoderLoader -> LTX2AudioDecoder, chains the
  vocoder) + build_torch_vocoder (VocoderLoader -> LTX2Vocoder); registered the 'audio_vae'/'vocoder'
  cuda component kinds; stamped their checkpoint subfolders (_WAN21_SUBFOLDERS).
- Fix: TorchGemma.encode_av must pass output_hidden_states=True — the 2.3 connector's SEPARATE audio
  projection lives in hidden_states[0] only then (gemma.py:703); without it the audio text fell back to
  the video embedding (4096 vs 2048 -> audio cross-attn shape mismatch).
- GPU-verified T2VS: video (3,33,256,384, std 0.68) + audio (stereo 2x61920 @24kHz, nonzero, std 0.059).
- test_torch_backend: cuda component kinds now include audio_vae + vocoder. CPU suite green (202+2).
2026-07-04 17:15:34 +00:00
SolitaryThinker 1d23399d81 [feat] v2 LTX-2.3 T2VS: single-stage joint audio+video card/loop/program (CPU-verified)
Makes LTX-2.3 a first-class, faithful card (was wrongly merged into the single-stage base):
- LTX23DenoiseLoop (loop.py): single-pass joint A/V denoise — one DiT forward per step cross-attends
  video<->audio via the adapter's (v_vel,a_vel) return; full-res video latent + a [8,T,16] audio latent;
  distilled few-step schedule (BASE_SIGMAS). Video-only when no audio requested.
- build_ltx2_3_card (model_id 'ltx2.3-distilled' — the name now correctly names the REAL 2.3): 5
  components incl. audio_vae (AudioDecoder) + vocoder (required_for t2vs, optional_for t2v); caps
  T2V + T2VS. build_ltx2_3_program: dual-connector text-encode -> joint denoise -> video + audio decode.
- registry.py routes FastVideo/LTX-2.3-Distilled-Diffusers -> this card (split from the base entry).
- Toy support: ToyTextEncoder.encode_av (separate video/audio text), ToyDiT joint A/V (audio now a
  keyword-only arg so positional  callers like the talker are unaffected), channel-agnostic
  ToyAudioVAE (np.resize identity for the existing 2-stage T2VS).
- CPU-verified: toy T2VS -> video+audio (8 steps), T2V -> video-only; CPU suite green (202+2).
GPU audio-VAE/vocoder loaders + the real-T2VS GPU verify are the next step.
2026-07-04 17:15:34 +00:00
SolitaryThinker fefcd415ff [feat] v2 LTX-2 adapters: A/V foundation (joint DiT forward + dual text connector + audio decode)
Foundation for the LTX-2.3 T2VS port (card/loop/program wiring + GPU verify to follow):
- TorchLTX2DiT.__call__ gains an optional joint audio path: pass audio_latent[8,T,16] + audio_text and it
  feeds audio_hidden_states/audio_encoder_hidden_states/audio_timestep/audio_sigma in ONE forward
  (LTX-2.3 cross-attends video<->audio) and returns (video_velocity, audio_velocity). Video-only call is
  byte-for-byte unchanged (audio_latent=None).
- TorchGemma.encode_av returns the SEPARATE (video_text, audio_text) projections from the 2.3 connector
  (video=last_hidden_state, audio=hidden_states[0]); 2.0 returns them equal.
- TorchLTX2AudioVAE (AudioDecoder -> Vocoder -> waveform@24kHz) + TorchLTX2Vocoder wrapper.
Additive + backward-compatible; CPU suite unaffected (adapters are GPU-lazy).
2026-07-04 17:15:34 +00:00
SolitaryThinker 7ac2ff0d1c [refactor] v2: shared model registry (HF-id primary + arch fallback) for all entrypoints
Per review: dispatch should be a directly-mapped HF-string -> card registry (like fastvideo), shared
by every entrypoint (VideoGenerator + a future CLI / server), not buried in the generator.

- New v2/registry.py mirrors fastvideo's fastvideo/registry.py hybrid resolution: (1) exact HF repo id
  in an explicit ModelEntry registry [PRIMARY — correct per-model card/capabilities, and the only way to
  split same-architecture capability variants like Wan2.1 T2V vs the i2v 'InP' 1.3B], (2) short repo-name
  match, (3) architecture inference [FALLBACK — local paths / unregistered repos]. resolve(model_path[,
  root]) is the single shared entry point.
- video_generator.py: moved _read_arch_signature/_select_builders into the registry; from_config now
  calls resolve() — registered ids resolve with no config read, else arch inference on a cheap *.json
  snapshot. Reconciles the earlier 'no brittle table' refactor with the 'map hf string -> card' ask: one
  clean registry + a fallback, not three coupled structures.
- Verified: registry resolves exact-id / short-name / arch-fallback / unregistered correctly; CPU suite
  green (202+2); wan21 GPU smoke generates via the new resolve path.
2026-07-04 17:15:34 +00:00
SolitaryThinker fa9c58b419 [fix] v2: name LTX-2 cards by architecture (2-stage vs single-stage) + Wan2.1 is T2V-only
Addresses the 'how is ltx2 separate from ltx2.3' confusion + a wrong capability:

- LTX-2 cards renamed by ARCHITECTURE (the version labels did not map to it): build_ltx2_card model_id
  'ltx2.3-distilled' -> 'ltx2-2stage-distilled' (two-stage base->upsample->refine; serves the
  upsampler-having FastVideo/LTX2-Distilled-Diffusers); build_ltx2_base_card 'ltx2.base' ->
  'ltx2-single-stage' (one loop; serves Davids048 base + the single-stage FastVideo/LTX-2.3-Distilled,
  which has NO spatial_upsampler). Dispatch already splits on has_spatial_upsampler. Updated the
  model-id refs in the mini's tests/examples.
- Wan2.1 base is T2V-only: dropped the wrong Capability.IMAGE_TO_VIDEO + narrowed components'
  required_for to {t2v} (i2v is the separate InP variant; v2 has no i2v path yet). build_wan21_card is
  shared by wan21 + wan2.2-ti2v; the A14B card was already T2V-only.
- CPU suite green (202 passed, 2 skipped).
2026-07-04 17:15:34 +00:00
SolitaryThinker 0662b42510 [docs] v2: LTX-2 base/2.3 GPU-verified on rebuilt x86 stack + remaining-port mechanisms
- LTX-2 base (Davids048) and LTX-2.3-Distilled both generate real video (inter-frame motion 4.5 / 6.6)
  via the single-stage base card — moved to Working (7 models now verified).
- Environment: the aarch64 venv was rebuilt for x86 (torch 2.11.0+cu128) + re-validated (CPU suite green,
  wan21 + LTX-2 base/2.3 generate).
- Remaining Wan+LTX-2 ports documented with concrete mechanisms: Wan2.2-i2v (SigLIP image_encoder +
  VAE-encode first-frame + concat-mask -> larger-in_channels i2v DiT), TurboWan (RCMScheduler consistency
  loop), Lucy-Edit (Wan v2v via VideoVAEEncodingStage), FastWan (VSA, env-blocked: no nvcc).
2026-07-04 17:15:34 +00:00
SolitaryThinker ae6d8085de [docs] v2: LTX-2.3 example + roadmap (A14B offload working, base/2.3 ported, env status)
- v2_basic_ltx2_3_distilled.py: LTX-2.3-Distilled routes to the single-stage base card (no
  spatial_upsampler) via the arch dispatch; pass few steps for the distilled schedule.
- V2_PORTING_STATUS.md: A14B moved to Working (CPU expert offload, 60GB peak); LTX-2 base + 2.3 added
  (code-complete, GPU re-verify pending); Environment-status note on the mid-session aarch64->x86 host
  reschedule that blocks GPU re-verify.
2026-07-04 17:15:34 +00:00
SolitaryThinker d0648ba8d9 [feat] v2: Wan2.2-A14B MoE CPU offload (fits 1 GPU) + LTX-2 base/2.3 single-stage port
Within the bounded Wan+LTX-2 scope:

- Wan2.2-A14B MoE now GENERATES on a single 80GB GPU via CPU offload (TorchWanDiT offload_group):
  the two 14B experts live on CPU and only the active one is swapped onto the GPU at the boundary-
  timestep transition (a single swap, not per-step). GPU-verified: 60GB peak (vs 79GB OOM), produced
  wan22_a14b_lion.mp4 (17x480x832, std 54.7, motion 8.36). Single-expert Wan stays resident.

- LTX-2 base (single-stage) port: build_ltx2_base_card + build_ltx2_base_program reuse the LTX-2
  adapters at FULL latent res with a request-driven many-step flow-match (LTX2DenoiseLoop full_res/
  request_steps/base_flow_sigmas; distilled base/refine path preserved via False defaults). The SAME
  single-stage card serves LTX-2.3-Distilled (also single-stage: no spatial_upsampler) — dispatched by
  the new has_spatial_upsampler discriminator in _select_builders. v2_basic_ltx2.py added; VideoGenerator
  gains shutdown() for API parity.

- Fixes: from_config 'os' scoping (shadowed module import); LTX-2 upsampler per_channel_statistics
  source (the AE's .decoder, not the top-level module).

Verification status: A14B offload, the upsampler, and the arch-dispatch refactor were GPU-verified
earlier this session. LTX-2 base/2.3 are CPU-verified (cards/programs build, dispatch routes, schedule
correct); their GPU smoke tests were pending when the box was rescheduled aarch64->x86 mid-session,
which broke the aarch64 venv (numpy/torch unrunnable on x86) — GPU re-verify blocked on the env.
2026-07-04 17:15:34 +00:00
SolitaryThinker e10828346f [feat] v2: architecture-driven dispatch + real LTX-2 upsampler + Wan2.2-A14B MoE card
Two reviewer asks + the next Wan port, all within the bounded Wan+LTX-2 scope:

1. Architecture-driven dispatch (replaces the HF-id table + substring fallback): from_config reads the
   checkpoint's pipeline/transformer/VAE class names (+ z_dim, transformer_2) and picks the v2 card via
   _select_builders — mirroring fastvideo's get_pipeline_config_cls_from_name. Resolves local paths /
   renamed repos / new distilled variants of a known arch with no table edits, and cleanly REJECTS
   FastWan (detected by WanDMDPipeline) with a precise message instead of a confusing load crash.

2. Real LTX-2 spatial upsampler (was a nearest-neighbor np.repeat stand-in): new 'upsampler' component
   kind -> TorchLTX2Upsampler wraps the real LTX2LatentUpsampler and applies the repo's upsample_video
   (un_normalize via the VAE decoder's per_channel_statistics -> learned 2x upsample -> normalize). CPU
   keeps ToyUpsampler (np.repeat) via the factory terminal, so the program calls
   component('spatial_upsampler').upsample(...) on both backends with no device branch. GPU-verified:
   9x512x768, std 70.5, motion 6.51.

3. Wan2.2-T2V-A14B MoE card (build_wan22_a14b_card): two WanTransformer3DModel experts +
   BoundaryTimestepRouting @0.875. GPU-verified that both experts denoise; OOMs in VAE decode on one
   80GB GPU (~70GB resident) — upstream offloads the DiT for MoE; documented as offload-blocked.

CPU suite green (202 passed, 2 skipped). FastWan root-caused (non-strict load of VSA gate_compress +
VSA not built); roadmap (V2_PORTING_STATUS.md) updated with the bounded scope + per-model status.
2026-07-04 17:15:34 +00:00
SolitaryThinker 655f362cf4 [feat] v2 port: Wan2.2-TI2V-5B (T2V) — 4th GPU-verified model
- Wan2.2-TI2V-5B reuses the Wan adapters (WanTransformer3DModel / AutoencoderKLWan / UMT5); deltas
  are the higher-compression VAE geometry (z_dim=48, 16x spatial, 4x temporal) and 480p flow-shift 5.0.
  The DiT forward accepts a scalar timestep (1D path), so no per-frame expand_timesteps for pure t2v.
- WanDenoiseLoop / build_wan21_card gain optional geometry params (latent_channels/spatial_ratio/
  temporal_ratio) defaulting to Wan2.1 (16/8/4) -> wan21 path unchanged; build_wan22_ti2v_card sets
  48/16/4. Registered as family 'wan2.2-ti2v' in VideoGenerator; v2_basic_wan2_2_ti2v.py added (T2V).
- Verified: 25x448x768 mp4, std 62.3, inter-frame motion 4.89 (coherent). CPU suite green (202+2skip).
- Corrected the now-disproven FastWan='wan21 reuse' mapping (its DMD checkpoint to_gate_compress param
  mapping differs); roadmap updated (TI2V-5B working; A14B MoE + I2V remain).
2026-07-04 17:15:34 +00:00
SolitaryThinker 3541e81d66 [feat] v2 VideoGenerator: convenience API (from_pretrained/generate_video) + porting roadmap
- VideoGenerator gains the convenience surface most basic examples use: from_pretrained(model,
  num_gpus/*_cpu_offload/...) and generate_video(prompt, sampling_param=, **kwargs) -> result, on top
  of the typed from_config/generate. Accepts SamplingParam.
- v2 examples matching the upstream convenience-API examples for the verified models: v2_basic.py
  (Wan2.1) and v2_basic_self_forcing_causal.py (SF-causal).
- V2_PORTING_STATUS.md: honest per-family roadmap. Working: wan21, wan_causal, ltx2-distilled. Each
  further model needs per-model work (FastWan: WanDMD to_gate_compress param mapping; TurboWan: RCM
  consistency sampler; Wan2.2: MoE card; LTX2 base/i2v; new families: new cards/adapters; gated Flux2 /
  local GEN3C / audio StableAudio / interactive MatrixGame blocked in this env).
2026-07-04 17:15:34 +00:00
SolitaryThinker f51497ee6d [feat] v2 VideoGenerator: typed fastvideo.api entrypoint over the v2 engine
Mirrors fastvideo.entrypoints.VideoGenerator (from_config(GeneratorConfig) -> generate(
GenerationRequest) -> GenerationResult.video_path), reusing the OFFICIAL fastvideo.api config classes
so a basic_dmd_new_api.py-style script differs only by importing VideoGenerator from v2.
- model_path -> v2 card registry (Wan2.1 / FastWan -> wan21; SFWan -> wan_causal; LTX2 -> ltx2);
  snapshot_download + stamp_wan21_checkpoints + Engine(cuda) + program; generate maps SamplingConfig
  -> DiffusionParams -> make_request -> eng.run, saves the [C,T,H,W] decode as an mp4.
- Lazy v2.__getattr__ keeps 'import v2' torch-free (verified) so the CPU mini stays green (202+2).
- examples/inference/basic/v2_basic_new_api.py runs all three GPU models through this API.
- Single-GPU, resident, TORCH_SDPA (EngineConfig offload/num_gpus>1/VSA accepted for parity, not applied).
verified: wan21 from_config->generate->mp4 (frames (5,256,256,3) uint8, video_path written).
2026-07-04 17:15:34 +00:00
SolitaryThinker d8af2e60d2 [feat] v2 ltx2: two-stage distilled GPU bring-up (LTX2Transformer3DModel 18.88B)
Official FastVideo/LTX2-Distilled-Diffusers. New torch_ltx2.py adapters (build_torch_* dispatch on
class):
- TorchLTX2DiT: patchify-internal; per-token timestep ones(B,tok,1)*sigma (sigma direct) + per-sample
  video_sigma; DiT predicts x0 so the adapter returns velocity=(x_t-x0)/sigma for the v2 flow-match step.
- TorchLTX2VAE: CausalVideoAutoencoder decode (internal per-channel un_normalize).
- TorchGemma: LTX2GemmaTextEncoderModel (Gemma + feature-extractor + connectors) -> last_hidden_state.
- ltx2 loop: real 128-ch latent geometry on cuda (32x spatial / 8x temporal; half-res base, 2x upsample).
e2e two-stage (8+3 steps) -> coherent, high-quality video (surfers at sunset), (3,9,512,768). NOTE:
still uses the v2 program's np.repeat upsampler between stages (the refine regenerates from noise so
output is faithful-quality); real LTX2LatentUpsampler swap-in is a follow-up.
2026-07-04 17:15:34 +00:00
SolitaryThinker f79919ba8a [feat] v2 wan_causal: causal DiT (CausalWanTransformer3DModel) GPU bring-up
Official SF checkpoint wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers (reuses TorchWanVAE + TorchT5Encoder).
- TorchWanDiT detects the causal transformer: ignores the chunk_rollout loop's latent `context` (the
  real model conditions across chunks via an internal kv_cache, not a forward arg) -> dispatches to
  full-attention _forward_train, and passes a per-latent-frame timestep [B, num_frames] (the causal
  block asserts a per-frame temb), uniform per chunk.
- chunk_rollout real geometry on cuda (16ch; chunk_size latent frames; 8x spatial).
e2e: chunk_rollout over the SF student -> coherent video (cat in a garden), (3,21,480,832). Fidelity
gap (artifacts): the v2 loop's per-chunk few-step sampling != the official kv-cache streaming + SF
schedule (a follow-up).
2026-07-04 17:15:34 +00:00
SolitaryThinker 1594f8e6be [fix] v2 tests: make no-torch-import guards GPU-aware (skipif torch installed)
The cuda-availability probe imports torch by design; the no-torch-import invariant is only
verifiable when torch is absent. skipif torch installed -> green on GPU box (202 passed, 2 skipped),
still enforced in torchless CI.
2026-07-04 17:15:34 +00:00
SolitaryThinker 64cadaa0bf [feat] v2 wan21: backend-aware latent geometry + checkpoint stamping
- latent_shape(req, model): real Wan geometry (16ch; (T-1)//4+1, H/8, W/8) on the cuda backend;
  toy stand-in stays on accel/cpu.
- stamp_wan21_checkpoints(card, model_root): map a root (local dir or HF id) onto the 3 components'
  ComponentSpec.checkpoint; build_wan21_card(checkpoint_root=...) optional.
2026-07-04 17:15:34 +00:00
SolitaryThinker 750fc1245b [feat] v2 cuda backend: real fastvideo construction + risk A-E fixes (Wan2.1 verified on H100)
Take the written-not-run torch adapters to runs-and-generates on 1x H100 (aarch64):
- A: FastVideoArgs.from_kwargs(model_path=root) builds the real pipeline_config; single-GPU dist
  init; load each component from its subfolder; tokenizer from the sibling <root>/tokenizer.
- B/C: DiT forward wrapped in set_forward_context(attn_metadata=None) (SDPA dense path);
  timestep=sigma*1000 + bare-velocity output confirmed.
- D: VAE decode denormalizes z*std+mean; removed the double-mean (it re-added
  shift_factor==latents_mean) that washed out the video.
- E: UMT5 from config; text embeds zero-padded to text_len (Wan t5_postprocess_text) - the fix
  that took output from a dark blur to a coherent prompt-matching scene.
- Components run at native precision (DiT bf16, VAE/text fp32); checkpoint check before dist init.
2026-07-04 17:15:34 +00:00
SolitaryThinker 7725998b0c [docs] add v2/HANDOFF.md for GPU-side bring-up of the torch backend
Orientation + process doc for an agent on a GPU branch: the 6 commits already
landed, the files to touch, the gating tasks (Risk A FastVideoArgs/checkpoint),
the verification bar (CPU suite stays 204; GPU generation matches a reference via
the SSIM harness), commit/push rules (no Claude co-author; don't rewrite history;
wandb token referenced not embedded), and the gotchas. Points to
GPU_BRINGUP.md for the detailed checklist + risk table.
2026-07-04 17:15:34 +00:00
SolitaryThinker 4b61dedc43 [fix] correct GPU adapters against real fastvideo API (cross-check findings)
Adversarial cross-check of the written-not-run torch adapters against the real
fastvideo source confirmed the interface contracts (DiT returns bare velocity
tensor; timestep=sigma*1000; encode().mode() + bare decode; .last_hidden_state;
no fused solver kernel) but caught a wrong construction layer. Fixed in code:

- Construction: WanTransformer3DModel / AutoencoderKLWan have NO from_pretrained.
  Replace it with the real FastVideo loaders (TransformerLoader / VAELoader /
  TextEncoderLoader + TokenizerLoader, each load(model_path, fastvideo_args)).
  The loader resolves the class from the checkpoint config — so UMT5-vs-T5 is
  chosen correctly instead of hardcoded (was BLOCKER #1/#4/#5).
- Text encoder: wrap the forward in set_forward_context(...) — the (U)MT5
  attention reads global state via get_forward_context(); a bare call mis-encodes
  (was BLOCKER #2). Drop the wrong padding="max_length".
- VAE: apply latent normalization the DiT expects — (z-mean)*inv_std on encode,
  inverse on decode, with latents_std stored as its reciprocal; shift_factor
  before decode (was BLOCKER #3). Skipping it yields washed-out video, not error.

The remaining unknowns are genuinely box-dependent (FastVideoArgs fields,
shift_factor placement, exact tokenizer kwargs, FSDP) — GPU_BRINGUP.md reconciled
to mark what's now fixed-in-code vs what still needs the box. 204 CPU tests pass.
2026-07-04 17:15:34 +00:00
SolitaryThinker 1e58d90d02 [feat] real torch/CUDA backend (written-not-run) behind the cuda cells
Implement the GPU backend the substrate was built for: Platform.detect() ->
cuda resolves real torch adapters + torch solver ops instead of the numpy
rungs, with the existing loops/policies/scheduler/training unchanged.

WRITTEN-NOT-RUN: this environment has no GPU/torch, so the torch code is
grounded in the verbatim real fastvideo APIs (DiT forward signature confirmed
from source) but cannot be executed/verified here. It is gated available=False
(CPU mini stays green; importing the backends never imports torch), with every
on-box confirm point marked `# BRINGUP` and an ordered checklist in
platform/backends/GPU_BRINGUP.md.

- torch_adapters.py: TorchWanDiT / TorchWanVAE / TorchT5Encoder wrap the real
  module named by each card's load_id and bridge it to the mini's duck-typed
  surface (numpy<->torch at the boundary; loop math stays numpy fp32). DiT
  weight-surface (copy_from/blend_from/clone) for serving sync; mse_grad_step
  raises (GPU training is a separate workstream).
- torch_kernels.py: flow_match_step / flow_sde_step as plain torch elementwise.
  Grounded conclusion from the kernel audit: fastvideo-kernel ships NO fused
  solver kernel (only attention/norm/quant primitives), so the cuda solver is
  torch, registered at arch generic with an honest source string.
- torch_cuda.py: rewritten as lazy trampolines (torch imported only inside
  builder/kernel bodies). Adds the missing vae + text_encoder cuda components
  (they'd otherwise silently fall back to the toy) and corrects the dishonest
  "fastvideo-kernel:flow_*" labels.
- ComponentSpec.checkpoint: the weights source for the torch adapter (risk A;
  the one field the cards didn't carry). Empty on the CPU toys.
- 7 CPU-verifiable wiring tests: honest registration/sources, torch-free import,
  cuda-resolves-real-cells-not-toy, build-fails-loudly-without-torch.

204 tests pass (CPU). The torch path needs a GPU box to verify (GPU_BRINGUP.md).
2026-07-04 17:15:34 +00:00
SolitaryThinker 47f7a04e09 [feat] static-buffer capture form for the cudagraph step body (Path A)
Close the loudest deferred gap from the cudagraph audit: the capturable step
now binds its I/O to address-stable static buffers (modeling real CUDA static
I/O buffers), instead of allocating fresh arrays per call.

- StaticWorkspace: address-stable buffers allocated once per capture key; bind()
  copies the current step's inputs in place via np.copyto, which RAISES on a
  shape/dtype mismatch — turning the weak peak_activation_bytes proxy into a real
  key-soundness backstop (a step whose shape doesn't fit can't replay an
  incompatible graph). Output written into a static buffer too.
- WorkPlan.graph_fn / graph_inputs: a capturable step exposes its deterministic
  op-structure as graph_fn(model, workspace) reading EVERY per-step input (latent,
  sigmas, conditioning, scale) from the workspace — never from closure over
  per-step data — plus the dict of current values. Loops without both stay on the
  eager path. wan21 factors a shared _velocity() so graph_fn and the eager run
  stay bit-identical.
- Capturer dispatch captures/replays via graph_fn against the keyed workspace;
  the workspace is shared per key on the instance. Correct under the engine's
  synchronous step execution (bind+graph_fn atomic per dispatch, output returned
  as a copy) — proven by the batch-of-N interleave gate running two same-key
  requests through the shared workspace bit-identically. A concurrent/multi-stream
  executor would need a per-stream pool (documented).
- 2 new tests (no-static-form eager-break, static-buffer shape-mismatch raises);
  workspace collision test split into bytes-proxy vs shape-backstop.

197 tests pass.
2026-07-04 17:15:34 +00:00
SolitaryThinker 9b8838834c [feat] piecewise CUDA-graph capture/replay at the step boundary (Path A)
Wire the capture/replay lifecycle into the driven-loop step boundary — the
other half of Path A (hand-fused kernels behind the registry + piecewise
cudagraphs, no compiler). Models and tests the correctness-critical control
logic; replay re-runs the current step thunk (CPU models the lifecycle, not the
GPU speedup).

- GraphCapturer on the instance (cross-request cache), wired into
  RuntimeLoopContext.execute and gated by LoopSpec.graph_capture ==
  "breakable_cudagraph". Capture key = (device, arch, loop, shape_sig,
  resident-weight-versions, graph_key). Eager-break for non-capturable (SDE) and
  interceptor-overridden steps. Never executes a stored thunk, so interleaved
  requests can't smear state.
- WorkPlan.capturable / graph_key. wan21 sets capturable=not sde and folds
  compute dtype (shape_sig.dtype) + CFG branch set + expert + scheduler-precision
  into the key — closing a cross-precision key-collision corruption path a
  non-fp32 build would otherwise hit on a real GPU (audit finding).
- Version-in-key auto-invalidation + real eviction: set_weights_version evicts
  the synced component's graphs (duck-typed, so card/ imports no runtime),
  preventing a GPU graph leak across FlowGRPO's per-iteration syncs.
- register_kernel gains the workspace_bytes capture-safety contract (declared in
  the matrix; cuda cells declare real scratch, numpy reference is 0).
- 11 tests: capture-once/replay-many, eager-break (SDE + override), capture ≡
  pure-eager bit-identical, recapture on shape + weight-version change with
  eviction, eager-loop gating, accel-backend capture, and capturer unit tests
  (key discrimination, eager-break, workspace-collision safety net, eviction).

Honestly deferred (bite a real GPU, not the CPU tests): the static-buffer
refactor of the step body, admission budgeting of capture cost (GRAPH_CAPTURE),
and per-card opt-in beyond wan21 — all documented in cudagraph.py + README.

195 tests pass.
2026-07-04 17:15:34 +00:00
SolitaryThinker 634f0828a2 [feat] route all diffusion loops + RL recompute through the kernel table
Finish the kernel seam across the board so the platform's KernelTable is the
universal solver-dispatch path, not just wan21.

- Loops: ltx2, wan_causal, adapters, adaptive now resolve the flow-match solver
  via model.platform.kernels.get(FLOW_MATCH_STEP) instead of importing the numpy
  sampler directly (wan21 already did). On CPU this is bit-identical (the cpu
  kernel IS the old function); a GPU/accel backend now overrides every loop.
- RL: the FlowGRPO log-prob recompute in unified_rl / joint_multi_rl /
  workflow_rl dispatches FLOW_SDE_STEP through the platform, pinned to the SAME
  kernel the rollout used (C2 kernel-pinning — otherwise the PPO ratio biases on
  a real GPU where rollout and recompute kernels could differ).
- accel backend: add a kind-generic AccelComponent wrapper and override the vae
  component too (text_encoder left unregistered to keep the device->cpu fallback
  demonstrated), closing the "only dit is overridden" gap.
- tests: vae override assertion; a second-loop (wan-causal chunk rollout) parity
  oracle proving accel == cpu bit-identical beyond wan21. README scope updated.

184 tests pass.
2026-07-04 17:15:34 +00:00
SolitaryThinker 2f044c02dd [feat] multi-backend dispatch substrate (device/arch/kernel registries)
Add v2/platform/: the (recipe, runtime) backend membrane that lets CPU, GPU,
and other devices coexist behind one dispatch substrate.

- Two tuple-keyed registries: COMPONENTS(kind, device, variant) for
  weight-bearing components and KERNELS(op, device, arch, variant) for
  stateless primitives, each with an availability predicate and an enumerable
  manifest (declared-but-unavailable cells listed without importing torch).
- Platform: detected (device, arch) owning the device + arch fallback chains
  and a per-platform cached KernelTable; detect() is honest (CPU/numpy unless
  torch+CUDA are actually present). Arch fallback is monotonic — only degrades
  to older/portable archs, never a newer binary-incompatible one.
- Seams wired: ModelInstance.component() -> platform.build_component(spec,self)
  with spec.factory as the numpy terminal rung (existing cards untouched); the
  wan21 denoise thunk dispatches solver ops through model.platform.kernels.
- Three backends: cpu (numpy terminal + parity oracle), accel (pure-python
  stand-in proving cross-device resolution, arch fallback, and the oracle),
  torch_cuda (declared-but-unavailable; no faked GPU).
- test_platform.py: 16 tests — detection, terminal rung, arch-fallback walk +
  monotonicity, device precedence, variant fallback, component override +
  per-kind device fallback, the parity oracle (accel == cpu, bit_identical via
  the C1 ladder), and matrix enumeration without importing torch.

Scope is honestly bounded in the README/docstrings: only wan21-denoise routes
through the kernel table and only the dit kind is overridden today (each a
one-line adoption); the torch/CUDA path and a cudagraph workspace-safety
contract are declared/deferred, not implemented. 183 tests pass.
2026-07-04 17:15:34 +00:00
SolitaryThinker 431f4daddb [feat] Adapter plane, non-linear workflows, RL→distill flywheel
Three more capabilities on distinct untested surfaces (167 tests pass, no new runtime primitive).

A — Adapter plane (§9.19): one base + swappable LoRA/ControlNet adapters, selected per request
(DiffusionParams.adapters); AdapterDenoiseLoop applies each active adapter's velocity delta. Per-request
selection changes output, multi-LoRA composes, ControlNet conditions on a control image, mixed-adapter
requests interleave without smearing, hot-swap changes generation, cache key partitions by adapter stack
(the adapter_versions field, previously declared-only). ToyLoRA/ToyControlNet; models/adapters/. 6 tests.

B — Non-linear workflows (§9.17): ParallelWorkflow (fan-out: one input → N models → merged) and
BestOfNWorkflow (generate N → score with the served reward card → return best; inference-time scaling).
The shapes a linear chain can't express. 4 tests.

D — RL→distill flywheel (§9.18): run_flywheel RL-improves the base (NFT), then distills FROM the RL'd model
(DMD2 teacher = RL'd policy) into a faster card, recording the base→rl→distilled provenance chain in
RecipeSpec.parents. The distilled student is measurably closer to the RL'd teacher than the base; the
distilled card serves few-step. training/flywheel.py. 4 tests.

designv4 §9.17–§9.19 + layout/closing/counts (167 tests, 29 files).
2026-07-04 17:15:34 +00:00
SolitaryThinker 6220d02746 [feat] #8b speculative (draft-verify) decoding — exact + lower-latency AR
The last audit stress test. A cheap draft model proposes K tokens, the target verifies
them in one batched step, and SpeculativeARLoop accepts the matching prefix + one target
correction — a variable accepted-length per round (a ragged AR loop the model owns).

- Exactness: the emitted sequence equals the target's OWN greedy decode for any draft
  quality (every accepted token is one the target would produce; the correction is the
  target's token) — the speedup is free.
- Speedup scales with accept rate: draft-agree 0.3→1x, 0.7→3x, 1.0→4x=K tokens/round
  (fewer verify_rounds, the expensive model's latency steps, for the same output).
- Two components (draft + target) co-scheduled on one resident instance; each round an
  AR_TOKEN WorkUnit.

models/speculative/ (loop+card+program); backend ToyTargetModel/ToyDraftModel with a
shared length-dependent target formula (no degenerate fixed point). 5 tests; full suite
153 passed. designv4 §9.16 + counts (153 tests, 26 files).
2026-07-04 17:15:34 +00:00
SolitaryThinker 9489c6c1dd [feat] Five more stress tests: LTX-2 A/V, weight-sync, served reward, cache-dit, nested workflows
The remaining design_v3 probes (all except 8b speculative decoding). All fit with no new runtime
primitive (148 tests pass).

#6 LTX-2 joint audio+video (§9.11) — LTX-2 declared an audio_vae required_for t2vs but never used
   it; now a single 2-stage denoise carries a synchronized audio latent (conditioned on video),
   applies per-modality CFG (guidance_per_modality), and decodes via video VAE + audio VAE → video +
   audio. Gated on requesting audio, so the T2V path is byte-identical (existing tests untouched).
   ToyAudioVAE; build_ltx2_av_program. 5 tests.

#4 Live weight-sync under in-flight serving (§9.14) — WeightSyncController makes the freeze → drain →
   transfer → bump version + invalidate → resume lifecycle explicit. Tests: a mid-flight swap corrupts
   (the hazard); draining first leaves the in-flight request bit-identical to baseline while a
   post-sync request reflects new weights; transformer-only sync, so the frozen text-encoder cache
   survives. The RL flywheel's hardest correctness. 3 tests.

#5 Reward-model-as-a-served-card (§9.15) — a reward model is a card (scorer + a score loop emitting
   REWARD_BATCH units); ServedRewardScorer drop-in-replaces the numpy scorer so any RL method becomes
   RLHF/RLAIF with no method change. ToyRewardModel; models/reward/. 4 tests.

#7 Content-adaptive control flow (§9.12) — CacheDiTDenoiseLoop (isolated WanDenoiseLoop subclass)
   reuses the cached velocity when predictions barely change (cache-dit skip) and early-exits on
   convergence — variable step count; interleave parity holds across ragged loops. models/adaptive/. 4 tests.

#8a Nested workflows (§9.13) — a workflow stage can invoke another workflow (engine.run routes ids);
   requires/validate recurse; cycles caught at registration + a run-time guard (engine._wf_running).
   build_t2i_i2v_extend_workflow. 5 tests.

designv4 §9.11–§9.15 + falsifier/layout/closing updates (148 tests, 25 files). Also removed a
pre-existing unused import in ltx2/loop.py.
2026-07-04 17:15:34 +00:00
SolitaryThinker 69c9871154 [examples] Add v2_examples/{training,omni,workflows}/ — runnable examples
Three more example folders alongside inference/, all CPU/numpy, self-contained
(sys.path bootstrap), public API only, every script verified to run green.

training/ (7) — one per method, all via the uniform method.train_step seam:
  01 finetune · 02 dmd2 distillation · 03 diffusion_nft (likelihood-free RL,
  samples from the old policy, feature-cache reuse) · 04 self-forcing (causal
  chunk_rollout) · 05 joint LM+generator RL (UniRL; joint + prompt-only) ·
  06 N-way joint RL (per_expert vs shared credit) · 07 end-to-end workflow RL
  (T2I+I2V from one final-video reward).

omni/ (4) — 01 Cosmos3 (reason→joint denoise, shared MoT) · 02 BAGEL
  (text→image, shared MoT; scheduler prices both WorkUnit kinds) · 03 Qwen-Omni
  (thinker→talker→vocoder, three separate experts, text+audio) · 04 interleave
  parity across AR + diffusion loop types.

workflows/ (2) — 01 cross-model T2I→I2V workflow (image provably conditions the
  video) · 02 workflow as a first-class servable (requires/validate, address by
  id, register_workflows catalog, WorkflowRegistry).

Each folder has a README indexing its scripts.
2026-07-04 17:15:34 +00:00
SolitaryThinker 18dd295e8d [examples] Add v2_examples/inference/ — runnable Wan2.1 inference examples
Five self-contained, runnable scripts (CPU/numpy) for the Wan2.1-1.3B card on the
v2 runtime, each bootstrapping the repo onto sys.path so they run from anywhere:

- 01_basic_t2v.py                  minimal path: build engine → T2V request → run → video
- 02_params_and_reproducibility.py DiffusionParams knobs + seeded bit-identical reproducibility
- 03_streaming.py                  per-denoise-step preview chunks (OutputSpec stream)
- 04_concurrent_interleaved.py     step-interleaved batching + interleave parity gate + cache reuse
- 05_async_serving.py              AsyncEngine: concurrent generate, event stream, step-boundary cancel

+ README.md indexing them. All five run green; use only the public API.
2026-07-04 17:15:34 +00:00
SolitaryThinker 4c333e0509 [docs] Update v2/README to designv4 + current scope (127 tests)
- v2/README.md: point to designv4.md as the unified design (design_v3 as
  north star); refresh the scope table (joint/N-way/workflow RL, Qwen-Omni
  cascade, cross-model Workflow, tiled VAE co-scheduling, WorldModelSession),
  package layout (program/Workflow, runtime/session, the 7 methods, new model
  dirs), the demonstrated-stress-tests list (§9.3–§9.10), and counts (49→127,
  20 files). Sessions moved out of "out of scope"; WebRTC wire stays out.
- designv4.md: drop two intermediate absolute suite totals (milestone "91/97
  passed") in favor of "zero regressions" so the only absolute count is the
  current 127 (intro + layout) — no stale numbers.
2026-07-04 17:15:34 +00:00
SolitaryThinker 32a7a6b87b [feat] Three stress tests: interactive sessions, workflow RL, heterogeneous co-scheduling
Targets the three design_v3 claims that were most load-bearing AND least
exercised (sessions/realtime, training-plane boundary, the WorkUnit-generality
falsifier). All fit with no new runtime primitive (127 tests pass).

1. Interactive world-model session (runtime/session.py, §9.8) — the Session
   plane had ZERO coverage. WorldModelSession drives the causal chunk_rollout
   loop as a long-lived session: persistent cross-request world state on the
   Session.kv_handle, frame streaming, transactional step-boundary cancellation
   (a cancelled act leaves the world resumable), no cross-session smearing.
   Added only a continuation seam to the chunk loop (init seeds context from a
   world_context slot; default empty = unchanged one-shot path). 5 tests.

2. End-to-end RL over a cross-model workflow (training/methods/workflow_rl.py,
   §9.9) — trains BOTH flux-t2i and wan-i2v from ONE final-video reward. Rolls
   out the whole workflow with SDE capture in both instances; the same final
   advantage drives FlowGRPO PPO on each stage's transformer; two WeightSyncPlans
   on two instances. The earlier model (T2I) is trained by a reward on the final
   video — end-to-end credit across a model boundary — proven causal by a control
   (constant reward => zero advantage => nothing moves). 4 tests.

3. Heterogeneous WorkUnit co-scheduling (models/tiled/, §9.10) — the §17
   falsifier. VAETileLoop makes VAE decode a loop of VAE_TILE units; tiling is
   exact (== one-shot, C0), and VAE_TILE + DIFFUSION_STEP pipelines interleave
   bit-identically and co-run in one batch. Validates the mechanism; the
   economic half (does it pay) stays a port-time measurement. 4 tests.

designv4 §9.8–§9.10 + falsifier/layout/closing updates (127 tests, 20 files).
2026-07-04 17:15:34 +00:00
SolitaryThinker 58223c0c41 [feat] Register cross-model workflows as first-class named servables
Answers "what's the right way to name/register custom pipelines like T2I→I2V":
treat a Workflow like a card — a stable namespaced id in the same servable
namespace, declared dependencies, and a two-level registry. No new concepts;
mirrors how cards are registered (and vllm-omni's pipeline_registry).

- program/workflow.py: Workflow gains `requires` (the cards it composes, derived
  from stages) and `validate(engine)` (fail-fast if a required card is absent,
  P7). New WorkflowRegistry: declarative workflow_id -> builder catalog for
  out-of-tree/ad hoc use.
- runtime/engine.py: `_workflows` registry + register_workflow (validates deps,
  rejects id collision with a model_id) + serves(); engine.run routes a request
  whose model_id is a workflow to workflow.run — addressable exactly like a model.
  Single-model hot path untouched.
- runtime/async_engine.py + serving/server.py: serves() and /models include
  workflows (discoverable as servables).
- models/__init__.py: declarative `_WORKFLOWS` catalog (the cross-model analog of
  _BUILDERS) + register_workflows() helper; build_image_video_engine now registers
  the workflow too. Adding a custom pipeline = one catalog line.
- Naming convention: dotted/namespaced workflow_id (`image_video.t2i_i2v`),
  distinct from kebab model ids, collision-checked. Renamed from `t2i_then_i2v`.
- tests (+5, 12 total in the file): addressable by id, requires/validate, id
  collision, registry catalog, register_workflows helper. Full suite 114 passed.
- designv4 §9.6: the naming & registration convention documented.
2026-07-04 17:15:34 +00:00
SolitaryThinker b254d1affe [feat] Cross-model T2I→I2V workflow + N-way joint RL over arbitrary experts
Two more pipelines stress-testing the design, plus a BAGEL-placement note in
designv4. Both fit with no new runtime primitive (109 tests pass).

Pipeline 1 — cross-model T2I→I2V (program/workflow.py, models/image_video/):
- Realizes ProgramKind.WORKFLOW as a thin multi-instance orchestrator ABOVE the
  engine (the hot path stays single-instance). A Program composes one model's
  loops; a Workflow chains full engine.run calls across distinct cards, threading
  artifacts. (LTX-2 already covers same-card multi-stage; cross-model — FLUX→Wan
  — is the new capability the single-instance runner can't express.)
- flux-t2i (text→image) and wan-i2v (text+image→video) cards; the I2V program
  folds the conditioning image into text_embeds so WanDenoiseLoop is unchanged.
- Each model keeps its own interleave-parity guarantee (crossing instances is a
  Workflow boundary, not a loop step). 7 tests incl. video-depends-on-image.

Pipeline 2 — N-way joint RL (training/methods/joint_multi_rl.py, models/multi_expert/):
- JointMultiExpertRL generalizes UnifiedRLMethod (N=2) to N refiner LMs + a
  generator: one reward → one group advantage → N token-PG updates + 1 FlowGRPO
  PPO update, N+1 independent WeightSyncPlans. Proves the substrate was already
  N-ready (card holds N components/loops; per-component weight-sync; dict grad
  targets) — only the method body looped over two; now it loops over a list.
- credit="per_expert" learns all N cleanly; credit="shared" (faithful to UniRL)
  works but is noisier — the honest multi-agent credit-assignment result, a
  reward-shaping choice, not a substrate limit. 6 tests (N=1,3,4; prompt-only).

Fix — flow_sde_ml_velocity (loop/sampler.py): the toy FlowGRPO generator update
targeted the velocity the model already produced (a no-op once guidance_scale=1
was set for the C2 identity; the unified generator moved only on ~1e-7 noise).
The correct PG surrogate targets the max-likelihood velocity of the realized
sample — nonzero at ratio==1. Both UniRL and N-way generators now learn for real;
the C2 ratio==1 identity still holds (measured before the update).

BAGEL: MoT/shared-weight (one transformer on both loops), same row as Cosmos3;
real BAGEL's co-resident experts are expressible via the expert-routing policy
(partial sharing) — captured in designv4 §2.3.
2026-07-04 17:15:34 +00:00
SolitaryThinker c7e0a8e894 [feat] Add Qwen-Omni thinker→talker→vocoder model (3 experts, 3 loops)
Ports vllm-omni's canonical qwen2_5_omni omni-speech cascade as a v2 card:
a third weight-sharing topology — three disjoint experts (thinker, talker,
vocoder) on three loop types (ar_decode → ar_decode → audio_decode) in one
request, with chained cross-stage conditioning and streaming codec→waveform.
vllm-omni runs these as three opaque request-scheduled stages; v2 makes every
thinker token, talker token, and vocoder chunk a runtime-visible WorkUnit.

- models/backend.py: ToyTalker (a genuinely distinct AR expert, weight-salted)
  + ToyVocoder (streaming code2wav: codec tokens → waveform chunks).
- models/omni/vocoder_loop.py: VocoderLoop filling the pre-declared
  LoopKind.AUDIO_DECODE / WorkUnitKind.AUDIO_CHUNK slot.
- models/omni/ar_loop.py: ARDecodeLoop gains a configurable prompt_slot so two
  chained AR loops don't collide on the prefill slot (thinker vs talker).
- models/qwen_omni/: card (3 experts/3 loops) + program (tokenize → thinker →
  emit_text → thinker→talker full-payload hand-off → talker → talker→vocoder →
  vocoder → emit_audio). Cross-stage hand-offs are explicit Program nodes, the
  model-native form of vllm-omni's custom_process_input_func.
- _enums.py: Capability.TEXT_TO_SPEECH.
- tests/test_thinker_talker.py: 6 tests incl. three-loop interleave parity,
  cascade conditioning, AUDIO_CHUNK streaming. Full suite 97 passed.

designv4.md: §2.3 topology table extended to four topologies; new §9.5 on the
cascade; reference-synthesis + package layout updated.
2026-07-04 17:15:34 +00:00
SolitaryThinker b3ddf6014d [feat] UniRL/PromptRL joint LM+generator RL stress test + designv4
Stress-tests the v2 Card/Loop/Program design with a UniRL/PromptRL-style
joint RL recipe: a prompt-refiner LM expert and a flow generator expert,
two separate experts driven by two loop types in one request, both updated
simultaneously from a single RL reward.

- loop/sampler.py: flow_sde_step_with_logprob — FlowGRPO SDE rollout sampler
  (per-step Gaussian log-prob), distinct from the deterministic ODE serve step.
- request/params.py: gated sde_rollout/sde_noise_scale on DiffusionParams so
  the serve path stays byte-identical (default ODE).
- models/wan21/loop.py: gated SDE-rollout capture in WanDenoiseLoop.advance.
- models/backend.py: ToyPromptRefiner — a real REINFORCE categorical policy
  (the Qwen role), separate weights from the generator.
- models/unified/: the unified card+program — two disjoint experts (llm +
  transformer) on ar_decode + diffusion_denoise; the topological opposite of
  the Cosmos3 MoT card, same vocabulary.
- training/methods/unified_rl.py: joint GRPO — one reward -> group advantage
  -> LM token policy gradient + DiT FlowGRPO PPO; two LRs; prompt-only/joint
  flag; reuses the shared diffusion loop for rollout.
- training/weight_sync.py: WeightSyncPlan gains a component scope so the two
  experts version + cache-invalidate independently (LM sync never flushes the
  frozen text-encoder feature cache).
- tests/test_unified_rl.py: 9 tests incl. likelihood-based C2 identity, the
  two-loop interleave parity gate, joint vs prompt-only. Full suite 91 passed.

designv4.md: unified design doc reflecting v2 as built+tested, with the joint
RL stress test as the validating case study (the design held — new card +
new method, no new runtime primitive).
2026-07-04 17:15:34 +00:00
SolitaryThinker a9e5f6ee7a [refactor] rename package mini_fastvideo → v2
Directory rename (git mv, history preserved) plus rewrite of all references: absolute imports in
tests, the zero-dep runner, docstrings, comments, and the README. No behavior change.

Run: python3 -m pytest v2/tests/ -q ; python3 v2/run_tests.py ; python3 -m v2.examples
2026-07-04 17:15:34 +00:00
SolitaryThinker 7467076d72 [fix] mini-fastvideo serving: address adversarial-review findings (capacity/credit leaks, robustness)
Review confirmed the core bets (concurrent disaggregation is bit-identical, design conformance holds,
Dynamo genuinely optional, engine stays step-scheduled). Fixes for the untested failure paths:

- HIGH: pool capacity (RolePool.in_flight) no longer leaks when a disaggregated request is cancelled
  or errors mid-occupancy — DisaggregatedRunner.close() releases the occupied pool and AsyncEngine._run
  calls it in a finally (a cancel on a capacity-1 denoiser no longer bricks the pool).
- credit flow-control: cross-pool transfer wraps acquire/release in try/finally (no credit leak on a
  failing transfer); slot is re-homed only on a successful fetch.
- AsyncEngine: duplicate in-flight request_id is rejected (was a deadlock); submit() cancels the driver
  task when the consumer abandons the stream (client disconnect → no orphaned compute); bounded
  per-request history (no unbounded _events/_states/_results/_runners growth).
- cancellation is common-path on the OFFLINE path too (cancel check at the top of every runner.tick()).
- HTTP server: read timeout (slowloris guard → 408), body-size cap (→ 413), invalid Content-Length
  (→ 400), explicit StreamReader limit, and aclose() of the SSE generator on client disconnect.
- build_deployment_card no longer aliases one mutable CostModel across replica cards (dataclasses.replace),
  so online calibration of one worker's cost doesn't mutate another's.
- video-job tasks tracked (not fire-and-forget); server.close() cancels/drains them; jobs dict bounded;
  fleet affinity map bounded.

6 regression tests added for these paths. 82 tests pass (pytest + zero-dep runner).
2026-07-04 17:15:33 +00:00
SolitaryThinker 01dc0c3377 [feat] mini-fastvideo serving + fleet (our own version, Dynamo-optional)
Builds the full serving layer the design files specify, instead of deferring it to Dynamo:

- transport/ (§7.3): pluggable Connectors (in-proc zero-copy / SHM-fake copy) with chunk_ready
  readiness (vllm-omni) AND credit-based flow control (sglang-omni Relay); KVConnector protocol shape;
  TransferManifest.
- runtime/ (§6, §13; plan M3/M4): AsyncEngine — request queue, lifecycle state machine
  (waiting→running→completed/cancelled/failed), live AsyncIterator[OmniEvent] streaming, common-path
  cancellation, step-level concurrency. RolePool + DisaggregatedRunner (encoder→denoiser→decoder,
  capacity-aware dispatch, cross-pool transfers via connectors); disaggregated output is bit-identical
  to inline. No-progress detection (no busy-spin).
- deploy/ (§14, §6.3.5-6): DeploymentCard; OUR OWN LocalFleet (discovery, health/drain, least-loaded
  / cost-model / sticky-affinity routing) so we never rely on Dynamo; DynamoWorkerAdapter +
  FakeDynamoRuntime export the SAME card + cost model so Dynamo CAN front us — one object, two consumers.
- serving/ (§6.3.5, §12): framework-free stdlib-asyncio OpenAI server (our own version of the
  vllm-omni pattern): /v1/chat/completions (SSE), /v1/images/generations, /v1/videos (async job+poll)
  + /v1/videos/sync, /v1/models, /health, /metrics. A thin shim over the STEP-scheduled engine — the
  runtime-visible loop scheduler vllm-omni's request-scheduled opaque DIFFUSION stage lacks.

15 serving tests (real-socket HTTP+SSE via stdlib asyncio, disagg==inline, fleet routing, Dynamo
contract, cancellation); 76 tests pass total (pytest + zero-dep runner). ~7900 LOC.
2026-07-04 17:15:33 +00:00
SolitaryThinker f1dc587c74 [feat] mini-fastvideo phase 2: omni/MoT — Cosmos3 + canonical vllm-omni (BAGEL/lance)
One resident MoT instance runs BOTH an ar_decode loop and a diffusion_denoise loop on shared weights
(the §16 claim no DAG-of-engines can express), with both loops runtime-visible: the scheduler prices
ar_token AND diffusion_step WorkUnits — unlike vllm-omni's opaque DIFFUSION stage the scheduler never
sees inside.

- ARDecodeLoop: token decode until EOS/max_tokens, paged text-KV — the omni AR pathway (loop/§5).
- ToyMoTDiT: one module exposing an und pathway (ar_forward) AND a gen pathway (denoise __call__);
  ToyTokenizer. Binding both loops to one instance = shared weights, no duplication.
- models/cosmos3/: tokenize → reason(ar_decode) → pack(tokens→conditioning) → diffusion_denoise →
  vae_decode; sound_vae declared optional_for non-t2vs (the lazy-component P8 fix, not an env-var hack).
- models/bagel/: the canonical vllm-omni model — generate_text(ar_decode) → generate_image(diffusion),
  text+image outputs, both loops step-scheduled.
- The diffusion loop is WanDenoiseLoop reused (one loop definition bound to the MoT module).
- build_omni_engine() + an omni worked example; 7 omni tests (shared-instance, both-kinds-scheduled,
  interleave parity across loop types, lazy sound_vae). 61 tests pass (pytest + zero-dep runner).
2026-07-04 17:15:33 +00:00
SolitaryThinker 098bcf014a [fix] mini-fastvideo: address adversarial-review findings (admission liveness + §7.1 cache key)
- Admission fails fast with AdmissionInfeasible on infeasible/deadlocked reservations instead of a
  10M-iteration busy-spin: no-progress detection in run_to_completion/run_interleaved via a real
  progress token, plus feasibility pre-checks (need > pool capacity).
- Compute budget is now a refundable concurrency gate (release() refunds spent), not a
  never-refunded lifetime cap that silently deadlocks.
- §7.1: the text-encoder feature CacheKey carries adapter_versions + precision (no stale serve across
  te-LoRA stacks); per-component weight versions mean a transformer-only RL weight sync no longer
  flushes the frozen text-encoder cache (component-scoped invalidation, not wholesale).
- Interleave gate flags symmetric-empty output instead of passing it vacuously.
- skipped_steps counted only when the override is actually consumed; BatchScheduler wired for round
  batch-accounting (metric renamed stepped_units); ResidualCache.get cleanup; stream chunks carry a
  latent preview payload; dead progress-vars removed.
- 5 regression tests added for the previously-untested paths. 54 tests pass (pytest + zero-dep runner).
2026-07-04 17:15:33 +00:00
SolitaryThinker 270fae959d [feat] mini-fastvideo: model-native runtime per design_v3 (Wan2.1/LTX2 + 4 training methods)
A scoped, CPU-testable realization of design_v3.md — the architecture where the atomic unit
is a typed (recipe, runtime) ModelCard, the model owns loop semantics while the runtime owns
loop lifecycle, one resident instance runs many loops, and training records behavior on the
same loops it serves.

Implements:
- card/ loop/ runtime/ cache/ memory/ parallel/ parity/ extend/ program/ request/ training/
  spanning design_v3 §4-§13: ModelCard + validate(); driven loops (init/next/advance/finalize);
  step-interleaving Engine with reservation-before-admission + per-class caches keyed by CacheKey;
  the C0-C4 consistency ladder + the non-negotiable batch-of-N interleave parity gate.
- Inference: Wan2.1-1.3B (T2V), LTX2.3 (two-stage distilled, shared transformer), Wan-causal
  (chunk rollout + slab-KV streaming).
- Training (Wan2.1-1.3B), each driving the SAME loops the engine serves: finetune (flow-match),
  DMD2 (teacher/critic distribution matching), DiffusionNFT (likelihood-free C2, samples from the
  decay-blended old policy, group-relative advantages, shared-prompt cache reuse), self-forcing
  (causal chunk loop). The engine never imports training (the §10 dependency rule, grep-verified).

numpy-only core (no torch/GPU here); heavy Wan/LTX forwards are deterministic toy stand-ins with
lazy torch-adapter seams (ComponentSpec.load_id/factory) for a GPU box. 49 tests pass via pytest
and a zero-dependency runner; interleave parity verified (and a buggy module-global interceptor
provably breaks it). Omni-ready spine (ar_decode/chunk_step loop kinds, multi-loop instances,
LoopState.extension) for the phase-2 Cosmos3 + vllm-omni omni ports.

Run: python3 -m pytest mini_fastvideo/tests/ -q ; python3 -m mini_fastvideo.examples
2026-07-04 17:15:33 +00:00
SolitaryThinker 4a14c1afa3 update 2026-07-04 17:15:33 +00:00
SolitaryThinker f1c19050c3 design 2026-07-04 17:15:33 +00:00
375 changed files with 71279 additions and 2 deletions
+207
View File
@@ -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).
+2
View File
@@ -10,6 +10,8 @@ exclude: |
scripts/.*|
fastvideo/dataset/.*|
fastvideo/models/.*|
v2/(layers|attention|platforms|configs|distributed|models|logging_utils|third_party|hooks|api)/.*|
v2/(envs|logger|utils|version|forward_context|fastvideo_args)\.py|
^apps/dreamverse/web/.*|
examples/.*|
\.agents/.*|
+56
View File
@@ -0,0 +1,56 @@
# FastVideo — Design Philosophy
One page on *why* FastVideo is built the way it is. The full architecture, the as-built status, and the
forward roadmap live in **[`v2/README.md`](v2/README.md)** — this is the philosophy beneath it.
---
**A deployable model is a post-training artifact.** Unlike an LLM — where inference optimizes frozen weights
after the fact — a *usable* video/omni model is *created* by training: step distillation for latency, QAT for
precision, distillation + self-forcing for causal/world models. So every inference capability is a
**(recipe, runtime) pair**: the weights and the loop that produced-and-assumes them are one versioned object.
This is the source of the moat — whoever owns *both* sides of the pair owns the optimization frontier — and it
is why training and serving cannot be two systems.
**The work is loops, not `forward()`.** Denoise timesteps, AR decode, chunked rollout, VAE tiles, encoder
chunks, audio tokens, reward batches, optimizer steps, media chunks — video and omni inference is iteration. A
runtime that collapses everything to a single `forward` can't schedule, batch, cancel, stream, reserve memory
for, or capture the behavior of what actually runs. So loops are first-class, and they are **driven**: the
model describes the next step it needs, the runtime decides when and with whom it runs, the model folds the
result back. The model keeps content-adaptive control flow; the runtime keeps admission, batching, streaming,
and behavior capture. Per-request state lives in typed `LoopState`, never in module globals — so interleaving
requests through one model instance cannot smear state, by construction.
**The model is the center; everything else is a view over it.** A typed `ModelCard` owns components, loops,
the recipe, and the parity contract. Programs compose a card's loops into a task; Workflows compose cards into
pipelines; the scheduler runs the *steps* of all loops as `WorkUnit`s under one currency (predicted GPU-time,
because a bidirectional denoise step and an AR token are ~1000× apart and incommensurable in counts);
deployment places and routes; products stream artifacts. None of them define model semantics — they reference
the Model Plane. One resident instance can run many loop types on shared weights, which is what makes omni/MoT
native rather than a DAG that doubles weights.
**Correctness is a typed contract, not a hope.** Caches are correct by *key* — if a field can change output
semantics it is in the key, so reuse is partitioned, never blindly flushed. Parity between the train-forward
and the serve-forward is *measured* on a declared ladder (component → loop → behavioral → distribution →
artifact-quality), never assumed. And the non-negotiable gate is **interleave bit-parity**: N requests
interleaved at step granularity must be bit-identical to running them serially — the test the whole
loop-inversion bet lives or dies on.
**One substrate for inference, training, and RL.** The rollout forward *is* the serve forward plus capture —
same loop, same caches, same batcher, same numerics — so every serving optimization is automatically a rollout
optimization, and there is one numerics surface the ladder measures rather than a correction layer papering
over it. The engine doubles as the RL rollout engine under a strict rule: `training` consumes the engine; the
**engine never imports `training`**.
**Borrow aggressively; copy nothing as the core.** vLLM/SGLang scheduling, vLLM-Omni/SGLang-Omni omni serving,
Dynamo fleet orchestration, diffusers components, xDiT parallelism, TorchTitan mesh discipline,
verl-omni/miles RL lessons, ComfyUI workflows, Dreamverse/LiveKit sessions — each contributes a take, none is
the center. Deployment orchestration (Dynamo) sits *above* the engine, never inside it. Extensions are
versioned hook points, never monkeypatching. New frontier capabilities arrive as a card, a method, a loop, a
workflow, or a controller — **not a rewrite**.
> A model card is a (recipe, runtime) pair with a parity obligation. The model owns loop semantics; the runtime
> owns loop lifecycle. One resident instance runs many loops; one scheduler runs their steps in one currency.
> Caches are correct by key; parity is correct by test; the interleave gate is non-negotiable. Training records
> behavior on the same loops it serves. Deployment places and routes; products stream artifacts; neither defines
> the model.
@@ -0,0 +1,94 @@
# v2 porting status — fastvideo models → the v2 (recipe, runtime) substrate
Goal: every model in fastvideo's registry resolves through the **v2 `VideoGenerator`** / `Engine`
(typed `fastvideo.api` configs + the real torch backend) to a recipe that can construct and run it.
**Scope: ALL fastvideo models (achieved).** v2 now resolves **63/64** of fastvideo's registered HF ids
by exact id (PRIMARY), plus the architecture fallback for local/unregistered checkpoints. The single
remaining id — `FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers` — is **environment-blocked**: its VSA
(Sparse-Linear Attention) kernels require `nvcc` (not built in this bring-up). It arch-resolves to the
base Wan card but needs the VSA kernel build to run faithfully.
Dispatch is **architecture-driven** (`v2/registry.py`): exact HF id → short-name → architecture
inference from the checkpoint (pipeline / transformer / VAE class names + `z_dim`, `transformer_2`,
`spatial_upsampler`). Adding a model is one `_BUCKET_C` row (HF ids → builders + transformer class).
## The porting mechanism — self-contained recipe packages
Every net-new arch is a **self-contained recipe package** (`v2/recipes/<arch>/` = `card.py` `loop.py`
`program.py` [+ `sampler.py`] + an optional `v2/platform/backends/torch_<arch>.py` adapter). The card
declares its torch adapter via **`ComponentSpec.adapter="module:Class"`** (the `_explicit_adapter` seam in
`torch_backend.py`) instead of editing the shared `_make_dit`/`_make_vae`/`_make_text_encoder` dispatch —
so a port adds **only new files**, never touching shared code, and parallel ports never conflict. New
samplers/loops live in-package. Registration is one row in `v2/registry.py:_BUCKET_C`.
## Working today (GPU-verified, real video/audio) — committed on `v2`
| Official example(s) | Model | v2 card |
|---|---|---|
| `basic.py`, `basic_mps.py`, `basic_ray.py` | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | wan21 |
| `basic_self_forcing_causal.py` | `wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers` | wan_causal |
| `basic_ltx2_distilled.py` | `FastVideo/LTX2-Distilled-Diffusers` (2-stage + spatial upsampler) | ltx2 |
| `basic_wan2_2_ti2v.py` | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | wan2.2-ti2v |
| `basic_wan2_2.py` | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` (MoE + CPU expert offload) | wan2.2-a14b |
| `basic_ltx2.py` | `Davids048/LTX2-Base-Diffusers` | ltx2 base |
| `basic_ltx2_3_distilled.py` | `FastVideo/LTX-2.3-Distilled-Diffusers` (joint T2VS, video+audio) | ltx2.3-distilled |
Plus the **Wan2.1 i2v cluster** (Fun-1.3B-InP GPU-verified; I2V-14B-480P/720P + Wan2.2-I2V-A14B MoE reuse
the i2v card) — CLIP image-encoder + first-frame `[mask|cond]` → 36ch DiT.
## GPU bring-up results (real weights on H100 NVL, single-GPU, TORCH_SDPA)
**20 models generate real video/audio on GPU** — the 7 above + **13 of the newly-ported** archs, each run
end-to-end through the real `VideoGenerator` (resolve → stamp → CUDA load → generate). The rest are blocked
by a **fastvideo-shared-code / missing-kernel / HF-access** wall, NOT a v2 recipe bug (the v2 recipes are
faithful — e.g. cosmos25's DiT+VAE produced finite output; only its Qwen2.5-VL encoder hit a library
incompat). All ports also resolve + run end-to-end on the CPU toy backend (`test_bucket_c_ports.py`).
| GPU status | Models |
|---|---|
| ✅ **Verified** (real GPU output) | stable_audio (audio), matrixgame2, matrixgame3, gen3c, wan_fun_control, lucy_edit, hunyuangamecraft, hunyuan_video, hunyuan_video15, longcat (13.58B), sfwan22 (2×14B MoE, expert offload), lingbotworld (2×14B, offload), fastwan (TI2V-5B-FullAttn DMD) |
| 🚫 fastvideo/env-blocked | **cosmos25** (DiT+VAE ran; Qwen2.5-VL encoder → transformers 5.12.1 incompat in fastvideo); **kandinsky5** (fastvideo registry registers a bare `PipelineConfig`); **hyworld** (fastvideo DiT hardcodes `flash_attn`, not built); **turbowan** 1.3B/i2v + **fastwan** VSA-variants (SLA/VSA sparse-attn params + Triton kernels need nvcc) |
| 🚫 access-blocked (HF-gated) | cosmos2, flux2, sd35 (no HF token in this env) |
To unblock the env-blocked: build `fastvideo-kernel` (SLA/VSA Triton, needs nvcc); pin a fastvideo-compatible
`transformers` for the Qwen2.5-VL encoder; add a Kandinsky5 `PipelineConfig` + an SDPA fallback in the
hyworld DiT (all fastvideo-side / environment, not v2 recipe work).
## Newly ported (recipe details)
Each resolves through the registry AND runs end-to-end on the CPU toy backend via the public `Engine`
path (the `v2/tests/test_bucket_c_ports.py` regression guard), emitting the correct modality artifact.
**15 net-new architectures** (each a new `TorchComponent` adapter + recipe):
- **cosmos2** (Cosmos-Predict2-2B-Video2World) — EDM-Karras denoiser; new `CosmosDenoiseLoop` +
`build_karras_sigmas` (the reference port). **cosmos25** (Cosmos-Predict2.5 2B/14B) — flow-match,
per-frame plain-sigma timestep, Reason1/Qwen2.5-VL encoder. **gen3c** (GEN3C) — EDM + 82ch pose-buffer.
- **hunyuan_video** (+FastHunyuan) — reuses WanDenoiseLoop, dual LLaMA+CLIP encoders, Hunyuan VAE.
**hunyuan_video15** (480p/720p). **hunyuangamecraft**, **hyworld** — interactive (camera/action).
- **longcat** (T2V/I2V/VC). **kandinsky5** (5.0 T2V Lite).
- **sd35** (MMDiT, image, triple-encoder). **flux2** (dev/klein, MMDiT image). **stable_audio** (audio).
- **lingbotworld** (camera/Plucker), **matrixgame2**, **matrixgame3** — interactive world models.
**5 Wan-family variants** (reuse the Wan/Causal arch, new in-package sampler/loop/conditioning):
- **turbowan** — rCM few-step (faithful RCMScheduler port), 1.3B/14B T2V + I2V-A14B MoE.
- **lucy_edit** — v2v editor (video-VAE-encode node → 96ch DiT input). **wan_fun_control** — control input.
- **sfwan22** — Self-Forcing Wan2.2-A14B causal + MoE (i2v + t2v). **fastwan** — DMD 3-step (TI2V-5B-FullAttn
loadable; VSA-trained variants + non-strict `to_gate_compress` load are BRINGUP).
BRINGUP scope per port (documented in each package): GPU load/run; for interactive/world-model archs the
action/camera/memory conditioning needs a request-API extension (the t2v/degenerate path is what
CPU-verifies); video2world/i2v frame-replace conditioning is threaded but inert without conditioning inputs.
## Environment
v2 bring-up runs **single-GPU, resident, on the `TORCH_SDPA` backend** (no fastvideo-kernel / VSA / FP4).
The box has been rescheduled across hosts/arches/python versions mid-session; rebuild the venv for the
current arch when that happens: `uv venv --python 3.12 .venv`; comment out `fastvideo-kernel` in
`pyproject.toml`; `uv pip install -e ".[dev]"`. Source `/home/scratch.willlin_ent/.bringup_env`
(`HF_HOME=./.cache` on scratch, `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`). v2 CPU mini: 240 passed, 2 skipped.
## How to add a model to the v2 substrate
1. `v2/recipes/<arch>/` — card (declare adapters via `ComponentSpec.adapter`; per-model `SamplingDefaults`),
loop (reuse `WanDenoiseLoop`/`chunk_rollout` or a new in-package loop+sampler), program.
2. `v2/platform/backends/torch_<arch>.py` — a `TorchComponent` subclass (only the forward semantics) if the
arch is genuinely new; reuse `WanDiT`/`LTX2DiT`/`WanVAE`/`T5Encoder` via `load_id` when it isn't.
3. One row in `v2/registry.py:_BUCKET_C` (HF ids → builders; `transformer_cls` for the arch fallback, or
`""` for explicit-id-only capability variants of an existing arch).
4. CPU-verify: it resolves + runs on the toy backend (auto-covered by `test_bucket_c_ports.py`). Then GPU
bring-up (`stamp_*_checkpoints` → real weights) per BRINGUP notes.
+36
View File
@@ -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()
+29
View File
@@ -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()
+225
View File
@@ -0,0 +1,225 @@
# FastVideo Runtime — Aggressive Implementation Plan
**Companion to** `design.md` (v19) and `design_summary.md` · **Stance:** this plan trades interface stability for
speed. Where it deviates from design.md's conservative migration (§10), the deviation is flagged with **⚡**.
design.md remains the architectural authority; this is the execution order.
---
## 1. Rules of engagement
**We break, freely and early:**
- The public Python API: `generate_video(**kwargs)` and `SamplingParam` are **deleted**, not deprecated.
- Config schemas: `FastVideoArgs` (1,272 lines, 81 fields) stops being a public or threaded surface.
- `fastvideo/api/compat.py` (651 lines): **deleted in M1** ⚡ (design.md §6.6 shrinks it monotonically to Phase 5 —
that policy existed only to honor signatures we are now licensed to break).
- CLI flags, YAML schemas, package layout, `fastvideo.api` exports, ComfyUI node params, every example.
- In-repo dependents (`apps/dreamverse`, `comfyui/`, `examples/`, `scripts/`) get **fixed in the same PR train** —
we own them; no deprecation period, no shims.
**We never break, at any speed:**
- **Numerics.** Bit-identical loop parity and SSIM gates are not "conservative" — they are the definition of
correct. Aggression applies to interfaces, never to outputs.
- **Model coverage** for the families that matter (tier list in §6 — the tail is a decision, not a casualty).
- The frozen legacy `fastvideo/training/` stack (N2) and the bit-exact porting methodology (N3).
- External users get **batched breakage**: all user-visible breaks land in at most two releases (R1 = request/config
cut, R2 = engine default), each with a migration guide and a `fastvideo migrate` codemod — never a drip.
## 2. The sequencing argument (answering "fix omni request first, then separate the planes?")
**Yes to the first half. The second half should not be a project.** The three planes are not separated by moving
code into plane-named directories — today's monolithic stages would just get reshuffled and then rewritten when
loops invert. The planes are *born* from two cuts, and a third that is really a config change:
1. **The request-plane cut (M1)** — `OmniRequest` becomes the only currency crossing the boundary. Everything
behind it is implementation. This is your "fix the omni input and request first," and it goes first because it
is low-risk, it defines the vocabulary every later stage consumes, and it gets the user-facing pain over with
while the codebase is still familiar.
2. **Loop inversion (M2)** — this *is* the pipeline/execution plane separation. Once families expose
`init/step/finalize` step bodies, something other than the family must own iteration; that owner is the
executor, and the execution plane exists by construction. Before inversion there is nothing for an execution
plane to schedule — "separating" it would be an empty directory.
3. **The config cut (inside M1)** — the real coupling between planes today is `FastVideoArgs`: one 81-field object
threaded through entrypoints, pipelines, stages, and executors, mixing deploy-time, model-time, and
request-time concerns. Splitting it into `DeployConfig` / `ModelSpec` / `OmniRequest` (design.md §6.6's four
layers) is the single highest-leverage "separation" action, and it's schema work, not architecture work.
So the order is: **M1 request+config cut → M2 loop inversion (planes now exist) → M3 engine on top.** Plane
separation is the *outcome* of M1+M2, not a milestone.
## 3. Milestones
Timeline assumes 3–4 engineers on the runtime critical path. Overlap is deliberate; gates are not ⚡-able.
### M0 — Baselines, harness, enforcement (weeks 0–3, overlaps M1)
The license for everything aggressive afterward. Not skippable, not shrinkable. M0 does not block M1 (which
changes no numerics) — the only hard rule is **no family's M2 migration starts before its baseline exists**.
- Merge the `feat/cosmos3-reasoning` chain (design.md sizes this alone at 2–3 engineer-months — it runs as its own
track); seed SSIM references for the ~7 uncovered families.
- ParityAligner v0: record/compare named taps on *current* pipelines (it must exist before anything changes).
- **The enforcement package, on day one** (design.md §10 — the prior freeze was broken 19× for lack of exactly
this): CI path gates (reject new `fastvideo/training/` files now; reject new `DenoisingStage` subclasses once the
first M2 family lands), CODEOWNERS on the frozen and migrating paths, a named owner per milestone, and the
inflow rule — new model families land on the new abstractions from the first Wan/Flux2 landing onward.
- Announce the M1 freeze window for in-flight PRs touching `fastvideo/api/`, `fastvideo_args.py`, entrypoints.
*Gate: every tier-A family has a recorded SSIM + activation baseline; CI gates live.*
### M1 — The request-plane cut (weeks 1–4) → **breaking release R1**
The typed API is partway there: `VideoGenerator.generate(GenerationRequest)` is already the documented primary
entrypoint (`generate_video` carries a deprecation warning), and `fastvideo/entrypoints/openai/` already serves
`POST /v1/videos` and `POST /v1/images`. But the legacy path is still what's *used*: Dreamverse calls
`generate_video(**kwargs)` (`apps/dreamverse/dreamverse/video_generation.py:508`), as do ComfyUI and most
examples. M1 finishes the cut instead of bridging it:
- **`OmniRequest` / `OmniOutput` / `OmniEvent`**: evolve `api/schema.py`'s `GenerationRequest` in place — typed
modality parts, `TaskType`, per-model `ModelOptions` registered blocks (formalizing the `api/matrixgame2.py`
pattern), seeds/priority/streaming flags. `api/results.py`'s `Video*Event` types become `OmniEvent`.
- **Config: four layers, one owner each** (§6.6): extract `DeployConfig` (placement, parallelism axes, memory/
offload, compile, plugins) from the runtime third of `FastVideoArgs` + `EngineConfig`/`ParallelismConfig`;
`ModelSpec` manifest v0 (manifest-first component resolution; today's name-detectors as fallback);
`OmniRequest` absorbs every per-call field. CLI flags, OpenAI protocol models, and presets are **generated**
from the schema.
- **Delete** ⚡: `compat.py` (651), `sampling_param.py` (411), `generate_video()`, the `fastvideo.api` legacy
exports, `FastVideoArgs` as a *public* type. Internally it survives as a boundary-constructed shim for as long
as anything still receives it: migrated families drop it per-family in M2, but unmigrated tier-B stages
(`LegacyPipelineNode`) and the frozen `training/` stack (whose `TrainingArgs` subclasses it) carry it until M6 —
it dies as a type with the tail, not before.
- **Fix in-train**: ComfyUI nodes (legacy-API callers), all `examples/` (~75 files, mostly mechanical),
`scripts/`, docs. Ship `fastvideo migrate` (codemod: old kwargs/YAML → `OmniRequest`/`DeployConfig`).
- Internals unchanged: `ForwardBatch` is built *from* `OmniRequest` at the boundary; the executor and stages are
untouched in M1.
*Gate: all SSIM suites unchanged; Dreamverse + ComfyUI + examples green on the new surface; R1 notes + codemod
published.*
### M2 — Loop inversion (weeks 4–10): the pipeline plane is born
- `DenoiseLoop` / `ARDecodeLoop` with `init/step/finalize`; runtime owns iteration; custom-step escape hatch from
day one (the Cosmos3-port and self-forcing pattern is legitimate, §6.2.3).
- **Family order** (each lands step body + policies, and **deletes its legacy stage code in the same PR** ⚡ —
continuous deletion, no end-of-plan cliff): **Wan 2.1/2.2 + Flux2 first, jointly** — design.md's rationale
stands: together they exercise CFG variants, expert routing, chunk-KV, and the image path, so the step-body
contract freezes only after all four are exercised → Wan-causal (self-forcing student) → LTX-2 → HunyuanVideo →
Stable Audio → remaining image families. Unmigrated families keep running via `LegacyPipelineNode`.
- Policies: `CFGPolicy` (absorbs the 3 CFG copies), `AttnMetadataProvider`, `FlowShiftPolicy`, `PrecisionPolicy`.
- Extension core lands with the loop (it's why the loop is being rebuilt): observer bus, ParityAligner promoted to
per-request observer, Profiler/NaNWatch, and **cache-dit as the first interceptor** (retiring `enable_teacache`).
- `forward_context.py` off the *migrated* inference path (194 references across ~68 files today: ~8 importer files
in frozen `training/`, the rest spread across train/ models, tier-B inference stages, quantization, and tests);
the module survives as a shim for frozen `training/` **and unmigrated tier-B stages** until M6 — what M2
guarantees is that no migrated family and no new code touches it.
- **`train/` migrates per-family, immediately behind inference**: DMD2 and the landed DiffusionNFT (#1450) adopt
the shared step functions as each family's body lands — `rl/common/sampling.py`'s loop is deleted, #1396
grad-norm refs extended to the migrated methods (RL included).
*Gate, per family: old-vs-new loop bit-identical (ParityAligner) + SSIM + a recorded loop-overhead / batch-of-1
latency measurement (the baseline M3 gates against); for train/: seeded rollout latents identical, reward metrics
- grad-norms neutral. No family is ever dual-maintained.*
### M3 — Execution plane: engine + scheduler (weeks 8–14, overlaps M2) → **breaking release R2**
- `AsyncEngine` (queue, admission, cancellation-as-common-path, failure isolation); offline `VideoGenerator` keeps
its name, becomes a thin sync wrapper that can bypass the queue.
- `StepScheduler` v0: multiplexes denoise steps across requests in a pool; budget currency = **predicted GPU-time**
from a calibrated per-(model, phase, shape) cost table (the cost *model* matures later; the currency is right
from day one). Carries the `ARDecodeLoop` contract; AR batching itself waits for its workload (N5).
- CacheManager v0: per-request chunk-KV slabs behind `KVHandle`; CFG-parallel axis (2-branch in practice).
- **Dynamo stock worker** (registration, health/drain, cost metrics), retiring the locked
`dynamo/examples/diffusers/worker.py` pattern.
- **Dreamverse hard-cut** (per design.md Phase 2; the aggressive delta is doing it in one PR): `gpu_pool.py`,
queue, warmup, and stream relay deleted and replaced by engine-client calls; the duty-cycle concurrency study
runs on the result.
- Colocated weight-sync RPC + component-granular sleep/wake + `RolloutClient` (engine-client RL mode for #1450).
*Gate: serving load tests; batch-of-1 latency regression ≤ 2% vs the M2-recorded measurement; Dreamverse
single-session parity; RL engine-client seeded final-latent parity vs in-process; deploys under stock Dynamo.*
### M4 — Graphs, parallelism, multi-session (weeks 14–20)
- `PipelineSpec` graph IR: per-family pipeline classes shrink to **spec + step body + policies**
(`create_pipeline_stages()` retires); LTX-2 and Hunyuan15+SR land as real fan-out graphs.
- Role pools + connectors (port `multimodal_gen`'s disagg state machine); declarative stacked-parallelism axes
compiled to DeviceMesh; general cross-mesh `WeightSyncPlan`.
- ComfyUI workflow→spec compiler MVP (tier-1 ~20-node vocabulary) + weight/adapter fleet cache.
*Gate (design.md Phase 3's, in full): ≥2 Dreamverse sessions/GPU on the recorded duty-cycle trace, p95 within SLO
— this is also where the loop-inversion **falsifier** is evaluated (see §7); LTX-2 A/V full-fan-out end-to-end;
disaggregated-vs-colocated throughput benchmark; CPU-only topology validation suite; ComfyUI tier-1 workflows
compile and run with equivalence reports; spec-built pipelines SSIM-identical to M2 loop versions.*
### M5 — Omni/MoT native + RL hardening (weeks 20–30)
- Cosmos3 re-port onto specs: packed factored sequences, dual-pathway attention, reasoner paged KV, joint denoise,
world-model `ChunkRollout`; `/v1/chat/completions`; AR continuous batching arrives **with** this workload (N5).
- Consistency ladder enforced end-to-end: C1 default in CI, C2 bitwise mode for goldens, Behavior Record opt-in;
first GRPO-class method lands on the engine-client rollout path (log-prob drift becomes the gated metric).
*Gate: Cosmos3 150-test parity suite on the new runtime; reasoner pool efficiency — tokens/s/GPU at target
concurrent denoise throughput, with the ≥10×-vs-re-prefill sanity floor; C1 drift ≈ 0 on a Wan RL run with the
drift dashboard live.*
### M6 — The tail and the precondition (week 30+)
Continuous deletion (M1/M2) shrinks the final phase but does not eliminate it: what remains by M5 is the tier-B
tail on `LegacyPipelineNode` and the frozen `training/` stack — which is a *live consumer* of
`ComposedPipelineBase` and `forward_context`, so its retirement is the precondition, exactly as design.md Phase 5
states. M6 = execute the §6 tail decision (migrate or deprecate each tier-B family), retire `training/` per the
checklist, then delete `ComposedPipelineBase`, the legacy `DenoisingStage`, `forward_context.py`,
`FastVideoArgs`/`TrainingArgs`, and `RayDistributedExecutor` together. **4 loop copies → 1.**
## 4. Breakage manifest (user-visible)
| Release | What breaks | Replacement | Aid |
|---|---|---|---|
| **R1** (M1) | `generate_video(prompt, **kwargs)`, `SamplingParam`, `FastVideoArgs` as public type, `fastvideo.api` legacy exports, CLI flag names, YAML config schema, streaming event types (`Video*Event` → `OmniEvent`, `schema_version`'d from day one) | `VideoGenerator.generate(OmniRequest)`, `DeployConfig`, generated CLI/protocol, `OmniEvent` | `fastvideo migrate` codemod, migration guide, R0 pinned |
| **R2** (M3) | Default execution path becomes the engine (offline bypass preserved); server lifecycle (queue/admission semantics, job states) | `AsyncEngine` | guide; `OmniEvent` schema unchanged from R1 |
| after R2 | nothing user-visible — M4/M5 are additive | — | — |
## 5. Deviations from design.md §10, stated honestly
| design.md | this plan | why it's safe now |
|---|---|---|
| Phase 0 keeps `VideoGenerator`/CLI signatures; `compat.py` shrinks to Phase 5 | M1 breaks signatures, deletes `compat.py` ⚡ | the only argument for the shim was signature stability — explicitly revoked |
| Legacy code deleted at Phase 5 | per-family deletion at parity, M2 onward ⚡ | parity gate is per-family anyway; carrying dead code to a final phase only invites the 19×-broken-freeze failure mode |
| Phases strictly sequential | M2/M3 overlap ⚡ | the engine consumes step bodies, not finished families; the step-body contract freezes at the Wan+Flux2 landing |
| Phases −1 through 4 sized at 36–54 engineer-months | ~21–28 engineer-months (3–4 eng × 30 wks) ⚡ | the delta is real deleted work — no compat maintenance, no adapter upkeep, no dual-stack carry — plus M2/M3 overlap; treat 30 weeks as the aggressive case and 36–40 as the planning case |
| Unchanged | parity/SSIM gates (restored in full at every milestone), enforcement package (CI path gates, CODEOWNERS, inflow rule — now at M0), train/RL migration timing (design.md Phase 1 already migrates NFT), Dreamverse hard-cut (Phase 2 already prescribes it), N2/N3/N5, cost-model currency, Dynamo asks + fallbacks, schema versioning | aggression budget is spent on interfaces only |
## 6. Decisions needed before M0
1. **Tier the model zoo.** Tier A (migrated, coverage guaranteed): Wan 2.1/2.2, Wan-causal/self-forcing, LTX-2,
Flux2, HunyuanVideo, Stable Audio, Cosmos3 (contingent on the M0 merge — it is not on `main` today), image
families. Tier B (runs on `LegacyPipelineNode` until someone claims it, candidate for deprecation at M6):
gen3c, matrixgame2/3, longcat, the rest. **Approve or edit the split** — it bounds M2.
2. **Release framing.** R1 as `v0.3.0` (pre-1.0 semantics, loud notes) vs holding breaks for a `v1.0` story.
Recommendation: `v0.3.0` now — waiting taxes every milestone.
3. **Freeze windows.** M1 freezes `api/`/args/entrypoints PRs ~2 weeks; M2 freezes per-family stage PRs while that
family migrates (days each). Needs maintainer sign-off.
4. **Staffing.** Critical path is M2's per-family step bodies — parallelizable per family after the Wan+Flux2
reference lands. 3–4 engineers ≈ 30 weeks to M5 in the aggressive case (design.md's own sizing implies 36–40
weeks at the same staffing — see §5); 2 engineers ≈ stretch ~1.5×. The Cosmos3-chain merge (M0) is its own
2–3 engineer-month track and should be staffed separately from the runtime critical path.
## 7. Risks specific to the aggressive posture
- **In-flight PR collisions** with layout/schema moves → freeze windows (above) + landing schema cuts at
milestone *starts*, not ends.
- **Community churn at R1** (ComfyUI users, script users) → codemod covers the mechanical 90%; the 10% that isn't
mechanical (kwargs with changed semantics) is enumerated in the guide; previous version stays pinned and
installable.
- **Parity harness becomes the bottleneck** — every aggressive deletion is licensed by it. Mitigation: it is the
*first* deliverable (M0), and per-family migration PRs are template-driven (record → port → compare → delete).
- **Overlap risk (M2/M3)**: the engine team building against a moving step-body contract → the contract
(`init/step/finalize` + `StepResult`) freezes at the *first* family (Wan), enforced by the same schema-version
discipline as external surfaces.
- **The known unknown**: loop inversion at scheduler granularity has no production precedent (design.md §1). The
falsifier stands, on design.md §11.6's schedule: the M3 duty-cycle study *publishes the targets*; the falsifier
is **evaluated at the M4 gate** — if step-level multiplexing can't beat request-level serialization on real
Dreamverse traces, StepScheduler retreats to request-level dispatch and the loop contract keeps only its
streaming/preemption seams, with no family code changing — step bodies and the M1/M2 cuts retain full value.
+3 -2
View File
@@ -223,8 +223,9 @@ follow_imports = "silent"
skip = "./data,./wandb,apps/fastvideo_studio/package-lock.json,apps/performance_dashboard/frontend/package-lock.json,*/_vendored/*"
# "tread" matches daVinci-MagiHuman's acronym "TReAD" (Token Routing and
# Early Drop). codespell lowercases ignore-words entries, so the single
# lowercase form silences all case variants.
ignore-words-list = "tread,passt"
# lowercase form silences all case variants. "mot" = Mixture-of-Transformers
# (MoT); "clen" = a Content-Length local; "te" = a text-embeds local.
ignore-words-list = "tread,passt,mot,clen,te"
[tool.ruff]
# Allow lines to be as long as 120.
+267
View File
@@ -0,0 +1,267 @@
# Adversarial Review of `design.md` (v12) — FastVideo Next-Generation Inference Runtime
**Date:** 2026-06-11
**Method:** Multi-agent adversarial review. 9 fact-check agents verified 70 concrete claims against the repo, the local reference checkouts (`cosmos-framework/`, `dynamo/`, `vllm-omni/`, `~/sglang`, `~/vllm`, `~/miles`, `~/verl-omni`, `~/diffusers`, `~/torchtitan`, `~/xDiT`, `~/ComfyUI`, `~/sglang-omni`, `~/cosmos-rl`), and GitHub. 9 attack lenses (abstractions, scheduler/perf, memory/cache, training/RL, strategy, migration, internal consistency, omissions, external borrowings) plus a completeness critic raised 70 findings; every finding went to a refute-by-default verifier. 36 findings were refuted; this document contains only the 34 that survived (1 critical, 26 major — consolidated below where lenses converged — 7 minor), plus fact-check corrections.
---
## Verdict
The architecture survives its strongest attacks — loop inversion's expressibility, the typed-state hybrid, the N1/N5 scope discipline, the clean-room GPL posture, and the C2-for-batch-1-video argument all held under refutation attempts. What does not survive is:
1. **The migration plan**, which consumes its own substrate two phases before building it and rests on a "frozen legacy stack" premise this repo has already empirically falsified.
2. **Two load-bearing factual errors** about reference systems (vLLM's BlockPool page sizes, diffusers' loop ownership) that each drove a recorded design decision.
3. **A family of undesigned failure/memory/trust paths** that the multiplexing bet itself creates. One is critical.
---
## Critical
### C1. No failure-isolation or cancellation semantics for the multiplexed pool — the blast-radius problem the architecture itself creates
**Where:** §6.3.1; absent from §12.
Today one request per pool means one request's CUDA error is its own problem. Step-multiplexing changes the failure class categorically: a mid-step OOM/illegal-access/NaN from one request poisons the CUDA context and desyncs in-flight NCCL collectives for *every* co-scheduled tenant on the pool, including resident Dreamverse session caches. The doc designs none of the machinery: no SPMD-consistent abort broadcast (the dual of its scheduling broadcast), no request-fatal vs pool-fatal classification, no pool re-init + cache-invalidation policy, no partial-artifact semantics for fan-out graphs. "OOM" and request cancellation appear nowhere in 1799 lines; "abort" appears once (RL stragglers).
Ordinary cancellation is also missing — and vibe directing makes abandoning in-flight generations the *common* path. Worse, Phase 2 retires Dreamverse's `gpu_pool.py`, which today has a working sentinel-fd worker-death watch (`gpu_pool.py:542-586`), into engine-client calls — a reliability regression for the flagship customer if the gate ships as written. vLLM v1, the doc's own scheduler template, needed first-class machinery for exactly this (`ENGINE_CORE_DEAD`, `EngineDeadError`, `abort_requests`). Risk 4 covers only scheduling-decision divergence; the long-job-resilience known-gap is single-job-framed.
The abort path shapes the StepScheduler loop, the worker RPC surface, and CacheManager handle lifetimes — it must be designed *with* Phase 2, and by the doc's own standard ("absence reads as a decision"), this absence is an oversight.
---
## Major — reference-system misreads that drove recorded decisions
### M1. The single-BlockPool CacheManager rests on a property vLLM explicitly does not have: per-group page sizes
**Where:** §6.3.2 lines 555-559.
The sentence asserts two mutually exclusive properties. vLLM's one-pool/no-fragmentation guarantee exists *only because* physical bytes-per-block are uniform across all groups: `kv_cache_utils.py` asserts a single page size (`get_uniform_page_size`), and its docstring says verbatim that breaking this "is non-trivial due to memory fragmentation concerns." Groups differ only in tokens-per-block at equal byte size; the unification mechanism inflates the smaller group's `block_size`.
Apply that to FastVideo's groups: a text-KV page (~64 KB/layer) vs a latent-frame slab (9.6–32 MB/layer for 1.3B/14B causal Wan) is a 150–500× ratio — unification means a 500-token reasoner prompt strands a multi-MB slab per layer-group. The one vLLM path with multiple page sizes (DeepseekV4) statically partitions capacity at startup over a single global block-id free list, which is harmless when group demand is token-coupled (every token passes through all layer groups) but wasteful exactly when demand is workload-decoupled — FastVideo's regime, where text-KV and chunk-KV demand vary independently with request mix.
Since this misread is what reversed the two-pool sketch (recorded at line 280), the decision rests on a false premise: either chunk-KV stays uniformly fine-paged (losing the slab semantics the MoT "falls out naturally" story depends on), or the two-pool design returns and needs its own fragmentation/deadlock argument.
### M2. diffusers Modular is not loop inversion — the "strongest external validation" of the keystone doesn't validate it
**Where:** §5 line 277, §6.2.3 lines 426-428.
`LoopSequentialPipelineBlocks.__call__` raises `NotImplementedError`; every concrete family hand-writes `for i, t in enumerate(timesteps)` inside its own blocking wrapper (`wan/denoise.py:434`, `stable_diffusion_xl/denoise.py:701` — SDXL ships four such wrappers, the subclass forest again). The iteration is block-owned, invisible to any runtime — no init/step/finalize, no external driver, none of the properties §6.2.2 says inversion exists for (scheduling, interleaving, preemption, streaming, fair sharing). In scheduling terms it is the current `DenoisingStage` with a refactored body — i.e., it validates the Guiders/policy pillar but as evidence for inversion it is *equally consistent with the alternative the design rejects* ("keep loops in stages, make bodies pluggable"). The class also carries an explicit experimental warning.
Consequence: no surveyed system — vLLM, sglang, multimodal_gen, diffusers — implements runtime-owned diffusion iteration at scheduler granularity. Loop inversion is the design's most novel element with zero production precedent, and risk 3 (which admits novelty only for the hybrid AR+denoise slice) should say so instead of borrowing validation the reference doesn't provide.
### M3. Cost-currency scheduling drops the memory half of vLLM's admission — and memory is never a scheduling resource anywhere in the design
**Where:** §6.3.1 (lines 476-547), §6.3.2; two lenses converged here.
vLLM's token budget is not a prediction — it is an exact cap checked *in the same loop as memory admission* (`allocate_slots` per request, preempt on allocation failure; activation memory separately bounded by a profiled worst case). The design takes the accounting structure, swaps the currency for a *forecast* (predicted GPU-time), and drops the memory dimension entirely: latents, conditioning sets, CFG duplicates, and activation peaks live in `RequestState`, explicitly outside the CacheManager, and nothing bounds how many concurrent LoopStates a pool admits — for a workload the doc itself calls memory-bound (line 499). Two items that each fit alone can jointly OOM, and a GPU-seconds currency cannot see it; combined with C1, that OOM is a pool-wide event. "Preemption only at step boundaries" never defines what happens to a preempted request's multi-GB resident state (offload? drop-and-resume-from-LoopState? — different economics from KV recompute).
Related internal contradiction, verified: cost is "static and known at admission... a table lookup" (line 538), but the same cost model is cache-dit-aware (line 520) — DBCache skip decisions are runtime data-dependent residual comparisons, unknowable at admission.
**Fix:** the budget needs a memory axis (resident-state + peak-activation per schedulable item), admission needs a memory planner over RequestState, and preemption semantics must be specified. The Phase-2 "≥2 sessions per GPU" gate rests on unaccounted memory until then.
### M4. Punica cannot express ComfyUI LoRA semantics
**Where:** §9.4 lines 1376-1380 (also §6.3.2 lines 575-579). *Verifier rated minor-to-major; grouped here with the borrowings cluster.*
vLLM's `LoRARequest` carries one `lora_int_id` and no strength field; scaling is baked into `lora_b` at registration; the Punica wrapper maps one adapter index per token. ComfyUI traffic — the workload §9.4 names — is N stacked LoRAs per request with continuous user-set `strength_model` *and* `strength_clip`, routinely tweaked per generation. Pushing that through Punica means registering each (ordered-set, strengths) tuple as a synthetic concatenated adapter: near-zero cache-hit rate across strength tweaks, registration churn in the stacked GPU weight slots, and concatenated ranks colliding with `max_lora_rank`. "Strictly better than hot-swap-only" is unsupported without a composition layer that doesn't exist anywhere, including in vLLM.
---
## Major — execution-plane gaps
### M5. MoT mode multiplexing has no parallelism answer
**Where:** §6.3.1 lines 502-503 vs §6.3.4; Phase 4 gate.
The "mode multiplexer" claim assumes both loop types share one static pool layout (`parallel: [dp, cfg, sp, tp]`), but their optimal layouts are disjoint: denoise wants SP+CFG; AR decode is sequence-length-1 — SP has nothing to shard and CFG doesn't exist. On a `[cfg(2), sp(4)]` 8-GPU pool the reasoner either runs replicated (1/8 useful work, paged KV duplicated 8×) or needs TP — and TP-everywhere regresses the bread-and-butter denoise workload on the flagship pool. Per-phase re-layout of the same resident weights is not expressible in the §6.3.4 spec (one static stack per pool), and resharding machinery exists only for train↔rollout weight sync (§8.6). §6.3.1's own jumbo-step mitigation (split cost classes across pools) is structurally unavailable for MoT — AR steps and denoise steps are the same weights — so concurrent reasoner token latency is gated by indivisible 50–500 ms denoise steps.
A workable resolution exists (AR continuous batching data-parallel across the cfg×sp weight-replica axes onto TP subgroups, plus §6.3.3 per-pathway TP, plus routing pure-REASON traffic to differently-shaped pools), but the doc never states one, and the Phase-4 gate ("reasoner ≥10× faster than re-prefill") is measured against an O(n²) strawman baseline that certifies nothing about pool efficiency. Risk 3's "prototype early in Phase 4" defers a *design contradiction*, not an implementation unknown.
### M6. The engine's own multi-node story is unstated, and the Ray executor silently disappears
**Where:** N1 line 143, §6.3.5 line 673, §6.0 line 304.
Whether one worker pool may span nodes is a load-bearing decision the doc never makes — Dynamo routes *between* workers; it does not own the NCCL mesh *inside* one. If pools are single-node by fiat, SP degree caps at ~8 GPUs, directly contradicting line 543's jumbo-step mitigation ("shrink jumbo step wall-time with SP"), capping MoT model scale — and `RayDistributedExecutor`, today's shipping multi-node path, is silently dropped: it appears in the §3.1 diagram and then never again in §6, §10, §11, or §12 (violating the plan's own "every phase deletes or freezes what it replaces" discipline). If pools may span nodes, the engine owns cross-node collective bring-up, NCCL-timeout fault domains, and a multi-node health/drain contract — none designed, and C1's recovery problem becomes a multi-node recovery problem. Either answer changes Phase 2/3 scope. "Node-group" appears once, undefined.
### M7. Policies carry per-request mutable state with no state-scoping contract — and the doc contradicts itself on when policies are resolved
**Where:** §6.2.3 lines 412-417 vs §6.2.2 line 387 vs risk 2 line 1621; §6.4 lines 837-840.
The doc says policies are resolved at pipeline build (lines 412-413; risk 2: "resolved to bound methods at build time") *and* in `DenoiseLoop.init` (line 387) — a genuine contradiction on a load-bearing contract. It matters: AdaptiveGateCFG — a named CFGPolicy example and the Wan2.2 worked-example default — is per-request mutable state in shipped code (`denoising.py:338-343, 507-551`: `delta_cached`, `delta_cached_model_id`, gate counters). Build-time-resolved singletons mean request A's cached CFG delta gets applied to request B the moment Phase 2 interleaving lands — silent quality corruption no Phase-2 gate (load tests, latency budget) can catch. This is the *exact* failure mode §6.4 cites to justify interceptor state scoping ("silently corrupts under concurrent requests") — the contract was designed for the plugin tier and forgotten for the policy tier, which sits on a hotter path. Cheap fix (policy state into LoopState, same as plugins), but it must be in the spec.
### M8. The six-policy taxonomy does not factor the shipped step bodies — no step skeleton or cross-policy interaction contract is defined
**Where:** §6.2.3 (policy table, line 424 claim); §6.2.2 lines 386-389; §6.4 lines 837-844.
The proposed step is three phases (forward → CFG combine → scheduler step); the shipped loops need ~six, with dependencies that cross policy boundaries. Verified examples:
- **Cosmos** conditioning-frame injection consumes the *sampler's* EDM coefficients, applies per-CFG-branch both pre-forward (input mix) and post-forward (x0 clamp), and the CFG combine runs in x0 space — ConditioningInjector × Sampler × CFGPolicy interleaved inside each branch, unownable by any one of them (`denoising.py:845-933`).
- **TI2V** clamps latents *after* `scheduler.step` — a post-step constraint with no policy slot (`denoising.py:570-573`).
- **Cosmos2.5** builds per-frame timestep vectors with a conditioned-frame override and re-clamps GT every step pre-forward.
- **CausalDMD** renoises between steps choosing `add_noise` vs `add_noise_high` by expert boundary — Sampler × ExpertRouting (`causal_denoising.py:268-301`).
- **AdaptiveGateCFG** must observe ExpertRouting's switch to invalidate its delta (today an inline `id(current_model)` check) — yet no channel for one policy to observe another is defined anywhere.
- **LTX2** guidance is 1–4 runtime-decided passes whose branches alter the network via forward kwargs (`skip_cross_modal_attn`, `skip_video/audio_self_attn_blocks`) — colliding with BlockInterceptor's domain in a way the "two block-skippers conflict" pre-flight check cannot see, and breaking §6.4's per-CFG-branch state scoping, which assumes a fixed cond/uncond branch vocabulary (`ltx2_denoising.py:503-605, 620-631`).
None of the six policies covers prediction-space conversion, per-token timestep construction, post-step latent constraints, inter-step renoising, or chunk-boundary refresh. The fix is not abandoning policies — the Sampler registry is the natural home for some of this, and composition still strips the duplicated offload/attn-metadata/autocast/trajectory plumbing — but the design needs the fixed step skeleton with ordered, typed extension points and an explicit policy-interaction contract, worked through Cosmos2.5 and LTX2 *in the doc*. Until then, "a new model contributes policies + a graph spec; it does not edit shared loop code" (line 424) is asserted, not demonstrated.
### M9. OmniRequest cannot parameterize multi-loop graphs
**Where:** §6.1 lines 318-334; §6.6 line 905; worked examples (c)(d) lines 943-950.
One flat `SamplingParams` + one flat `DiffusionParams` per request, while the design's own flagship examples are multi-loop graphs needing per-node knobs: LTX-2's refine loop has its own step count and guidance scale *today* as first-class fields (`fastvideo_args.py:204-205`, threaded through `compat.py` and `dynamo/examples/diffusers/worker.py:201-203`); a thinker and talker need different `max_tokens`/`temperature`/`stop`. No request→graph-node parameter binding is defined anywhere; the only escape hatch is line 905's per-model `ModelOptions` blocks — i.e., the `ltx2_*` field-leakage pattern the doc indicts at P3, with a type wrapper, regenerated into the OpenAI/CLI views that derive from the request schema (line 907). Needs a real decision — parameters keyed by graph-node id, or per-node override blocks validated against the PipelineSpec — made in Phase 0, because that schema ships first and external consumers build against it.
---
## Major — caches and weights
### M10. No feature-cache invalidation story under LoRA hot-swap — te-LoRAs make the embedding cache serve stale embeddings in the workflow cloud
**Where:** §6.3.2 lines 570-574 vs §9.4 lines 1349-1380.
The only invalidation rule in the document is RL `update_weights` → `reset()`. But ComfyUI-grade LoRAs routinely patch the *text encoder* alongside the DiT (`comfy/lora.py` maintains `lora_te/lora_te1/lora_te2` key maps; `load_lora_for_models` takes a separate `strength_clip`), so a content-hash-keyed embedding cache returns embeddings computed under the wrong adapter state the moment two workflows share a prompt but differ in te-LoRA stacks — silent wrong output in the exact product (§9.4 "exact mode") whose trust claim is reproducibility. §11.8 even makes cross-request embedding reuse load-bearing as the radix-cache substitute. And once Punica-style batched multi-LoRA lands, requests with different adapter stacks coexist concurrently on one pool, so the cache must be key-*partitioned* by (encoder identity × adapter set × strengths), not flushed — a different design from the `EncoderCacheManager` reset() semantics being adopted, which come from a world where encoders are never patched per request. The key schema needs a weight-state epoch / adapter-set hash as a mandatory component, decided before Phase 3.
### M11. Checkpoint/LoRA patching mutates pool-shared weights — a pool-quiescing barrier the StepScheduler has no vocabulary for
**Where:** §9.4 lines 1371-1380 vs §6.3.1 and §6.0 line 299.
Components are "one resident copy per worker pool"; patch/unpatch mutates that copy, which is global to every loop interleaved on the pool — yet step-interleaving is the engine's core Phase-2 value. Two interleaved loops requiring different patch states cannot coexist, so every cross-group transition is a drain barrier: finish in-flight steps, apply/undo `W += scale·BA` across 14–28 GB shard-consistently across TP/SP ranks (ComfyUI keeps weight backups for the undo — 2× weight memory or a CPU→GPU restore at PCIe seconds), re-admit. Under workflow-cloud traffic (long-tail checkpoints, per-request adapter stacks), transition frequency is the whole game — and the §6.3.1 cost model (lines 516-521) has no weight-state-transition term, no notion of weight state as schedulable state, and no quiesce-vs-queue policy, even though transition cost is exactly what A1 checkpoint-affinity routing must weigh. The §8.6 safe-point-swap pattern shows the doc knows the shape but never applies it here. §9.4 calls this "the one real new subsystem"; §12 carries no risk entry for it.
---
## Major — training/RL
### M12. "Step bodies are plain tensor programs, so autograd composes" is contradicted by the distillation code the substrate must absorb
**Where:** §8.2 lines 1016-1019; §6.2.2; §6.3.2.
Self-forcing does not "drive `DenoiseLoop.step`": its rollout samples per-block exit indices broadcast across ranks, runs no-grad steps to the exit, runs exactly *one* grad-enabled forward, then a separate no-grad `store_kv=True` context-caching pass with context noise, gated by `start_gradient_frame`. None of this fits `init/step/finalize` + `StepResult(done, emit)` without grad-gating flags, per-step cache-write control, and per-block exit policies — training-only surface in substrate code, or the method keeps its own loop and the "3 copies → 1" dedup claim dies for the hardest case. The KV path needs grad/AC-aware semantics the engine pool lacks: today's causal model snapshots KV indices whenever `torch.is_grad_enabled()` so activation-checkpoint recompute doesn't double-advance the cache (`wan_causal.py:119-120,405-431`), and never recycles blocks mid-rollout — while §6.3.2 specs vLLM-style out-of-window block recycling, and §8.5's own profile taxonomy says "training forward … *no caches*," showing the grad+KV case was never designed. §8.3 explicitly stakes the architecture on ChunkKVPool serving self-forcing training.
(Note: the related forward-context-backward attack was refuted — the Phase-1 retirement of the global plus explicit metadata passing *helps* autograd composition. The surviving residue is the grad-window/cache-mode design above.)
### M13. Behavior Record cost is understated ~1.5 orders of magnitude for its own flagship case (MoE diffusion)
**Where:** §8.5 lines 1156-1160; §5 miles row line 1088.
The miles ~60 MB/sample figure is per-token routing, one forward per generated token. Diffusion re-routes the *entire packed sequence at every denoise step, twice under CFG*: the record is steps × CFG × tokens × MoE-layers × top_k. For a Cosmos3-class request (Qwen3-VL-MoE config: 60 experts, top_k 4, ~24 sparse layers via `decoder_sparse_step=1`, ~50K packed tokens, 35-50 steps × 2 branches) that is ~1.3–1.9 GB/sample int32 — ~20–30 GB per 16-sample GRPO group, before latents. "Cheap because trajectory capture is already an OutputSpec feature" conflates plumbing cost with byte cost; at these sizes the Record forces a buffering/transport/storage design (GB-scale trajectories through connectors from disaggregated rollout fleets) that appears nowhere — not in §8.7's TrajectoryBuffer, not in §12, not in the known-gaps list. (The RNG-draws sub-claim was refuted: seeded generators in a single shared loop reproduce draws; uint8 expert IDs also cut 4×. The routing-record problem stands.)
### M14. The omni-RL pilot is a Phase-4 deliverable with no objective design
**Where:** §8.7 lines 1236-1240; §10 Phase 4.
The section establishes *expressibility* (one trajectory, two segment types — true, and a real structural advantage over engine-per-stage stacks) and quietly upgrades it to a deliverable without posing the algorithm problem:
- **Scale mismatch:** token log-probs are O(1–10) nats over 10²–10³ tokens; per-step diffusion SDE log-probs are Gaussian densities over 10⁶–10⁷ latent dims — any joint clipped-ratio objective needs principled per-segment normalization that none of the cited recipes (FlowGRPO/DanceGRPO/NFT/AIPO/GSPO) provides; get it wrong and one modality silently dominates the shared trunk.
- **Credit assignment:** the reasoner influences video reward only through *sampled discrete tokens* re-entering as conditioning — a non-differentiable boundary, so token segments get sparse trajectory-level REINFORCE signal while denoise segments get dense per-step ratios, both updating shared attention-trunk weights, with no interference analysis.
- **Reasoning regression:** RL-updating the und pathway on video-reward-correlated signal risks degrading its reasoning; reference-model KL anchoring for hybrid episodes is never mentioned.
The entire treatment is the phrase "optimized with mixed objectives," and §12's 15 open questions contain nothing on it — for the capability marketed as "the capability nobody else has." Either it gets an algorithm sketch and an open-question entry with an owner, or the Phase-4 item should be demoted from "pilot" to "trajectory capture demonstrated."
---
## Major — the migration plan (the weakest section)
### M15. The "frozen legacy stack" premise is empirically false in this very repo
**Where:** lines 5, 110, 1026; §11.4; risk 5.
The anti-third-stack defense is a declared freeze plus intent to delete — and this repo has already run that experiment and it failed within weeks. Verified from git: `fastvideo/train/` landed 2026-03-09 (#1159); since then **19 commits modified the "frozen" `fastvideo/training/`**, including a *brand-new* `cosmos2_5_training_pipeline.py` added to the legacy stack on 2026-05-11 (#1227) — **nine days after `training/AGENTS.md` explicitly forbade adding new models there**, and eleven days after the same model landed in `train/` (#1224). World-model training (#1179) and LongCat finetuning (#1244) also landed in the frozen stack in May; EMA bugfixes as recently as June 8-9; `AGENTS.md` still calls `training/` "authoritative for shipped models."
The doc invokes the training/-vs-train/ "lesson" but proposes nothing mechanically different from what was tried: no CI gate rejecting new files under legacy paths, no codeowner veto, no named owner per family, no calendar date for Phase 5. "Phase 5 is a scheduled deletion, not an aspiration" (risk 5) — but nothing in the document is scheduled. Under the same model-port pressure that broke the training/ freeze (measurably higher on the inference side), this freeze breaks the same way. Name the enforcement mechanism that did not exist last time, or the deprecation commitment is the prior failure restated with more confidence.
### M16. Phase dependency inversion: Phases 1–2 consume the substrate Phase 4 builds
**Where:** §10 lines 1406-1446 vs §6.3.1 lines 487-489, §6.3.2; three lenses converged on this.
Phase 1 migrates causal Wan ("exercises chunk-KV"); Phase 2 ships "AR continuous batching" — which §6.3.1 *constitutively defines* as "(continuous batching; paged KV; chunked prefill)"; the CacheManager owning both lands in Phase 4, and risk 3 even defers the StepScheduler+KVPool prototype to "early in Phase 4," contradicting Phase 2. Compounding it: **no AR-pathway model exists on the new runtime before the Phase-4 Cosmos3 re-port** (Wan-causal is chunked denoise, not token AR; thinkers/talkers are Phase 4), so Phase 2's headline deliverable has neither a cache backing nor a workload — and none of Phase 2's gates (lines 1427-1430) tests AR batching.
The Phase-1 half is softenable: an interim per-request chunk-KV behind the unchanged `KVHandle` seam, with a Phase-4 allocator swap, is normal incremental staging — but the doc never states this, and its own "no third stack / every phase deletes what it replaces" principle cuts against unstated throwaway implementations. Fix structurally: pull a CacheManager v0 (chunk-KV slabs + minimal paged text-KV) into Phases 1–2, or move AR batching to Phase 4 and rewrite the Phase-2 gate to what it actually exercises.
### M17. Phase 4 re-ports a baseline that is not on main, and the plan schedules neither its merge nor its rebase
**Where:** §10 Phase 0 line 1405, Phase 4 lines 1439-1446; §1 lines 42-49; Appendix.
`fastvideo/pipelines/basic/cosmos3/` on main contains only `__pycache__` — the design's forcing function exists solely as the unmerged 5-branch stacked chain (`feat/cosmos3-tier-a-port` → … → `feat/cosmos3-reasoning`). Phase 0's "Cosmos3 audio leaves `batch.extra`" cannot execute against main: it presupposes the chain is merged (a major-model review effort the plan never schedules) or means maintaining the migration on a side branch, continuously rebased across the most churn-heavy refactors in the repo's history (ForwardBatch→RequestState, loop inversion, executor→engine) — months of conflict-resolution work, unowned and unsized, on the artifact whose 150/150 bit-exactness is the design's proudest credential and whose parity suite the Phase-4 gate requires ("every phase ships green" cannot apply to a suite that is not in the tree). The plan sequences other in-flight work explicitly (`fastvideo/api/` in Phase 0, PR #1438 in Phase 1) but skips this. Needs an explicit merge milestone before Phase 0 touches the port.
### M18. G5's enforcement instrument has holes: ~6-7 shipped families have no SSIM test, and the CI-cost mitigation is incoherent for substrate PRs
**Where:** G5 lines 128-129; Phase 0 gate line 1405; risk 6.
`fastvideo/tests/ssim/` covers ~14 of 20+ families. Cosmos(2/2.5), Hunyuan, Hunyuan15(+SR), HYWorld, MagiHuman, Waypoint, and MatrixGame-v1 have no SSIM test — "all SSIM suites unchanged" passes *vacuously* for roughly a third of shipped pipelines, exactly the ones sitting on the shared loop being refactored. And risk 6's "gated to touched families" mitigation is designed for model-local PRs; Phases 0–2 are by construction not model-local — the ForwardBatch adapter, loop inversion, and executor replacement sit under every family, so "touched families" = all of them on precisely the riskiest PRs. Either substrate PRs run the full GPU matrix (a cost the plan should budget — SSIM runs on Modal L40S today) or gating quietly degrades to sampling, which is how regressions slip through. Needs: a reference-seeding work item before Phase 1, or G5 restated as "zero regression for the SSIM-covered subset," plus a stated per-phase GPU-CI budget.
### M19. Phase 5's deletion milestone breaks the "frozen and untouched" legacy training/ stack
**Where:** lines 5-6, 144, 1026-1027 vs Phase 5 line 1448.
The frozen stack is a live consumer of exactly the code Phase 5 deletes: `fastvideo/training/training_pipeline.py:39` imports `ComposedPipelineBase`/`ForwardBatch`/`LoRAPipeline`, holds `validation_pipeline: ComposedPipelineBase`, and its validation instantiates real legacy pipelines that run the legacy `DenoisingStage`; `distillation_pipeline.py:31` likewise. So Phase 5 cannot remove `ComposedPipelineBase` and `DenoisingStage` while leaving `training/` untouched — either the deletion milestone hollows to "delete except what legacy training/ needs" (the old path never dies — the very smell being fixed) or the scope statement is false and `training/` breaks on this plan's schedule. Relatedly, "loop inversion makes the step functions the single shared implementation" is arithmetically 3→2, not 3→1: the legacy inlined copies are out of scope forever. The doc needs an explicit answer: what happens to `fastvideo/training/` at Phase 5?
### M20. "Retire `fastvideo/forward_context.py` (Phase 1)" is infeasible as scheduled
**Where:** §6.3.3 lines 618-621; Phase 1 lines 1412-1414; vs N2/N4; Appendix line 1791.
194 references across ~50 files. The global is read inside `fastvideo/attention/layer.py` — the shared Attention module on *every* family's hot path — and set in 27 places inside the frozen `training/` stack (8 module-level imports). Phase 1 migrates only Wan+Flux2; the other ~16 families run "unmodified" behind the legacy adapter (N4) and still set the global. So in Phase 1 the file cannot be deleted (touches the frozen stack, violating N2; breaks every unmigrated family), and `attention/layer.py` must serve both worlds simultaneously — a dual-sourcing branch in the hottest shared layer, undesigned. The honest description: Phase 1 *adds a second context mechanism beside the global*, and the global survives until Phase 5 at the earliest — where the deliverables list never mentions it. Appendix A states "retired Phase 1" as accomplished fact. Rewrite as "new-path-only StageContext; `forward_context` frozen for legacy consumers; deletion gated on Phase 5," and design the dual-mechanism cost.
### M21. §10 is a dependency ordering, not a plan — no timeline, no staffing, no sizing, and no policy for the ~1-2 new model ports per month that arrive during the migration
**Where:** §10; N4 line 153; risk 1.
The scope — typed I/O, loop inversion + policies, extension system, async engine + StepScheduler + online-calibrated cost model, four-class CacheManager, PackedSeq/MoT layers, declarative parallelism compiler, workflow compiler, RL layer, Dynamo contract, config collapse — is plainly multi-engineer-years, with zero dates, headcount, per-phase sizing, or owners; "by Phase 2" decision deadlines (§11.1, risks 7/15) are unanchored because Phase 2 is not a date.
The sharper, unanswered problem is **inflow**: git shows ~1–2 new families landing per month (Flux2 Klein and Lucy Edit on 2026-06-09 alone; MatrixGame3 05-27; MagiHuman 05-12; Stable Audio 05-01; Gen3C 04-01…). Over multi-quarter Phases 0–4, another 10–15 models arrive, and the doc never says what they target: land them on legacy abstractions and the Phase-5 tail grows faster than phases retire it (negative net migration velocity); force them onto the new stack and every port blocks on machinery that doesn't exist until Phase 1/3/4. Either answer materially changes the plan; choosing neither means the terminal state recedes indefinitely. Minimum fix: per-phase engineer-month estimates, a named owner per phase, a calendar target for Phase 5, and an explicit "new ports target the new stack starting at Phase X" rule with its porting-velocity cost stated.
---
## Major — product/trust surfaces
### M22. Per-request plugin enablement is an unsandboxed third-party-code and noisy-neighbor surface; only workflow JSON is named untrusted
**Where:** §6.4 lines 859-861 vs §12 input-hardening gap lines 1693-1695.
Entry-point plugins execute arbitrary code inside the serving engine, and the doc makes their selection part of the *request* (`diffusion.plugins=[{"name": "cache_dit", "Fn": 8, "Bn": 8}]`) in the same engine pitched as a multi-tenant cloud — and since the OpenAI protocol is *generated from the request schema* (lines 907-908), the field derives into the public API with no carve-out. Consequences forcing a design change: (a) **correctness** — a caller can attach a distribution-altering interceptor to a request the product has labeled "exact mode" (the §9.4 trust claim), or pass unvalidated kwargs into third-party code; (b) **isolation** — a `needs_eager` observer on one request drops compile/cudagraph capture for scopes shared with co-scheduled tenants (line 809), a noisy-neighbor vector with no cost attribution anywhere in the metrics design; (c) **supply chain** — entry-point resolution imports whatever package claims the name. The needed contract: enablement/allowlisting at DeployConfig scope only; requests merely parameterize pre-enabled plugins against per-plugin validated schemas; plugin overhead attributed per-request in the cost model. §12's input-hardening gap names only workflow JSON — a categorically different surface.
### M23. No versioning or stability contract for the serialized schemas shipped to external consumers mid-migration
**Where:** §6.4 line 861; §6.6 lines 920-927; §10; open question 12.
By Phase 3 there are at least four externally consumed serialized surfaces: hub-published ModelSpec manifests (interchange with diffusers' `modular_model_index.json` — a format co-owned with an external party), compiled-workflow PipelineSpecs (content-hash-keyed in the weight-fleet cache — schema changes silently change hashes and invalidate fleet affinity), the OmniEvent streaming schema (Dreamverse's frontend; proposed as Dynamo ask A3's wire format), and per-model ModelOptions blocks. Phase 4 then lands PackedSeq, session-scoped inputs, and the Cosmos3 re-port — guaranteed churn after consumers exist. The migration plan gates *behavior* at every phase (SSIM, parity, load) and gates *interfaces* at none; the only versioning commitment in the document is hook-point names (open question 12 is scoped to hook points). Without per-surface decisions now — `schema_version` fields, frozen-vs-experimental tiers per phase, a deprecation window — Phase 4 either breaks published artifacts or gets paralyzed by accidental freezing. G5 protects only the Python `VideoGenerator` call.
---
## Minor (confirmed)
1. **ForwardBatch has 111 fields, not ~250** (AST-verified; stated twice, lines 33/188). P3 survives at 111, but the headline metric is inflated 2.3× in a doc that brands its pain points "evidence-backed" — it invites discounting of the numbers that *do* verify exactly (1381 lines and 35 probes both check out).
2. **"Prediction is a table lookup" vs the design's own flagship features** (§6.3.1 vs §6.4): DBCache/FBCache/TaylorSeer decide per step from runtime residual similarity — a stochastic per-step cost multiplier unknowable at admission; VSA tile selection is content-dependent; and AR decode lengths are unbounded (the doc concedes vLLM "must guess decode lengths," then silently exempts its own AR group).
3. **Worked example (g) is internally contradictory**: cache-dit + C1 + "identical trajectories" are pairwise incompatible under §8.5's own `distribution_altering` contract (§8.7 states the rule correctly: cache acceleration is C0). Matters because (g) is the template PR #1438 is told to target in Phase 1.
4. **The Phase-2 Dreamverse gate is untestable as written**: at ~4.55 s GPU-saturating per 5 s clip (line 1263), "≥2 concurrent sessions per GPU at unchanged segment latency" is only passable under an unstated think-time/collision-rate assumption — the gate can be passed or failed at will by choosing the test's session behavior. More broadly, no quantitative multiplexing target (sessions/GPU under a stated load profile, GPU-utilization, cost/clip) exists anywhere, so there is no way to conclude after Phase 2 whether step-level scheduling earned its complexity over the §11.6-rejected simpler design.
5. **The exec summary launders Dynamo contingencies into outcomes** (line 75: "each with a fallback — so Dynamo fronts both production serving and RL rollout fleets"): the body is honest (A1–A7 with fallbacks; §11.9; §12.15), but the asks are unfiled RFCs on an NVIDIA-governed roadmap; A5's own fallback "weakens fleet-scale async RL," and if A3 misses Phase 2, Dreamverse ships on the direct-WebSocket bypass and the production-hardened fallback becomes permanent — the exact "permanent workaround" dynamic §11.9 claims the direct relationship avoids. Ask-sequencing (§12.15) has no owner or decision dates.
6. **diffusers as "convergent validation" cuts both ways** (see M2): its four-wrappers-per-family shape is the subclass forest again; the citation supports the rejected alternative as well as the chosen one.
7. **Punica/ComfyUI LoRA semantics gap** — see M4.
---
## Fact-check corrections
70 concrete claims were checked; **none was fabricated**; 13 need correction. Everything else verified, including the claims most likely to be embellished: vLLM RFC #42770 (author/date/content/two-tier resolution), PR #42304 **merged** 2026-05-16 with `VLLM_USE_BREAKABLE_CUDAGRAPH`, vllm-omni RFC #4084, the Thinking Machines numbers (80/1000 unique outputs, divergence at token 103, 26s→42s, KL results), the Dynamo worker's `asyncio.Lock`, cache-dit, the cosmos-framework MoT details (PackedAttentionMoT, MoTDecoderLayer, ReasonerKVCache, MoE gen-MLP), miles/verl-omni/sglang-omni mechanics, sglang's cache-dit monkeypatch scars, and `enable_teacache` genuinely having no consumer.
| # | design.md says | Reality |
|---|---|---|
| 1 | "1381-line `DenoisingStage`" (lines 34, 201) | 1381 is the **file**; the class is ~670 lines (47–715) plus 6 subclasses in-file. The 35-probe count is exact for the file. |
| 2 | "~250-field ForwardBatch" (33, 188) | **111 fields** (whole file incl. TrainingBatch/PreprocessBatch: ~153). |
| 3 | "19 denoising-stage classes" (201) | **22** model/variant classes (+ base = 23); the list omits Magi-class and two other same-category stages predating the doc. |
| 4 | "Cosmos2.5 clamping … hardcoded in the shared loop" (201) | Clamping lives in the `Cosmos25DenoisingStage` **subclass**; the Wan2.2 expert switch (`denoising.py:229-235, 352-376`) and TI2V inline VAE encode (`:239-268, 399-404, 570-572`) are in the shared loop as claimed. |
| 5 | `SamplingParam` "~170 fields" (887) | **75**. The ~170 figure belongs to TrainingArgs (90 own + 81 inherited = 171). |
| 6 | `FastVideoArgs` "~96 fields" (885) | **81** (TrainingArgs subclassing claim correct). |
| 7 | "TP and SP (Ulysses/ring)" (192) | Main is **Ulysses-only** (`all_to_all_4D`); no ring-attention SP is wired into FastVideo. |
| 8 | CFG "3 copies: `stages/conditioning.py` vs …" (993) | Right count, wrong citation: the inference-stack copy is in `denoising.py`, not `conditioning.py`. |
| 9 | ComfyUI "~45 `comfy_extras` packs", "90+ blueprints" (1335-1339) | **117** packs (matching nodes.py's 117-entry registration list); **80** in-tree blueprints (the larger library ships via the registry). 64 core nodes, 39 API providers, GPL-3.0, FIFO-no-batching all verify. |
| 10 | kv-router events "`{sequence_hash, block_hash, removed}`" (707, A1 733-739) | Paraphrase: actual shape is `KvCacheEventData::Stored{parent_hash, blocks[{block_hash, tokens_hash}]}` / `Removed` / `Cleared` (`protocols.rs:627-646`). Token-prefix-derived keying verifies. |
| 11 | miles TIS clamp "to `[0.5, 2.0]`" (1086) | Configurable `[tis_clip_low, tis_clip]`, CLI defaults [0, 2.0]; the 0.5/2.0 pair comes from the MIS example config (`mis.yaml`). |
| 12 | sglang-omni "`DllmScheduler` for a DiT talker" (269) | DllmScheduler serves the **LLaDA2-Uni thinker** (diffusion-LLM); the DiT talker is Ming-Omni's, on a different scheduler. |
| 13 | `_iter_packed_batches` under `model/vfm/` (236); §11.3's claim that the port's "own status notes" list reasoning-KV/batching/streaming/prefix-reuse as "missing for production" | Lives at `cosmos_framework/inference/inference.py:66`. PORT_STATUS.md confirms 150/150 but contains no such missing-for-production list — that framing is the design doc's own and should not be attributed to the port's status notes. |
---
## Attacks that failed (the doc survives these)
The refute-by-default verifiers killed 36 findings, several of them attacks a hostile reviewer would lead with — worth knowing they don't land:
- **ChunkRollout/DenoiseLoop nesting is expressible** in the stated Stage/LoopStage/StepResult contracts ("one solver step / one token / one chunk" + composition).
- **N1 vs engine-internal pools** is consistent on a careful read (N1 is about datacenter orchestration; §6.3.5 states the reconciliation).
- **The trainer-scope line (N2 vs §8)** is drawn consistently — N2's own text enumerates exactly what §8 changes.
- **G6 vs the ≤2% Phase-2 gate** is goal-vs-acceptance-gate, not contradiction (Phase 1 is gated bit-identical).
- **The clean-room GPL posture holds**: sampler/scheduler math (DPM-Solver, Karras sigmas, flow-match shift) is published outside GPL sources.
- **C2 for the video denoise path is fine**: batch-1 fixed shapes are trivially batch-invariant — the doc's own analysis at lines 1145-1147 is correct; the AR/image/sharding exposures are correctly identified there too.
- **Self-forcing's cross-chunk gradients truncate by construction** (KV written under `no_grad` on detached context), so the engine KV pool is not blocked the way one might fear — the surviving residue is M12's grad-window cache mode.
- **"Every phase deletes or freezes something" survives audit** at the phase-deliverable level (the failures are the specific items in M19/M20).
- **The tier-1 ComfyUI vocabulary claim survives** blueprint-corpus measurement under the doc's actual claim (curated canonical workflows, not top-N node frequency).
- **The sglang reconvergence deferral** is substantively defended in §11.1 with reasons valid under either outcome.
- **WeightSyncPlan's "literal no-op"** is correctly scoped to colocated same-layout in the doc's own sentence; FSDP-vs-TP/SP is explicitly routed to in-place reshard.
---
## Ranked recommendations
1. **Design the abort/cancellation/OOM path with Phase 2** (C1) **and add memory as a budget axis with admission planning and preemption semantics** (M3). These two are the soundness conditions of the multiplexing bet; everything else in the execution plane sits on them.
2. **Re-derive §6.3.2 from the real vLLM constraint** (M1). The two-pool→one-pool reversal was made on a false premise; either accept uniform page bytes (and redesign the slab story) or bring back two pools with an explicit fragmentation/deadlock argument.
3. **Fix the migration plan's three structural defects**: CacheManager v0 into Phases 1–2 or AR batching out of Phase 2 (M16); a merge milestone for the cosmos3 chain before Phase 0 touches it (M17); a new-port inflow rule plus a freeze-enforcement mechanism that did not exist last time — CI path gate, codeowners, a date (M15, M21). Also reconcile Phase 5 with the frozen `training/` stack (M19) and restate the `forward_context` retirement honestly (M20).
4. **Specify the step skeleton and the policy contracts** — ordered, typed extension points; a policy state-scoping rule (state in LoopState, like plugins); a policy-observation channel — and work the mapping through Cosmos2.5 and LTX2 in the doc (M7, M8). Decide per-node request parameter binding in Phase 0 (M9).
5. **Give MoT a stated parallelism answer** (M5) and make the single-pool-spans-nodes decision explicit, including the fate of `RayDistributedExecutor` (M6).
6. **Close the workflow-cloud trust/correctness holes before Phase 3**: adapter-aware feature-cache keys (M10), weight-state transitions as a scheduled, costed operation (M11), DeployConfig-scoped plugin allowlisting (M22), per-surface schema stability tiers (M23), and an honest assessment of Punica's fit (M4).
7. **Right-size the RL claims**: design the grad+KV cache mode or scope self-forcing out of the shared loop (M12); budget the Behavior Record at real byte counts (M13); demote the omni-RL pilot or give it an objective sketch and an owner (M14); fix worked example (g).
8. **Reclassify loop inversion as unprecedented at scheduler granularity** in risk 3 and drop the diffusers "validation" (M2). The bet may still be right — but it should be made with open eyes, and the parity-gate plan is then carrying more weight than the doc admits.
9. **Correct the thirteen numbers above before circulating.** The doc's credibility rests on its "evidence-backed" brand; ~250-vs-111 is the kind of error that makes a reader re-check everything else — and most of everything else checks out.
+3
View File
@@ -0,0 +1,3 @@
__pycache__/
*.pyc
*.pyo
+105
View File
@@ -0,0 +1,105 @@
# Handoff — GPU bring-up of the v2 torch backend
**For: an agent on a GPU box, branched from `will/mini-fastvideo`.**
**Your job:** take the *written-not-run* `cuda` backend to *runs-and-generates*, then commit + push.
Everything below is committed on `will/mini-fastvideo` and CPU-tested (**204 tests pass**). The torch
path was authored on a machine with **no GPU and no torch**, so it is grounded in the real
`fastvideo` APIs and cross-checked against the source, but **never executed**. That's what you finish.
---
## 0. Orientation (read these first, in order)
1. **`v2/README.md`** — what the whole v2 mini is (the `(recipe, runtime)` runtime; "architecture is
real, kernels are toys"). The "Honest scope" paragraph says exactly what's wired.
2. **`v2/platform/backends/GPU_BRINGUP.md`** — *your checklist*: the ordered 10-step bring-up + the
risk table (A–G), each tied to a `# BRINGUP` marker in the source. **This handoff is orientation +
process; GPU_BRINGUP.md is the work.**
3. This file — the meta-instructions (verify bar, commit/push, gotchas).
## 1. What's already done (commits on this branch)
```
d6d0580a [fix] correct GPU adapters against real fastvideo API (cross-check findings)
b8d78f40 [feat] real torch/CUDA backend (written-not-run) behind the cuda cells
27791b51 [feat] static-buffer capture form for the cudagraph step body (Path A)
ae6a170d [feat] piecewise CUDA-graph capture/replay at the step boundary (Path A)
9308d87e [feat] route diffusion loops through the kernel table
7490c590 [feat] multi-backend dispatch substrate (device/arch/kernel registries)
```
The dispatch substrate (two tuple-keyed registries `COMPONENTS(kind,device,variant)` +
`KERNELS(op,device,arch,variant)`, a detected `Platform`, numpy terminal + parity oracle), the
universal kernel seam (diffusion loops go through `model.platform.kernels`), the
piecewise cudagraph lifecycle, and the torch backend cells are all in place. On a GPU box,
`Platform.detect()` returns a `cuda` platform and resolves the torch cells instead of the numpy toys —
**the inference loops/policies/scheduler are unchanged**; only the resolved implementations differ.
## 2. The files you'll touch
| File | What it is |
|---|---|
| `v2/platform/backends/torch_adapters.py` | `TorchWanDiT` / `TorchWanVAE` / `TorchT5Encoder` — wrap the real `fastvideo.models.*` (named by each card's `load_id`) to the mini's duck-typed surface. Built via the real FastVideo loaders. |
| `v2/platform/backends/torch_kernels.py` | torch `flow_match_step` / `flow_sde_step` (plain elementwise — there is **no** fused solver kernel in fastvideo-kernel; don't look for one). |
| `v2/platform/backends/torch_cuda.py` | registers the `cuda` cells as lazy trampolines (torch imported only inside builder bodies). |
| `v2/card/specs.py` | `ComponentSpec.checkpoint` — the per-component weights source (empty on toys; **you fill it in**). |
The surface the adapters must honor (what the loops call):
`dit(latent, text_embed, sigma) -> velocity` · `vae.decode(latent)` / `vae.encode(video)` ·
`text_encoder.encode(text)`. The CPU toys in `v2/models/backend.py` are the reference behavior.
## 3. Your task (the gating items — full detail in GPU_BRINGUP.md)
1. **Env:** install `torch` + the parent `fastvideo` package + weights. (`fastvideo` source lives at
`/Users/willlin/src/FastVideo`.)
2. **Risk A — the one blocking gap:** the builders call `_load_via_fastvideo(...)` → the real loaders
need a **`FastVideoArgs`**, which `_fastvideo_args(spec)` builds minimally from `spec.checkpoint`.
Confirm/extend its fields (model config, precision, parallelism). And stamp `ComponentSpec.checkpoint`
onto the wan21 card — a tiny helper that maps a model root onto the three components is the cleanest
way (the toy cards leave it `""`).
3. **Work the risk list (A–G in GPU_BRINGUP.md).** The *interface* contracts were cross-checked as
matching (DiT returns bare velocity; `timestep=sigma*1000`; `encode().mode()`; `.last_hidden_state`;
no fused solver kernel) — confirm them numerically. The *construction* layer was fixed (real loaders,
`set_forward_context`, latent normalization, UMT5-from-config). What's left is box-dependent:
`FastVideoArgs` fields, `shift_factor` placement/sign, exact tokenizer kwargs, FSDP sharding.
4. **Bring up in order:** build each component in isolation → one DiT step → one solver step → VAE
decode → full t2v → SDE stochastic sampling → cudagraph capture (last).
## 4. The verification bar (how you know it's right)
- **CPU suite must stay green:** `python3 -m pytest v2/ -q` → still **204 passed**. The torch path is
gated `available=False` off-GPU; importing the backends must never import torch. If you break either,
you broke the substrate. (`v2/tests/test_torch_backend.py` pins these.)
- **Parity oracle is the spec:** the substrate's whole point is that a real backend matches the numpy
reference on the consistency ladder. On GPU, compare a full generation against a known-good fastvideo
output — use the parent repo's SSIM regression harness (`fastvideo/tests/ssim/`). Target C4 (SSIM /
artifact quality); component/trajectory parity (C0/C1) is bit-level vs the reference pipeline.
- **Don't trust "it ran" — trust "it matched."** A wrong `timestep` scale or `shift_factor` produces
plausible-but-wrong video, not a crash (risks B/D). Diff against a reference, don't eyeball.
## 5. Commit + push
- **You are on a GPU branch** (branched from `will/mini-fastvideo`). Commit your bring-up fixes there,
focused by concern (e.g. one commit per confirmed risk), in the existing style (`[fix]`/`[feat] …`).
- **NEVER add Claude as a co-author** (repo policy, `/Users/willlin/src/.claude/CLAUDE.md`).
- **Do not rewrite or force-push** the six commits above — build on top.
- When the CPU suite is green **and** a GPU generation matches the reference, **push your branch.**
- If you launch inference with wandb logging enabled, log in with the token in the project
`CLAUDE.md` (`/Users/willlin/src/.claude/CLAUDE.md`) — **do not paste it into any committed file.**
## 6. Gotchas (don't relearn these the hard way)
- **No fused solver kernel exists.** `fastvideo-kernel` ships only attention/norm/quant primitives;
the cuda `flow_match_step`/`flow_sde_step` are plain torch by design. Don't hunt for a `.cu` solver.
- **`from_pretrained` is not the loader.** `WanTransformer3DModel`/`AutoencoderKLWan` have none — the
real path is the `*Loader().load(model_path, fastvideo_args)` classes in
`fastvideo/models/loader/component_loader.py`. The loader resolves the class from the checkpoint
config (this is what makes UMT5-vs-T5 correct without hardcoding).
- **T5 needs `set_forward_context`.** A bare encoder forward reads stale/None global context.
- **The loop surface stays numpy** for bring-up; adapters marshal numpy↔torch at the boundary. A
torch-native surface (latent on-device through forward→combine→solver) is the **perf follow-up**
(Risk G), not bring-up — don't rewrite `cfg.combine`/`precision.cast`/the samplers yet.
- **cudagraph capture is last.** The wan21 loop declares `breakable_cudagraph`; the v2 capturer models
the lifecycle with a numpy `StaticWorkspace`. Capturing a real `torch.cuda.CUDAGraph` is GPU-only
work and the riskiest step — leave it until inference is verified.
+148
View File
@@ -0,0 +1,148 @@
# FastVideo v2 - Inference Runtime Scope
**Status:** source of truth for `v2/`.
`v2` is the model-native inference runtime for FastVideo. It owns model cards,
programs, loops, runtime execution, serving, cache/memory policy, backend dispatch,
compile/cudagraph integration, and inference parity checks.
`v2` does **not** own training, finetuning, distillation, RL, optimizer steps, or
checkpoint production. Training remains in the existing FastVideo stacks:
- `fastvideo/train/` - the new modular trainer.
- `fastvideo/training/` - the legacy shipped training pipelines.
`v2` may record how a checkpoint was produced through `RecipeSpec` metadata
(`method`, `parents`, `assumes_loop`, `assumes_precision`), because inference must
know which runtime loop and precision policy a post-training checkpoint expects.
That metadata is provenance, not a v2 training API.
## Design Center
FastVideo v2 is video-generation first. Wan/LTX-style diffusion video inference
is the baseline path, and unified models such as BAGEL/Cosmos3 are first-class:
one resident model may run AR, diffusion, VAE, and codec loops in one request.
Audio and TTS are supported as additional modalities on the same stage/loop
model, not as the reason to build a separate universal serving framework.
The core abstraction should stay small:
- `ModelCard` declares the resident components, loops, capabilities, precision,
caches, and sampling defaults a checkpoint needs.
- `Program` is an ordered list of typed nodes passing values through named
slots. It is not a general DAG, Walk graph, or declarative control-flow IR.
- `Loop` owns model semantics. The runtime drives the loop, handles admission,
cancellation, streaming, cache access, and backend dispatch.
Do not add a public contract field until a runtime path consumes it. Future
optimizations such as richer stage placement, multi-GPU transport, or paged KV
should start behind a concrete Wan/BAGEL/Cosmos/Qwen use case and graduate only
after they simplify at least two model recipes.
## Scope
In scope:
- Python inference entrypoint through `v2.VideoGenerator`.
- Typed `ModelCard` declarations for components, loops, capabilities, parity,
sampling defaults, precision, and checkpoint layout.
- Driven inference loops such as diffusion denoise, AR decode, causal/world
continuation, VAE/audio decode, and multi-stage programs.
- Runtime execution through `Engine` and `AsyncEngine`.
- Serving through OpenAI-compatible HTTP/SSE surfaces and deployment cards.
- Backend dispatch through the CPU toy backend, accelerator stand-ins, and the
real torch/CUDA backend.
- Inference acceleration features such as FP8/NVFP4 loading, Sage/Flash/SDPA
attention backend selection, `torch.compile`, cudagraph capture, cache policy,
and component placement.
- Inference parity and regression tests.
Out of scope:
- Training methods, optimizers, loss functions, RL rewards, rollout trainers,
weight-sync training loops, and behavior records for policy updates.
- Training examples under `v2_examples/`.
- Any CLI/API that advertises v2 as a trainer.
- A universal graph runtime, Walk/state-machine authoring layer, or parallelism
vocabulary that is not consumed by the current inference runtime.
## Core Model
The atomic inference artifact is a `(recipe, runtime)` pair:
- `RecipeSpec` records what the weights assume: parent checkpoints, post-training
method name, required loop, and required precision.
- `ModelCard` declares the runtime surface: components, loops, capabilities,
caches, precision, parallelism, sampling defaults, and checkpoint manifest.
- `Program` composes component nodes and loop nodes into a user-facing task as
an ordered named-slot stage list.
- `ModelInstance` is the resident loaded card with shared components, caches,
weight versions, and optional captured graphs.
This keeps post-training artifacts serveable without making `v2` responsible for
creating them.
## Execution Model
Loops are model-owned state machines:
```python
state = loop.init(req, model, ctx)
while True:
plan = loop.next(state)
if isinstance(plan, Done):
break
result = ctx.execute(plan)
state = loop.advance(state, result)
return loop.finalize(state)
```
The loop owns semantics. The runtime owns execution, admission, cancellation,
streaming, cache access, graph capture, and backend dispatch.
Serving is pooled run-to-completion. `AsyncEngine` bounds concurrency by pool
slots; each request runs its program to completion. The synchronous `Engine` is
the offline path used by tests and `VideoGenerator`.
## Package Layout
```text
v2/
video_generator.py public inference facade
registry.py model id -> card builder registry
core/
card/ ModelCard, specs, ModelInstance
loop/ loop contracts, driver, sampler, policies
program/ task programs and workflows
request/ request params, tasks, outputs, sessions
parity/ inference parity helpers
parallel/ named parallel plans
recipes/ model-specific cards, loops, and programs
runtime/ Engine, AsyncEngine, cache, memory, cudagraph, transport
serving/ HTTP/SSE server and deployment adapters
platform/ backend/device/kernel dispatch
_vendor/ vendored FastVideo model/loader/config pieces for inference
tests/ v2 inference/runtime/serving/parity tests
```
There is intentionally no `v2/training/` package.
## Current Inference Path
The torch backend builds real components from stamped checkpoint paths, keeps
components in eval mode, and dispatches inference through the same cards and loops
used by the CPU tests. Wan2.1 T2V inference is the primary real path today.
Wan/FastWan inference supports:
- real Wan component loading through vendored component loaders,
- FP8 post-load quantization for `FastVideo/FastWan-QAD-FP8-1.3B`,
- attention backend selection, including SageAttention when installed,
- `torch.compile` for inference DiT modules,
- on-device latent residency for cards that set `device_io=True`.
## Boundary Rule
If a change adds training behavior, it belongs in `fastvideo/train/` or
`fastvideo/training/`, not in `v2/`. If inference needs to consume the result of
that training, add or update a v2 card, loop, registry entry, checkpoint loader,
sampling defaults, and inference tests.
+89
View File
@@ -0,0 +1,89 @@
"""v2 - the FastVideo inference runtime (see v2/README.md).
> A model card is a (recipe, runtime) pair with a parity obligation.
> The model owns loop semantics; the runtime owns loop lifecycle.
> One resident instance runs many loops; one scheduler runs their steps in one currency.
> Caches are correct by key; parity is correct by test.
v2 is inference-only. Training, finetuning, distillation, RL, and optimizer
loops belong to ``fastvideo/train`` or ``fastvideo/training``. v2 only records
checkpoint provenance in recipe metadata so inference can bind weights to the
right loop and precision policy.
The core is numpy-only and CPU-testable; heavy Wan/LTX neural forwards become lazy torch
adapters (see ``v2/platform/backends/``) that are off the test path.
"""
from __future__ import annotations
from v2.core.enums import (
Capability,
ConsistencyLevel,
ExecutionProfile,
LoopKind,
WorkUnitKind,
)
from v2.core.card import (
CapabilityMatrix,
ComponentSpec,
LoopSpec,
ModelCard,
ModelInstance,
ParitySpec,
RecipeSpec,
load_card,
)
from v2.core.program import ComponentNode, ModelLoopNode, Program, ProgramKind, when_opt, when_task
from v2.core.request import (
DiffusionParams,
Output,
Request,
SamplingParams,
Session,
TaskType,
make_request,
)
from v2.runtime import AsyncEngine, Engine
__version__ = "0.2.0"
__all__ = [
"ModelCard",
"ComponentSpec",
"LoopSpec",
"RecipeSpec",
"ParitySpec",
"CapabilityMatrix",
"ModelInstance",
"load_card",
"Engine",
"AsyncEngine",
"Program",
"ProgramKind",
"ComponentNode",
"ModelLoopNode",
"when_task",
"when_opt",
"Request",
"Session",
"Output",
"make_request",
"TaskType",
"SamplingParams",
"DiffusionParams",
"LoopKind",
"WorkUnitKind",
"ConsistencyLevel",
"ExecutionProfile",
"Capability",
"VideoGenerator",
"__version__",
]
def __getattr__(name: str):
# Lazy: the GPU entrypoint imports torch / fastvideo, so resolve it only on access — plain
# ``import v2`` (and the CPU-only mini) stay torch-free.
if name == "VideoGenerator":
from v2.video_generator import VideoGenerator
return VideoGenerator
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
+1
View File
@@ -0,0 +1 @@
"""Vendored fastvideo code (copied, standalone). Internal layout mirrors upstream for diffing; v2-native code must not edit these ad hoc."""
+28
View File
@@ -0,0 +1,28 @@
"""Slim vendored API config surface for the v2 VideoGenerator.
Only the inference-config dataclasses (schema) + result types are vendored. The fastvideo
parser / presets / overrides modules are intentionally NOT vendored — they pull the fastvideo
pipeline runtime, which v2 replaces. See v2/README.md (vendoring)."""
from __future__ import annotations
from v2._vendor.api.results import GenerationResult
from v2._vendor.api.schema import (
CompileConfig,
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
SamplingConfig,
)
__all__ = [
"CompileConfig",
"EngineConfig",
"GenerationRequest",
"GeneratorConfig",
"OffloadConfig",
"OutputConfig",
"SamplingConfig",
"GenerationResult",
]
+16
View File
@@ -0,0 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
class ConfigValidationError(ValueError):
"""Validation error that keeps track of the nested config path."""
def __init__(self, path: str, message: str):
self.path = path
self.message = message
super().__init__(str(self))
def __str__(self) -> str:
if self.path:
return f"{self.path}: {self.message}"
return self.message
+15
View File
@@ -0,0 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from v2._vendor.api.sampling_param import SamplingParam
@dataclass
class MatrixGame2SamplingParam(SamplingParam):
height: int = 352
width: int = 640
num_frames: int = 57
fps: int = 25
guidance_scale: float = 1.0
num_inference_steps: int = 3
negative_prompt: str | None = None
+233
View File
@@ -0,0 +1,233 @@
# SPDX-License-Identifier: Apache-2.0
"""Track which GenerationRequest fields the user explicitly provided.
When translating a GenerationRequest into a legacy SamplingParam we must
distinguish user-provided values (which should override model defaults)
from schema defaults (which should NOT override model defaults).
The mechanism: a single ``_fastvideo_explicit_paths`` set stored on the
root ``GenerationRequest``. It holds dotted leaf paths (e.g.
``"sampling.guidance_scale"``) the user has touched, either via raw
config at bind time or via attribute assignment at runtime. A patched
``__setattr__`` on the request dataclass types records assignments into
this set.
The set holds leaf paths only. Nested dataclass or mapping assignments
are flattened to their leaves at record time.
"""
from __future__ import annotations
from collections.abc import Callable, Mapping
import dataclasses
from typing import Any, cast
from v2._vendor.api.schema import (
ContinuationState,
GenerationPlan,
GenerationRequest,
InputConfig,
OutputConfig,
PlannedStage,
RequestRuntimeConfig,
RunConfig,
SamplingConfig,
ServeConfig,
)
EXPLICIT_PATHS_ATTR = "_fastvideo_explicit_paths"
_TRACKING_ROOT_ATTR = "_fastvideo_request_tracking_root"
_TRACKING_PATH_ATTR = "_fastvideo_request_tracking_path"
_TRACKING_PATCHED_ATTR = "_fastvideo_request_tracking_patched"
_TRACKED_REQUEST_TYPES = (
GenerationRequest,
InputConfig,
SamplingConfig,
RequestRuntimeConfig,
OutputConfig,
ContinuationState,
PlannedStage,
GenerationPlan,
)
def bind_generation_request_raw(
request: GenerationRequest,
raw: Mapping[str, Any] | None,
) -> GenerationRequest:
"""Install explicit-path tracking on *request*.
*raw* is the parsed config dict (YAML/JSON/kwargs); every leaf key
in it becomes an explicit path. Subsequent attribute assignments on
*request* or its nested dataclasses are recorded automatically via a
patched ``__setattr__``.
"""
_ensure_request_tracking()
# Disable recording while we walk the tree to install roots.
object.__setattr__(request, EXPLICIT_PATHS_ATTR, None)
_set_tracking_roots(request, request, "")
paths: set[str] = set()
_record_value_paths(raw or {}, "", paths)
object.__setattr__(request, EXPLICIT_PATHS_ATTR, paths)
return request
def bind_run_config_raw(
config: RunConfig,
raw: Mapping[str, Any],
) -> RunConfig:
request_raw = raw.get("request")
if isinstance(request_raw, Mapping):
bind_generation_request_raw(config.request, request_raw)
else:
bind_generation_request_raw(config.request, {})
return config
def bind_serve_config_raw(
config: ServeConfig,
raw: Mapping[str, Any],
) -> ServeConfig:
default_request_raw = raw.get("default_request")
if isinstance(default_request_raw, Mapping):
bind_generation_request_raw(config.default_request, default_request_raw)
else:
bind_generation_request_raw(config.default_request, {})
return config
def get_explicit_paths(request: GenerationRequest) -> frozenset[str]:
"""Return a snapshot of the explicit paths set on *request*."""
paths = getattr(request, EXPLICIT_PATHS_ATTR, None)
if isinstance(paths, set | frozenset):
return frozenset(paths)
return frozenset()
def reset_tracking_roots(request: GenerationRequest) -> None:
"""Re-install tracking roots after a deepcopy or manual clone.
The paths set itself deepcopies correctly; we only need to repoint
the tracking root on nested dataclasses at the new root.
"""
_ensure_request_tracking()
_set_tracking_roots(request, request, "")
# ---------------------------------------------------------------------------
# Path recording
# ---------------------------------------------------------------------------
def _record_value_paths(
value: Any,
prefix: str,
out: set[str],
) -> None:
"""Add every leaf path under *value* to *out*.
A leaf is any terminal value (non-dataclass, non-mapping, or empty
mapping/dataclass). ``prefix`` is the dotted path at which *value*
sits. When called with an empty ``prefix`` (the root), leaves are
recorded at their own key.
"""
if dataclasses.is_dataclass(value) and not isinstance(value, type):
dc_fields = dataclasses.fields(value)
if not dc_fields:
if prefix:
out.add(prefix)
return
for field in dc_fields:
child = getattr(value, field.name)
path = f"{prefix}.{field.name}" if prefix else field.name
_record_value_paths(child, path, out)
return
if isinstance(value, Mapping):
if not value:
if prefix:
out.add(prefix)
return
for key, child in value.items():
path = f"{prefix}.{key}" if prefix else key
_record_value_paths(child, path, out)
return
if prefix:
out.add(prefix)
# ---------------------------------------------------------------------------
# __setattr__ patching
# ---------------------------------------------------------------------------
def _ensure_request_tracking() -> None:
for config_type in _TRACKED_REQUEST_TYPES:
_patch_tracking_setattr(config_type)
def _patch_tracking_setattr(config_type: type[Any]) -> None:
if getattr(config_type, _TRACKING_PATCHED_ATTR, False):
return
original_setattr = cast(
Callable[[Any, str, Any], None],
config_type.__setattr__,
)
field_names = {field.name for field in dataclasses.fields(config_type)}
def _tracking_setattr(self: Any, name: str, value: Any) -> None:
if name.startswith("_fastvideo_") or name not in field_names:
original_setattr(self, name, value)
return
original_setattr(self, name, value)
root = getattr(self, _TRACKING_ROOT_ATTR, None)
if root is None:
return
paths = getattr(root, EXPLICIT_PATHS_ATTR, None)
if not isinstance(paths, set):
return
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
path = f"{prefix}.{name}" if prefix else name
# Wholesale dataclass replacement: install roots on the new
# instance so its future mutations are tracked too.
if dataclasses.is_dataclass(value) and not isinstance(value, type):
_set_tracking_roots(root, value, path)
_record_value_paths(value, path, paths)
type.__setattr__(config_type, "__setattr__", _tracking_setattr)
setattr(config_type, _TRACKING_PATCHED_ATTR, True)
# ---------------------------------------------------------------------------
# Tree walk to set tracking root/path on nested dataclasses
# ---------------------------------------------------------------------------
def _set_tracking_roots(
root: GenerationRequest,
obj: Any,
prefix: str,
) -> None:
if not dataclasses.is_dataclass(obj) or isinstance(obj, type):
return
object.__setattr__(obj, _TRACKING_ROOT_ATTR, root)
object.__setattr__(obj, _TRACKING_PATH_ATTR, prefix)
for field in dataclasses.fields(obj):
child = getattr(obj, field.name)
child_path = f"{prefix}.{field.name}" if prefix else field.name
if dataclasses.is_dataclass(child) and not isinstance(child, type):
_set_tracking_roots(root, child, child_path)
__all__ = [
"EXPLICIT_PATHS_ATTR",
"bind_generation_request_raw",
"bind_run_config_raw",
"bind_serve_config_raw",
"get_explicit_paths",
"reset_tracking_roots",
]
+173
View File
@@ -0,0 +1,173 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
from collections.abc import Mapping
from v2._vendor.api.schema import ContinuationState
@dataclass
class GenerationResult:
prompt: str | None = None
prompt_index: int | None = None
samples: Any | None = None
frames: Any | None = None
audio: Any | None = None
audio_sample_rate: int | None = None
size: tuple[int, int, int] | None = None
generation_time: float | None = None
logging_info: Any | None = None
trajectory: Any | None = None
trajectory_timesteps: Any | None = None
trajectory_decoded: Any | None = None
video_path: str | None = None
peak_memory_mb: float | None = None
state: ContinuationState | None = None
extra: dict[str, Any] = field(default_factory=dict)
@classmethod
def from_legacy_result(
cls,
result: Mapping[str, Any],
) -> GenerationResult:
prompt = result.get("prompt")
if prompt is None:
prompt = result.get("prompts")
extra = {
key: value
for key, value in result.items() if key not in {
"prompt",
"prompt_index",
"prompts",
"samples",
"frames",
"audio",
"audio_sample_rate",
"size",
"generation_time",
"logging_info",
"trajectory",
"trajectory_timesteps",
"trajectory_decoded",
"video_path",
"peak_memory_mb",
"state",
}
}
return cls(
prompt=prompt,
prompt_index=result.get("prompt_index"),
samples=result.get("samples"),
frames=result.get("frames"),
audio=result.get("audio"),
audio_sample_rate=result.get("audio_sample_rate"),
size=result.get("size"),
generation_time=result.get("generation_time"),
logging_info=result.get("logging_info"),
trajectory=result.get("trajectory"),
trajectory_timesteps=result.get("trajectory_timesteps"),
trajectory_decoded=result.get("trajectory_decoded"),
video_path=result.get("video_path"),
peak_memory_mb=result.get("peak_memory_mb"),
state=result.get("state"),
extra=extra,
)
def to_legacy_dict(self) -> dict[str, Any]:
result = {
"prompts": self.prompt,
"samples": self.samples,
"frames": self.frames,
"audio": self.audio,
"audio_sample_rate": self.audio_sample_rate,
"size": self.size,
"generation_time": self.generation_time,
"logging_info": self.logging_info,
"trajectory": self.trajectory,
"trajectory_timesteps": self.trajectory_timesteps,
"trajectory_decoded": self.trajectory_decoded,
"video_path": self.video_path,
"peak_memory_mb": self.peak_memory_mb,
}
if self.prompt_index is not None:
result["prompt_index"] = self.prompt_index
result["prompt"] = self.prompt
if self.state is not None:
result["state"] = self.state
result.update(self.extra)
return result
# Alias the canonical result type; matches the public docs.
VideoResult = GenerationResult
@dataclass
class VideoProgressEvent:
"""Per-step progress event emitted by :meth:`VideoGenerator.generate_async`.
Consumers treat these as best-effort telemetry; ``total_steps`` is
the count the pipeline reported at the start of the run, not a
rolling estimate.
"""
step: int
total_steps: int
stage: str = "denoise"
"""Logical stage name (``denoise`` | ``refine`` | ``decode`` | …)."""
@dataclass
class VideoPartialEvent:
"""Chunk of decoded frames ready for streaming.
Emitted only on the streaming path; the aggregated code path never
yields partials. ``frames`` is a numpy ``(N, H, W, 3)`` uint8
ndarray; ``index`` is a monotonic chunk index starting at 0.
"""
frames: Any
index: int
@dataclass
class VideoFinalEvent:
"""Terminal event carrying the generated video and metadata.
Exactly one ``VideoFinalEvent`` is emitted per request. When
``request.output.return_state`` is True the event also carries the
:class:`ContinuationState` the caller needs to resume.
"""
video_bytes: bytes | None = None
tensor: Any | None = None
frames: Any | None = None
metadata: dict[str, Any] = field(default_factory=dict)
continuation_state: ContinuationState | None = None
result: VideoResult | None = None
"""The full :class:`VideoResult` for callers that want everything.
Streaming consumers typically only care about ``frames`` /
``continuation_state``; keeping the full result here avoids a
second code path."""
VideoEvent = VideoProgressEvent | VideoPartialEvent | VideoFinalEvent
"""Union of every event :meth:`VideoGenerator.generate_async` yields.
Consumers match by ``isinstance`` rather than ``type`` so subclasses
(e.g. a future ``VideoAudioSegmentEvent``) slot in without breaking
existing code."""
__all__ = [
"GenerationResult",
"VideoEvent",
"VideoFinalEvent",
"VideoPartialEvent",
"VideoProgressEvent",
"VideoResult",
]
+411
View File
@@ -0,0 +1,411 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import copy
from dataclasses import dataclass, field, fields
from typing import TYPE_CHECKING, Any
from v2._vendor.logger import init_logger
from v2._vendor.utils import StoreBoolean
if TYPE_CHECKING:
from v2._vendor.api.schema import ContinuationState
logger = init_logger(__name__)
@dataclass
class SamplingParam:
"""
Sampling parameters for video generation.
"""
# All fields below are copied from ForwardBatch
data_type: str = "video"
# Image inputs
image_path: str | None = None
pil_image: Any | None = None
# Video inputs
video_path: str | None = None
# Optional pre-generated diffusion latents. Used by parity/debug harnesses
# and advanced callers that need deterministic latent reuse.
latents: Any | None = None
# Action control inputs (Matrix-Game)
mouse_cond: Any | None = None # Shape: (B, T, 2)
keyboard_cond: Any | None = None # Shape: (B, T, K)
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
# Camera/action control inputs (GameCraft)
camera_states: Any | None = None # Plücker coordinates [B, T_video, 6, H, W]
camera_trajectory: str | None = None
action_list: list[str] | None = None
action_speed_list: list[float] | None = None
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
# Camera control inputs (LingBotWorld)
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
# Refine inputs (LongCat 480p->720p upscaling)
# Path-based refine (load stage1 video from disk, e.g. MP4)
refine_from: str | None = None # Path to stage1 video (480p output from distill)
t_thresh: float = 0.5 # Threshold for timestep scheduling in refinement
spatial_refine_only: bool = False # If True, only spatial (no temporal doubling)
num_cond_frames: int = 0 # Number of conditioning frames
# In-memory refine input (for two-stage pipeline where stage1 frames are already in memory)
# This mirrors LongCat's demo where a list of frames (e.g. np.ndarray or PIL.Image)
# is passed directly to the refinement pipeline instead of reloading from disk.
stage1_video: Any | None = None
# Text inputs
prompt: str | list[str] | None = None
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
max_sequence_length: int | None = None
prompt_path: str | None = None
output_path: str = "outputs/"
output_video_name: str | None = None
# Batch info
num_videos_per_prompt: int = 1
seed: int = 1024
# Original dimensions (before VAE scaling)
num_frames: int = 125
height: int = 720
width: int = 1280
height_sr: int = 1072
width_sr: int = 1920
fps: int = 24
# Denoising parameters
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
sigmas: list[float] | None = None
# TeaCache parameters
enable_teacache: bool = False
# GEN3C camera control
trajectory_type: str | None = None
movement_distance: float | None = None
camera_rotation: str | None = None
# LTX-2 multi-modal CFG and STG.
# Class-level defaults match the *distilled* LTX-2 schedule
# (mirrors ``FastVideo-internal/.../LTX2DistilledSamplingParam``):
# the distilled model expects neutral guidance scales — modality 1,
# rescale 0, STG 0 — and explicit-CFG callers (full LTX-2) opt back
# in by selecting the ``LTX2_BASE`` preset, which overrides these
# to mod=3.0 / rescale=0.7 / stg=1.0 in its ``defaults`` dict.
# cfg_scale defaults stay at 1.0 (CFG off) so
# ``ForwardBatch.__post_init__`` doesn't force CFG on non-LTX-2
# models that never override these fields.
ltx2_cfg_scale_video: float = 1.0
ltx2_cfg_scale_audio: float = 1.0
ltx2_modality_scale_video: float = 1.0
ltx2_modality_scale_audio: float = 1.0
ltx2_rescale_scale: float = 0.0
ltx2_stg_scale_video: float = 0.0
ltx2_stg_scale_audio: float = 0.0
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
# LTX-2 image / video / continuation conditioning. These flow from
# generate_video(...) kwargs through ``sampling_param.update(kwargs)``
# onto the ForwardBatch fields of the same name. ``ltx2_image_crf``
# gates the conditioning-image H.264 re-encode; the streaming
# session controller passes ``ltx2_image_crf=0.0`` because it
# conditions on already-decoded VAE-quality frames.
ltx2_images: list[tuple[str, int, float]] | None = None
ltx2_image_crf: float = 33.0
ltx2_conditioning_latent_stage1: Any | None = None
ltx2_conditioning_latent_stage2: Any | None = None
ltx2_video_conditions: list[tuple[list[str], int, float]] | None = None
# Stable Audio (T2A): clip start/end in seconds. Honored by
# `StableAudioConditioningStage` + `StableAudioDecodingStage`. Other
# families ignore them.
audio_start_in_s: float | None = None
audio_end_in_s: float | None = None
# Stable Audio audio-to-audio (variation):
# `init_audio` -- a path or `[B, C, samples]` waveform at the model
# sample rate; the pipeline encodes it via the VAE
# and uses it as the starting latent.
# `init_audio_strength` -- 0..1, higher = closer to the reference
# (matches the convention of Stability's
# commercial Stable Audio 2.0 UI). 1.0 ~=
# VAE round-trip, 0.0 ~= plain T2A.
# `init_noise_level` -- legacy raw `sigma_max` override (0.3..500,
# higher = more freedom). Kept for callers
# that already use it; prefer `init_audio_strength`.
init_audio: Any = None
init_audio_strength: float | None = None
init_noise_level: float | None = None
# Stable Audio inpainting (RePaint-style): `inpaint_audio` is the
# reference clip, `inpaint_mask` is a [samples] tensor in {0, 1} where
# 1 means *keep the reference* and 0 means *regenerate*.
inpaint_audio: Any = None
inpaint_mask: Any = None
# Continuation state carried across streaming/multi-segment calls.
continuation_state: ContinuationState | None = None
# When True, the pipeline returns a ContinuationState on the result so
# the caller can resume from the generated segment.
return_continuation_state: bool = False
# Misc
save_video: bool = True
return_frames: bool = True
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
def check_sampling_param(self):
if self.prompt_path and not self.prompt_path.endswith(".txt"):
raise ValueError("prompt_path must be a txt file")
def update(self, source_dict: dict[str, Any]) -> None:
valid_fields = {f.name for f in fields(self)}
unknown = [key for key in source_dict if key not in valid_fields]
if unknown:
raise ValueError(f"{type(self).__name__}.update() received unknown field(s): "
f"{sorted(unknown)}. All kwargs must correspond to declared "
f"SamplingParam fields. If a kwarg is meant to flow into "
f"ForwardBatch.extra (e.g. LTX2 audio conditioning), route it "
f"via VideoGenerator._BATCH_EXTRA_PASSTHROUGH_KEYS instead.")
for key, value in source_dict.items():
setattr(self, key, value)
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str) -> SamplingParam:
sampling_param = cls._from_preset(model_path)
if sampling_param is not None:
return sampling_param
logger.warning(
"Couldn't find a preset for %s."
" Using the default sampling param.",
model_path,
)
return cls()
@classmethod
def _from_preset(
cls,
model_path: str,
) -> SamplingParam | None:
"""Build a SamplingParam from preset defaults.
Returns ``None`` when no preset is configured for
*model_path*, letting the caller fall back to the legacy
subclass lookup.
"""
from v2.registry import get_preset_selection
try:
preset_name, model_family = get_preset_selection(model_path)
except (ValueError, RuntimeError):
return None
if preset_name is None or model_family is None:
return None
from v2._vendor.api.presets import get_preset
preset = get_preset(preset_name, model_family)
sp = cls()
valid_fields = {f.name for f in fields(cls)}
for key, value in preset.defaults.items():
if key in valid_fields:
setattr(sp, key, copy.deepcopy(value))
sp.__post_init__()
return sp
@staticmethod
def add_cli_args(parser: Any) -> Any:
"""Add CLI arguments for SamplingParam fields"""
parser.add_argument(
"--prompt",
type=str,
default=SamplingParam.prompt,
help="Text prompt for video generation",
)
parser.add_argument(
"--negative-prompt",
type=str,
default=SamplingParam.negative_prompt,
help="Negative text prompt for video generation",
)
parser.add_argument(
"--prompt-path",
type=str,
default=SamplingParam.prompt_path,
help="Path to a text file containing the prompt",
)
parser.add_argument(
"--output-path",
type=str,
default=SamplingParam.output_path,
help="Path to save the generated video",
)
parser.add_argument(
"--output-video-name",
type=str,
default=SamplingParam.output_video_name,
help="Name of the output video",
)
parser.add_argument(
"--num-videos-per-prompt",
type=int,
default=SamplingParam.num_videos_per_prompt,
help="Number of videos to generate per prompt",
)
parser.add_argument(
"--seed",
type=int,
default=SamplingParam.seed,
help="Random seed for generation",
)
parser.add_argument(
"--num-frames",
type=int,
default=SamplingParam.num_frames,
help="Number of frames to generate",
)
parser.add_argument(
"--height",
type=int,
default=SamplingParam.height,
help="Height of generated video",
)
parser.add_argument(
"--width",
type=int,
default=SamplingParam.width,
help="Width of generated video",
)
parser.add_argument(
"--fps",
type=int,
default=SamplingParam.fps,
help="Frames per second for saved video",
)
parser.add_argument(
"--num-inference-steps",
type=int,
default=SamplingParam.num_inference_steps,
help="Number of denoising steps",
)
parser.add_argument(
"--guidance-scale",
type=float,
default=SamplingParam.guidance_scale,
help="Classifier-free guidance scale",
)
parser.add_argument(
"--guidance-rescale",
type=float,
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--boundary-ratio",
type=float,
default=SamplingParam.boundary_ratio,
help="Boundary timestep ratio",
)
parser.add_argument(
"--save-video",
action="store_true",
default=SamplingParam.save_video,
help="Whether to save the video to disk",
)
parser.add_argument(
"--no-save-video",
action="store_false",
dest="save_video",
help="Don't save the video to disk",
)
parser.add_argument(
"--return-frames",
action="store_true",
default=False,
help="Whether to return the raw frames",
)
parser.add_argument(
"--image-path",
type=str,
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--video-path",
type=str,
default=SamplingParam.video_path,
help="Path to input video for video-to-video generation",
)
parser.add_argument(
"--refine-from",
type=str,
default=SamplingParam.refine_from,
help="Path to stage1 video for refinement (LongCat 480p->720p)",
)
parser.add_argument(
"--t-thresh",
type=float,
default=SamplingParam.t_thresh,
help="Threshold for timestep scheduling in refinement (default: 0.5)",
)
parser.add_argument(
"--spatial-refine-only",
action=StoreBoolean,
default=SamplingParam.spatial_refine_only,
help="Only perform spatial super-resolution (no temporal doubling)",
)
parser.add_argument(
"--num-cond-frames",
type=int,
default=SamplingParam.num_cond_frames,
help="Number of conditioning frames for refinement",
)
parser.add_argument(
"--moba-config-path",
type=str,
default=None,
help="Path to a JSON file containing V-MoBA specific configurations.",
)
parser.add_argument(
"--return-trajectory-latents",
action="store_true",
default=SamplingParam.return_trajectory_latents,
help="Whether to return the trajectory",
)
parser.add_argument(
"--return-trajectory-decoded",
action="store_true",
default=SamplingParam.return_trajectory_decoded,
help="Whether to return the decoded trajectory",
)
return parser
@dataclass
class CacheParams:
cache_type: str = "none"
+307
View File
@@ -0,0 +1,307 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
@dataclass
class ServerConfig:
host: str = "0.0.0.0"
port: int = 8000
output_dir: str = "outputs/"
@dataclass
class ParallelismConfig:
tp_size: int = -1
sp_size: int = -1
hsdp_replicate_dim: int = 1
hsdp_shard_dim: int = -1
dist_timeout: int | None = None
@dataclass
class OffloadConfig:
dit: bool = True
dit_layerwise: bool = True
text_encoder: bool = True
image_encoder: bool = True
vae: bool = True
pin_cpu_memory: bool = True
@dataclass
class CompileConfig:
"""Typed ``torch.compile`` configuration.
``backend``/``fullgraph``/``mode``/``dynamic`` are the four most
common ``torch.compile`` knobs. ``extras`` holds any remaining
``torch.compile`` kwargs (e.g. ``options``, ``disable``).
The ``enabled`` switch covers the DiT transformer path (including
``transformer_2`` and the LTX-2 stage-2 ``transformer_refine``).
Per-component flags below are independent overlays — set to ``True``
to compile that component, ``None`` to leave it eager. Each
``*_kwargs`` dict overrides the master ``backend``/``fullgraph``/
``mode``/``dynamic``/``extras`` for that component when non-empty;
leaving it empty inherits the master kwargs.
"""
enabled: bool = False
backend: str | None = None
fullgraph: bool | None = None
mode: str | None = None
dynamic: bool | None = None
extras: dict[str, Any] = field(default_factory=dict)
text_encoder_enabled: bool | None = None
vae_enabled: bool | None = None
audio_vae_enabled: bool | None = None
dit_kwargs: dict[str, Any] = field(default_factory=dict)
text_encoder_kwargs: dict[str, Any] = field(default_factory=dict)
vae_kwargs: dict[str, Any] = field(default_factory=dict)
audio_vae_kwargs: dict[str, Any] = field(default_factory=dict)
@dataclass
class QuantizationConfig:
text_encoder_quant: str | None = None
transformer_quant: str | None = None
@dataclass
class EngineConfig:
num_gpus: int = 1
execution_backend: Literal["mp", "ray"] = "mp"
parallelism: ParallelismConfig = field(default_factory=ParallelismConfig)
offload: OffloadConfig = field(default_factory=OffloadConfig)
compile: CompileConfig = field(default_factory=CompileConfig)
enable_stage_verification: bool = True
use_fsdp_inference: bool = False
disable_autocast: bool = False
quantization: QuantizationConfig | None = None
@dataclass
class ComponentConfig:
config_root: str | None = None
pipeline_config_path: str | None = None
text_encoder_weights: str | None = None
transformer_weights: str | None = None
transformer_2_weights: str | None = None
vae_weights: str | None = None
upsampler_weights: str | None = None
lora_path: str | None = None
override_pipeline_cls_name: str | None = None
override_transformer_cls_name: str | None = None
@dataclass
class PipelineSelection:
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
preset: str | None = None
preset_version: int | None = None
components: ComponentConfig = field(default_factory=ComponentConfig)
vae_tiling: bool | None = None
"""Tile-based VAE decode. ``None`` keeps the model's default."""
preset_overrides: dict[str, Any] = field(default_factory=dict)
experimental: dict[str, Any] = field(default_factory=dict)
@dataclass
class GeneratorConfig:
model_path: str
revision: str | None = None
trust_remote_code: bool = False
engine: EngineConfig = field(default_factory=EngineConfig)
pipeline: PipelineSelection = field(default_factory=PipelineSelection)
@dataclass
class InputConfig:
prompt_path: str | None = None
image_path: str | list[str] | None = None
video_path: str | list[str] | None = None
pil_image: Any | None = None
pose: str | None = None
mouse_cond: Any | None = None
keyboard_cond: Any | None = None
grid_sizes: Any | None = None
c2ws_plucker_emb: Any | None = None
refine_from: str | None = None
stage1_video: Any | None = None
@dataclass
class SamplingConfig:
num_videos_per_prompt: int = 1
seed: int = 1024
num_frames: int = 125
height: int = 720
width: int = 1280
height_sr: int = 1072
width_sr: int = 1920
fps: int = 24
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
guidance_rescale: float = 0.0
true_cfg_scale: float | None = None
boundary_ratio: float | None = None
sigmas: list[float] | None = None
@dataclass
class RequestRuntimeConfig:
enable_teacache: bool = False
return_trajectory_latents: bool = False
return_trajectory_decoded: bool = False
@dataclass
class OutputConfig:
output_path: str = "outputs/"
output_video_name: str | None = None
save_video: bool = True
return_frames: bool = True
return_state: bool = False
@dataclass
class ContinuationState:
kind: str
payload: dict[str, Any]
@dataclass
class PlannedStage:
name: str
kind: str
source: str | None = None
overrides: dict[str, Any] = field(default_factory=dict)
@dataclass
class GenerationPlan:
stages: list[PlannedStage]
final_stage: str | None = None
@dataclass
class GenerationRequest:
prompt: str | list[str] | None = None
negative_prompt: str | None = None
inputs: InputConfig = field(default_factory=InputConfig)
sampling: SamplingConfig = field(default_factory=SamplingConfig)
runtime: RequestRuntimeConfig = field(default_factory=RequestRuntimeConfig)
output: OutputConfig = field(default_factory=OutputConfig)
stage_overrides: dict[str, Any] = field(default_factory=dict)
state: ContinuationState | None = None
plan: GenerationPlan | None = None
extensions: dict[str, Any] = field(default_factory=dict)
@dataclass
class RunConfig:
generator: GeneratorConfig
request: GenerationRequest
@dataclass
class WarmupConfig:
enabled: bool = True
prompt: str = ("A cinematic drone shot over coastal cliffs at sunrise, "
"golden light, gentle ocean waves, ultra detailed")
timeout_seconds: int = 2400
@dataclass
class GpuPoolConfig:
num_workers: int | None = None
enable_audio_reencode: bool = True
conditioning_num_frames: int = 9
conditioning_end_offset: int = 0
@dataclass
class PromptEnhancerConfig:
enabled: bool = False
provider: Literal["cerebras", "groq"] = "cerebras"
model: str = "gpt-oss-120b"
timeout_ms: int = 20000
system_prompt_dir: str | None = None
@dataclass
class PromptSafetyConfig:
enabled: bool = False
classifier_path: str | None = None
@dataclass
class StreamingConfig:
session_timeout_seconds: int = 300
generation_segment_cap: int = 6
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
warmup: WarmupConfig = field(default_factory=WarmupConfig)
pool: GpuPoolConfig = field(default_factory=GpuPoolConfig)
prompt: PromptEnhancerConfig = field(default_factory=PromptEnhancerConfig)
safety: PromptSafetyConfig = field(default_factory=PromptSafetyConfig)
@dataclass
class ServeConfig:
"""Typed serve config loaded from ``fastvideo serve --config``.
``default_request`` is a full :class:`GenerationRequest` — the same type
clients POST to ``/v1/videos``. At request time the server merges it into
the incoming body as the operator-pinned baseline.
Important nuance: only fields the operator **explicitly wrote** in the
serve YAML/JSON count as defaults. Although the in-memory object is
fully populated (schema defaults fill every unset field), the merge
walks ``_fastvideo_explicit_paths`` — populated during parse — so
unset fields are *not* forced onto requests. Per-request precedence:
body (client-explicit) > default_request (operator-explicit)
> hardcoded fallback (e.g. ``fps=24``)
See :func:`v2._vendor.api.compat.explicit_request_updates` for the
projection and ``entrypoints/openai/video_api.py::_build_generation_kwargs``
for the merge.
"""
generator: GeneratorConfig
server: ServerConfig = field(default_factory=ServerConfig)
default_request: GenerationRequest = field(default_factory=GenerationRequest)
streaming: StreamingConfig | None = None
__all__ = [
"CompileConfig",
"ComponentConfig",
"ContinuationState",
"EngineConfig",
"GenerationPlan",
"GenerationRequest",
"GeneratorConfig",
"GpuPoolConfig",
"InputConfig",
"OffloadConfig",
"OutputConfig",
"ParallelismConfig",
"PipelineSelection",
"PlannedStage",
"PromptEnhancerConfig",
"PromptSafetyConfig",
"QuantizationConfig",
"RequestRuntimeConfig",
"RunConfig",
"SamplingConfig",
"ServeConfig",
"ServerConfig",
"StreamingConfig",
"WarmupConfig",
]
+58
View File
@@ -0,0 +1,58 @@
# `fastvideo/attention/` — Attention Backends
**Generated:** 2026-05-02
Backend registry + selector wrapping FlashAttn / SageAttn / SageAttn3 / SDPA / VSA / VMoBA / SLA / BSA.
## Layout
```
attention/
├── __init__.py # Exports DistributedAttention, LocalAttention, get_attn_backend
├── layer.py # DistributedAttention, DistributedAttention_VSA, LocalAttention
├── selector.py # get_attn_backend (cached) + env-var override
├── backends/
│ ├── abstract.py # AttentionBackend / AttentionMetadata / AttentionMetadataBuilder
│ ├── flash_attn.py # FA2/FA3
│ ├── sage_attn.py # SageAttention v1
│ ├── sage_attn3.py # SageAttention v3
│ ├── sdpa.py # torch SDPA fallback
│ ├── video_sparse_attn.py # VSA (paper: Video Sparse Attention)
│ ├── vmoba.py # Video-MoBA
│ ├── sla.py # Sliding-window (STA)
│ └── bsa_attn.py # Block-sparse
└── utils/
├── flash_attn_cute.py
└── flash_attn_no_pad.py
```
## Selection Order
`get_attn_backend()` resolves via:
1. Env-var override `FASTVIDEO_ATTENTION_BACKEND` (see `STR_BACKEND_ENV_VAR` in `fastvideo/utils.py`).
2. Per-platform default from `fastvideo/platforms/`.
3. Heuristic fallback to SDPA.
The result is `@lru_cache`d. Tests that need a specific backend must use the
`global_force_attn_backend(...)` context manager from `selector.py`, never set
the env var mid-process.
## Adding a Backend
1. Subclass `AttentionBackend` in `backends/<name>.py`.
2. Implement `AttentionMetadata` + `AttentionMetadataBuilder` for the new path.
3. Register the enum value in `fastvideo/platforms/interface.py` (`AttentionBackendEnum`).
4. Wire string → class resolution in `selector.py`.
5. Verify the new backend works with `DistributedAttention` (sequence parallel)
and `LocalAttention` (single-rank). If it cannot support SP, document the
gap in the backend file's module docstring.
## Anti-Patterns
- Calling `torch.nn.functional.scaled_dot_product_attention` directly inside a
model's forward — go through `DistributedAttention` / `LocalAttention`.
- Reading `os.environ[STR_BACKEND_ENV_VAR]` from arbitrary call sites. Use
`get_env_variable_attn_backend()`.
- Caching backend instances per-module. The selector cache is process-wide; do
not duplicate it.
+16
View File
@@ -0,0 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
from v2._vendor.attention.backends.abstract import (AttentionBackend, AttentionMetadata, AttentionMetadataBuilder)
from v2._vendor.attention.layer import (DistributedAttention, DistributedAttention_VSA, LocalAttention)
from v2._vendor.attention.selector import get_attn_backend
__all__ = [
"DistributedAttention",
"LocalAttention",
"DistributedAttention_VSA",
"AttentionBackend",
"AttentionMetadata",
"AttentionMetadataBuilder",
# "AttentionState",
"get_attn_backend",
]
+177
View File
@@ -0,0 +1,177 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/backends/abstract.py
from abc import ABC, abstractmethod
from dataclasses import dataclass, field, fields
from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar
if TYPE_CHECKING:
pass
import torch
class AttentionBackend(ABC):
"""Abstract class for attention backends."""
# For some attention backends, we allocate an output tensor before
# calling the custom op. When piecewise cudagraph is enabled, this
# makes sure the output tensor is allocated inside the cudagraph.
accept_output_buffer: bool = False
@staticmethod
@abstractmethod
def get_name() -> str:
raise NotImplementedError
@staticmethod
@abstractmethod
def get_impl_cls() -> type["AttentionImpl"]:
raise NotImplementedError
@staticmethod
@abstractmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
# @staticmethod
# @abstractmethod
# def get_state_cls() -> Type["AttentionState"]:
# raise NotImplementedError
# @classmethod
# def make_metadata(cls, *args, **kwargs) -> "AttentionMetadata":
# return cls.get_metadata_cls()(*args, **kwargs)
@staticmethod
@abstractmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
@dataclass
class AttentionMetadata:
"""Attention metadata for prefill and decode batched together."""
# Current step of diffusion process
current_timestep: int
VSA_sparsity: float = field(default=0.0, kw_only=True)
def __getattr__(self, name: str) -> Any:
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
def asdict_zerocopy(self, skip_fields: set[str] | None = None) -> dict[str, Any]:
"""Similar to dataclasses.asdict, but avoids deepcopying."""
if skip_fields is None:
skip_fields = set()
# Note that if we add dataclasses as fields, they will need
# similar handling.
return {field.name: getattr(self, field.name) for field in fields(self) if field.name not in skip_fields}
T = TypeVar("T", bound=AttentionMetadata)
class AttentionMetadataBuilder(ABC, Generic[T]):
"""Abstract class for attention metadata builders."""
@abstractmethod
def __init__(self) -> None:
"""Create the builder, remember some configuration and parameters."""
raise NotImplementedError
@abstractmethod
def prepare(self) -> None:
"""Prepare for one batch."""
raise NotImplementedError
@abstractmethod
def build(
self,
**kwargs: Any,
) -> AttentionMetadata:
"""Build attention metadata with on-device tensors."""
raise NotImplementedError
class AttentionLayer(Protocol):
_k_scale: torch.Tensor
_v_scale: torch.Tensor
_k_scale_float: float
_v_scale_float: float
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
kv_cache: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
...
class AttentionImpl(ABC, Generic[T]):
@abstractmethod
def __init__(
self,
num_heads: int,
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
raise NotImplementedError
def preprocess_qkv(self, qkv: torch.Tensor, attn_metadata: T) -> torch.Tensor:
"""Preprocess QKV tensor before performing attention operation.
Default implementation returns the tensor unchanged.
Subclasses can override this to implement custom preprocessing
like reshaping, tiling, scaling, or other transformations.
Called AFTER all_to_all for distributed attention
Args:
qkv: The query-key-value tensor
attn_metadata: Metadata for the attention operation
Returns:
Processed QKV tensor
"""
return qkv
def postprocess_output(
self,
output: torch.Tensor,
attn_metadata: T,
) -> torch.Tensor:
"""Postprocess the output tensor after the attention operation.
Default implementation returns the tensor unchanged.
Subclasses can override this to implement custom postprocessing
like untiling, scaling, or other transformations.
Called BEFORE all_to_all for distributed attention
Args:
output: The output tensor from the attention operation
attn_metadata: Metadata for the attention operation
Returns:
Postprocessed output tensor
"""
return output
@abstractmethod
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: T,
) -> torch.Tensor:
raise NotImplementedError
@@ -0,0 +1,125 @@
# SPDX-License-Identifier: Apache-2.0
import importlib
import sys
from collections.abc import Callable
from pathlib import Path
import torch
from v2._vendor.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
_project_root = Path(__file__).resolve().parent.parent.parent.parent
_kernel_root = _project_root / "fastvideo-kernel"
_kernel_python_root = _kernel_root / "python"
_attn_qat_infer: Callable[..., torch.Tensor] | None = None
_attn_qat_infer_import_attempted = False
def _ensure_kernel_paths() -> None:
for path in (_project_root, _kernel_root, _kernel_python_root):
path_str = str(path)
if path_str not in sys.path:
sys.path.insert(0, path_str)
def _get_attn_qat_infer() -> Callable[..., torch.Tensor] | None:
global _attn_qat_infer
global _attn_qat_infer_import_attempted
if _attn_qat_infer_import_attempted:
return _attn_qat_infer
_attn_qat_infer_import_attempted = True
_ensure_kernel_paths()
try:
# Prefer the in-repo kernel implementation during local development.
_attn_qat_infer = importlib.import_module("attn_qat_infer").sageattn_blackwell
except ImportError:
_attn_qat_infer = None
return _attn_qat_infer
def is_attn_qat_infer_available() -> bool:
return _get_attn_qat_infer() is not None
class AttnQatInferBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@staticmethod
def get_name() -> str:
return "ATTN_QAT_INFER"
@staticmethod
def get_impl_cls() -> type["AttnQatInferImpl"]:
return AttnQatInferImpl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder[AttentionMetadata]"]:
raise NotImplementedError
class AttnQatInferImpl(AttentionImpl[AttentionMetadata]):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
dropout_p = extra_impl_args.get("dropout_p", 0.0)
if dropout_p > 0:
raise NotImplementedError(f"attn_qat_infer does not support dropout (got dropout_p={dropout_p}). "
"The QAT inference kernel applies no stochastic dropout.")
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
attn_qat_infer = _get_attn_qat_infer()
if attn_qat_infer is None:
raise ImportError("attn_qat_infer is not available. Please ensure the "
"attn_qat_infer kernel package is installed.")
query = query.transpose(1, 2).contiguous()
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
output = attn_qat_infer(
query,
key,
value,
attn_mask=None,
is_causal=self.causal,
sm_scale=self.softmax_scale,
)
return output.transpose(1, 2).contiguous()
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
import importlib
import sys
from collections.abc import Callable
from pathlib import Path
import torch
from v2._vendor.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
_project_root = Path(__file__).resolve().parent.parent.parent.parent
_kernel_root = _project_root / "fastvideo-kernel"
_kernel_python_root = _kernel_root / "python"
_attn_qat_train_attention: Callable[..., torch.Tensor] | None = None
_attn_qat_train_import_attempted = False
def _ensure_kernel_paths() -> None:
for path in (_project_root, _kernel_root, _kernel_python_root):
path_str = str(path)
if path_str not in sys.path:
sys.path.insert(0, path_str)
def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
global _attn_qat_train_attention
global _attn_qat_train_import_attempted
if _attn_qat_train_import_attempted:
return _attn_qat_train_attention
_attn_qat_train_import_attempted = True
_ensure_kernel_paths()
try:
_attn_qat_train_attention = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train").attention
except ImportError:
_attn_qat_train_attention = None
return _attn_qat_train_attention
def is_attn_qat_train_available() -> bool:
return _get_attn_qat_train_attention() is not None
def attn_qat_train(q_BLHD: torch.Tensor,
k_BLHD: torch.Tensor,
v_BLHD: torch.Tensor,
is_causal: bool = False,
sm_scale: float | None = None) -> torch.Tensor:
attention = _get_attn_qat_train_attention()
if attention is None:
raise ImportError("fastvideo_kernel.triton_kernels.attn_qat_train is not available. "
"Please ensure the FastVideo kernel package is installed.")
q_BHLD = q_BLHD.permute(0, 2, 1, 3).contiguous()
k_BHLD = k_BLHD.permute(0, 2, 1, 3).contiguous()
v_BHLD = v_BLHD.permute(0, 2, 1, 3).contiguous()
use_qat_qkv_backward = True
smooth_k = False
warp_specialize = True
is_qat = True
two_level_quant_p_sage3 = False
fake_quant_p_bwd = True
use_high_prec_o = True
smooth_q = False
if sm_scale is None:
sm_scale = 1.0 / (q_BHLD.shape[-1]**0.5)
use_global_sf_qkv = False
use_global_sf_p = False
o_BHLD = attention(
q_BHLD,
k_BHLD,
v_BHLD,
is_causal,
sm_scale,
use_qat_qkv_backward,
smooth_k,
warp_specialize,
is_qat,
two_level_quant_p_sage3,
fake_quant_p_bwd,
use_high_prec_o,
smooth_q,
use_global_sf_p,
use_global_sf_qkv,
)
return o_BHLD.permute(0, 2, 1, 3).contiguous()
class AttnQatTrainBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 96, 128, 160, 192, 224, 256]
@staticmethod
def get_name() -> str:
return "ATTN_QAT_TRAIN"
@staticmethod
def get_impl_cls() -> type["AttnQatTrainImpl"]:
return AttnQatTrainImpl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder[AttentionMetadata]"]:
raise NotImplementedError
class AttnQatTrainImpl(AttentionImpl[AttentionMetadata]):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
dropout_p = extra_impl_args.get("dropout_p", 0.0)
if dropout_p > 0:
raise NotImplementedError(f"attn_qat_train does not support dropout (got dropout_p={dropout_p}). "
"The QAT training kernel applies no stochastic dropout.")
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
return attn_qat_train(query, key, value, is_causal=self.causal, sm_scale=self.softmax_scale)
+740
View File
@@ -0,0 +1,740 @@
# SPDX-License-Identifier: Apache-2.0
"""
Bidirectional Sparse Attention (BSA) backend for FastVideo.
Pure-PyTorch reference implementation from:
"Bidirectional Sparse Attention for Faster Video Diffusion Training"
(arXiv:2509.01085)
BSA sparsifies both queries (pruning redundant tokens per block) and
key-value pairs (keeping only relevant KV blocks per query block).
This is a training-free inference backend: it works with any model
trained with full attention by applying BSA sparsity at inference time.
"""
import functools
import math
from dataclasses import dataclass
from typing import Any
import torch
import torch.nn.functional as F
from v2._vendor.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from v2._vendor.distributed import get_sp_group
from v2._vendor.logger import init_logger
try:
from v2._vendor.attention.utils.flash_attn_no_pad import (
flash_attn_varlen_func_impl, )
FLASH_ATTN_AVAILABLE = True
except ImportError:
flash_attn_varlen_func_impl = None
FLASH_ATTN_AVAILABLE = False
logger = init_logger(__name__)
BSA_TILE_SIZE = (4, 4, 4)
# ---------------------------------------------------------------------------
# Cached index helpers (same pattern as VSA)
# ---------------------------------------------------------------------------
@functools.lru_cache(maxsize=10)
def get_tile_partition_indices(
dit_seq_shape: tuple[int, int, int],
tile_size: tuple[int, int, int],
device: torch.device,
) -> torch.LongTensor:
"""Map raster-order tokens to tile-contiguous order."""
T, H, W = dit_seq_shape
ts, hs, ws = tile_size
indices = torch.arange(T * H * W, device=device, dtype=torch.long).reshape(T, H, W)
ls = []
for t in range(math.ceil(T / ts)):
for h in range(math.ceil(H / hs)):
for w in range(math.ceil(W / ws)):
ls.append(indices[
t * ts:min(t * ts + ts, T),
h * hs:min(h * hs + hs, H),
w * ws:min(w * ws + ws, W),
].flatten())
return torch.cat(ls, dim=0)
@functools.lru_cache(maxsize=10)
def get_reverse_tile_partition_indices(
dit_seq_shape: tuple[int, int, int],
tile_size: tuple[int, int, int],
device: torch.device,
) -> torch.LongTensor:
"""Inverse mapping: tile-contiguous order back to raster order."""
return torch.argsort(get_tile_partition_indices(dit_seq_shape, tile_size, device))
# ---------------------------------------------------------------------------
# BSA core operations
# ---------------------------------------------------------------------------
def _prune_queries(
q_blocks: torch.Tensor,
keep_ratio: float,
) -> tuple[torch.Tensor, torch.Tensor, int]:
"""
Prune redundant query tokens within each block.
Scores tokens by cosine similarity to the block center.
Keeps the LEAST similar (most informative) tokens.
Args:
q_blocks: [B, N_heads, N_blocks, block_size, D]
keep_ratio: fraction of tokens to keep
Returns:
sparse_q: [B, N_heads, N_blocks, keep_size, D]
keep_indices: [B, N_heads, N_blocks, keep_size]
keep_size: int
"""
B, H, N, S, D = q_blocks.shape
keep_size = max(1, int(S * keep_ratio))
if keep_size >= S:
idx = torch.arange(S, device=q_blocks.device)
idx = idx.view(1, 1, 1, S).expand(B, H, N, S)
return q_blocks, idx, S
center_idx = S // 2
center = q_blocks[:, :, :, center_idx:center_idx + 1, :]
q_norm = F.normalize(q_blocks, dim=-1)
c_norm = F.normalize(center, dim=-1)
similarity = (q_norm * c_norm).sum(dim=-1) # [B, H, N, S]
# lowest similarity = most distinctive = keep
_, indices = similarity.topk(keep_size, dim=-1, largest=False)
indices, _ = indices.sort(dim=-1)
idx_expand = indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
sparse_q = torch.gather(q_blocks, 3, idx_expand)
return sparse_q, indices, keep_size
def _select_kv_blocks(
sparse_q: torch.Tensor,
k_blocks: torch.Tensor,
cumulative_threshold: float,
min_kv_blocks: int,
) -> torch.Tensor:
"""
Dynamically select KV blocks for each query block.
Mean-pools to block level, computes block attention scores,
admits blocks in descending order until cumulative mass
exceeds threshold.
Args:
sparse_q: [B, H, N, Sq, D]
k_blocks: [B, H, N, Sk, D]
cumulative_threshold: e.g. 0.9
min_kv_blocks: minimum blocks to keep
Returns:
kv_mask: [B, H, N, N] boolean
"""
B, H, N, _, D = sparse_q.shape
q_repr = sparse_q.mean(dim=3)
k_repr = k_blocks.mean(dim=3)
scores = torch.matmul(q_repr, k_repr.transpose(-1, -2)) / (D**0.5)
block_attn = F.softmax(scores, dim=-1)
sorted_attn, sorted_idx = block_attn.sort(dim=-1, descending=True)
cumsum = sorted_attn.cumsum(dim=-1)
keep_sorted = torch.ones_like(cumsum, dtype=torch.bool)
keep_sorted[..., 1:] = cumsum[..., :-1] < cumulative_threshold
min_mask = torch.zeros_like(keep_sorted)
min_mask[..., :min(min_kv_blocks, N)] = True
keep_sorted = keep_sorted | min_mask
kv_mask = torch.zeros_like(block_attn, dtype=torch.bool)
kv_mask.scatter_(-1, sorted_idx, keep_sorted)
return kv_mask
def _compute_sparse_attention(
sparse_q: torch.Tensor,
k_blocks: torch.Tensor,
v_blocks: torch.Tensor,
kv_mask: torch.Tensor,
) -> torch.Tensor:
"""
Compute attention for each query block against selected KV blocks.
Handles per-batch and per-head KV masks correctly.
Uses flash_attn_varlen_func when available on GPU.
Falls back to pure-PyTorch reference on CPU.
Args:
sparse_q: [B, H, N, Sq, D]
k_blocks: [B, H, N, Sk, D]
v_blocks: [B, H, N, Sk, D]
kv_mask: [B, H, N, N] boolean (per-batch, per-head)
Returns:
output: [B, H, N, Sq, D]
"""
if FLASH_ATTN_AVAILABLE and sparse_q.is_cuda:
return _compute_sparse_attention_flash(sparse_q, k_blocks, v_blocks, kv_mask)
else:
return _compute_sparse_attention_reference(sparse_q, k_blocks, v_blocks, kv_mask)
def _compute_sparse_attention_reference(
sparse_q: torch.Tensor,
k_blocks: torch.Tensor,
v_blocks: torch.Tensor,
kv_mask: torch.Tensor,
) -> torch.Tensor:
"""Pure-PyTorch fallback with per-batch, per-head mask support."""
B, H, N, Sq, D = sparse_q.shape
output = torch.zeros_like(sparse_q)
for b in range(B):
for h in range(H):
for qb in range(N):
selected = kv_mask[b, h, qb] # [N] boolean
sel_idx = selected.nonzero(as_tuple=True)[0]
if sel_idx.shape[0] == 0:
continue
# [num_sel * Sk, D]
sel_k = k_blocks[b, h, sel_idx].reshape(-1, D)
sel_v = v_blocks[b, h, sel_idx].reshape(-1, D)
q = sparse_q[b, h, qb] # [Sq, D]
scores = torch.matmul(q, sel_k.transpose(-1, -2)) / (D**0.5)
weights = F.softmax(scores, dim=-1)
output[b, h, qb] = torch.matmul(weights, sel_v)
return output
def _compute_sparse_attention_flash(
sparse_q: torch.Tensor,
k_blocks: torch.Tensor,
v_blocks: torch.Tensor,
kv_mask: torch.Tensor,
) -> torch.Tensor:
"""
FlashAttention implementation with per-batch, per-head mask support.
Strategy: check if all heads share the same mask. If so, use a single
FlashAttention call per batch (fast path). If not, process each head
separately (correct path).
Args:
sparse_q: [B, H, N, Sq, D]
k_blocks: [B, H, N, Sk, D]
v_blocks: [B, H, N, Sk, D]
kv_mask: [B, H, N, N] boolean
Returns:
output: [B, H, N, Sq, D]
"""
B, H, N, Sq, D = sparse_q.shape
Sk = k_blocks.shape[3]
device = sparse_q.device
output = torch.zeros_like(sparse_q)
for b in range(B):
# Check if all heads share the same mask for this batch element
# Compare each head's mask to head 0's mask
head0_mask = kv_mask[b, 0] # [N, N]
all_heads_same = all(torch.equal(kv_mask[b, h], head0_mask) for h in range(1, H))
if all_heads_same:
# Fast path: all heads share the same mask, single FA call
_flash_attn_single_mask(
sparse_q[b],
k_blocks[b],
v_blocks[b],
head0_mask,
output[b],
H,
N,
Sq,
Sk,
D,
device,
)
else:
# Per-head path: process each head individually
for h in range(H):
head_mask = kv_mask[b, h] # [N, N]
# Process single head: squeeze head dim, run FA, put back
_flash_attn_single_head(
sparse_q[b, h],
k_blocks[b, h],
v_blocks[b, h],
head_mask,
output,
b,
h,
N,
Sq,
Sk,
D,
device,
)
return output
def _flash_attn_single_mask(
sparse_q_b: torch.Tensor, # [H, N, Sq, D]
k_blocks_b: torch.Tensor, # [H, N, Sk, D]
v_blocks_b: torch.Tensor, # [H, N, Sk, D]
mask: torch.Tensor, # [N, N] boolean
output_b: torch.Tensor, # [H, N, Sq, D] (modified in-place)
H: int,
N: int,
Sq: int,
Sk: int,
D: int,
device: torch.device,
) -> None:
"""Run FlashAttention for all heads sharing the same KV mask."""
q_list = []
k_list = []
v_list = []
cu_seqlens_q = [0]
cu_seqlens_k = [0]
active_blocks = []
for qb in range(N):
selected = mask[qb] # [N] boolean
sel_idx = selected.nonzero(as_tuple=True)[0]
if sel_idx.shape[0] == 0:
continue
active_blocks.append(qb)
num_kv_tokens = sel_idx.shape[0] * Sk
# [H, Sq, D] -> [Sq, H, D]
q_block = sparse_q_b[:, qb].permute(1, 0, 2)
q_list.append(q_block)
# [H, num_sel, Sk, D] -> [num_kv_tokens, H, D]
sel_k = k_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
sel_v = v_blocks_b[:, sel_idx].permute(1, 2, 0, 3).reshape(num_kv_tokens, H, D)
k_list.append(sel_k)
v_list.append(sel_v)
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
if not q_list:
return
flat_q = torch.cat(q_list, dim=0)
flat_k = torch.cat(k_list, dim=0)
flat_v = torch.cat(v_list, dim=0)
# Compute max_seqlen_k from the Python list before moving to GPU to
# avoid a `.item()` round-trip that would force a host/device sync.
max_seqlen_q = Sq
max_seqlen_k = max(b - a for a, b in zip(cu_seqlens_k[:-1], cu_seqlens_k[1:], strict=False))
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
orig_dtype = flat_q.dtype
compute_dtype = orig_dtype
if compute_dtype not in (torch.float16, torch.bfloat16):
compute_dtype = torch.bfloat16
flat_q = flat_q.to(compute_dtype)
flat_k = flat_k.to(compute_dtype)
flat_v = flat_v.to(compute_dtype)
flat_out = flash_attn_varlen_func_impl(
flat_q,
flat_k,
flat_v,
cu_seqlens_q_t,
cu_seqlens_k_t,
max_seqlen_q,
max_seqlen_k,
causal=False,
)
if compute_dtype != orig_dtype:
flat_out = flat_out.to(orig_dtype)
idx = 0
for qb in active_blocks:
block_out = flat_out[idx:idx + Sq] # [Sq, H, D]
output_b[:, qb] = block_out.permute(1, 0, 2) # [H, Sq, D]
idx += Sq
def _flash_attn_single_head(
sparse_q_bh: torch.Tensor, # [N, Sq, D]
k_blocks_bh: torch.Tensor, # [N, Sk, D]
v_blocks_bh: torch.Tensor, # [N, Sk, D]
mask: torch.Tensor, # [N, N] boolean
output: torch.Tensor, # [B, H, N, Sq, D] (modified in-place)
b: int,
h: int,
N: int,
Sq: int,
Sk: int,
D: int,
device: torch.device,
) -> None:
"""Run FlashAttention for a single head with its own KV mask."""
q_list = []
k_list = []
v_list = []
cu_seqlens_q = [0]
cu_seqlens_k = [0]
active_blocks = []
for qb in range(N):
selected = mask[qb]
sel_idx = selected.nonzero(as_tuple=True)[0]
if sel_idx.shape[0] == 0:
continue
active_blocks.append(qb)
num_kv_tokens = sel_idx.shape[0] * Sk
# [Sq, D] -> [Sq, 1, D] (single head)
q_block = sparse_q_bh[qb].unsqueeze(1)
q_list.append(q_block)
# [num_sel, Sk, D] -> [num_kv_tokens, 1, D]
sel_k = k_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
sel_v = v_blocks_bh[sel_idx].reshape(num_kv_tokens, 1, D)
k_list.append(sel_k)
v_list.append(sel_v)
cu_seqlens_q.append(cu_seqlens_q[-1] + Sq)
cu_seqlens_k.append(cu_seqlens_k[-1] + num_kv_tokens)
if not q_list:
return
flat_q = torch.cat(q_list, dim=0)
flat_k = torch.cat(k_list, dim=0)
flat_v = torch.cat(v_list, dim=0)
# Compute max_seqlen_k from the Python list before moving to GPU to
# avoid a `.item()` round-trip that would force a host/device sync.
max_seqlen_q = Sq
max_seqlen_k = max(b - a for a, b in zip(cu_seqlens_k[:-1], cu_seqlens_k[1:], strict=False))
cu_seqlens_q_t = torch.tensor(cu_seqlens_q, dtype=torch.int32, device=device)
cu_seqlens_k_t = torch.tensor(cu_seqlens_k, dtype=torch.int32, device=device)
orig_dtype = flat_q.dtype
compute_dtype = orig_dtype
if compute_dtype not in (torch.float16, torch.bfloat16):
compute_dtype = torch.bfloat16
flat_q = flat_q.to(compute_dtype)
flat_k = flat_k.to(compute_dtype)
flat_v = flat_v.to(compute_dtype)
flat_out = flash_attn_varlen_func_impl(
flat_q,
flat_k,
flat_v,
cu_seqlens_q_t,
cu_seqlens_k_t,
max_seqlen_q,
max_seqlen_k,
causal=False,
)
if compute_dtype != orig_dtype:
flat_out = flat_out.to(orig_dtype)
idx = 0
for qb in active_blocks:
block_out = flat_out[idx:idx + Sq] # [Sq, 1, D]
output[b, h, qb] = block_out.squeeze(1) # [Sq, D]
idx += Sq
def _reconstruct_pruned(
sparse_output: torch.Tensor,
keep_indices: torch.Tensor,
block_size: int,
) -> torch.Tensor:
"""
Scatter sparse output back to full block size.
Pruned positions get nearest kept token's output.
Handles per-batch, per-head indices correctly.
Args:
sparse_output: [B, H, N, keep_size, D]
keep_indices: [B, H, N, keep_size]
block_size: original tokens per block
Returns:
full_output: [B, H, N, block_size, D]
"""
B, H, N, keep_size, D = sparse_output.shape
device = sparse_output.device
if keep_size >= block_size:
return sparse_output
full_output = torch.zeros(B, H, N, block_size, D, device=device, dtype=sparse_output.dtype)
# Scatter kept tokens
idx_expand = keep_indices.unsqueeze(-1).expand(-1, -1, -1, -1, D)
full_output.scatter_(3, idx_expand, sparse_output)
# Fill pruned positions with nearest kept token (vectorized)
all_pos = torch.arange(block_size, device=device)
for b in range(B):
for h in range(H):
for n in range(N):
kept = keep_indices[b, h, n] # [keep_size]
# Distance from every position to every kept position
dists = (all_pos.view(-1, 1) - kept.view(1, -1)).abs()
nearest_local_idx = dists.argmin(dim=1) # [block_size]
# Identify pruned positions
is_pruned = torch.ones(block_size, dtype=torch.bool, device=device)
is_pruned[kept] = False
pruned_indices = is_pruned.nonzero(as_tuple=True)[0]
if pruned_indices.numel() > 0:
src_indices = nearest_local_idx[pruned_indices]
full_output[b, h, n, pruned_indices] = sparse_output[b, h, n, src_indices]
return full_output
# ---------------------------------------------------------------------------
# FastVideo backend classes
# ---------------------------------------------------------------------------
class BSAAttentionBackend(AttentionBackend):
accept_output_buffer: bool = False
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@staticmethod
def get_name() -> str:
return "BSA_ATTN"
@staticmethod
def get_impl_cls() -> type["BSAAttentionImpl"]:
return BSAAttentionImpl
@staticmethod
def get_metadata_cls() -> type["BSAAttentionMetadata"]:
return BSAAttentionMetadata
@staticmethod
def get_builder_cls() -> type["BSAAttentionMetadataBuilder"]:
return BSAAttentionMetadataBuilder
@dataclass
class BSAAttentionMetadata(AttentionMetadata):
current_timestep: int
dit_seq_shape: tuple[int, int, int]
total_seq_length: int
num_blocks: int
block_size: int
tile_partition_indices: torch.LongTensor
reverse_tile_partition_indices: torch.LongTensor
# BSA-specific config
query_keep_ratio: float
kv_cumulative_threshold: float
min_kv_blocks: int
class BSAAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build(
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
device: torch.device,
bsa_query_keep_ratio: float = 0.5,
bsa_kv_cumulative_threshold: float = 0.9,
bsa_min_kv_blocks: int = 4,
**kwargs: dict[str, Any],
) -> "BSAAttentionMetadata":
# Ensure patching does not drop tokens silently.
assert all(r % p == 0 for r, p in zip(raw_latent_shape, patch_size, strict=False)), (
"raw_latent_shape must be divisible by patch_size for BSA", )
dit_seq_shape = (
raw_latent_shape[0] // patch_size[0],
raw_latent_shape[1] // patch_size[1],
raw_latent_shape[2] // patch_size[2],
)
total_seq_length = math.prod(dit_seq_shape)
block_size = math.prod(BSA_TILE_SIZE)
# Require exact tiling to avoid reshape failures later.
assert all(d % t == 0 for d, t in zip(dit_seq_shape, BSA_TILE_SIZE, strict=False)), (
"dit_seq_shape must be divisible by BSA_TILE_SIZE", )
num_blocks = total_seq_length // block_size
tile_partition_indices = get_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
reverse_tile_partition_indices = get_reverse_tile_partition_indices(dit_seq_shape, BSA_TILE_SIZE, device)
return BSAAttentionMetadata(
current_timestep=current_timestep,
dit_seq_shape=dit_seq_shape,
total_seq_length=total_seq_length,
num_blocks=num_blocks,
block_size=block_size,
tile_partition_indices=tile_partition_indices,
reverse_tile_partition_indices=reverse_tile_partition_indices,
query_keep_ratio=bsa_query_keep_ratio,
kv_cumulative_threshold=bsa_kv_cumulative_threshold,
min_kv_blocks=bsa_min_kv_blocks,
)
class BSAAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.prefix = prefix
self.num_heads = num_heads
self.head_size = head_size
if num_kv_heads is not None and num_kv_heads != num_heads:
raise ValueError("BSA backend does not support grouped-query attention")
if causal:
raise ValueError("BSA backend is bidirectional; causal=True is unsupported")
if softmax_scale is not None:
expected_scale = 1.0 / math.sqrt(self.head_size)
if not math.isclose(softmax_scale, expected_scale, rel_tol=1e-4, abs_tol=1e-5):
raise ValueError("softmax_scale must be default (1/sqrt(d)) for BSA")
try:
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
except (AssertionError, RuntimeError):
self.sp_size = 1
def preprocess_qkv(
self,
qkv: torch.Tensor,
attn_metadata: BSAAttentionMetadata,
) -> torch.Tensor:
"""Reorder tokens from raster order to tile-contiguous order."""
# qkv: [B, L, num_heads, D]
return qkv[:, attn_metadata.tile_partition_indices]
def postprocess_output(
self,
output: torch.Tensor,
attn_metadata: BSAAttentionMetadata,
) -> torch.Tensor:
"""Reorder tokens from tile-contiguous order back to raster order."""
return output[:, attn_metadata.reverse_tile_partition_indices]
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: BSAAttentionMetadata,
) -> torch.Tensor:
"""
BSA attention forward pass.
Input tensors are already in tile-contiguous order from preprocess_qkv.
Args:
query: [B, L, num_heads, D] (tile-ordered)
key: [B, L, num_heads, D] (tile-ordered)
value: [B, L, num_heads, D] (tile-ordered)
attn_metadata: BSA metadata
Returns:
output: [B, L, num_heads, D] (tile-ordered)
"""
B, L, H, D = query.shape
block_size = attn_metadata.block_size
num_blocks = attn_metadata.num_blocks
assert num_blocks * block_size == L, "Sequence length must match tiling"
# Reshape to [B, H, L, D] for attention computation
q = query.transpose(1, 2).contiguous() # [B, H, L, D]
k = key.transpose(1, 2).contiguous()
v = value.transpose(1, 2).contiguous()
# Reshape into blocks: [B, H, num_blocks, block_size, D]
q_blocks = q.view(B, H, num_blocks, block_size, D)
k_blocks = k.view(B, H, num_blocks, block_size, D)
v_blocks = v.view(B, H, num_blocks, block_size, D)
# --- Query sparsification ---
sparse_q, keep_indices, keep_size = _prune_queries(q_blocks, attn_metadata.query_keep_ratio)
# --- KV block selection ---
kv_mask = _select_kv_blocks(
sparse_q,
k_blocks,
attn_metadata.kv_cumulative_threshold,
attn_metadata.min_kv_blocks,
)
# --- Sparse attention ---
sparse_output = _compute_sparse_attention(sparse_q, k_blocks, v_blocks, kv_mask)
# --- Reconstruct pruned positions ---
full_output = _reconstruct_pruned(sparse_output, keep_indices, block_size)
# Reshape back: [B, H, num_blocks, block_size, D] -> [B, H, L, D] -> [B, L, H, D]
hidden_states = full_output.view(B, H, L, D).transpose(1, 2)
return hidden_states
+341
View File
@@ -0,0 +1,341 @@
# SPDX-License-Identifier: Apache-2.0
import os
import torch
import torch.nn.functional as F
from dataclasses import dataclass
try:
from v2._vendor.attention.utils.flash_attn_cute import flash_attn_func
fa_version = "4"
except ImportError:
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 no longer have a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "3"
except ImportError:
from flash_attn import flash_attn_func as flash_attn_2_func
flash_attn_func = flash_attn_2_func
fa_version = "2"
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
# already a registered torch.library custom op, so dynamo treats it as a
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
# breaks the graph at the call site (observed: wanvideo.py self-attn,
# once per layer every step), which fragments the compiled region and
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
# opaque-but-traceable node. The kernel still runs eager inside the op
# (correct — flash-attn must run eager); only dynamo's treatment of the
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
if fa_version in ("2", "3"):
_fa_default = flash_attn_func
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
# call shape used by FlashAttentionImpl.forward's default branch
# (see `flash_attn_func_compilable(...)` call site below). The
# masked/no-pad and varlen / cross-attn paths use different
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
# are intentionally out of scope for this PR — wrapping them is a
# natural follow-up. The wrapper's signature is the contract: any
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
# return_attn_probs, ...) raises TypeError at the call site, so
# silent loss of kwargs is not a failure mode.
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
del softmax_scale, causal
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
# same dtype/device as q (head dim taken from v).
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Autograd carve-out. The custom op above registers a forward + fake
# kernel but NO backward (register_autograd), so it is opaque to
# autograd. Inference runs under no_grad / inference_mode and routes
# through the traceable custom op — that is the torch.compile win, and
# the only path this PR claims. Training backprops through attention,
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
# (itself an autograd.Function, so backward is correct) at the cost of a
# dynamo graph break on the training path — i.e. pre-PR behavior, no
# regression. Full autograd parity for the custom op (mirroring the FP4
# cute template) is a tracked follow-up.
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
return torch.ops.v2._flash_attn_default_forward(q, k, v, softmax_scale, causal)
elif fa_version == "4":
# FA4 path: `flash_attn_func` is already a torch.library custom op
# (registered in `v2._vendor.attention.utils.flash_attn_cute`), so a
# passthrough is enough — no extra registration needed.
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
else:
# Defensive: the probe above only ever sets fa_version to "2", "3",
# or "4"; an unexpected value means an import/probe regression and
# we want a loud error at import, not a silent NameError later.
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
f"'2', '3', or '4' from the import probe above.")
from v2._vendor.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
logger.info("Using FlashAttention-%s backend", fa_version)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
# The FP4 path uses a dedicated custom_op wrapper (flash_attn_fp4_func) so that
# torch.compile treats the CuTeDSL kernel as an opaque boundary.
try:
from v2._vendor.attention.utils.flash_attn_cute import flash_attn_fp4_func
_FA4_FP4_AVAILABLE = True
except ImportError:
flash_attn_fp4_func = None
_FA4_FP4_AVAILABLE = False
def _nvfp4_quantize_for_fa4(tensor_4d: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]:
"""Quantize a (batch, seqlen, nheads, headdim) BF16 tensor to FP4.
Returns:
fp4_tensor: torch.float4_e2m1fn_x2, shape (batch, seqlen_padded, nheads, headdim//2)
where seqlen_padded is seqlen rounded up to multiple of 128.
Caller should slice [:, :orig_seqlen] before passing to FA4.
sf_tensor: torch.uint8, shape (32, 4, rest_m, 4, rest_k, nheads, batch) with stride[3]=1
"""
from flashinfer.quantization import nvfp4_quantize, SfLayout
batch, seqlen, nheads, headdim = tensor_4d.shape
sf_vec_size = 16
# Pad seqlen to multiple of 128 (required by nvfp4_quantize layout_128x4)
tile_m = 128
seqlen_padded = (seqlen + tile_m - 1) // tile_m * tile_m
if seqlen_padded != seqlen:
tensor_4d = F.pad(tensor_4d, (0, 0, 0, 0, 0, seqlen_padded - seqlen))
# Quantize with nheads squashed into K dimension so M=batch*seqlen (divisible by 128)
# and K=nheads*headdim. This ensures 128-row SF tiles align with seqlen boundaries.
t2d = tensor_4d.reshape(batch * seqlen_padded, nheads * headdim)
one = torch.ones(1, device=t2d.device, dtype=torch.float32)
fp4_data, sf_data = nvfp4_quantize(t2d, one, sfLayout=SfLayout.layout_128x4, do_shuffle=False)
# FP4 data: (batch*seqlen, nheads*headdim/2) → (batch, seqlen, nheads, headdim/2)
fp4_tensor = (fp4_data.reshape(batch, seqlen_padded, nheads,
headdim // 2).view(torch.int8).view(torch.float4_e2m1fn_x2))
# SF layout conversion: nvfp4_quantize layout_128x4 → FA4 MMA layout
# layout_128x4 buffer: [mTile, kTile, 32, 4, 4]
# FA4 expects: (32, 4, rest_m, 4, rest_k, nheads, batch) with stride[3]=1
atom_m0, atom_m1, atom_k = 32, 4, 4
rest_m = seqlen_padded // tile_m
sf_k_per_head = headdim // sf_vec_size # 8 for headdim=128
rest_k = sf_k_per_head // atom_k # 2
total_m_tiles = batch * rest_m
total_k_tiles = (nheads * sf_k_per_head) // atom_k
sf_swizzled = sf_data.reshape(total_m_tiles, total_k_tiles, atom_m0, atom_m1, atom_k)
sf_decomposed = sf_swizzled.reshape(batch, rest_m, nheads, rest_k, atom_m0, atom_m1, atom_k)
sf_canonical = sf_decomposed.permute(0, 2, 1, 3, 4, 5, 6).contiguous()
sf_mma = sf_canonical.permute(4, 5, 2, 6, 3, 1, 0)
return fp4_tensor, sf_mma
class FlashAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
def get_name() -> str:
return "FLASH_ATTN"
@staticmethod
def get_impl_cls() -> type["FlashAttentionImpl"]:
return FlashAttentionImpl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
def _key_padding_mask_from_attn_mask(attn_mask: torch.Tensor, key_len: int) -> torch.Tensor:
# Normalize attn_mask to [B, key_len] where True means valid token.
if attn_mask.dim() == 4:
attn_mask = attn_mask[:, 0, 0, :]
elif attn_mask.dim() == 3:
attn_mask = attn_mask[:, 0, :]
elif attn_mask.dim() != 2:
raise ValueError(f"Unsupported attn_mask shape for FLASH_ATTN: {attn_mask.shape}")
# SDPA additive mask convention: valid=0, masked=-inf/large negative.
key_padding_mask = attn_mask if attn_mask.dtype == torch.bool else attn_mask >= 0
if key_padding_mask.shape[-1] != key_len:
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
f"expected {key_len}, got {key_padding_mask.shape[-1]}")
return key_padding_mask
@dataclass
class FlashAttnMetadata(AttentionMetadata):
current_timestep: int
attn_mask: torch.Tensor | None = None
class FlashAttnMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
current_timestep: int,
attn_mask: torch.Tensor,
) -> FlashAttnMetadata:
return FlashAttnMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
class FlashAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.nvfp4_fa4 = extra_impl_args.get("nvfp4_fa4", False) or os.environ.get("FASTVIDEO_NVFP4_FA4", "0") == "1"
if self.nvfp4_fa4:
cap = torch.cuda.get_device_capability()
assert cap in [(10, 0), (10, 3)], (f"NVFP4 FA4 requires Blackwell (sm100a/sm103a), got sm{cap[0]}{cap[1]}")
assert _FA4_FP4_AVAILABLE, ("NVFP4 FA4 requires flash-attention-fp4 (flash_attn.cute). "
"Install via instructions in docs/inference/optimizations.md")
logger.info("NVFP4 FA4 enabled for FlashAttentionImpl (quant_qk only)")
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: FlashAttnMetadata,
):
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
from v2._vendor.attention.utils.flash_attn_no_pad import (
flash_attn_no_pad,
flash_attn_varlen_qk_no_pad,
)
attn_mask = attn_metadata.attn_mask
# flash_attn_no_pad packs q/k/v as one tensor and assumes equal q/k
# sequence lengths. Cross-attention can violate this.
if query.shape[1] != key.shape[1]:
query_padding_mask = torch.ones(
(query.shape[0], query.shape[1]),
dtype=torch.bool,
device=query.device,
)
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, key.shape[1]).to(device=key.device)
return flash_attn_varlen_qk_no_pad(
query,
key,
value,
query_padding_mask=query_padding_mask,
key_padding_mask=key_padding_mask,
causal=self.causal,
dropout_p=0.0,
softmax_scale=self.softmax_scale,
)
qkv = torch.stack([query, key, value], dim=2)
attn_mask_padded = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0), value=True)
output = flash_attn_no_pad(qkv, attn_mask_padded, causal=False, dropout_p=0, softmax_scale=None)
elif self.nvfp4_fa4:
output = self._forward_nvfp4(query, key, value)
else:
# Route through the compilable wrapper so dynamo sees a
# registered op (no graph break) for FA2/FA3; identical
# kernel + numerics, op runs eager internally.
output = flash_attn_func_compilable(
query, # type: ignore[no-untyped-call]
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal,
)
return output
def _forward_nvfp4(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor) -> torch.Tensor:
"""FP4 flash attention with quantized Q and K, BF16 V."""
orig_seqlen_q = query.shape[1]
orig_seqlen_k = key.shape[1]
# Quantize Q/K to FP4 (internally pads to multiple of 128 for SF layout)
q_fp4, q_sf = _nvfp4_quantize_for_fa4(query)
k_fp4, k_sf = _nvfp4_quantize_for_fa4(key)
# Pass original seqlen to FA4 — the kernel handles non-multiple-of-128
# via boundary masking. FP4/SF data is padded to 128-multiple but FA4
# only attends to orig_seqlen positions, avoiding softmax bias on padding.
q_fp4 = q_fp4[:, :orig_seqlen_q]
k_fp4 = k_fp4[:, :orig_seqlen_k]
output = flash_attn_fp4_func(
q_fp4,
k_fp4,
value,
q_sf,
k_sf,
softmax_scale=self.softmax_scale,
causal=self.causal,
)
if isinstance(output, tuple):
output = output[0]
return output
@@ -0,0 +1,64 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from sageattention import sageattn
from v2._vendor.attention.backends.abstract import ( # FlashAttentionMetadata,
AttentionBackend, AttentionImpl, AttentionMetadata)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
class SageAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
def get_name() -> str:
return "SAGE_ATTN"
@staticmethod
def get_impl_cls() -> type["SageAttentionImpl"]:
return SageAttentionImpl
# @staticmethod
# def get_metadata_cls() -> Type["AttentionMetadata"]:
# return FlashAttentionMetadata
class SageAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
output = sageattn(
query,
key,
value,
# since input is (batch_size, seq_len, head_num, head_dim)
tensor_layout="NHD",
is_causal=self.causal)
return output
@@ -0,0 +1,70 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from sageattn3 import sageattn3_blackwell
from v2._vendor.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
class SageAttention3Backend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128, 256]
@staticmethod
def get_name() -> str:
return "SAGE_ATTN_THREE"
@staticmethod
def get_impl_cls() -> type["SageAttention3Impl"]:
return SageAttention3Impl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
raise NotImplementedError
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
raise NotImplementedError
# @staticmethod
# def get_metadata_cls() -> Type["AttentionMetadata"]:
# return FlashAttentionMetadata
class SageAttention3Impl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
output = output.transpose(1, 2)
return output
+95
View File
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from dataclasses import dataclass
from v2._vendor.attention.backends.abstract import ( # FlashAttentionMetadata,
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
class SDPABackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [32, 64, 96, 128, 160, 192, 224, 256]
@staticmethod
def get_name() -> str:
return "SDPA"
@staticmethod
def get_impl_cls() -> type["SDPAImpl"]:
return SDPAImpl
# @staticmethod
# def get_metadata_cls() -> Type["AttentionMetadata"]:
# return FlashAttentionMetadata
@dataclass
class SDPAMetadata(AttentionMetadata):
current_timestep: int
attn_mask: torch.Tensor | None = None
class SDPAMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
current_timestep: int,
attn_mask: torch.Tensor,
) -> SDPAMetadata:
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
class SDPAImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.causal = causal
self.softmax_scale = softmax_scale
self.dropout = extra_impl_args.get("dropout_p", 0.0)
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: SDPAMetadata,
) -> torch.Tensor:
# transpose to bs, heads, seq_len, head_dim
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
and hasattr(attn_metadata, "attn_mask")) else None
attn_kwargs = {
"attn_mask": attn_mask,
"dropout_p": self.dropout,
"is_causal": self.causal,
"scale": self.softmax_scale
}
if query.shape[1] != key.shape[1]:
attn_kwargs["enable_gqa"] = True
output = torch.nn.functional.scaled_dot_product_attention(query, key, value, **attn_kwargs)
output = output.transpose(1, 2)
return output
+561
View File
@@ -0,0 +1,561 @@
# SPDX-License-Identifier: Apache-2.0
# SLA (Sparse-Linear Attention) backend for FastVideo
# Adapted from TurboDiffusion SLA implementation
#
# Copyright (c) 2025 by SLA team.
# Citation:
# @article{zhang2025sla,
# title={SLA: Beyond Sparsity in Diffusion Transformers via Fine-Tunable Sparse-Linear Attention},
# author={Jintao Zhang and Haoxu Wang and Kai Jiang and Shuo Yang and Kaiwen Zheng and
# Haocheng Xi and Ziteng Wang and Hongzhou Zhu and Min Zhao and Ion Stoica and
# Joseph E. Gonzalez and Jun Zhu and Jianfei Chen},
# journal={arXiv preprint arXiv:2509.24006},
# year={2025}
# }
from dataclasses import dataclass
from typing import Any
from collections.abc import Callable
import torch
import torch.nn as nn
import torch.nn.functional as F
import triton
import triton.language as tl
from v2._vendor.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
from fastvideo_kernel.triton_kernels.sla_triton import _attention
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
# ============================================================================
# SLA Utility functions (moved from sla_kernels/utils.py)
# ============================================================================
@triton.jit
def compress_kernel(
X,
XM,
L: tl.constexpr,
D: tl.constexpr,
BLOCK_L: tl.constexpr,
):
idx_l = tl.program_id(0)
idx_bh = tl.program_id(1)
offs_l = idx_l * BLOCK_L + tl.arange(0, BLOCK_L)
offs_d = tl.arange(0, D)
x_offset = idx_bh * L * D
xm_offset = idx_bh * ((L + BLOCK_L - 1) // BLOCK_L) * D
x = tl.load(X + x_offset + offs_l[:, None] * D + offs_d[None, :], mask=offs_l[:, None] < L)
nx = min(BLOCK_L, L - idx_l * BLOCK_L)
x_mean = tl.sum(x, axis=0, dtype=tl.float32) / nx
tl.store(XM + xm_offset + idx_l * D + offs_d, x_mean.to(XM.dtype.element_ty))
def mean_pool(x: torch.Tensor, BLK: int) -> torch.Tensor:
"""Mean pool tensor along sequence dimension with block size BLK."""
assert x.is_contiguous()
B, H, L, D = x.shape
L_BLOCKS = (L + BLK - 1) // BLK
x_mean = torch.empty((B, H, L_BLOCKS, D), device=x.device, dtype=x.dtype)
grid = (L_BLOCKS, B * H)
compress_kernel[grid](x, x_mean, L, D, BLK)
return x_mean
def get_block_map(
q: torch.Tensor,
k: torch.Tensor,
topk_ratio: float,
BLKQ: int = 64,
BLKK: int = 64,
) -> tuple[torch.Tensor, torch.Tensor, int]:
"""Compute sparse block map for attention based on QK similarity.
Args:
q: Query tensor of shape (B, H, L, D)
k: Key tensor of shape (B, H, L, D)
topk_ratio: Ratio of key blocks to attend to (0-1)
BLKQ: Query block size
BLKK: Key block size
Returns:
sparse_map: Binary mask of shape (B, H, num_q_blocks, num_k_blocks)
lut: Top-k indices of shape (B, H, num_q_blocks, topk)
topk: Number of key blocks selected
"""
arg_k = k - torch.mean(k, dim=-2, keepdim=True) # smooth-k technique from SageAttention
pooled_qblocks = mean_pool(q, BLKQ)
pooled_kblocks = mean_pool(arg_k, BLKK)
pooled_score = pooled_qblocks @ pooled_kblocks.transpose(-1, -2)
K = pooled_score.shape[-1]
topk = min(K, int(topk_ratio * K))
lut = torch.topk(pooled_score, topk, dim=-1, sorted=False).indices
sparse_map = torch.zeros_like(pooled_score, dtype=torch.int8)
sparse_map.scatter_(-1, lut, 1)
return sparse_map, lut, topk
# ============================================================================
# SLA Backend classes
# ============================================================================
class SLAAttentionBackend(AttentionBackend):
"""Sparse-Linear Attention backend."""
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@staticmethod
def get_name() -> str:
return "SLA_ATTN"
@staticmethod
def get_impl_cls() -> type["SLAAttentionImpl"]:
return SLAAttentionImpl
@staticmethod
def get_metadata_cls() -> type["SLAAttentionMetadata"]:
return SLAAttentionMetadata
@staticmethod
def get_builder_cls() -> type["SLAAttentionMetadataBuilder"]:
return SLAAttentionMetadataBuilder
@dataclass
class SLAAttentionMetadata(AttentionMetadata):
"""Metadata for SLA attention."""
current_timestep: int
topk_ratio: float = 0.5 # Ratio of key blocks to attend to
class SLAAttentionMetadataBuilder(AttentionMetadataBuilder):
"""Builder for SLA attention metadata."""
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build(
self,
current_timestep: int,
topk_ratio: float = 0.5,
**kwargs: dict[str, Any],
) -> SLAAttentionMetadata:
return SLAAttentionMetadata(
current_timestep=current_timestep,
topk_ratio=topk_ratio,
)
class SLAAttentionImpl(AttentionImpl, nn.Module):
"""SLA attention implementation with learnable linear projection.
This implementation combines sparse attention with linear attention,
using a learnable projection to blend the outputs. The sparse attention
uses a block-sparse pattern determined by QK similarity.
Args:
num_heads: Number of attention heads
head_size: Dimension of each head
topk_ratio: Ratio of key blocks to attend to (0-1), default 0.5
feature_map: Feature map for linear attention ('softmax', 'elu', 'relu')
BLKQ: Query block size for sparse attention
BLKK: Key block size for sparse attention
use_bf16: Whether to use bfloat16 for computation
"""
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool = False,
softmax_scale: float | None = None,
num_kv_heads: int | None = None,
prefix: str = "",
# SLA-specific parameters - matched to TurboDiffusion defaults
topk_ratio: float = 0.1, # TurboDiffusion uses topk=0.1
feature_map: str = "softmax",
BLKQ: int = 128, # TurboDiffusion uses BLKQ=128
BLKK: int = 64, # TurboDiffusion uses BLKK=64
use_bf16: bool = True,
**extra_impl_args,
) -> None:
nn.Module.__init__(self)
self.num_heads = num_heads
self.head_size = head_size
self.softmax_scale = softmax_scale if softmax_scale else head_size**-0.5
self.causal = causal
self.prefix = prefix
# SLA-specific config
self.topk_ratio = topk_ratio
self.BLKQ = BLKQ
self.BLKK = BLKK
self.dtype = torch.bfloat16 if use_bf16 else torch.float16
# Learnable linear projection for combining sparse + linear attention
self.proj_l = nn.Linear(head_size, head_size, dtype=torch.float32)
# Feature map for linear attention
# Type annotation for callables
self.feature_map_q: Callable[[torch.Tensor], torch.Tensor]
self.feature_map_k: Callable[[torch.Tensor], torch.Tensor]
if feature_map == "elu":
self.feature_map_q = lambda x: F.elu(x) + 1
self.feature_map_k = lambda x: F.elu(x) + 1
elif feature_map == "relu":
self.feature_map_q = F.relu
self.feature_map_k = F.relu
elif feature_map == "softmax":
self.feature_map_q = lambda x: F.softmax(x, dim=-1)
self.feature_map_k = lambda x: F.softmax(x, dim=-1)
else:
raise ValueError(f"Unknown feature map: {feature_map}")
self._init_weights()
def _init_weights(self) -> None:
"""Initialize projection weights to zero for residual-like behavior."""
with torch.no_grad():
nn.init.zeros_(self.proj_l.weight)
nn.init.zeros_(self.proj_l.bias) # type: ignore[arg-type]
def _calc_linear_attention(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
) -> torch.Tensor:
"""Compute linear attention: (Q @ K^T @ V) / normalizer.
Args:
q: Query tensor (B, H, L, D) after feature map
k: Key tensor (B, H, L, D) after feature map
v: Value tensor (B, H, L, D)
Returns:
Linear attention output (B, H, L, D)
"""
kvsum = k.transpose(-1, -2) @ v # (B, H, D, D)
ksum = torch.sum(k, dim=-2, keepdim=True) # (B, H, 1, D)
return (q @ kvsum) / (1e-5 + (q * ksum).sum(dim=-1, keepdim=True))
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
"""Forward pass for SLA attention.
Input tensors are in FastVideo format: (B, L, H, D)
Internally converted to SLA format: (B, H, L, D)
Args:
query: Query tensor (B, L, H, D)
key: Key tensor (B, L, H, D)
value: Value tensor (B, L, H, D)
attn_metadata: Attention metadata
Returns:
Output tensor (B, L, H, D)
"""
original_dtype = query.dtype
# Convert from FastVideo format (B, L, H, D) to SLA format (B, H, L, D)
q = query.transpose(1, 2).contiguous()
k = key.transpose(1, 2).contiguous()
v = value.transpose(1, 2).contiguous()
# Get topk ratio from metadata if available
topk_ratio = self.topk_ratio
if hasattr(attn_metadata, 'topk_ratio'):
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
# Compute block-sparse attention pattern
sparse_map, lut, real_topk = get_block_map(q, k, topk_ratio=topk_ratio, BLKQ=self.BLKQ, BLKK=self.BLKK)
# Convert to compute dtype
q = q.to(self.dtype)
k = k.to(self.dtype)
v = v.to(self.dtype)
# Sparse attention
o_s = _attention.apply(q, k, v, sparse_map, lut, real_topk, self.BLKQ, self.BLKK)
# Linear attention with feature maps. Note: softmax / elu / relu
# are elementwise and preserve layout, so the inputs are already
# contiguous from the transpose-contiguous above — no need to
# call .contiguous() again here.
q_linear = self.feature_map_q(q).to(self.dtype)
k_linear = self.feature_map_k(k).to(self.dtype)
o_l = self._calc_linear_attention(q_linear, k_linear, v)
# Project linear attention output and combine
with torch.amp.autocast('cuda', dtype=self.dtype):
o_l = self.proj_l(o_l)
# Combine sparse and linear outputs
output = (o_s + o_l).to(original_dtype)
# Convert back to FastVideo format (B, L, H, D)
output = output.transpose(1, 2)
return output
# Check if spas_sage_attn is available for SageSLA
SAGESLA_ENABLED = True
try:
import spas_sage_attn._qattn as qattn
import spas_sage_attn._fused as fused
from spas_sage_attn.utils import get_vanilla_qk_quant, block_map_lut_triton
except ImportError:
SAGESLA_ENABLED = False
SAGE2PP_ENABLED = True
try:
from spas_sage_attn._qattn import qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold
except ImportError:
SAGE2PP_ENABLED = False
class SageSLAAttentionBackend(AttentionBackend):
"""Quantized Sparse-Linear Attention backend using SageAttention kernels."""
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@staticmethod
def get_name() -> str:
return "SAGE_SLA_ATTN"
@staticmethod
def get_impl_cls() -> type["SageSLAAttentionImpl"]:
return SageSLAAttentionImpl
@staticmethod
def get_metadata_cls() -> type["SLAAttentionMetadata"]:
return SLAAttentionMetadata
@staticmethod
def get_builder_cls() -> type["SLAAttentionMetadataBuilder"]:
return SLAAttentionMetadataBuilder
def _get_cuda_arch(device_index: int) -> str:
"""Get CUDA architecture string for the given device."""
major, minor = torch.cuda.get_device_capability(device_index)
return f"sm{major}{minor}"
class SageSLAAttentionImpl(AttentionImpl, nn.Module):
"""SageSLA attention implementation using quantized SageAttention kernels.
This uses INT8 quantization for Q/K and FP8 for V to achieve better performance
while maintaining accuracy. Requires spas_sage_attn package.
Args:
num_heads: Number of attention heads
head_size: Dimension of each head (must be 64 or 128)
topk_ratio: Ratio of key blocks to attend to (0-1), default 0.5
feature_map: Feature map for linear attention ('softmax', 'elu', 'relu')
use_bf16: Whether to use bfloat16 for computation
"""
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool = False,
softmax_scale: float | None = None,
num_kv_heads: int | None = None,
prefix: str = "",
# SageSLA-specific parameters
topk_ratio: float = 0.5,
feature_map: str = "softmax",
use_bf16: bool = True,
**extra_impl_args,
) -> None:
nn.Module.__init__(self)
if not SAGESLA_ENABLED:
raise ImportError("SageSLA requires spas_sage_attn. "
"Install with: uv pip install git+https://github.com/thu-ml/SpargeAttn.git")
assert head_size in [64, 128], f"SageSLA requires head_size in [64, 128], got {head_size}"
self.num_heads = num_heads
self.head_size = head_size
self.softmax_scale = softmax_scale if softmax_scale else head_size**-0.5
self.causal = causal
self.prefix = prefix
# SageSLA-specific config
self.topk_ratio = topk_ratio
self.dtype = torch.bfloat16 if use_bf16 else torch.float16
# Learnable linear projection for combining sparse + linear attention
self.proj_l = nn.Linear(head_size, head_size, dtype=torch.float32)
# Feature map for linear attention
# Type annotation for callables
self.feature_map_q: Callable[[torch.Tensor], torch.Tensor]
self.feature_map_k: Callable[[torch.Tensor], torch.Tensor]
if feature_map == "elu":
self.feature_map_q = lambda x: F.elu(x) + 1
self.feature_map_k = lambda x: F.elu(x) + 1
elif feature_map == "relu":
self.feature_map_q = F.relu
self.feature_map_k = F.relu
elif feature_map == "softmax":
self.feature_map_q = lambda x: F.softmax(x, dim=-1)
self.feature_map_k = lambda x: F.softmax(x, dim=-1)
else:
raise ValueError(f"Unknown feature map: {feature_map}")
self._init_weights()
def _init_weights(self) -> None:
"""Initialize projection weights to zero for residual-like behavior."""
with torch.no_grad():
nn.init.zeros_(self.proj_l.weight)
nn.init.zeros_(self.proj_l.bias) # type: ignore[arg-type]
def _calc_linear_attention(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
) -> torch.Tensor:
"""Compute linear attention: (Q @ K^T @ V) / normalizer."""
kvsum = k.transpose(-1, -2) @ v
ksum = torch.sum(k, dim=-2, keepdim=True)
return (q @ kvsum) / (1e-5 + (q * ksum).sum(dim=-1, keepdim=True))
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
"""Forward pass for SageSLA attention with quantized kernels.
Input tensors are in FastVideo format: (B, L, H, D)
Args:
query: Query tensor (B, L, H, D)
key: Key tensor (B, L, H, D)
value: Value tensor (B, L, H, D)
attn_metadata: Attention metadata
Returns:
Output tensor (B, L, H, D)
"""
original_dtype = query.dtype
# Convert from FastVideo format (B, L, H, D) to SLA format (B, H, L, D)
q = query.transpose(1, 2).contiguous()
k = key.transpose(1, 2).contiguous()
v = value.transpose(1, 2).contiguous()
# Get topk ratio from metadata if available
topk_ratio = self.topk_ratio
if hasattr(attn_metadata, 'topk_ratio'):
topk_ratio = attn_metadata.topk_ratio # type: ignore[union-attr]
# Determine block sizes based on GPU architecture
arch = _get_cuda_arch(q.device.index)
if arch == "sm90":
BLKQ, BLKK = 64, 128
else:
BLKQ, BLKK = 128, 64
# Compute block-sparse attention pattern
sparse_map, lut, real_topk = get_block_map(q, k, topk_ratio=topk_ratio, BLKQ=BLKQ, BLKK=BLKK)
# Convert to compute dtype
q = q.to(self.dtype)
k = k.to(self.dtype)
v = v.to(self.dtype)
# ========== SPARGE QUANTIZED ATTENTION ==========
km = k.mean(dim=-2, keepdim=True)
headdim = q.size(-1)
scale = 1.0 / (headdim**0.5)
# Quantize Q, K to INT8
q_int8, q_scale, k_int8, k_scale = get_vanilla_qk_quant(q, k, km, BLKQ, BLKK)
lut_triton, valid_block_num = block_map_lut_triton(sparse_map)
# Quantize V to FP8
b, h_kv, kv_len, head_dim = v.shape
padded_len = (kv_len + 127) // 128 * 128
v_transposed_permutted = torch.empty((b, h_kv, head_dim, padded_len), dtype=v.dtype, device=v.device)
fused.transpose_pad_permute_cuda(v, v_transposed_permutted, 1)
v_fp8 = torch.empty(v_transposed_permutted.shape, dtype=torch.float8_e4m3fn, device=v.device)
v_scale = torch.empty((b, h_kv, head_dim), dtype=torch.float32, device=v.device)
fused.scale_fuse_quant_cuda(v_transposed_permutted, v_fp8, v_scale, kv_len, 2.25, 1)
# Sparse attention with quantized kernels
o_s = torch.empty_like(q)
if arch == "sm90":
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_sm90(
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num, q_scale, k_scale, v_scale, 1, False, 1, scale)
else:
pvthreshold = torch.full((q.shape[-3], ), 1e6, dtype=torch.float32, device=q.device)
if SAGE2PP_ENABLED:
qk_int8_sv_f8_accum_f16_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num, pvthreshold, q_scale, k_scale, v_scale, 1,
False, 1, scale, 0)
else:
qattn.qk_int8_sv_f8_accum_f32_block_sparse_attn_inst_buf_fuse_v_scale_with_pv_threshold(
q_int8, k_int8, v_fp8, o_s, lut_triton, valid_block_num, pvthreshold, q_scale, k_scale, v_scale, 1,
False, 1, scale, 0)
# ========== END SPARGE ==========
# Linear attention with feature maps (see SLAAttentionImpl.forward
# for why .contiguous() is unnecessary here).
q_linear = self.feature_map_q(q).to(self.dtype)
k_linear = self.feature_map_k(k).to(self.dtype)
o_l = self._calc_linear_attention(q_linear, k_linear, v)
# Project linear attention output and combine
with torch.amp.autocast('cuda', dtype=self.dtype):
o_l = self.proj_l(o_l)
# Combine sparse and linear outputs
output = (o_s + o_l).to(original_dtype)
# Convert back to FastVideo format (B, L, H, D)
output = output.transpose(1, 2)
return output
@@ -0,0 +1,319 @@
# SPDX-License-Identifier: Apache-2.0
import functools
import math
from dataclasses import dataclass
import torch
try:
from fastvideo_kernel import video_sparse_attn
except ImportError:
video_sparse_attn = None
try:
from fastvideo_kernel import video_sparse_attn_bshd
except ImportError:
video_sparse_attn_bshd = None
from typing import Any
from v2._vendor.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from v2._vendor.distributed import get_sp_group
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
# VSA tile shape. The tile volume picks the kernel path automatically in
# forward(): (4,4,4)=64 -> existing TK/Triton path (default, unchanged);
# (4,8,8)=256 -> FA4 CuTe block-sparse attention fastpath (Blackwell).
VSA_TILE_SIZE = (4, 4, 4)
@functools.lru_cache(maxsize=10)
def get_tile_partition_indices(
dit_seq_shape: tuple[int, int, int],
tile_size: tuple[int, int, int],
device: torch.device,
) -> torch.LongTensor:
T, H, W = dit_seq_shape
ts, hs, ws = tile_size
indices = torch.arange(T * H * W, device=device, dtype=torch.long).reshape(T, H, W)
ls = []
for t in range(math.ceil(T / ts)):
for h in range(math.ceil(H / hs)):
for w in range(math.ceil(W / ws)):
ls.append(indices[t * ts:min(t * ts + ts, T), h * hs:min(h * hs + hs, H),
w * ws:min(w * ws + ws, W)].flatten())
index = torch.cat(ls, dim=0)
return index
@functools.lru_cache(maxsize=10)
def get_reverse_tile_partition_indices(
dit_seq_shape: tuple[int, int, int],
tile_size: tuple[int, int, int],
device: torch.device,
) -> torch.LongTensor:
return torch.argsort(get_tile_partition_indices(dit_seq_shape, tile_size, device))
@functools.lru_cache(maxsize=10)
def construct_variable_block_sizes(
dit_seq_shape: tuple[int, int, int],
num_tiles: tuple[int, int, int],
device: torch.device,
) -> torch.LongTensor:
"""
Compute the number of valid (non‑padded) tokens inside every
(ts_t × ts_h × ts_w) tile after padding ‑‑ flattened in the order
(t‑tile, h‑tile, w‑tile) that `rearrange` uses.
Returns
-------
torch.LongTensor # shape: [∏ full_window_size]
"""
# unpack
t, h, w = dit_seq_shape
ts_t, ts_h, ts_w = VSA_TILE_SIZE
n_t, n_h, n_w = num_tiles
def _sizes(dim_len: int, tile: int, n_tiles: int) -> torch.LongTensor:
"""Vector with the size of each tile along one dimension."""
sizes = torch.full((n_tiles, ), tile, dtype=torch.int, device=device)
# size of last (possibly partial) tile
remainder = dim_len - (n_tiles - 1) * tile
sizes[-1] = remainder if remainder > 0 else tile
return sizes
t_sizes = _sizes(t, ts_t, n_t) # [n_t]
h_sizes = _sizes(h, ts_h, n_h) # [n_h]
w_sizes = _sizes(w, ts_w, n_w) # [n_w]
# broadcast‑multiply to get voxels per tile, then flatten
block_sizes = (
t_sizes[:, None, None] # [n_t, 1, 1]
* h_sizes[None, :, None] # [1, n_h, 1]
* w_sizes[None, None, :] # [1, 1, n_w]
).reshape(-1) # [n_t * n_h * n_w]
return block_sizes
@functools.lru_cache(maxsize=10)
def get_non_pad_index(
variable_block_sizes: torch.LongTensor,
max_block_size: int,
):
n_win = variable_block_sizes.shape[0]
device = variable_block_sizes.device
starts_pad = torch.arange(n_win, device=device) * max_block_size
index_pad = starts_pad[:, None] + torch.arange(max_block_size, device=device)[None, :]
index_mask = torch.arange(max_block_size, device=device)[None, :] < variable_block_sizes[:, None]
return index_pad[index_mask]
class VideoSparseAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 128]
@staticmethod
def get_name() -> str:
return "VIDEO_SPARSE_ATTN"
@staticmethod
def get_impl_cls() -> type["VideoSparseAttentionImpl"]:
return VideoSparseAttentionImpl
@staticmethod
def get_metadata_cls() -> type["VideoSparseAttentionMetadata"]:
return VideoSparseAttentionMetadata
@staticmethod
def get_builder_cls() -> type["VideoSparseAttentionMetadataBuilder"]:
return VideoSparseAttentionMetadataBuilder
@dataclass
class VideoSparseAttentionMetadata(AttentionMetadata):
current_timestep: int
dit_seq_shape: list[int]
num_tiles: list[int]
total_seq_length: int
tile_partition_indices: torch.LongTensor
reverse_tile_partition_indices: torch.LongTensor
variable_block_sizes: torch.LongTensor
non_pad_index: torch.LongTensor
# Precomputed fancy index that fuses ``x[:, non_pad_index][:, reverse_tile_partition_indices]``
# in postprocess_output(). Avoids materializing the intermediate
# ``[B, len(non_pad_index), H, D]`` tensor on every layer.
untile_combined_index: torch.LongTensor
# Per-step shared padded buffer used by tile(). Inference can reuse this
# across VSA layers, but training disables it so activation checkpointing
# can release the large tiled QKVG scratch tensor after each attention call.
tile_buf: torch.Tensor | None = None
cache_tile_buf: bool = True
class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
VSA_sparsity: float,
device: torch.device,
cache_tile_buf: bool = True,
**kwargs: dict[str, Any],
) -> VideoSparseAttentionMetadata:
patch_size = patch_size
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
raw_latent_shape[2] // patch_size[2])
num_tiles = (math.ceil(dit_seq_shape[0] / VSA_TILE_SIZE[0]), math.ceil(dit_seq_shape[1] / VSA_TILE_SIZE[1]),
math.ceil(dit_seq_shape[2] / VSA_TILE_SIZE[2]))
total_seq_length = math.prod(dit_seq_shape)
tile_partition_indices = get_tile_partition_indices(dit_seq_shape, VSA_TILE_SIZE, device)
reverse_tile_partition_indices = get_reverse_tile_partition_indices(dit_seq_shape, VSA_TILE_SIZE, device)
variable_block_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device)
non_pad_index = get_non_pad_index(variable_block_sizes, math.prod(VSA_TILE_SIZE))
untile_combined_index = non_pad_index[reverse_tile_partition_indices]
return VideoSparseAttentionMetadata(
current_timestep=current_timestep,
dit_seq_shape=dit_seq_shape, # type: ignore
VSA_sparsity=VSA_sparsity, # type: ignore
num_tiles=num_tiles, # type: ignore
total_seq_length=total_seq_length, # type: ignore
tile_partition_indices=tile_partition_indices, # type: ignore
reverse_tile_partition_indices=reverse_tile_partition_indices,
variable_block_sizes=variable_block_sizes,
non_pad_index=non_pad_index,
untile_combined_index=untile_combined_index,
cache_tile_buf=cache_tile_buf)
class VideoSparseAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.prefix = prefix
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
def tile(self, x: torch.Tensor, attn_metadata: VideoSparseAttentionMetadata) -> torch.Tensor:
"""Tile ``x`` into ``attn_metadata.tile_buf`` and return it.
The returned tensor aliases the per-metadata buffer and is only
valid until the next ``tile()`` / ``preprocess_qkv`` call on the
same ``attn_metadata``. Callers must consume (or copy) the
result before invoking another VSA layer with the same metadata.
Today both call sites materialize copies via
``.transpose(...).contiguous()`` inside ``forward()``, so the
contract holds; future callers must preserve it.
"""
num_tiles = attn_metadata.num_tiles
t_padded_size = num_tiles[0] * VSA_TILE_SIZE[0]
h_padded_size = num_tiles[1] * VSA_TILE_SIZE[1]
w_padded_size = num_tiles[2] * VSA_TILE_SIZE[2]
target_shape = (x.shape[0], t_padded_size * h_padded_size * w_padded_size, x.shape[-2], x.shape[-1])
if not attn_metadata.cache_tile_buf:
buf = torch.zeros(target_shape, device=x.device, dtype=x.dtype)
buf[:, attn_metadata.non_pad_index] = x[:, attn_metadata.tile_partition_indices]
return buf
# Reuse the per-step buffer stashed on metadata (lazily allocated
# on the first VSA layer's call within a denoising step). Pad
# positions are zero from the initial torch.zeros and never
# written to. Scoping to metadata makes reuse safe across
# concurrent requests and keeps the "pad positions are zero"
# invariant trivially true: ``non_pad_index`` is fixed within
# a single metadata instance.
buf = attn_metadata.tile_buf
if (buf is None or buf.shape != target_shape or buf.dtype != x.dtype or buf.device != x.device):
buf = torch.zeros(target_shape, device=x.device, dtype=x.dtype)
attn_metadata.tile_buf = buf
buf[:, attn_metadata.non_pad_index] = x[:, attn_metadata.tile_partition_indices]
return buf
def untile(self, x: torch.Tensor, untile_combined_index: torch.LongTensor) -> torch.Tensor:
# Single fancy index using precomputed combined indices; avoids
# the intermediate ``[B, len(non_pad_index), H, D]`` tensor that
# the two-step ``x[:, non_pad_index][:, reverse_tile_partition_indices]``
# would allocate on every layer.
return x[:, untile_combined_index]
def preprocess_qkv(
self,
qkv: torch.Tensor,
attn_metadata: VideoSparseAttentionMetadata,
) -> torch.Tensor:
"""Tile QKV; aliasing contract: see ``tile()``."""
return self.tile(qkv, attn_metadata)
def postprocess_output(
self,
output: torch.Tensor,
attn_metadata: VideoSparseAttentionMetadata,
) -> torch.Tensor:
return self.untile(output, attn_metadata.untile_combined_index)
def forward( # type: ignore[override]
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
gate_compress: torch.Tensor,
attn_metadata: VideoSparseAttentionMetadata,
) -> torch.Tensor:
VSA_sparsity = attn_metadata.VSA_sparsity
block_elements = math.prod(VSA_TILE_SIZE)
cur_topk = math.ceil((1 - VSA_sparsity) * (attn_metadata.total_seq_length / block_elements))
# 256-element tiles auto-route to the FA4 CuTe BSHD fastpath, which
# consumes [B, S, H, D] directly -- skip the transpose round-trip.
if block_elements == 256 and video_sparse_attn_bshd is not None:
return video_sparse_attn_bshd(query,
key,
value,
attn_metadata.variable_block_sizes,
attn_metadata.variable_block_sizes,
cur_topk,
block_size=VSA_TILE_SIZE,
compress_attn_weight=gate_compress)
if video_sparse_attn is None:
raise NotImplementedError("video_sparse_attn is not installed")
# Default 64-element-tile path (unchanged): BHSD round-trip.
query = query.transpose(1, 2).contiguous()
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
gate_compress = gate_compress.transpose(1, 2).contiguous()
return video_sparse_attn(query,
key,
value,
attn_metadata.variable_block_sizes,
attn_metadata.variable_block_sizes,
cur_topk,
block_size=VSA_TILE_SIZE,
compress_attn_weight=gate_compress).transpose(1, 2)
+202
View File
@@ -0,0 +1,202 @@
# SPDX-License-Identifier: Apache-2.0
import re
from dataclasses import dataclass
import torch
from einops import rearrange
from fastvideo_kernel import (moba_attn_varlen, process_moba_input, process_moba_output)
from v2._vendor.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
class VMOBAAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_name() -> str:
return "VMOBA_ATTN"
@staticmethod
def get_impl_cls() -> type["VMOBAAttentionImpl"]:
return VMOBAAttentionImpl
@staticmethod
def get_metadata_cls() -> type["VideoMobaAttentionMetadata"]:
return VideoMobaAttentionMetadata
@staticmethod
def get_builder_cls() -> type["VideoMobaAttentionMetadataBuilder"]:
return VideoMobaAttentionMetadataBuilder
@dataclass
class VideoMobaAttentionMetadata(AttentionMetadata):
current_timestep: int
temporal_chunk_size: int
temporal_topk: int
spatial_chunk_size: tuple[int, int]
spatial_topk: int
st_chunk_size: tuple[int, int, int]
st_topk: int
moba_select_mode: str
moba_threshold: float
moba_threshold_type: str
patch_resolution: list[int]
first_full_step: int = 12
first_full_layer: int = 0
# temporal_layer -> spatial_layer -> st_layer
temporal_layer: int = 1
spatial_layer: int = 1
st_layer: int = 1
class VideoMobaAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
temporal_chunk_size: int,
temporal_topk: int,
spatial_chunk_size: tuple[int, int],
spatial_topk: int,
st_chunk_size: tuple[int, int, int],
st_topk: int,
moba_select_mode: str = 'threshold',
moba_threshold: float = 0.25,
moba_threshold_type: str = 'query_head',
device: torch.device | None = None,
first_full_layer: int = 0,
first_full_step: int = 12,
temporal_layer: int = 1,
spatial_layer: int = 1,
st_layer: int = 1,
**kwargs,
) -> VideoMobaAttentionMetadata:
if device is None:
device = torch.device("cpu")
assert raw_latent_shape[0] % patch_size[0] == 0 and raw_latent_shape[1] % patch_size[
1] == 0 and raw_latent_shape[2] % patch_size[
2] == 0, f"spatial patch_resolution {raw_latent_shape} should be divisible by patch_size {patch_size}"
patch_resolution = [t // pt for t, pt in zip(raw_latent_shape, patch_size, strict=False)]
return VideoMobaAttentionMetadata(
current_timestep=current_timestep,
temporal_chunk_size=temporal_chunk_size,
temporal_topk=temporal_topk,
spatial_chunk_size=spatial_chunk_size,
spatial_topk=spatial_topk,
st_chunk_size=st_chunk_size,
st_topk=st_topk,
moba_select_mode=moba_select_mode,
moba_threshold=moba_threshold,
moba_threshold_type=moba_threshold_type,
patch_resolution=patch_resolution,
first_full_layer=first_full_layer,
first_full_step=first_full_step,
temporal_layer=temporal_layer,
spatial_layer=spatial_layer,
st_layer=st_layer,
)
class VMOBAAttentionImpl(AttentionImpl):
def __init__(self,
num_heads,
head_size,
softmax_scale,
causal=False,
num_kv_heads=None,
prefix="",
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
from flash_attn.bert_padding import pad_input
self.pad_input = pad_input
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
if not match:
raise ValueError(f"Invalid prefix: {prefix}")
return int(match.group(1))
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: VideoMobaAttentionMetadata,
) -> torch.Tensor:
"""
query: [B, L, H, D]
key: [B, L, H, D]
value: [B, L, H, D]
attn_metadata: AttentionMetadata
"""
batch_size, sequence_length, num_heads, head_dim = query.shape
# select chunk type according to layer idx:
loop_layer_num = attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer
assert self.layer_idx is not None, "VMoBA attention requires layer_idx to be set"
moba_layer = self.layer_idx - attn_metadata.first_full_layer
moba_chunk_size: int | tuple[int, int] | tuple[int, int, int]
if moba_layer % loop_layer_num < attn_metadata.temporal_layer:
moba_chunk_size = attn_metadata.temporal_chunk_size
moba_topk = attn_metadata.temporal_topk
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer:
moba_chunk_size = attn_metadata.spatial_chunk_size
moba_topk = attn_metadata.spatial_topk
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer:
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
else:
raise ValueError(f"Invalid MoBA layer selection for layer {moba_layer}")
query, chunk_size = process_moba_input(query, attn_metadata.patch_resolution, moba_chunk_size)
key, chunk_size = process_moba_input(key, attn_metadata.patch_resolution, moba_chunk_size)
value, chunk_size = process_moba_input(value, attn_metadata.patch_resolution, moba_chunk_size)
max_seqlen = query.shape[1]
indices_q = torch.arange(0, query.shape[0] * query.shape[1], device=query.device)
cu_seqlens = torch.arange(0,
query.shape[0] * query.shape[1] + 1,
query.shape[1],
dtype=torch.int32,
device=query.device)
query = rearrange(query, "b s ... -> (b s) ...")
key = rearrange(key, "b s ... -> (b s) ...")
value = rearrange(value, "b s ... -> (b s) ...")
# current_timestep=attn_metadata.current_timestep
hidden_states = moba_attn_varlen(
query,
key,
value,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=chunk_size,
moba_topk=moba_topk,
select_mode=attn_metadata.moba_select_mode,
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = self.pad_input(hidden_states, indices_q, batch_size, sequence_length)
hidden_states = process_moba_output(hidden_states, attn_metadata.patch_resolution, moba_chunk_size)
return hidden_states
+287
View File
@@ -0,0 +1,287 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import torch.nn as nn
from v2._vendor.attention.selector import backend_name_to_enum, get_attn_backend
from v2._vendor.distributed.communication_op import (sequence_model_parallel_all_gather,
sequence_model_parallel_all_to_all_4D)
from v2._vendor.distributed.parallel_state import (get_sp_parallel_rank, get_sp_world_size)
from v2._vendor.forward_context import ForwardContext, get_forward_context
from v2._vendor.platforms import AttentionBackendEnum
from v2._vendor.utils import get_compute_dtype
from v2._vendor.layers.rotary_embedding import _apply_rotary_emb
class DistributedAttention(nn.Module):
"""Distributed attention layer.
"""
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
else:
self.softmax_scale = softmax_scale
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = get_compute_dtype()
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
# Register attn_impl as submodule if it has learnable parameters (e.g., SLA's proj_l)
# This ensures its parameters are included in state_dict() for saving/loading
if isinstance(self.attn_impl, nn.Module):
self.add_module('attn_impl', self.attn_impl)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
self.backend = backend_name_to_enum(attn_backend.get_name())
self.dtype = dtype
@torch.compiler.disable
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
original_seq_len: int | None = None,
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
Args:
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
original_seq_len (int): Original (unpadded) full sequence length
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
- o (torch.Tensor): Output tensor after attention for the main sequence
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
"""
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
batch_size, _, num_heads, _ = q.shape
local_rank = get_sp_parallel_rank()
world_size = get_sp_world_size()
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
# Stack QKV
qkv = torch.cat([q, k, v], dim=0) # [3*batch, seq_len, num_heads, head_dim]
# Redistribute heads across sequence dimension
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
# After all-to-all, each rank has the full sequence but only a subset of heads.
# Trim away SP padding for attention compute, then pad back before returning.
original_seq_len = original_seq_len or qkv.shape[1]
pad_seq_len = qkv.shape[1] - original_seq_len
qkv = qkv[:, :original_seq_len, :, :]
if freqs_cis is not None:
cos, sin = freqs_cis
qkv[:batch_size * 2] = _apply_rotary_emb(qkv[:batch_size * 2], cos, sin, is_neox_style=False)
# Apply backend-specific preprocess_qkv
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
# Concatenate with replicated QKV if provided
if replicated_q is not None:
assert replicated_k is not None and replicated_v is not None
replicated_qkv = torch.cat([replicated_q, replicated_k, replicated_v],
dim=0) # [3, seq_len, num_heads, head_dim]
heads_per_rank = num_heads // world_size
replicated_qkv = replicated_qkv[:, :, local_rank * heads_per_rank:(local_rank + 1) * heads_per_rank]
qkv = torch.cat([qkv, replicated_qkv], dim=1)
q, k, v = qkv.chunk(3, dim=0)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
# Redistribute back if using sequence parallelism
replicated_output = None
if replicated_q is not None:
split_idx = original_seq_len
replicated_output = output[:, split_idx:]
output = output[:, :split_idx]
# TODO: make this asynchronous
replicated_output = sequence_model_parallel_all_gather(replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_seq_len))
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
return output, replicated_output
class DistributedAttention_VSA(DistributedAttention):
"""Distributed attention layer with VSA support.
"""
@torch.compiler.disable
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
original_seq_len: int,
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
gate_compress: torch.Tensor | None = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
Args:
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
original_seq_len (int): Original (unpadded) full sequence length
gate_compress (torch.Tensor): Gate compress tensor [batch_size, seq_len, num_heads, head_dim]
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
- o (torch.Tensor): Output tensor after attention for the main sequence
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
"""
# Check text tokens are not supported for VSA now
assert replicated_q is None and replicated_k is None and replicated_v is None, "Replicated QKV is not supported for VSA now"
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
batch_size, seq_len, num_heads, head_dim = q.shape
# Stack QKV
qkvg = torch.cat([q, k, v, gate_compress], dim=0) # [4*batch, seq_len, num_heads, head_dim]
# Redistribute heads across sequence dimension
# Before: [4*batch, shard_seq_len, num_heads, head_dim]
# After: [4*batch, full_seq_len, shard_num_heads, head_dim]
qkvg = sequence_model_parallel_all_to_all_4D(qkvg, scatter_dim=2, gather_dim=1)
# After all-to-all, each rank has the full sequence but only a subset of heads
pad_seq_len = qkvg.shape[1] - original_seq_len
qkvg = qkvg[:, :original_seq_len, :, :]
if freqs_cis is not None:
cos, sin = freqs_cis
qkvg[:batch_size * 2] = _apply_rotary_emb(qkvg[:batch_size * 2], cos, sin, is_neox_style=False)
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
output = self.attn_impl.forward(q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
# Redistribute back if using sequence parallelism
replicated_output = None
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_seq_len))
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
return output, replicated_output
class LocalAttention(nn.Module):
"""Attention layer.
"""
def __init__(self,
num_heads: int,
head_size: int,
num_kv_heads: int | None = None,
softmax_scale: float | None = None,
causal: bool = False,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
else:
self.softmax_scale = softmax_scale
if num_kv_heads is None:
num_kv_heads = num_heads
dtype = get_compute_dtype()
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
self.backend = backend_name_to_enum(attn_backend.get_name())
self.dtype = dtype
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""
Apply local attention between query, key and value tensors.
Args:
q (torch.Tensor): Query tensor of shape [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor of shape [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor of shape [batch_size, seq_len, num_heads, head_dim]
Returns:
torch.Tensor: Output tensor after local attention
"""
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
if freqs_cis is not None:
cos, sin = freqs_cis
q = _apply_rotary_emb(q, cos, sin, is_neox_style=False)
k = _apply_rotary_emb(k, cos, sin, is_neox_style=False)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return output
+154
View File
@@ -0,0 +1,154 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py
import os
from collections.abc import Generator
from contextlib import contextmanager
from functools import cache
from typing import cast
import torch
import v2._vendor.envs as envs
from v2._vendor.attention.backends.abstract import AttentionBackend
from v2._vendor.logger import init_logger
from v2._vendor.platforms import AttentionBackendEnum
from v2._vendor.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
def backend_name_to_enum(backend_name: str) -> AttentionBackendEnum | None:
"""
Convert a string backend name to a _Backend enum value.
Returns:
* _Backend: enum value if backend_name is a valid in-tree type
* None: otherwise it's an invalid in-tree type or an out-of-tree platform is
loaded.
"""
assert backend_name is not None
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
None
def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
Returns:
* _Backend enum value if an override is specified
* None otherwise
'''
backend_name = os.environ.get(STR_BACKEND_ENV_VAR)
return (None if backend_name is None else backend_name_to_enum(backend_name))
# Global state allows a particular choice of backend
# to be forced, overriding the logic which auto-selects
# a backend based on system & workload configuration
# (default behavior if this variable is None)
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: AttentionBackendEnum | None = None
def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None:
'''
Force all attention operations to use a specified backend.
Passing `None` for the argument re-enables automatic
backend selection.,
Arguments:
* attn_backend: backend selection (None to revert to auto)
'''
global forced_attn_backend
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
'''
return forced_attn_backend
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
@cache
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
#
# THIS SELECTION OVERRIDES THE FASTVIDEO_ATTENTION_BACKEND
# ENVIRONMENT VARIABLE.
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: AttentionBackendEnum | None = (get_global_forced_attn_backend())
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
else:
# Check the environment variable and override if specified
backend_by_env_var: str | None = envs.FASTVIDEO_ATTENTION_BACKEND
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
# get device-specific attn_backend
from v2._vendor.platforms import current_platform
if selected_backend not in supported_attention_backends:
selected_backend = None
attention_cls = current_platform.get_attn_backend_cls(selected_backend, head_size, dtype)
if not attention_cls:
raise ValueError(f"Invalid attention backend for {current_platform.device_name}")
return cast(type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
@contextmanager
def global_force_attn_backend_context_manager(attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
'''
Globally force a FastVideo attention backend override within a
context manager, reverting the global attention backend
override to its prior state upon exiting the context
manager.
Arguments:
* attn_backend: attention backend to force
Returns:
* Generator
'''
# Save the current state of the global backend override (if any)
original_value = get_global_forced_attn_backend()
# Globally force the new backend override
global_force_attn_backend(attn_backend)
# Yield control back to the enclosed code block
try:
yield
finally:
# Revert the original global backend override, if any
global_force_attn_backend(original_value)
@@ -0,0 +1,327 @@
from __future__ import annotations
import torch
if torch.cuda.is_available():
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
else:
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
def _check_dropout(dropout_p: float) -> None:
if dropout_p != 0.0:
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
out, lse = _flash_attn_fwd(
q,
k,
v,
softmax_scale=softmax_scale,
causal=causal,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=1,
pack_gqa=None,
)
return out, lse
@torch.library.register_fake("fastvideo::_flash_attn_cute_forward")
def _flash_attn_cute_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del k, softmax_scale, causal, deterministic
batch, seqlen_q, nheads = q.shape[:3]
out = q.new_empty(batch, seqlen_q, nheads, v.shape[-1])
lse = q.new_empty(batch, nheads, seqlen_q, dtype=torch.float32)
return out, lse
def _flash_attn_cute_setup_context(ctx: torch.autograd.function.FunctionCtx, inputs, output) -> None:
q, k, v, softmax_scale, causal, deterministic = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse)
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.deterministic = deterministic
def _flash_attn_cute_backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse = ctx.saved_tensors
dq, dk, dv = _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
softcap=0.0,
window_size_left=None,
window_size_right=None,
deterministic=ctx.deterministic,
)
return dq, dk, dv, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_cute_forward",
_flash_attn_cute_backward,
setup_context=_flash_attn_cute_setup_context,
)
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_varlen_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_varlen_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
out, lse = _flash_attn_fwd(
q,
k,
v,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
softmax_scale=softmax_scale,
causal=causal,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=1,
pack_gqa=None,
)
return out, lse
@torch.library.register_fake("fastvideo::_flash_attn_cute_varlen_forward")
def _flash_attn_cute_varlen_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
softmax_scale: float | None,
causal: bool,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del k, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, softmax_scale
del causal
del deterministic
total_q, nheads = q.shape[:2]
out = q.new_empty(total_q, nheads, v.shape[-1])
lse = q.new_empty(nheads, total_q, dtype=torch.float32)
return out, lse
def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionCtx, inputs, output) -> None:
(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
softmax_scale,
causal,
deterministic,
) = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k)
ctx.max_seqlen_q = max_seqlen_q
ctx.max_seqlen_k = max_seqlen_k
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.deterministic = deterministic
def _flash_attn_cute_varlen_backward(
ctx: torch.autograd.function.FunctionCtx,
grad_out: torch.Tensor,
grad_lse: torch.Tensor | None,
):
del grad_lse
q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors
dq, dk, dv = _flash_attn_bwd(
q,
k,
v,
out,
grad_out,
lse,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
softcap=0.0,
window_size_left=None,
window_size_right=None,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=ctx.max_seqlen_q,
max_seqlen_k=ctx.max_seqlen_k,
deterministic=ctx.deterministic,
)
return dq, dk, dv, None, None, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_cute_varlen_forward",
_flash_attn_cute_varlen_backward,
setup_context=_flash_attn_cute_varlen_setup_context,
)
def flash_attn_func(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
causal: bool = False,
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
_check_dropout(dropout_p)
out, _ = torch.ops.v2._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
return out
# ---------------------------------------------------------------------------
# FP4 (NVFP4 block-scaled) variant
# ---------------------------------------------------------------------------
# The FP4 path needs the mSFQ/mSFK scale-factor tensors that the regular
# wrapper does not expose. We register a separate custom op so that
# torch.compile can treat the kernel as an opaque boundary (the underlying
# CuTeDSL kernel uses cuda.CUstream which dynamo cannot trace).
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_fp4_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_cute_fp4_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
sfq: torch.Tensor,
sfk: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
out, _ = _flash_attn_fwd(
q,
k,
v,
softmax_scale=softmax_scale,
causal=causal,
window_size_left=None,
window_size_right=None,
softcap=0.0,
num_splits=1,
pack_gqa=None,
mSFQ=sfq,
mSFK=sfk,
)
return out
@torch.library.register_fake("fastvideo::_flash_attn_cute_fp4_forward")
def _flash_attn_cute_fp4_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
sfq: torch.Tensor,
sfk: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
del k, sfq, sfk, softmax_scale, causal
# q is FP4 packed: shape (batch, seqlen, nheads, headdim/2). Output is in
# V's dtype with full headdim.
batch, seqlen_q, nheads = q.shape[:3]
return v.new_empty(batch, seqlen_q, nheads, v.shape[-1])
def flash_attn_fp4_func(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
sfq: torch.Tensor,
sfk: torch.Tensor,
softmax_scale: float | None = None,
causal: bool = False,
) -> torch.Tensor:
"""FP4 (NVFP4 block-scaled) flash attention. q/k are FP4-packed; v is BF16."""
return torch.ops.v2._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
def flash_attn_varlen_func(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens_q: torch.Tensor,
cu_seqlens_k: torch.Tensor,
max_seqlen_q: int,
max_seqlen_k: int,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
causal: bool = False,
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
_check_dropout(dropout_p)
out, _ = torch.ops.v2._flash_attn_cute_varlen_forward(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
softmax_scale,
causal,
deterministic,
)
return out
@@ -0,0 +1,182 @@
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results there from are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
from typing import Any
import torch
from einops import rearrange
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
def _resolve_flash_attn_varlen_func() -> Any:
try:
from v2._vendor.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
return flash_attn_varlen_func_cute
except ImportError:
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
return flash_attn_varlen_func_interface
except ImportError:
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
return flash_attn_varlen_func_flash
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
def flash_attn_no_pad(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool = False,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
deterministic: bool = False,
) -> torch.Tensor:
batch_size = qkv.shape[0]
seqlen = qkv.shape[1]
nheads = qkv.shape[-2]
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads)
output_unpad = flash_attn_varlen_qkvpacked_func(
x_unpad,
cu_seqlens,
max_s,
dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
output = rearrange(
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"),
indices,
batch_size,
seqlen,
),
"b s (h d) -> b s h d",
h=nheads,
)
return output
def flash_attn_no_pad_v3(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool = False,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
deterministic: bool = False,
) -> torch.Tensor:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_v3, )
if flash_attn_varlen_func_v3 is None:
raise ImportError("FlashAttention V3 backend not available")
batch_size, seqlen, _, nheads, head_dim = qkv.shape
query, key, value = qkv.unbind(dim=2)
query_unpad, indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
key_padding_mask)
key_unpad, _, cu_seqlens_k, _, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
value_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
output_unpad = flash_attn_varlen_func_v3(
query_unpad,
key_unpad,
value_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_q,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
output = rearrange(
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"),
indices,
batch_size,
seqlen,
),
"b s (h d) -> b s h d",
h=nheads,
)
return output
def flash_attn_varlen_qk_no_pad(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool = False,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
deterministic: bool = False,
) -> torch.Tensor:
batch_size, q_seqlen, nheads, _ = query.shape
query_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
query_padding_mask)
key_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
value_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
output_unpad = flash_attn_varlen_func_impl(
query_unpad,
key_unpad,
value_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
output = rearrange(
pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"),
q_indices,
batch_size,
q_seqlen,
),
"b s (h d) -> b s h d",
h=nheads,
)
return output
+53
View File
@@ -0,0 +1,53 @@
# `fastvideo/configs/` — Config-Driven Model Registry
**Generated:** 2026-05-02
Two layers of dataclass configs feed every pipeline: **arch configs** (what the model is) and **pipeline configs** (how to run it).
## Layout
```
configs/
├── configs.py # Dataset / loader enums (DatasetType, VideoLoaderType)
├── utils.py # update_config_from_args, shallow_asdict helpers
├── backend/ # Attention backend defaults
├── models/
│ ├── base.py # ModelConfig ABC
│ ├── dits/ # DiTConfig per model (wanvideo, ltx2, hunyuan, ...)
│ ├── vaes/ # VAEConfig per model
│ ├── encoders/ # EncoderConfig (t5, clip, llama, qwen2_5, gemma, siglip, ...)
│ ├── upsamplers/ # UpsamplerConfig (hunyuan15)
│ └── audio/ # Audio-model configs (ltx2_audio_vae, ...)
├── pipelines/
│ ├── base.py # PipelineConfig ABC + (de)serialization
│ └── <model>.py # Concrete configs (HunyuanConfig, WanT2V480PConfig, ...)
└── *.json # Frozen reference configs for shipped models
```
## How Configs Hook Into the Registry
`fastvideo/registry.py` imports every concrete `PipelineConfig` and exposes
`get_pipeline_config_cls_from_name(...)`. Adding a new pipeline config requires:
1. Subclass `PipelineConfig` in `pipelines/<model>.py`.
2. Reference its component arch configs (DiT / VAE / encoder / upsampler).
3. Add the import + name mapping in `fastvideo/registry.py`.
Configs that do not appear in `registry.py` are unreachable from `VideoGenerator`.
## Arch vs Pipeline — Where Does This Field Go?
| Field type | Lives on |
|-----------|----------|
| Architecture constants (hidden dim, num heads, layer count) | `configs/models/<role>/<model>.py` |
| Default sampling params (steps, cfg, shift, fps) | `configs/pipelines/<model>.py` |
| Runtime overrides (precision, sp_size, tp_size, attention backend) | `configs/pipelines/base.py` defaults + CLI flags via `fastvideo_args.py` |
| `param_names_mapping` for HF → FastVideo state-dict | Arch config (lives with the model definition) |
If a knob is tunable per inference call → `SamplingParam`, not `PipelineConfig`.
## Anti-Patterns
- Hard-coding architecture constants inside model classes — always read from the arch config.
- Using `argparse` directly here. Configs deserialize from dicts via `update_config_from_args`.
- Importing from `fastvideo.pipelines` here. Configs are the lower layer; the dependency is one-way.
View File
@@ -0,0 +1,16 @@
{
"temporal_chunk_size": 2,
"temporal_topk": 2,
"spatial_chunk_size": [4, 13],
"spatial_topk": 6,
"st_chunk_size": [4, 4, 13],
"st_topk": 18,
"moba_select_mode": "topk",
"moba_threshold": 0.25,
"moba_threshold_type": "query_head",
"first_full_layer": 0,
"first_full_step": 12,
"temporal_layer": 1,
"spatial_layer": 1,
"st_layer": 1
}
@@ -0,0 +1,16 @@
{
"temporal_chunk_size": 2,
"temporal_topk": 3,
"spatial_chunk_size": [3, 4],
"spatial_topk": 20,
"st_chunk_size": [4, 6, 4],
"st_topk": 15,
"moba_select_mode": "threshold",
"moba_threshold": 0.25,
"moba_threshold_type": "query_head",
"first_full_layer": 0,
"first_full_step": 12,
"temporal_layer": 1,
"spatial_layer": 1,
"st_layer": 1
}
+213
View File
@@ -0,0 +1,213 @@
import dataclasses
from enum import Enum
from typing import Any, Optional
from v2._vendor.configs.utils import update_config_from_args
from v2._vendor.logger import init_logger
from v2._vendor.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
class DatasetType(str, Enum):
"""
Enumeration for different dataset types.
"""
HF = "hf"
MERGED = "merged"
@classmethod
def from_string(cls, value: str) -> "DatasetType":
"""Convert string to DatasetType enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid dataset type: {value}. Must be one of: {', '.join([m.value for m in cls])}") from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [dataset_type.value for dataset_type in cls]
class VideoLoaderType(str, Enum):
"""
Enumeration for different video loaders.
"""
TORCHCODEC = "torchcodec"
TORCHVISION = "torchvision"
@classmethod
def from_string(cls, value: str) -> "VideoLoaderType":
"""Convert string to VideoLoader enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid video loader: {value}. Must be one of: {', '.join([m.value for m in cls])}") from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [video_loader.value for video_loader in cls]
@dataclasses.dataclass
class PreprocessConfig:
"""Configuration for preprocessing operations."""
# Model and dataset configuration
model_path: str = ""
dataset_path: str = ""
dataset_type: DatasetType = DatasetType.HF
dataset_output_dir: str = "./output"
# Dataloader configuration
dataloader_num_workers: int = 1
preprocess_video_batch_size: int = 2
# Saver configuration
samples_per_file: int = 64
flush_frequency: int = 256
# Video processing parameters
video_loader_type: VideoLoaderType = VideoLoaderType.TORCHCODEC
max_height: int = 480
max_width: int = 848
num_frames: int = 163
video_length_tolerance_range: float = 2.0
train_fps: int = 30
speed_factor: float = 1.0
drop_short_ratio: float = 1.0
do_temporal_sample: bool = False
# Model configuration
training_cfg_rate: float = 0.0
with_audio: bool = False
# framework configuration
seed: int = 42
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser, prefix: str = "preprocess") -> FlexibleArgumentParser:
"""Add preprocessing configuration arguments to the parser."""
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
preprocess_args = parser.add_argument_group("Preprocessing Arguments")
# Model & Dataset
preprocess_args.add_argument(f"--{prefix_with_dot}model-path",
type=str,
default=PreprocessConfig.model_path,
help="Path to the model for preprocessing")
preprocess_args.add_argument(f"--{prefix_with_dot}dataset-path",
type=str,
default=PreprocessConfig.dataset_path,
help="Path to the dataset directory for preprocessing")
preprocess_args.add_argument(f"--{prefix_with_dot}dataset-type",
type=str,
choices=DatasetType.choices(),
default=PreprocessConfig.dataset_type.value,
help="Type of the dataset")
preprocess_args.add_argument(f"--{prefix_with_dot}dataset-output-dir",
type=str,
default=PreprocessConfig.dataset_output_dir,
help="The output directory where the dataset will be written.")
# Dataloader
preprocess_args.add_argument(
f"--{prefix_with_dot}dataloader-num-workers",
type=int,
default=PreprocessConfig.dataloader_num_workers,
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.")
preprocess_args.add_argument(f"--{prefix_with_dot}preprocess-video-batch-size",
type=int,
default=PreprocessConfig.preprocess_video_batch_size,
help="Batch size (per device) for the training dataloader.")
# Saver
preprocess_args.add_argument(f"--{prefix_with_dot}samples-per-file",
type=int,
default=PreprocessConfig.samples_per_file,
help="Number of samples per output file")
preprocess_args.add_argument(f"--{prefix_with_dot}flush-frequency",
type=int,
default=PreprocessConfig.flush_frequency,
help="How often to save to parquet files")
# Video processing parameters
preprocess_args.add_argument(f"--{prefix_with_dot}video-loader-type",
type=str,
choices=VideoLoaderType.choices(),
default=PreprocessConfig.video_loader_type.value,
help="Type of the video loader")
preprocess_args.add_argument(f"--{prefix_with_dot}max-height",
type=int,
default=PreprocessConfig.max_height,
help="Maximum height for video processing")
preprocess_args.add_argument(f"--{prefix_with_dot}max-width",
type=int,
default=PreprocessConfig.max_width,
help="Maximum width for video processing")
preprocess_args.add_argument(f"--{prefix_with_dot}num-frames",
type=int,
default=PreprocessConfig.num_frames,
help="Number of frames to process")
preprocess_args.add_argument(f"--{prefix_with_dot}video-length-tolerance-range",
type=float,
default=PreprocessConfig.video_length_tolerance_range,
help="Video length tolerance range")
preprocess_args.add_argument(f"--{prefix_with_dot}train-fps",
type=int,
default=PreprocessConfig.train_fps,
help="Training FPS")
preprocess_args.add_argument(f"--{prefix_with_dot}speed-factor",
type=float,
default=PreprocessConfig.speed_factor,
help="Speed factor for video processing")
preprocess_args.add_argument(f"--{prefix_with_dot}drop-short-ratio",
type=float,
default=PreprocessConfig.drop_short_ratio,
help="Ratio for dropping short videos")
preprocess_args.add_argument(f"--{prefix_with_dot}do-temporal-sample",
action=StoreBoolean,
default=PreprocessConfig.do_temporal_sample,
help="Whether to do temporal sampling")
# Model Training configuration
preprocess_args.add_argument(f"--{prefix_with_dot}training-cfg-rate",
type=float,
default=PreprocessConfig.training_cfg_rate,
help="Training CFG rate")
preprocess_args.add_argument(f"--{prefix_with_dot}with-audio",
action=StoreBoolean,
default=PreprocessConfig.with_audio,
help="Whether to extract and encode audio")
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
type=int,
default=PreprocessConfig.seed,
help="Seed for random number generator")
return parser
@classmethod
def from_kwargs(cls, kwargs: dict[str, Any]) -> Optional["PreprocessConfig"]:
"""Create PreprocessConfig from keyword arguments."""
if 'dataset_type' in kwargs and isinstance(kwargs['dataset_type'], str):
kwargs['dataset_type'] = DatasetType.from_string(kwargs['dataset_type'])
if 'video_loader_type' in kwargs and isinstance(kwargs['video_loader_type'], str):
kwargs['video_loader_type'] = VideoLoaderType.from_string(kwargs['video_loader_type'])
preprocess_config = cls()
if not update_config_from_args(preprocess_config, kwargs, prefix="preprocess", pop_args=True):
return None
return preprocess_config
def check_preprocess_config(self) -> None:
if self.dataset_path == "":
raise ValueError("dataset_path must be set for preprocess mode")
if self.samples_per_file <= 0:
raise ValueError("samples_per_file must be greater than 0")
if self.flush_frequency <= 0:
raise ValueError("flush_frequency must be greater than 0")
+47
View File
@@ -0,0 +1,47 @@
{
"embedded_cfg_scale": 6,
"flow_shift": 17,
"dit_cpu_offload": false,
"disable_autocast": false,
"precision": "bf16",
"vae_precision": "fp32",
"vae_tiling": true,
"vae_sp": true,
"vae_config": {
"load_encoder": false,
"load_decoder": true,
"tile_sample_min_height": 256,
"tile_sample_min_width": 256,
"tile_sample_min_num_frames": 16,
"tile_sample_stride_height": 192,
"tile_sample_stride_width": 192,
"tile_sample_stride_num_frames": 12,
"blend_num_frames": 4,
"use_tiling": true,
"use_temporal_tiling": true,
"use_parallel_tiling": true
},
"dit_config": {
"prefix": "Hunyuan",
"quant_config": null
},
"text_encoder_precisions": [
"fp16",
"fp16"
],
"text_encoder_configs": [
{
"prefix": "llama",
"quant_config": null,
"lora_config": null
},
{
"prefix": "clip",
"quant_config": null,
"lora_config": null,
"num_hidden_layers_override": null,
"require_post_norm": null
}
],
"enable_torch_compile": false
}
+17
View File
@@ -0,0 +1,17 @@
from v2._vendor.configs.models.base import ModelConfig
from v2._vendor.configs.models.dits.base import DiTConfig
from v2._vendor.configs.models.encoders.base import EncoderConfig
from v2._vendor.configs.models.vaes.base import VAEConfig
from v2._vendor.configs.models.audio import (LTX2AudioDecoderConfig, LTX2AudioEncoderConfig, LTX2VocoderConfig)
from v2._vendor.configs.models.upsamplers.base import UpsamplerConfig
__all__ = [
"ModelConfig",
"VAEConfig",
"DiTConfig",
"EncoderConfig",
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
"UpsamplerConfig",
]
@@ -0,0 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
from v2._vendor.configs.models.audio.ltx2_audio_vae import (
LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
)
__all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -0,0 +1,28 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 audio VAE and vocoder configuration.
"""
from dataclasses import dataclass, field
from v2._vendor.configs.models.base import ArchConfig, ModelConfig
@dataclass
class LTX2AudioArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=list)
@dataclass
class LTX2AudioEncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(architectures=["LTX2AudioEncoder"]))
@dataclass
class LTX2AudioDecoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(architectures=["LTX2AudioDecoder"]))
@dataclass
class LTX2VocoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(architectures=["LTX2Vocoder"]))
+68
View File
@@ -0,0 +1,68 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field, fields
from typing import Any
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
# 1. ArchConfig contains all fields from diffuser's/transformer's config.json (i.e. all fields related to the architecture of the model)
# 2. ArchConfig should be inherited & overridden by each model arch_config
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
stacked_params_mapping: list[tuple[str, str, str]] = field(
default_factory=list) # mapping from huggingface weight names to custom names
@dataclass
class ModelConfig:
# Every model config parameter can be categorized into either ArchConfig or everything else
# Diffuser/Transformer parameters
arch_config: ArchConfig = field(default_factory=ArchConfig)
# FastVideo-specific parameters here
def __getattr__(self, name):
# Only called if 'name' is not found in ModelConfig directly
if hasattr(self.arch_config, name):
return getattr(self.arch_config, name)
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
def __getstate__(self):
# Return a dictionary of attributes to pickle
# Convert to dict and exclude any problematic attributes
state = self.__dict__.copy()
return state
def __setstate__(self, state):
# Restore instance attributes from the unpickled state
self.__dict__.update(state)
# This should be used only when loading from transformers/diffusers
def update_model_arch(self, source_model_dict: dict[str, Any]) -> None:
arch_config = self.arch_config
valid_fields = {f.name for f in fields(arch_config)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(arch_config, key, value)
if hasattr(arch_config, "__post_init__"):
arch_config.__post_init__()
def update_model_config(self, source_model_dict: dict[str, Any]) -> None:
assert "arch_config" not in source_model_dict, "Source model config shouldn't contain arch_config."
valid_fields = {f.name for f in fields(self)}
for key, value in source_model_dict.items():
if key in valid_fields:
setattr(self, key, value)
else:
logger.warning("%s does not contain field '%s'!", type(self).__name__, key)
raise AttributeError(f"Invalid field: {key}")
if hasattr(self, "__post_init__"):
self.__post_init__()
@@ -0,0 +1,19 @@
from v2._vendor.configs.models.dits.cosmos import CosmosVideoConfig
from v2._vendor.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from v2._vendor.configs.models.dits.flux_2 import Flux2Config
from v2._vendor.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from v2._vendor.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from v2._vendor.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from v2._vendor.configs.models.dits.longcat import LongCatVideoConfig
from v2._vendor.configs.models.dits.ltx2 import LTX2VideoConfig
from v2._vendor.configs.models.dits.magi_human import MagiHumanVideoConfig
from v2._vendor.configs.models.dits.stable_audio import StableAudioConfig
from v2._vendor.configs.models.dits.wanvideo import WanVideoConfig
from v2._vendor.configs.models.dits.hyworld import HYWorldConfig
from v2._vendor.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
]
+71
View File
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Any
from v2._vendor.configs.models.base import ArchConfig, ModelConfig
from v2._vendor.layers.quantization import QuantizationConfig
from v2._vendor.platforms import AttentionBackendEnum
@dataclass
class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
# When True, the denoising stage casts text/prompt embeddings to the DiT's
# working dtype before the diffusion loop. Flux2 requires this (BFL casts ctx
# to bf16 before denoising); models with fp32 text encoders (Wan, Hunyuan15,
# SD3.5) leave it False to preserve full-precision embeddings.
cast_prompt_embeds_to_dit_dtype: bool = False
_supported_attention_backends: tuple[AttentionBackendEnum,
...] = (AttentionBackendEnum.SAGE_ATTN, AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN, AttentionBackendEnum.SAGE_ATTN_THREE,
AttentionBackendEnum.ATTN_QAT_INFER,
AttentionBackendEnum.ATTN_QAT_TRAIN, AttentionBackendEnum.SLA_ATTN,
AttentionBackendEnum.SAGE_SLA_ATTN)
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
in_channels: int = 0
out_channels: int = 0
exclude_lora_layers: list[str] = field(default_factory=list)
boundary_ratio: float | None = None
def __post_init__(self) -> None:
if not self._compile_conditions:
self._compile_conditions = self._fsdp_shard_conditions.copy()
@dataclass
class DiTConfig(ModelConfig):
arch_config: DiTArchConfig = field(default_factory=DiTArchConfig)
# FastVideoDiT-specific parameters
prefix: str = ""
quant_config: QuantizationConfig | None = None
@staticmethod
def add_cli_args(parser: Any, prefix: str = "dit-config") -> Any:
"""Add CLI arguments for DiTConfig fields"""
parser.add_argument(
f"--{prefix}.prefix",
type=str,
dest=f"{prefix.replace('-', '_')}.prefix",
default=DiTConfig.prefix,
help="Prefix for the DiT model",
)
parser.add_argument(
f"--{prefix}.quant-config",
type=str,
dest=f"{prefix.replace('-', '_')}.quant_config",
default=None,
help="Quantization configuration for the DiT model",
)
return parser
+81
View File
@@ -0,0 +1,81 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_transformer_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class CosmosArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embed\.(.*)$": r"patch_embed.\1",
r"^time_embed\.time_proj\.(.*)$": r"time_embed.time_proj.\1",
r"^time_embed\.t_embedder\.(.*)$": r"time_embed.t_embedder.\1",
r"^time_embed\.norm\.(.*)$": r"time_embed.norm.\1",
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"transformer_blocks.\1.attn1.norm_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"transformer_blocks.\1.attn1.norm_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.norm_q\.(.*)$": r"transformer_blocks.\1.attn2.norm_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.norm_k\.(.*)$": r"transformer_blocks.\1.attn2.norm_k.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
r"^norm_out\.(.*)$": r"norm_out.\1",
r"^proj_out\.(.*)$": r"proj_out.\1",
})
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.ff\.(.*)$": r"transformer_blocks.\1.ff.\2",
})
# Cosmos-specific config parameters based on transformer_cosmos.py
# in_channels includes the condition_mask channel (16 latent + 1 cond = 17)
in_channels: int = 17
out_channels: int = 16
num_attention_heads: int = 16
attention_head_dim: int = 128
num_layers: int = 28
mlp_ratio: float = 4.0
text_embed_dim: int = 1024
adaln_lora_dim: int = 256
max_size: tuple[int, int, int] = (128, 240, 240)
patch_size: tuple[int, int, int] = (1, 2, 2)
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0)
concat_padding_mask: bool = True
extra_pos_embed_type: str | None = None
qk_norm: str = "rms_norm"
eps: float = 1e-6
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels
@dataclass
class CosmosVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=CosmosArchConfig)
prefix: str = "Cosmos"
+160
View File
@@ -0,0 +1,160 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_transformer_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class Cosmos25ArchConfig(DiTArchConfig):
"""Configuration for Cosmos 2.5 architecture (MiniTrainDIT)."""
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
# Remove "net." prefix and map official structure to FastVideo
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
r"^net\.x_embedder\.proj\.1\.(.*)$": r"patch_embed.proj.\1",
# Time embedding: net.t_embedder.1.linear_1.weight -> time_embed.t_embedder.linear_1.weight
r"^net\.t_embedder\.1\.linear_1\.(.*)$": r"time_embed.t_embedder.linear_1.\1",
r"^net\.t_embedder\.1\.linear_2\.(.*)$": r"time_embed.t_embedder.linear_2.\1",
# Time embedding norm: net.t_embedding_norm.weight -> time_embed.norm.weight
# Note: This also handles _extra_state if present
r"^net\.t_embedding_norm\.(.*)$": r"time_embed.norm.\1",
# Cross-attention projection (optional): net.crossattn_proj.0.weight -> crossattn_proj.0.weight
r"^net\.crossattn_proj\.0\.weight$": r"crossattn_proj.0.weight",
r"^net\.crossattn_proj\.0\.bias$": r"crossattn_proj.0.bias",
# Transformer blocks: net.blocks.N -> transformer_blocks.N
# Self-attention (self_attn -> attn1)
r"^net\.blocks\.(\d+)\.self_attn\.q_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
r"^net\.blocks\.(\d+)\.self_attn\.k_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
r"^net\.blocks\.(\d+)\.self_attn\.v_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
r"^net\.blocks\.(\d+)\.self_attn\.output_proj\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\.weight$": r"transformer_blocks.\1.attn1.norm_q.weight",
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\.weight$": r"transformer_blocks.\1.attn1.norm_k.weight",
# RMSNorm _extra_state keys (internal PyTorch state, will be recomputed automatically)
r"^net\.blocks\.(\d+)\.self_attn\.q_norm\._extra_state$":
r"transformer_blocks.\1.attn1.norm_q._extra_state",
r"^net\.blocks\.(\d+)\.self_attn\.k_norm\._extra_state$":
r"transformer_blocks.\1.attn1.norm_k._extra_state",
# Cross-attention (cross_attn -> attn2)
r"^net\.blocks\.(\d+)\.cross_attn\.q_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.k_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.v_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.output_proj\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\.weight$": r"transformer_blocks.\1.attn2.norm_q.weight",
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\.weight$": r"transformer_blocks.\1.attn2.norm_k.weight",
# RMSNorm _extra_state keys for cross-attention
r"^net\.blocks\.(\d+)\.cross_attn\.q_norm\._extra_state$":
r"transformer_blocks.\1.attn2.norm_q._extra_state",
r"^net\.blocks\.(\d+)\.cross_attn\.k_norm\._extra_state$":
r"transformer_blocks.\1.attn2.norm_k._extra_state",
# MLP: net.blocks.N.mlp.layer1 -> transformer_blocks.N.mlp.fc_in
r"^net\.blocks\.(\d+)\.mlp\.layer1\.(.*)$": r"transformer_blocks.\1.mlp.fc_in.\2",
r"^net\.blocks\.(\d+)\.mlp\.layer2\.(.*)$": r"transformer_blocks.\1.mlp.fc_out.\2",
# AdaLN-LoRA modulations: net.blocks.N.adaln_modulation_* -> transformer_blocks.N.adaln_modulation_*
# These are now at the block level, not inside norm layers
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.1\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_self_attn.1.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_self_attn\.2\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_self_attn.2.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.1\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_cross_attn.1.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_cross_attn\.2\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_cross_attn.2.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.1\.(.*)$": r"transformer_blocks.\1.adaln_modulation_mlp.1.\2",
r"^net\.blocks\.(\d+)\.adaln_modulation_mlp\.2\.(.*)$": r"transformer_blocks.\1.adaln_modulation_mlp.2.\2",
# Layer norms: net.blocks.N.layer_norm_* -> transformer_blocks.N.norm*.norm
r"^net\.blocks\.(\d+)\.layer_norm_self_attn\._extra_state$":
r"transformer_blocks.\1.norm1.norm._extra_state",
r"^net\.blocks\.(\d+)\.layer_norm_cross_attn\._extra_state$":
r"transformer_blocks.\1.norm2.norm._extra_state",
r"^net\.blocks\.(\d+)\.layer_norm_mlp\._extra_state$": r"transformer_blocks.\1.norm3.norm._extra_state",
# Final layer: net.final_layer.linear -> final_layer.proj_out
r"^net\.final_layer\.linear\.(.*)$": r"final_layer.proj_out.\1",
# Final layer AdaLN-LoRA: net.final_layer.adaln_modulation -> final_layer.linear_*
r"^net\.final_layer\.adaln_modulation\.1\.(.*)$": r"final_layer.linear_1.\1",
r"^net\.final_layer\.adaln_modulation\.2\.(.*)$": r"final_layer.linear_2.\1",
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
# - net.pos_embedder.* (seq, dim_spatial_range, dim_temporal_range) - These are computed dynamically
# in FastVideo's Cosmos25RotaryPosEmbed forward() method, so they don't need to be loaded.
# - net.accum_* keys (training metadata) - These are skipped during checkpoint loading.
})
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$": r"transformer_blocks.\1.mlp.\2",
})
# Cosmos 2.5 specific config parameters
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 16
attention_head_dim: int = 128 # 2048 / 16
num_layers: int = 28
mlp_ratio: float = 4.0
text_embed_dim: int = 1024
adaln_lora_dim: int = 256
use_adaln_lora: bool = True
max_size: tuple[int, int, int] = (128, 240, 240)
patch_size: tuple[int, int, int] = (1, 2, 2)
rope_scale: tuple[float, float, float] = (1.0, 3.0, 3.0) # T, H, W scaling
concat_padding_mask: bool = True
extra_pos_embed_type: str | None = None # "learnable" or None
# Note: Official checkpoint has use_crossattn_projection=True with 100K-dim input from Qwen 7B.
# When enabled, must provide 100,352-dim embeddings to match the projection layer in checkpoint.
use_crossattn_projection: bool = False
crossattn_proj_in_channels: int = 100352 # Qwen 7B embedding dimension
rope_enable_fps_modulation: bool = True
qk_norm: str = "rms_norm"
eps: float = 1e-6
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels
@dataclass
class Cosmos25_14BArchConfig(Cosmos25ArchConfig):
"""Configuration for Cosmos 2.5 14B architecture."""
num_attention_heads: int = 40
attention_head_dim: int = 128 # 5120 / 40
num_layers: int = 36
@dataclass
class Cosmos25VideoConfig(DiTConfig):
"""Configuration for Cosmos 2.5 video generation model."""
arch_config: DiTArchConfig = field(default_factory=Cosmos25ArchConfig)
prefix: str = "Cosmos25"
@dataclass
class Cosmos25_14BVideoConfig(DiTConfig):
"""Configuration for Cosmos 2.5 14B video generation model."""
arch_config: DiTArchConfig = field(default_factory=Cosmos25_14BArchConfig)
prefix: str = "Cosmos25"
+77
View File
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
# Copied and adapted from: https://github.com/sglang-ai/sglang
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
from v2._vendor.logger import init_logger
logger = init_logger(__name__)
@dataclass
class Flux2ArchConfig(DiTArchConfig):
"""Architecture configuration for Flux2 transformer model."""
cast_prompt_embeds_to_dit_dtype: bool = True
# Flux2-specific architecture parameters
patch_size: int = 1
in_channels: int = 64
out_channels: int | None = None
num_layers: int = 19 # Number of double-stream transformer blocks
num_single_layers: int = 38 # Number of single-stream transformer blocks
attention_head_dim: int = 128
num_attention_heads: int = 24
joint_attention_dim: int = 4096 # Dimension for text encoder output
timestep_guidance_channels: int = 256 # Dimension for timestep embedding
mlp_ratio: float = 3.0
axes_dims_rope: tuple[int, ...] = (32, 32, 32, 32) # RoPE dimensions per axis (match diffusers Flux2)
rope_theta: int = 2000 # Base frequency for RoPE (match diffusers Flux2)
eps: float = 1e-6
guidance_embeds: bool = True # Whether to use guidance embeddings
# When True, compute SwiGLU in fp32 inside ``ff_context`` only (bf16 noise mitigation).
ff_context_swiglu_fp32: bool = False
# Parameter name mapping for loading HuggingFace checkpoints
param_names_mapping: dict = field(default_factory=lambda: {
r"transformer\.(\w*)\.(.*)$": r"\1.\2",
})
def __post_init__(self) -> None:
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
def update_from_weight_keys(self, all_keys: set[str]) -> None:
"""Infer num_layers and num_single_layers from checkpoint weight keys so the model is built with the same number of blocks as the weights."""
if not all_keys:
return
num_layers = 0
num_single_layers = 0
for k in all_keys:
if "single_transformer_blocks." not in k and "transformer_blocks." in k:
parts = k.split("transformer_blocks.")[-1].split(".")
if parts[0].isdigit():
num_layers = max(num_layers, int(parts[0]) + 1)
if "single_transformer_blocks." in k:
parts = k.split("single_transformer_blocks.")[-1].split(".")
if parts[0].isdigit():
num_single_layers = max(num_single_layers, int(parts[0]) + 1)
if num_layers > 0:
self.num_layers = num_layers
logger.info("Inferred num_layers=%s from checkpoint keys", num_layers)
if num_single_layers > 0:
self.num_single_layers = num_single_layers
logger.info("Inferred num_single_layers=%s from checkpoint keys", num_single_layers)
if num_layers > 0 or num_single_layers > 0:
self.__post_init__()
@dataclass
class Flux2Config(DiTConfig):
"""Configuration for Flux2 transformer model."""
arch_config: DiTArchConfig = field(default_factory=Flux2ArchConfig)
prefix: str = "Flux"
+188
View File
@@ -0,0 +1,188 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_transformer_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class Gen3CArchConfig(DiTArchConfig):
"""Configuration for GEN3C architecture (VideoExtendGeneralDIT)."""
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_transformer_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
# Official GEN3C checkpoint key naming to FastVideo mapping.
# The official checkpoint uses nn.Sequential patterns like attn.to_q.0 (Linear)
# and attn.to_q.1 (RMSNorm), and layer1/layer2 for MLP.
#
# Patch embedding: net.x_embedder.proj.1.weight -> patch_embed.proj.weight
r"^net\.x_embedder\.proj\.1\.(.*)$": r"patch_embed.proj.\1",
# Time embedding: net.t_embedder.1.linear_*.weight -> time_embed.t_embedder.linear_*.weight
r"^net\.t_embedder\.0\.(.*)$": r"time_embed.time_proj.\1",
r"^net\.t_embedder\.1\.linear_1\.(.*)$": r"time_embed.t_embedder.linear_1.\1",
r"^net\.t_embedder\.1\.linear_2\.(.*)$": r"time_embed.t_embedder.linear_2.\1",
# Augment sigma embedding (GEN3C-specific)
r"^net\.augment_sigma_embedder\.0\.(.*)$": r"augment_sigma_embed.time_proj.\1",
r"^net\.augment_sigma_embedder\.1\.linear_1\.(.*)$": r"augment_sigma_embed.t_embedder.linear_1.\1",
r"^net\.augment_sigma_embedder\.1\.linear_2\.(.*)$": r"augment_sigma_embed.t_embedder.linear_2.\1",
# Affine embedding norm: net.affline_norm.weight -> affine_norm.weight
# Note: "affline" is a typo in the official GEN3C checkpoint (should be "affine")
r"^net\.affline_norm\.(.*)$": r"affine_norm.\1",
# Extra positional embeddings (learnable per-axis)
r"^net\.extra_pos_embedder\.pos_emb_t$": r"learnable_pos_embed.pos_emb_t",
r"^net\.extra_pos_embedder\.pos_emb_h$": r"learnable_pos_embed.pos_emb_h",
r"^net\.extra_pos_embedder\.pos_emb_w$": r"learnable_pos_embed.pos_emb_w",
# Transformer blocks: net.blocks.blockN -> transformer_blocks.N
# Official uses: block.attn.to_q.0 (Linear), block.attn.to_q.1 (QK RMSNorm)
#
# Self-attention (block index 0)
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_q\.1\.(.*)$":
r"transformer_blocks.\1.attn1.norm_q.\2",
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_k\.1\.(.*)$":
r"transformer_blocks.\1.attn1.norm_k.\2",
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
r"^net\.blocks\.block(\d+)\.blocks\.0\.block\.attn\.to_out\.0\.(.*)$":
r"transformer_blocks.\1.attn1.to_out.\2",
# AdaLN modulation for self-attention
r"^net\.blocks\.block(\d+)\.blocks\.0\.adaLN_modulation\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_self_attn.\2",
# Cross-attention (block index 1)
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_q\.1\.(.*)$":
r"transformer_blocks.\1.attn2.norm_q.\2",
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_k\.1\.(.*)$":
r"transformer_blocks.\1.attn2.norm_k.\2",
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_v\.0\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
r"^net\.blocks\.block(\d+)\.blocks\.1\.block\.attn\.to_out\.0\.(.*)$":
r"transformer_blocks.\1.attn2.to_out.\2",
# AdaLN modulation for cross-attention
r"^net\.blocks\.block(\d+)\.blocks\.1\.adaLN_modulation\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_cross_attn.\2",
# MLP (block index 2): layer1 -> fc_in, layer2 -> fc_out
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer1\.(.*)$": r"transformer_blocks.\1.mlp.fc_in.\2",
r"^net\.blocks\.block(\d+)\.blocks\.2\.block\.layer2\.(.*)$": r"transformer_blocks.\1.mlp.fc_out.\2",
# AdaLN modulation for MLP
r"^net\.blocks\.block(\d+)\.blocks\.2\.adaLN_modulation\.(.*)$":
r"transformer_blocks.\1.adaln_modulation_mlp.\2",
# Final layer: net.final_layer.linear -> final_layer.proj_out
r"^net\.final_layer\.linear\.(.*)$": r"final_layer.proj_out.\1",
# Final layer AdaLN: net.final_layer.adaLN_modulation -> final_layer.adaln_modulation
r"^net\.final_layer\.adaLN_modulation\.(.*)$": r"final_layer.adaln_modulation.\1",
# Note: The following keys from official checkpoint are NOT mapped and can be safely ignored:
# - net.pos_embedder.* (rope position embeddings computed dynamically)
# - net.accum_* keys (training metadata)
# - logvar.* (training-only module, not used in inference)
})
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$": r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$": r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$": r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$": r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$": r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.mlp\.(.*)$": r"transformer_blocks.\1.mlp.\2",
})
# GEN3C architecture parameters
# Base VAE latent channels
in_channels: int = 16
out_channels: int = 16
# Channels per 3D cache buffer: 16 (warped frame latent) + 16 (warped mask latent)
CHANNELS_PER_BUFFER: int = 32
# Number of 3D cache buffers
frame_buffer_max: int = 2
# Attention configuration (7B model: 32 heads x 128 dim = 4096 hidden)
num_attention_heads: int = 32
attention_head_dim: int = 128 # 4096 / 32
num_layers: int = 28
mlp_ratio: float = 4.0
# Text encoder configuration
text_embed_dim: int = 1024
# AdaLN-LoRA configuration
adaln_lora_dim: int = 256
use_adaln_lora: bool = True
# GEN3C-specific: augment sigma embedding for conditioning noise augmentation
# Note: The official GEN3C-Cosmos-7B checkpoint was trained without this
add_augment_sigma_embedding: bool = False
# Position embedding configuration
max_size: tuple[int, int, int] = (128, 240, 240) # T, H, W
patch_size: tuple[int, int, int] = (1, 2, 2)
rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0) # T, H, W scaling
# GEN3C uses learnable positional embeddings in addition to RoPE
extra_pos_embed_type: str = "learnable"
# Padding mask handling
concat_padding_mask: bool = True
# Cross-attention projection (not used in GEN3C 7B)
use_crossattn_projection: bool = False
# RoPE FPS modulation
rope_enable_fps_modulation: bool = True
# QK normalization
qk_norm: str = "rms_norm"
eps: float = 1e-6
# Affine embedding normalization
affine_emb_norm: bool = True
# Block format (THWBD for GEN3C compatibility)
block_x_format: str = "THWBD"
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels
# Calculate total input channels for patch embedding:
# - in_channels (16): VAE latent
# - condition_video_input_mask (1): Binary mask for conditioning frames
# - condition_video_pose (frame_buffer_max * 32): 3D cache buffers
# - padding_mask (1 if concat_padding_mask): Padding mask
self.buffer_channels = self.frame_buffer_max * self.CHANNELS_PER_BUFFER
self.total_input_channels = (
self.in_channels + # 16: VAE latent
1 + # 1: condition_video_input_mask
self.buffer_channels # 64: 3D cache buffers (2 * 32)
)
# padding_mask is added in build_patch_embed if concat_padding_mask=True
@dataclass
class Gen3CVideoConfig(DiTConfig):
"""Configuration for GEN3C video generation model."""
arch_config: DiTArchConfig = field(default_factory=Gen3CArchConfig)
prefix: str = "Gen3C"
@@ -0,0 +1,143 @@
# SPDX-License-Identifier: Apache-2.0
"""
Configuration for HunyuanGameCraft transformer model.
HunyuanGameCraft extends HunyuanVideo with:
1. CameraNet for camera/action conditioning
2. 33 input channels (16 latent + 16 gt_latent + 1 mask)
3. Mask-based conditioning for autoregressive generation
"""
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
def is_camera_net(n: str, m) -> bool:
return "camera_net" in n
@dataclass
class HunyuanGameCraftArchConfig(DiTArchConfig):
"""Architecture config for HunyuanGameCraft transformer."""
# Version field for compatibility with saved config.json
_fastvideo_version: str = "0.1.0"
# Camera net flag (for config.json compatibility)
camera_net: bool = True
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_refiner_block, is_camera_net])
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
# Parameter names mapping from official checkpoint to FastVideo naming
# GameCraft weights are already close to FastVideo format with minor adjustments
param_names_mapping: dict = field(
default_factory=lambda: {
# MLP naming: fc1 -> fc_in, fc2 -> fc_out
r"^(.*)\.img_mlp\.fc1\.(.*)$": r"\1.img_mlp.fc_in.\2",
r"^(.*)\.img_mlp\.fc2\.(.*)$": r"\1.img_mlp.fc_out.\2",
r"^(.*)\.txt_mlp\.fc1\.(.*)$": r"\1.txt_mlp.fc_in.\2",
r"^(.*)\.txt_mlp\.fc2\.(.*)$": r"\1.txt_mlp.fc_out.\2",
# Single block MLP naming
r"^single_blocks\.(\d+)\.mlp\.fc1\.(.*)$": r"single_blocks.\1.mlp.fc_in.\2",
r"^single_blocks\.(\d+)\.mlp\.fc2\.(.*)$": r"single_blocks.\1.mlp.fc_out.\2",
# Token refiner naming
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.(.*)$": r"txt_in.refiner_blocks.\1.\2",
# Vector in naming
r"^vector_in\.in_layer\.(.*)$": r"vector_in.fc_in.\1",
r"^vector_in\.out_layer\.(.*)$": r"vector_in.fc_out.\1",
# Time embedder naming
r"^time_in\.mlp\.0\.(.*)$": r"time_in.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$": r"time_in.mlp.fc_out.\1",
# Guidance embedder naming (if present)
r"^guidance_in\.mlp\.0\.(.*)$": r"guidance_in.mlp.fc_in.\1",
r"^guidance_in\.mlp\.2\.(.*)$": r"guidance_in.mlp.fc_out.\1",
# Final layer adaLN modulation
r"^final_layer\.adaLN_modulation\.1\.(.*)$": r"final_layer.adaLN_modulation.linear.\1",
# Refiner block MLP naming
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc1\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.refiner_blocks\.(\d+)\.mlp\.fc2\.(.*)$": r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
# Camera net weights are already correctly named
})
# Reverse mapping for saving checkpoints
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Model architecture parameters
# patch_size can be int or tuple - if tuple, it's [T, H, W]
patch_size: int | tuple[int, int, int] = 2
patch_size_t: int = 1
in_channels: int = 33 # 16 latent + 16 gt_latent + 1 mask
out_channels: int = 16
num_attention_heads: int = 24
attention_head_dim: int = 128
mlp_ratio: float = 4.0
num_layers: int = 20 # Double stream blocks
num_single_layers: int = 40 # Single stream blocks
num_refiner_layers: int = 2
rope_axes_dim: tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False # GameCraft doesn't use guidance
dtype: torch.dtype | None = None
text_embed_dim: int = 4096 # LLaMA-3 hidden size
pooled_projection_dim: int = 768 # CLIP pooled output dim
rope_theta: int = 256
qk_norm: str = "rms_norm"
# Camera net parameters
camera_in_channels: int = 6 # Plücker coordinates
camera_downscale_coef: int = 8
camera_out_channels: int = 16
# Layers to exclude from LoRA
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in", "camera_net"])
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = 16 # Output is 16 channels
# Convert patch_size list to tuple if needed (from JSON deserialization)
if isinstance(self.patch_size, list):
self.patch_size = tuple(self.patch_size)
# Convert rope_axes_dim list to tuple if needed
if isinstance(self.rope_axes_dim, list):
self.rope_axes_dim = tuple(self.rope_axes_dim)
@dataclass
class HunyuanGameCraftConfig(DiTConfig):
"""Full config for HunyuanGameCraft model."""
arch_config: DiTArchConfig = field(default_factory=HunyuanGameCraftArchConfig)
prefix: str = "HunyuanGameCraft"
@@ -0,0 +1,168 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HunyuanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_double_block, is_single_block, is_refiner_block])
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
# 4. Top-level time_text_embed mappings:
r"^time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.mlp.fc_in.\1",
r"^time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.mlp.fc_out.\1",
r"^time_text_embed\.guidance_embedder\.linear_1\.(.*)$":
r"guidance_in.mlp.fc_in.\1",
r"^time_text_embed\.guidance_embedder\.linear_2\.(.*)$":
r"guidance_in.mlp.fc_out.\1",
r"^time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"vector_in.fc_in.\1",
r"^time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"vector_in.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 6. single_transformer_blocks mapping:
r"^single_transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"single_blocks.\1.q_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"single_blocks.\1.k_norm.\2",
r"^single_transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (r"single_blocks.\1.linear1.\2", 0, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (r"single_blocks.\1.linear1.\2", 1, 4),
r"^single_transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (r"single_blocks.\1.linear1.\2", 2, 4),
r"^single_transformer_blocks\.(\d+)\.proj_mlp\.(.*)$": (r"single_blocks.\1.linear1.\2", 3, 4),
# Corrected: map proj_out to modulation.linear rather than a separate proj_out branch.
r"^single_transformer_blocks\.(\d+)\.proj_out\.(.*)$":
r"single_blocks.\1.linear2.\2",
r"^single_transformer_blocks\.(\d+)\.norm\.linear\.(.*)$":
r"single_blocks.\1.modulation.linear.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: int = 2
patch_size_t: int = 1
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 24
attention_head_dim: int = 128
mlp_ratio: float = 4.0
num_layers: int = 20
num_single_layers: int = 40
num_refiner_layers: int = 2
rope_axes_dim: tuple[int, int, int] = (16, 56, 56)
guidance_embeds: bool = False
dtype: torch.dtype | None = None
text_embed_dim: int = 4096
pooled_projection_dim: int = 768
rope_theta: int = 256
qk_norm: str = "rms_norm"
exclude_lora_layers: list[str] = field(default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = self.in_channels
@dataclass
class HunyuanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HunyuanVideoArchConfig)
prefix: str = "Hunyuan"
@@ -0,0 +1,150 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double_blocks" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HunyuanVideo15ArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block])
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
# 2. txt_in_2 mapping:
r"^context_embedder_2\.(.*)$":
r"txt_in_2.\1",
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
# 4. Top-level time_text_embed mappings:
r"^time_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_out.\1",
r"^time_embed\.timestep_embedder_r\.linear_1\.(.*)$":
r"time_in.timestep_embedder_r.mlp.fc_in.\1",
r"^time_embed\.timestep_embedder_r\.linear_2\.(.*)$":
r"time_in.timestep_embedder_r.mlp.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
in_channels: int = 65
out_channels: int = 32
num_attention_heads: int = 16
attention_head_dim: int = 128
num_layers: int = 54
num_refiner_layers: int = 2
mlp_ratio: float = 4.0
patch_size: int = 1
patch_size_t: int = 1
qk_norm: str = "rms_norm"
text_embed_dim: int = 3584
text_embed_2_dim: int = 1472
image_embed_dim: int = 1152
rope_theta: float = 256.0
rope_axes_dim: tuple[int, ...] = (16, 56, 56)
target_size: int = 640
task_type: str = "i2v"
use_meanflow: bool = False
exclude_lora_layers: list[str] = field(default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = self.out_channels
@dataclass
class HunyuanVideo15Config(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HunyuanVideo15ArchConfig)
prefix: str = "Hunyuan15"
+169
View File
@@ -0,0 +1,169 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HYWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block])
_compile_conditions: list = field(default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. txt_in submodules (text embedder, refiner blocks):
r"^txt_in\.t_embedder\.mlp\.0\.(.*)$": r"txt_in.t_embedder.mlp.fc_in.\1",
r"^txt_in\.t_embedder\.mlp\.2\.(.*)$": r"txt_in.t_embedder.mlp.fc_out.\1",
r"^txt_in\.c_embedder\.linear_1\.(.*)$": r"txt_in.c_embedder.fc_in.\1",
r"^txt_in\.c_embedder\.linear_2\.(.*)$": r"txt_in.c_embedder.fc_out.\1",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm1\.(.*)$": r"txt_in.refiner_blocks.\1.norm1.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm2\.(.*)$": r"txt_in.refiner_blocks.\1.norm2.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_proj\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 2. time_in mappings:
r"^time_in\.mlp\.0\.(.*)$": r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$": r"time_in.timestep_embedder.mlp.fc_out.\1",
# 3. action_in mappings:
r"^action_in\.mlp\.0\.(.*)$": r"action_in.mlp.fc_in.\1",
r"^action_in\.mlp\.2\.(.*)$": r"action_in.mlp.fc_out.\1",
# 4. byt5_in -> txt_in_2 mappings:
r"^byt5_in\.layernorm\.(.*)$": r"txt_in_2.norm.\1",
r"^byt5_in\.fc1\.(.*)$": r"txt_in_2.linear_1.\1",
r"^byt5_in\.fc2\.(.*)$": r"txt_in_2.linear_2.\1",
r"^byt5_in\.fc3\.(.*)$": r"txt_in_2.linear_3.\1",
# 5. cond_type_embedding -> cond_type_embed:
r"^cond_type_embedding\.(.*)$": r"cond_type_embed.\1",
# 6. vision_in -> image_embedder mappings:
r"^vision_in\.proj\.0\.(.*)$": r"image_embedder.norm_in.\1",
r"^vision_in\.proj\.1\.(.*)$": r"image_embedder.linear_1.\1",
r"^vision_in\.proj\.3\.(.*)$": r"image_embedder.linear_2.\1",
r"^vision_in\.proj\.4\.(.*)$": r"image_embedder.norm_out.\1",
# 7. double_blocks mapping:
r"^double_blocks\.(\d+)\.img_attn_q\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.img_attn_k\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.img_attn_v\.(.*)$": (r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.txt_attn_q\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.txt_attn_k\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.txt_attn_v\.(.*)$": (r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.img_mlp\.fc1\.(.*)$": r"double_blocks.\1.img_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.img_mlp\.fc2\.(.*)$": r"double_blocks.\1.img_mlp.fc_out.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc1\.(.*)$": r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc2\.(.*)$": r"double_blocks.\1.txt_mlp.fc_out.\2",
# 8. Final layer mapping:
r"^final_layer\.adaLN_modulation\.1\.(.*)$": r"final_layer.adaLN_modulation.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Parameters from HY-WorldPlay config.json (loaded from checkpoint)
patch_size: list | tuple | int = field(default_factory=lambda: [1, 1, 1])
# Base latent channels - will be expanded in __post_init__ if concat_condition=True
in_channels: int = 32
concat_condition: bool = True
out_channels: int = 32
hidden_size: int = 2048
heads_num: int = 16
mlp_width_ratio: float = 4.0
mlp_act_type: str = "gelu_tanh"
mm_double_blocks_depth: int = 54
mm_single_blocks_depth: int = 0
rope_dim_list: list | tuple = field(default_factory=lambda: [16, 56, 56])
qkv_bias: bool = True
qk_norm: bool | str = True
qk_norm_type: str = "rms"
guidance_embed: bool = False
use_meanflow: bool = False
text_projection: str = "single_refiner"
use_attention_mask: bool = True
text_states_dim: int = 3584
text_states_dim_2: int | None = None
text_pool_type: str | None = None
rope_theta: float = 256.0
attn_mode: str = "flash"
attn_param: str | None = None
glyph_byT5_v2: bool = True
vision_projection: str = "linear"
vision_states_dim: int = 1152
is_reshape_temporal_channels: bool = False
use_cond_type_embedding: bool = True
ideal_resolution: str = "480p"
ideal_task: str = "i2v"
task_type: str = "i2v"
exclude_lora_layers: list[str] = field(default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
# Convert HY-WorldPlay naming to FastVideo naming conventions
self.num_attention_heads: int = self.heads_num
self.attention_head_dim: int = self.hidden_size // self.heads_num
self.num_layers: int = self.mm_double_blocks_depth
self.num_single_layers: int = self.mm_single_blocks_depth
self.num_refiner_layers: int = 2 # Default for HYWorld
self.mlp_ratio: float = float(self.mlp_width_ratio)
self.text_embed_dim: int = self.text_states_dim
self.text_embed_2_dim: int = self.text_states_dim_2 if self.text_states_dim_2 else 1472
self.image_embed_dim: int = self.vision_states_dim
self.rope_axes_dim: tuple[int, ...] = tuple(self.rope_dim_list)
self.num_channels_latents: int = self.out_channels
self.target_size: int = 640
# Handle concat_condition: when True, actual in_channels = base * 2 + 1
# (base latent + condition latent + mask channel)
# config.json has base in_channels (32), but img_in needs full (65)
if self.concat_condition and self.in_channels == 32:
if self.is_reshape_temporal_channels:
self.in_channels = self.in_channels + self.in_channels // 2 + 1
else:
self.in_channels = self.in_channels * 2 + 1 # 32 * 2 + 1 = 65
# Handle patch_size (can be list/tuple or int)
if isinstance(self.patch_size, list | tuple):
self.patch_size_t: int = self.patch_size[0]
# assume square patch size for height and width
patch_size_hw: int = self.patch_size[1]
object.__setattr__(self, 'patch_size', patch_size_hw)
else:
self.patch_size_t = 1
# Convert qk_norm to string format
if isinstance(self.qk_norm, bool):
if self.qk_norm:
self.qk_norm = "rms_norm" if self.qk_norm_type == "rms" else self.qk_norm_type
else:
self.qk_norm = "none"
@dataclass
class HYWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HYWorldArchConfig)
prefix: str = "HYWorld"
@@ -0,0 +1,65 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class Kandinsky5ArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [
lambda n, m:
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
])
# Native FastVideo implementation uses the same parameter names as diffusers
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
# FastVideo uses MLP `fc_in/fc_out`.
param_names_mapping: dict = field(
default_factory=lambda: {
r"^(.*feed_forward)\.in_layer\.(weight|bias)$": r"\1.mlp.fc_in.\2",
r"^(.*feed_forward)\.out_layer\.(weight|bias)$": r"\1.mlp.fc_out.\2",
})
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Diffusers Kandinsky5Transformer3DModel config fields.
in_visual_dim: int = 4
in_text_dim: int = 3584
in_text_dim2: int = 768
time_dim: int = 512
out_visual_dim: int = 4
patch_size: tuple[int, int, int] = (1, 2, 2)
model_dim: int = 2048
ff_dim: int = 5120
num_text_blocks: int = 2
num_visual_blocks: int = 32
axes_dims: tuple[int, int, int] = (16, 24, 24)
visual_cond: bool = False
attention_type: str = "regular"
attention_causal: bool | None = None
attention_local: bool | None = None
attention_glob: bool | None = None
attention_window: int | None = None
attention_P: float | None = None
attention_wT: int | None = None
attention_wW: int | None = None
attention_wH: int | None = None
attention_add_sta: bool | None = None
attention_method: str | None = None
def __post_init__(self):
super().__post_init__()
head_dim = sum(self.axes_dims)
if self.model_dim % head_dim != 0:
raise ValueError(f"model_dim ({self.model_dim}) must be divisible by head_dim ({head_dim})")
self.hidden_size = self.model_dim
self.num_attention_heads = self.model_dim // head_dim
self.in_channels = self.in_visual_dim
self.out_channels = self.out_visual_dim
self.num_channels_latents = self.in_visual_dim
@dataclass
class Kandinsky5VideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=Kandinsky5ArchConfig)
prefix: str = "Kandinsky5"
@@ -0,0 +1,96 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LingBotWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^patch_embedding_wancamctrl\.(.*)$": r"patch_embedding_wancamctrl.proj.\1",
r"^c2ws_hidden_states_layer1\.(.*)$": r"c2ws_mlp.fc_in.\1",
r"^c2ws_hidden_states_layer2\.(.*)$": r"c2ws_mlp.fc_out.\1",
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.cam_injector_layer1\.(.*)$": r"blocks.\1.cam_conditioner.cam_injector.fc_in.\2",
r"^blocks\.(\d+)\.cam_injector_layer2\.(.*)$": r"blocks.\1.cam_conditioner.cam_injector.fc_out.\2",
r"^blocks\.(\d+)\.cam_scale_layer\.(.*)$": r"blocks.\1.cam_conditioner.cam_scale_layer.\2",
r"^blocks\.(\d+)\.cam_shift_layer\.(.*)$": r"blocks.\1.cam_conditioner.cam_shift_layer.\2",
r"^head\.modulation$": r"scale_shift_table",
r"^head\.head\.(.*)$": r"proj_out.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
lora_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class LingBotWorldVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=LingBotWorldArchConfig)
prefix: str = "Wan"
+133
View File
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
"""
LongCat Video DiT configuration for native FastVideo implementation.
"""
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
from v2._vendor.platforms import AttentionBackendEnum
def is_longcat_blocks(n: str, m) -> bool:
"""FSDP shard condition for LongCat transformer blocks."""
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LongCatVideoArchConfig(DiTArchConfig):
"""Architecture configuration for native LongCat Video DiT."""
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_longcat_blocks])
# Enable torch.compile for transformer blocks (major speedup!)
_compile_conditions: list = field(default_factory=lambda: [is_longcat_blocks])
# Parameter name mapping for weight conversion
param_names_mapping: dict = field(
default_factory=lambda: {
# Embedders
r"^x_embedder\.(.*)$": r"patch_embed.\1",
r"^t_embedder\.mlp\.0\.(.*)$": r"time_embedder.linear_1.\1",
r"^t_embedder\.mlp\.2\.(.*)$": r"time_embedder.linear_2.\1",
r"^y_embedder\.y_proj\.0\.(.*)$": r"caption_embedder.linear_1.\1",
r"^y_embedder\.y_proj\.2\.(.*)$": r"caption_embedder.linear_2.\1",
# Transformer blocks - AdaLN modulation
r"^blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$": r"blocks.\1.adaln_linear_1.\2",
# Transformer blocks - Normalization
r"^blocks\.(\d+)\.mod_norm_attn\.(.*)$": r"blocks.\1.norm_attn.\2",
r"^blocks\.(\d+)\.mod_norm_ffn\.(.*)$": r"blocks.\1.norm_ffn.\2",
r"^blocks\.(\d+)\.pre_crs_attn_norm\.(.*)$": r"blocks.\1.norm_cross.\2",
# Self-attention: QKV fused -> separate (will need splitting in converter)
# Original has attn.qkv.weight -> need to split into to_q, to_k, to_v
r"^blocks\.(\d+)\.attn\.qkv\.(.*)$": r"blocks.\1.self_attn.qkv_fused.\2", # Marker for splitting
r"^blocks\.(\d+)\.attn\.proj\.(.*)$": r"blocks.\1.self_attn.to_out.\2",
r"^blocks\.(\d+)\.attn\.q_norm\.(.*)$": r"blocks.\1.self_attn.q_norm.\2",
r"^blocks\.(\d+)\.attn\.k_norm\.(.*)$": r"blocks.\1.self_attn.k_norm.\2",
# Cross-attention
r"^blocks\.(\d+)\.cross_attn\.q_linear\.(.*)$": r"blocks.\1.cross_attn.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.kv_linear\.(.*)$":
r"blocks.\1.cross_attn.kv_fused.\2", # Marker for splitting
r"^blocks\.(\d+)\.cross_attn\.proj\.(.*)$": r"blocks.\1.cross_attn.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.q_norm\.(.*)$": r"blocks.\1.cross_attn.q_norm.\2",
r"^blocks\.(\d+)\.cross_attn\.k_norm\.(.*)$": r"blocks.\1.cross_attn.k_norm.\2",
# FFN (SwiGLU)
r"^blocks\.(\d+)\.ffn\.w1\.(.*)$": r"blocks.\1.ffn.w1.\2", # gate
r"^blocks\.(\d+)\.ffn\.w2\.(.*)$": r"blocks.\1.ffn.w2.\2", # down
r"^blocks\.(\d+)\.ffn\.w3\.(.*)$": r"blocks.\1.ffn.w3.\2", # up
# Final layer
r"^final_layer\.adaLN_modulation\.1\.(.*)$": r"final_layer.adaln_linear.\1",
r"^final_layer\.norm_final\.(.*)$": r"final_layer.norm.\1",
r"^final_layer\.linear\.(.*)$": r"final_layer.proj.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# LoRA parameter name mapping
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Model architecture parameters
hidden_size: int = 4096
depth: int = 48 # Number of transformer blocks
num_attention_heads: int = 32
attention_head_dim: int = 128 # hidden_size / num_attention_heads
in_channels: int = 16 # Latent space channels
out_channels: int = 16
num_channels_latents: int = 16
# Patch embedding
patch_size: tuple[int, int, int] = (1, 2, 2) # [T, H, W] - no temporal compression
# Text/caption embedding
caption_channels: int = 4096 # UMT5 d_model
# Timestep embedding
adaln_tembed_dim: int = 512
frequency_embedding_size: int = 256
# FFN
mlp_ratio: int = 4
# Attention backend support
_supported_attention_backends: tuple = field(default_factory=lambda: (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
))
# Text padding behavior
text_tokens_zero_pad: bool = True
# Block Sparse Attention (BSA)
enable_bsa: bool = False
bsa_params: dict | None = field(default_factory=lambda: {
"sparsity": 0.9375,
"cdf_threshold": None,
"chunk_3d_shape_q": [4, 4, 4],
"chunk_3d_shape_k": [4, 4, 4],
})
# LoRA exclusions
exclude_lora_layers: list[str] = field(default_factory=lambda: [])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
# Ensure attention_head_dim matches
self.attention_head_dim = self.hidden_size // self.num_attention_heads
@dataclass
class LongCatVideoConfig(DiTConfig):
"""Main configuration for LongCat Video DiT."""
arch_config: DiTArchConfig = field(default_factory=LongCatVideoArchConfig)
prefix: str = "longcat"
+137
View File
@@ -0,0 +1,137 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 Transformer configuration for native FastVideo integration.
"""
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
import re
def is_ltx2_blocks(name: str, _module) -> bool:
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
return res
@dataclass
class LTX2VideoArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 video transformer."""
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo).
# The ``to_gate_compress`` -> ``to_gate_logits`` rules for the LTX-2.3
# gated-attention path are inserted at the front of this dict in
# ``__post_init__`` only when ``apply_gated_attention=True``. Without
# that flag the target model has no ``to_gate_logits`` slot, *and* the
# same-named ``to_gate_compress`` already lives on the LTX-2.0 VSA-QAT
# gate path (plus it is in the default ``lora_target_modules`` list).
# Applying the rename unconditionally silently breaks LTX-2.0 VSA
# checkpoints and default-target LoRAs.
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.diffusion_model\.(.*)$": r"model.\1",
r"^diffusion_model\.(.*)$": r"model.\1",
r"^model\.(.*)$": r"model.\1",
r"^(.*)$": r"model.\1",
})
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Core transformer settings (defaults from LTX-2 metadata)
num_attention_heads: int = 32
attention_head_dim: int = 128
num_layers: int = 48
cross_attention_dim: int = 4096
caption_channels: int = 3840
norm_eps: float = 1e-6
attention_type: str = "default"
rope_type: str = "split"
double_precision_rope: bool = True
# LTX-2.3 gated extensions. All default OFF == LTX-2.0 behavior.
cross_attention_adaln: bool = False
caption_proj_before_connector: bool = False
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(default_factory=lambda: [20, 2048, 2048])
timestep_scale_multiplier: int = 1000
use_middle_indices_grid: bool = True
# Patchification (video-only path)
patch_size: tuple[int, int, int] = (1, 1, 1)
num_channels_latents: int = 128
in_channels: int | None = None
out_channels: int | None = None
# Audio defaults (reserved for joint AV ports)
audio_num_attention_heads: int = 32
audio_attention_head_dim: int = 64
audio_in_channels: int = 128
audio_out_channels: int = 128
audio_cross_attention_dim: int = 2048
audio_positional_embedding_max_pos: list[int] = field(default_factory=lambda: [20])
av_ca_timestep_scale_multiplier: int = 1
# LTX-2.3 gated self-attention (distinct from the VSA-QAT to_gate_compress
# gate). Default OFF == LTX-2.0 behavior.
apply_gated_attention: bool = False
# Text connector/feature extractor compatibility fields carried in some
# transformer configs (used by the LTX-2.3 text stack). Defaults match
# the LTX-2.0 connector layout.
caption_projection_first_linear: bool = True
caption_proj_input_norm: bool = True
caption_projection_second_linear: bool = True
connector_num_attention_heads: int = 30
connector_attention_head_dim: int = 128
connector_num_layers: int = 2
audio_connector_num_attention_heads: int = 30
audio_connector_attention_head_dim: int = 128
audio_connector_num_layers: int = 2
# STG perturbation block index differs across model versions.
# LTX-2.0 defaults to block 29; LTX-2.3 (caption_proj_before_connector)
# uses block 28. ``None`` resolves in __post_init__.
stg_block_idx: int | None = None
def __post_init__(self):
super().__post_init__()
patch_volume = self.patch_size[0] * self.patch_size[1] * self.patch_size[2]
if self.in_channels is None:
self.in_channels = self.num_channels_latents * patch_volume
if self.out_channels is None:
self.out_channels = self.in_channels
if self.stg_block_idx is None:
self.stg_block_idx = 28 if self.caption_proj_before_connector else 29
# LTX-2.3 stores the gated-attention weight under ``to_gate_compress``
# upstream; FastVideo's internal name is ``to_gate_logits``. Only
# enable the rename when the gated path is actually configured: the
# LTX-2.0 attention module's own ``to_gate_compress`` parameter
# (created when the backend is ``VIDEO_SPARSE_ATTN``) and the default
# ``to_gate_compress`` LoRA target both share the upstream name, so
# an unconditional rename would silently retarget them. Inserted at
# the front so first-match-wins matching fires the rename before the
# generic prefix-strip rules.
if self.apply_gated_attention:
gate_rules = {
r"^model\.diffusion_model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^diffusion_model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^model\.(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
r"^(.*)\.to_gate_compress\.(.*)$": r"model.\1.to_gate_logits.\2",
}
self.param_names_mapping = {
**gate_rules,
**self.param_names_mapping,
}
@dataclass
class LTX2VideoConfig(DiTConfig):
"""Main configuration for LTX-2 transformer."""
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
prefix: str = "ltx2"
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: Apache-2.0
"""Architecture / model config for the daVinci-MagiHuman DiT.
The MagiHuman base DiT is a 15B-parameter single-stream transformer that
jointly denoises video, audio, and text tokens in one flat sequence. Layout
details verified against GAIR/daVinci-MagiHuman's base/ shards (2026-04-24).
This file captures only configuration. The module implementation lives in
fastvideo/models/dits/magi_human.py and the pipeline wiring in
fastvideo/pipelines/basic/magi_human/.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def _is_block_layer(n: str, m) -> bool:
# Match "block.layers.<idx>" — the FSDP shard boundary for MagiHuman.
parts = n.split(".")
return (len(parts) >= 3 and parts[0] == "block" and parts[1] == "layers" and str.isdigit(parts[2]))
@dataclass
class MagiHumanArchConfig(DiTArchConfig):
"""MagiHuman base DiT architecture constants.
**Scope contract:** fields here must match the `transformer/config.json`
emitted by `scripts/checkpoint_conversion/convert_magi_human_to_diffusers.py`
1:1, and both are sourced from the upstream Python reference
`inference/common/config.py::ModelConfig` (the HF root `config.json`
is empty so the Python source is canonical). Pipeline-level knobs
(VAE stride, fps, num_inference_steps, CFG scales, flow_shift,
t5_gemma_target_length) and data-proxy knobs (coords_style,
frame_receptive_field, ref_audio_offset, text_offset) live on
`MagiHumanBaseConfig`, NOT here.
`param_names_mapping` is intentionally empty: the FastVideo implementation
keeps the same module tree as the reference (`adapter.*`,
`block.layers.<i>.*`, `final_linear_{video,audio}.*`,
`final_norm_{video,audio}.*`), so converted weights load directly.
"""
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_block_layer])
# No renames needed — the FastVideo module mirrors the reference names.
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
# --- transformer shape ---
num_layers: int = 40
hidden_size: int = 5120
head_dim: int = 128
num_query_groups: int = 8 # num_heads_kv (GQA)
# --- modality channels ---
# video_in_channels = z_dim (48) * patch_size product (1*2*2=4), so the
# embedder receives 192 per token. text_in_channels is T5Gemma-9B's
# encoder hidden size.
video_in_channels: int = 192
audio_in_channels: int = 64
text_in_channels: int = 3584
# --- block-level architecture switches ---
# Sandwich MoE: first and last 4 layers have per-modality experts
# (video/audio/text), middle layers share a single set of weights.
mm_layers: tuple[int, ...] = (0, 1, 2, 3, 36, 37, 38, 39)
local_attn_layers: tuple[int, ...] = ()
gelu7_layers: tuple[int, ...] = (0, 1, 2, 3)
post_norm_layers: tuple[int, ...] = ()
enable_attn_gating: bool = True
activation_type: str = "swiglu7"
# --- DiT patching (upstream `ModelConfig`-equivalent; NOT the VAE
# stride, which is pipeline-level). ---
patch_size: tuple[int, int, int] = (1, 2, 2)
spatial_rope_interpolation: str = "extra"
# --- TReAD (token routing + early drop). Flattened from the upstream
# nested `tread_config` dict so it round-trips through
# `update_model_arch` cleanly. ---
tread_selection_rate: float = 0.5
tread_start_layer_idx: int = 2
tread_end_layer_idx: int = 25
# --- derived fields (populated in __post_init__) ---
num_attention_heads: int = 0 # hidden_size / head_dim
num_heads_kv: int = 0 # == num_query_groups
in_channels: int = 0 # mirror of video_in_channels (FastVideo contract)
out_channels: int = 0 # mirror of video_in_channels
def __post_init__(self) -> None:
super().__post_init__()
self.num_attention_heads = self.hidden_size // self.head_dim
self.num_heads_kv = self.num_query_groups
self.in_channels = self.video_in_channels
self.out_channels = self.video_in_channels
# num_channels_latents is the VAE latent z_dim (48 for Wan 2.2 TI2V-5B).
# We don't declare z_dim on the arch config (it's a VAE property),
# but we still set num_channels_latents for the BaseDiT contract.
self.num_channels_latents = 48
@dataclass
class MagiHumanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=MagiHumanArchConfig)
prefix: str = "magi_human"
@@ -0,0 +1,73 @@
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
@dataclass
class MatrixGame2WanVideoArchConfig(WanVideoArchConfig):
# Override param_names_mapping to remove patch_embedding transformation
# because Matrix-Game 2.0 checkpoints already have patch_embedding.proj format
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(?!proj\.)(.*)$": r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$": r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
})
action_config: dict = field(
default_factory=lambda: {
"blocks": list(range(15)),
"enable_mouse": True,
"enable_keyboard": True,
"heads_num": 16,
"hidden_size": 128,
"img_hidden_size": 1536,
"keyboard_dim_in": 4,
"keyboard_hidden_dim": 1024,
"mouse_dim_in": 2,
"mouse_hidden_dim": 1024,
"mouse_qk_dim_list": [8, 28, 28],
"patch_size": [1, 2, 2],
"qk_norm": True,
"qkv_bias": False,
"rope_dim_list": [8, 28, 28],
"rope_theta": 256,
"vae_time_compression_ratio": 4,
"windows_size": 3,
})
local_attn_size: int = -1
sink_size: int = 0
num_frames_per_block: int = 3
text_len: int = 512
text_dim: int = 0
image_dim: int = 1280
def _is_transformer_block(param_name: str, module: torch.nn.Module) -> bool:
return bool("blocks" in param_name and param_name.split(".")[-1].isdigit())
@dataclass
class MatrixGame2WanVideoConfig(WanVideoConfig):
arch_config: MatrixGame2WanVideoArchConfig = field(default_factory=MatrixGame2WanVideoArchConfig)
prefix: str = "Wan"
_compile_conditions: list = field(default_factory=lambda: [_is_transformer_block])
@@ -0,0 +1,79 @@
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
def _is_transformer_block(param_name: str, module: torch.nn.Module) -> bool:
return bool("blocks" in param_name and param_name.split(".")[-1].isdigit())
@dataclass
class MatrixGame3WanVideoArchConfig(WanVideoArchConfig):
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(weight|bias)$": r"patch_embedding.proj.\1",
r"^patch_embedding_wancamctrl\.(.*)$": r"camera_patch_embedding.proj.\1",
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^head\.head\.(.*)$": r"proj_out.\1",
r"^head\.modulation$": r"scale_shift_table",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
})
patch_size: tuple[int, int, int] = (1, 2, 2)
in_channels: int = 48
out_channels: int = 48
num_attention_heads: int = 24
attention_head_dim: int = 128
ffn_dim: int = 14336
num_layers: int = 30
text_len: int = 512
image_dim: int = 0
use_text_crossattn: bool = True
use_memory: bool = True
sigma_theta: float = 0.8
camera_embed_in_channels: int = 1536
action_config: dict = field(
default_factory=lambda: {
"blocks": list(range(15)),
"enable_mouse": True,
"enable_keyboard": True,
"heads_num": 16,
"hidden_size": 128,
"img_hidden_size": 3072,
"keyboard_dim_in": 6,
"keyboard_hidden_dim": 1024,
"mouse_dim_in": 2,
"mouse_hidden_dim": 1024,
"mouse_qk_dim_list": [8, 28, 28],
"patch_size": [1, 2, 2],
"qk_norm": True,
"qkv_bias": False,
"rope_dim_list": [8, 28, 28],
"rope_theta": 256,
"vae_time_compression_ratio": 4,
"windows_size": 3,
})
@dataclass
class MatrixGame3WanVideoConfig(WanVideoConfig):
arch_config: MatrixGame3WanVideoArchConfig = field(default_factory=MatrixGame3WanVideoArchConfig)
prefix: str = "Wan"
_compile_conditions: list = field(default_factory=lambda: [_is_transformer_block])
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class SD3Transformer2DArchConfig(DiTArchConfig):
# Diffusers SD3Transformer2DModel config fields.
sample_size: int = 128
patch_size: int = 2
num_layers: int = 24
attention_head_dim: int = 64
joint_attention_dim: int = 4096
caption_projection_dim: int = 1536
pooled_projection_dim: int = 2048
pos_embed_max_size: int = 384
dual_attention_layers: list[int] = field(default_factory=lambda: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12])
qk_norm: str = "rms_norm"
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 24
@dataclass
class SD3DiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=SD3Transformer2DArchConfig)
prefix: str = "sd3"
@@ -0,0 +1,76 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for the Stable Audio Open 1.0 DiT.
Note: the SA pipeline bypasses the standard `ComposedPipelineBase`
component loader because the published HF repo ships a single monolithic
`model.safetensors` (no Diffusers-style `model_index.json` or
per-subfolder layout). The arch fields and `param_names_mapping` here
document the architecture and key remap so the same conventions used by
the rest of the DiT family apply (FSDP shard conditions, supported
attention backends, future loader integrations) — they are not currently
consumed by `fastvideo/models/loader/fsdp_load.py` for SA.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
from v2._vendor.platforms import AttentionBackendEnum
def _is_transformer_layer(n: str, m) -> bool:
# Matches `transformer.layers.{i}` in the SA DiT module tree.
parts = n.split(".")
return (len(parts) >= 3 and parts[-3] == "transformer" and parts[-2] == "layers" and parts[-1].isdigit())
@dataclass
class StableAudioArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer])
# SA's checkpoint is `stable_audio_tools` raw format (not Diffusers),
# so the only remaps are: strip the `model.model.` host-pipeline
# prefix, and rename `nn.LayerNorm`'s `gamma`/`beta` to torch's
# canonical `weight`/`bias`. Linear / cross-attention naming already
# matches FastVideo's conventions, so no further remap is needed.
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.model\.(.*?)\.gamma$": r"\1.weight",
r"^model\.model\.(.*?)\.beta$": r"\1.bias",
r"^model\.model\.(.*)$": r"\1",
})
# SA only supports backends compatible with single-GPU LocalAttention.
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
# Architecture constants (from the published `model_config.json` for
# `stabilityai/stable-audio-open-1.0`).
io_channels: int = 64
embed_dim: int = 1536
depth: int = 24
num_attention_heads: int = 24
cond_token_dim: int = 768
global_cond_dim: int = 1536
project_cond_tokens: bool = False
project_global_cond: bool = True
# Set to "ln" to wrap attention Q/K in LayerNorm (used by
# `stable-audio-open-small`; absent in the 1.0 base).
qk_norm: str | None = None
def __post_init__(self) -> None:
super().__post_init__()
self.hidden_size = self.embed_dim
self.in_channels = self.io_channels
self.out_channels = self.io_channels
self.num_channels_latents = self.io_channels
self.attention_head_dim = self.embed_dim // self.num_attention_heads
@dataclass
class StableAudioConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=StableAudioArchConfig)
prefix: str = "StableAudio"
@@ -0,0 +1,97 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class WanVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$": r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$": r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class WanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanVideoArchConfig)
prefix: str = "Wan"
@@ -0,0 +1,21 @@
from v2._vendor.configs.models.encoders.base import (BaseEncoderOutput, EncoderConfig, ImageEncoderConfig,
TextEncoderConfig)
from v2._vendor.configs.models.encoders.clip import (CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
from v2._vendor.configs.models.encoders.llama import LlamaConfig
from v2._vendor.configs.models.encoders.t5 import T5Config, T5LargeConfig
from v2._vendor.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from v2._vendor.configs.models.encoders.siglip import SiglipVisionConfig
from v2._vendor.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
from v2._vendor.configs.models.encoders.gemma import LTX2GemmaConfig
from v2._vendor.configs.models.encoders.mistral3 import Mistral3TextConfig
from v2._vendor.configs.models.encoders.qwen3 import Qwen3TextConfig
from v2._vendor.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
StableAudioConditionerConfig)
from v2._vendor.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
]
@@ -0,0 +1,85 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Any
import torch
from v2._vendor.configs.models.base import ArchConfig, ModelConfig
from v2._vendor.layers.quantization import QuantizationConfig
from v2._vendor.platforms import AttentionBackendEnum
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=lambda: [])
_supported_attention_backends: tuple[AttentionBackendEnum,
...] = (AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
@dataclass
class TextEncoderArchConfig(EncoderArchConfig):
vocab_size: int = 0
hidden_size: int = 0
num_hidden_layers: int = 0
num_attention_heads: int = 0
pad_token_id: int = 0
eos_token_id: int = 0
text_len: int = 0
hidden_state_skip_layer: int = 0
decoder_start_token_id: int = 0
output_past: bool = True
scalable_attention: bool = True
tie_word_embeddings: bool = False
stacked_params_mapping: list[tuple[str, str, str]] = field(
default_factory=list) # mapping from huggingface weight names to custom names
tokenizer_kwargs: dict[str, Any] = field(default_factory=dict)
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
# When True, the tokenizer loader prefers AutoProcessor over AutoTokenizer
# for encoders whose tokenizer dir ships a processor_config.json (e.g. Flux2
# full's Mistral3 multimodal processor). Default False keeps every existing
# encoder on the historical AutoTokenizer path.
require_processor: bool = False
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
"truncation": True,
"max_length": self.text_len,
"return_tensors": "pt",
}
@dataclass
class ImageEncoderArchConfig(EncoderArchConfig):
pass
@dataclass
class BaseEncoderOutput:
last_hidden_state: torch.FloatTensor | None = None
pooler_output: torch.FloatTensor | None = None
hidden_states: tuple[torch.FloatTensor, ...] | None = None
attentions: tuple[torch.FloatTensor, ...] | None = None
attention_mask: torch.Tensor | None = None
@dataclass
class EncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=EncoderArchConfig)
prefix: str = ""
quant_config: QuantizationConfig | None = None
lora_config: Any | None = None
@dataclass
class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
is_chat_model: bool = False
treat_empty_as_dot: bool = False
@dataclass
class ImageEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=ImageEncoderArchConfig)
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (ImageEncoderArchConfig, ImageEncoderConfig, TextEncoderArchConfig,
TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embeddings")
@dataclass
class CLIPTextArchConfig(TextEncoderArchConfig):
vocab_size: int = 49408
hidden_size: int = 512
intermediate_size: int = 2048
projection_dim: int = 512
num_hidden_layers: int = 12
num_attention_heads: int = 8
max_position_embeddings: int = 77
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
dropout: float = 0.0
attention_dropout: float = 0.0
initializer_range: float = 0.02
initializer_factor: float = 1.0
pad_token_id: int = 1
bos_token_id: int = 49406
eos_token_id: int = 49407
text_len: int = 77
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_layer, _is_embeddings])
@dataclass
class CLIPVisionArchConfig(ImageEncoderArchConfig):
hidden_size: int = 768
intermediate_size: int = 3072
projection_dim: int = 512
num_hidden_layers: int = 12
num_attention_heads: int = 12
num_channels: int = 3
image_size: int = 224
patch_size: int = 32
hidden_act: str = "quick_gelu"
layer_norm_eps: float = 1e-5
dropout: float = 0.0
attention_dropout: float = 0.0
initializer_range: float = 0.02
initializer_factor: float = 1.0
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
@dataclass
class CLIPTextConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=CLIPTextArchConfig)
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = True
prefix: str = "clip"
@dataclass
class CLIPVisionConfig(ImageEncoderConfig):
arch_config: ImageEncoderArchConfig = field(default_factory=CLIPVisionArchConfig)
num_hidden_layers_override: int | None = 31
require_post_norm: bool | None = None
enable_scale: bool = False
is_causal: bool = False
prefix: str = "clip"
@dataclass
class WAN2_1ControlCLIPVisionConfig(CLIPVisionConfig):
num_hidden_layers_override: int | None = 31
require_post_norm: bool | None = False
enable_scale: bool = False
is_causal: bool = False
@@ -0,0 +1,75 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
def _is_feature_extractor_linear(n: str, m) -> bool:
# LTX-2.3 (caption_proj_before_connector) introduces separate
# video/audio feature extractor linears; keep the LTX-2.0 name too.
return (n.endswith("feature_extractor_linear") or n.endswith("video_feature_extractor_linear")
or n.endswith("audio_feature_extractor_linear"))
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embeddings_connector") or n.endswith("audio_embeddings_connector")
def _is_gemma_model(n: str, m) -> bool:
return "_gemma_model" in n
@dataclass
class LTX2GemmaArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
hidden_size: int = 3840
num_hidden_layers: int = 48
num_attention_heads: int = 30
text_len: int = 1024
pad_token_id: int = 0
eos_token_id: int = 2
gemma_model_path: str = ""
gemma_dtype: str = "bfloat16"
padding_side: str = "left"
feature_extractor_in_features: int = 3840 * 49
feature_extractor_out_features: int = 3840
# LTX-2.3 text-stack connector fields (default OFF == LTX-2.0 behavior).
video_feature_extractor_out_features: int | None = None
audio_feature_extractor_out_features: int | None = None
caption_proj_before_connector: bool = False
caption_projection_first_linear: bool = True
caption_proj_input_norm: bool = True
caption_projection_second_linear: bool = True
connector_num_attention_heads: int = 30
connector_attention_head_dim: int = 128
connector_num_layers: int = 2
# Separate audio connector geometry (None falls back to the video values).
audio_connector_num_attention_heads: int | None = None
audio_connector_attention_head_dim: int | None = None
audio_connector_num_layers: int | None = None
connector_positional_embedding_theta: float = 10000.0
connector_positional_embedding_max_pos: list[int] = field(default_factory=lambda: [4096])
connector_rope_type: str = "split"
connector_double_precision_rope: bool = False
connector_apply_gated_attention: bool = False
connector_num_learnable_registers: int | None = 128
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_feature_extractor_linear, _is_embeddings, _is_gemma_model])
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LTX2GemmaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=LTX2GemmaArchConfig)
prefix: str = "ltx2_gemma"
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (TextEncoderArchConfig, TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass
class LlamaArchConfig(TextEncoderArchConfig):
vocab_size: int = 32000
hidden_size: int = 4096
intermediate_size: int = 11008
num_hidden_layers: int = 32
num_attention_heads: int = 32
num_key_value_heads: int | None = None
hidden_act: str = "silu"
max_position_embeddings: int = 2048
initializer_range: float = 0.02
rms_norm_eps: float = 1e-6
use_cache: bool = True
pad_token_id: int = 0
bos_token_id: int = 1
eos_token_id: int = 2
pretraining_tp: int = 1
tie_word_embeddings: bool = False
rope_theta: float = 10000.0
rope_scaling: float | None = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: int | None = None
hidden_state_skip_layer: int = 2
text_len: int = 256
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0), # type: ignore
(".gate_up_proj", ".up_proj", 1), # type: ignore
])
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm])
@dataclass
class LlamaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=LlamaArchConfig)
prefix: str = "llama"
@@ -0,0 +1,38 @@
# SPDX-License-Identifier: Apache-2.0
"""Mistral3 text encoder configuration for full Flux2."""
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class Mistral3TextArchConfig(TextEncoderArchConfig):
"""Architecture config for the Mistral3 text encoder used by full Flux2."""
architectures: list[str] = field(default_factory=lambda: ["Mistral3ForConditionalGeneration"])
hidden_size: int = 5120
num_hidden_layers: int = 40
text_len: int = 512
output_hidden_states: bool = True
# Mistral3 (full Flux2) ships a multimodal processor; load via AutoProcessor.
require_processor: bool = True
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
"padding": "max_length",
"truncation": True,
"max_length": self.text_len,
"return_tensors": "pt",
}
@dataclass
class Mistral3TextConfig(TextEncoderConfig):
"""Top-level config for the Mistral3 full Flux2 text encoder."""
arch_config: TextEncoderArchConfig = field(default_factory=Mistral3TextArchConfig)
prefix: str = "mistral3"
is_chat_model: bool = True
@@ -0,0 +1,90 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (TextEncoderArchConfig, TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass
class Qwen2_5_VLArchConfig(TextEncoderArchConfig):
vocab_size: int = 152064
hidden_size: int = 8192
intermediate_size: int = 29568
num_hidden_layers: int = 80
num_attention_heads: int = 64
num_key_value_heads: int = 8
hidden_act: str = "silu"
max_position_embeddings: int = 32768
initializer_range: float = 0.02
rms_norm_eps: float = 1e-05
use_cache: bool = True
tie_word_embeddings: bool = False
rope_theta: float = 1000000.0
use_sliding_window: bool = False
sliding_window: int | None = 4096
max_window_layers: int = 80
layer_types: list = field(default_factory=list)
attention_dropout: float = 0.0
rope_scaling: dict | None = None
bos_token_id: int | None = None
eos_token_id: int | None = None
pad_token_id: int | None = None
vision_token_id: int = 151654
model_type: str = "qwen2_5_vl_text"
dtype: str = "bfloat16"
stacked_params_mapping: list[tuple[str, str, str
| int]] = field(default_factory=lambda: [
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0),
(".gate_up_proj", ".up_proj", 1),
])
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm])
def __post_init__(self):
super().__post_init__()
self.sliding_window = self.sliding_window if self.use_sliding_window else None
# for backward compatibility
if self.num_key_value_heads is None:
self.num_key_value_heads = self.num_attention_heads
if self.layer_types is None:
self.layer_types = [
"sliding_attention"
if self.sliding_window is not None and i >= self.max_window_layers else "full_attention"
for i in range(self.num_hidden_layers)
]
if self.rope_scaling is not None and "type" in self.rope_scaling:
if self.rope_scaling["type"] == "mrope":
self.rope_scaling["type"] = "default"
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
self.tokenizer_kwargs = {
"add_generation_prompt": True,
"tokenize": True,
"return_dict": True,
"max_length": 1000 + 108,
"truncation": True,
"return_tensors": "pt",
}
@dataclass
class Qwen2_5_VLConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=Qwen2_5_VLArchConfig)
prefix: str = "qwen2_5_vl"
is_chat_model: bool = True
treat_empty_as_dot: bool = True
@@ -0,0 +1,82 @@
# SPDX-License-Identifier: Apache-2.0
# Ported from SGLang: python/sglang/multimodal_gen/configs/models/encoders/qwen3.py
"""Qwen3 text encoder configuration for FastVideo diffusion models (e.g. Flux2 Klein)."""
from dataclasses import dataclass, field
from typing import Any
from v2._vendor.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
def _is_transformer_layer(n: str, m: Any) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m: Any) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m: Any) -> bool:
return n.endswith("norm")
@dataclass
class Qwen3TextArchConfig(TextEncoderArchConfig):
"""Architecture config for Qwen3 text encoder.
Qwen3 is similar to LLaMA but with QK-Norm (RMSNorm on Q and K before attention).
Used by Flux2 Klein.
"""
vocab_size: int = 151936
hidden_size: int = 2560
intermediate_size: int = 9728
num_hidden_layers: int = 36
num_attention_heads: int = 32
num_key_value_heads: int = 8
hidden_act: str = "silu"
max_position_embeddings: int = 40960
initializer_range: float = 0.02
rms_norm_eps: float = 1e-6
use_cache: bool = True
pad_token_id: int = 151643
bos_token_id: int = 151643
eos_token_id: int = 151645
tie_word_embeddings: bool = True
rope_theta: float = 1000000.0
rope_scaling: dict | None = None
attention_bias: bool = False
attention_dropout: float = 0.0
mlp_bias: bool = False
head_dim: int = 128
text_len: int = 512
output_hidden_states: bool = True # Klein needs hidden states from layers 9, 18, 27
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=lambda: [
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0),
(".gate_up_proj", ".up_proj", 1),
])
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_norm])
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
"padding": "max_length",
"truncation": True,
"max_length": self.text_len,
"return_tensors": "pt",
}
@dataclass
class Qwen3TextConfig(TextEncoderConfig):
"""Top-level config for Qwen3 text encoder."""
arch_config: TextEncoderArchConfig = field(default_factory=Qwen3TextArchConfig)
prefix: str = "qwen3"
is_chat_model: bool = True
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for Reason1 (Qwen2.5-VL) text encoder."""
from dataclasses import dataclass, field
from typing import Any
from v2._vendor.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig
@dataclass
class Reason1ArchConfig(TextEncoderArchConfig):
"""Architecture settings (defaults match Qwen2.5-VL-7B-Instruct)."""
architectures: list[str] = field(default_factory=lambda: ["Qwen2_5_VLForConditionalGeneration"])
model_type: str = "qwen2_5_vl"
vocab_size: int = 152064
hidden_size: int = 3584
num_hidden_layers: int = 28
num_attention_heads: int = 28
num_key_value_heads: int = 4
intermediate_size: int = 18944
text_len: int = 512
hidden_state_skip_layer: int = 0
bos_token_id: int = 151643
pad_token_id: int = 151643
eos_token_id: int = 151645
image_token_id: int = 151655
video_token_id: int = 151656
vision_token_id: int = 151654
vision_start_token_id: int = 151652
vision_end_token_id: int = 151653
vision_config: dict[str, Any] | None = None
rope_theta: float = 1000000.0
rope_scaling: dict[str, Any] | None = field(default_factory=lambda: {
"type": "mrope",
"mrope_section": [16, 24, 24]
})
max_position_embeddings: int = 128000
max_window_layers: int = 28
embedding_concat_strategy: str = "mean_pooling"
n_layers_per_group: int = 5
num_embedding_padding_tokens: int = 512
attention_dropout: float = 0.0
hidden_act: str = "silu"
initializer_range: float = 0.02
rms_norm_eps: float = 1e-6
use_sliding_window: bool = False
sliding_window: int = 32768
tie_word_embeddings: bool = False
use_cache: bool = False
output_hidden_states: bool = True
torch_dtype: str = "bfloat16"
_attn_implementation: str = "flash_attention_2"
@dataclass
class Reason1Config(TextEncoderConfig):
"""Reason1 text encoder config."""
arch_config: Reason1ArchConfig = field(default_factory=Reason1ArchConfig)
tokenizer_type: str = "Qwen/Qwen2.5-VL-7B-Instruct"
@@ -0,0 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
"""SigLIP vision encoder configuration for FastVideo."""
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (ImageEncoderArchConfig, ImageEncoderConfig)
@dataclass
class SiglipVisionArchConfig(ImageEncoderArchConfig):
"""Architecture configuration for SigLIP vision encoder.
Fields match the config.json from HuggingFace SigLIP checkpoints.
"""
# From config.json
architectures: list[str] = field(default_factory=lambda: ["SiglipVisionModel"])
attention_dropout: float = 0.0
dtype: str | None = None
hidden_act: str = "gelu_pytorch_tanh"
hidden_size: int = 1152
image_size: int = 384
intermediate_size: int = 4304
layer_norm_eps: float = 1e-6
model_type: str = "siglip_vision_model"
num_attention_heads: int = 16
num_channels: int = 3
num_hidden_layers: int = 27
patch_size: int = 14
# FastVideo specific - QKV fusion mapping
stacked_params_mapping: list = field(default_factory=lambda: [
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
@dataclass
class SiglipVisionConfig(ImageEncoderConfig):
"""Configuration for SigLIP vision encoder."""
arch_config: ImageEncoderArchConfig = field(default_factory=SiglipVisionArchConfig)
# FastVideo specific
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = False
prefix: str = "siglip"
@@ -0,0 +1,80 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for the Stable Audio Open 1.0 multi-conditioner.
The conditioner bundles three sub-conditioners — a T5 text encoder
(prompt) and two NumberConditioners (`seconds_start` / `seconds_total`)
— into the (cross_attn_cond, cross_attn_mask, global_embed) triple the
DiT consumes. The architecture is fully specified by the official
`stable_audio_tools` `MultiConditioner` config; the constants here
mirror that.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from v2._vendor.configs.models.base import ArchConfig
from v2._vendor.configs.models.encoders.base import (EncoderArchConfig, EncoderConfig)
def _default_configs() -> list[dict]:
"""Default = `stable-audio-open-1.0`'s three sub-conditioners."""
return [
{
"id": "prompt",
"type": "t5",
"config": {
"t5_model_name": "t5-base",
"max_length": 128
}
},
{
"id": "seconds_start",
"type": "number",
"config": {
"min_val": 0,
"max_val": 512
}
},
{
"id": "seconds_total",
"type": "number",
"config": {
"min_val": 0,
"max_val": 512
}
},
]
@dataclass
class StableAudioConditionerArchConfig(EncoderArchConfig):
architectures: list[str] = field(default_factory=lambda: ["StableAudioMultiConditioner"])
# Shared embedding width across all sub-conditioners (T5 last-hidden
# dim and NumberEmbedder feature dim both = `cond_dim`).
cond_dim: int = 768
# Sub-conditioner identifiers. Order in `cross_attention_cond_ids`
# is the concat order for the cross-attn token sequence; order in
# `global_cond_ids` is the concat order for the global FiLM-style
# embedding.
cross_attention_cond_ids: tuple[str, ...] = ("prompt", "seconds_start", "seconds_total")
global_cond_ids: tuple[str, ...] = ("seconds_start", "seconds_total")
# Per-sub-conditioner spec list (mirrors upstream
# `model_config.json.model.conditioning.configs`). Each entry is
# `{"id": ..., "type": "t5"|"number", "config": {...}}`. The default
# matches `stable-audio-open-1.0`; SA-small overrides via the
# `conditioner/config.json` shipped in the converted repo.
configs: list = field(default_factory=_default_configs)
# Match official `stable_audio_tools/models/conditioners.py:334`:
# T5 is loaded directly in fp16.
t5_dtype: str = "float16"
@dataclass
class StableAudioConditionerConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=StableAudioConditionerArchConfig)
prefix: str = "stable_audio_conditioner"
+106
View File
@@ -0,0 +1,106 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (TextEncoderArchConfig, TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "block" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("shared")
def _is_final_layernorm(n: str, m) -> bool:
return n.endswith("final_layer_norm")
@dataclass
class T5ArchConfig(TextEncoderArchConfig):
vocab_size: int = 32128
d_model: int = 512
d_kv: int = 64
d_ff: int = 2048
num_layers: int = 6
num_decoder_layers: int | None = None
num_heads: int = 8
relative_attention_num_buckets: int = 32
relative_attention_max_distance: int = 128
dropout_rate: float = 0.1
layer_norm_epsilon: float = 1e-6
initializer_factor: float = 1.0
feed_forward_proj: str = "relu"
dense_act_fn: str = ""
is_gated_act: bool = False
is_encoder_decoder: bool = True
use_cache: bool = True
pad_token_id: int = 0
eos_token_id: int = 1
classifier_dropout: float = 0.0
text_len: int = 512
dtype: str | None = None
gradient_checkpointing: bool = False
# Extra fields present in upstream HF T5Config but unused by FastVideo's
# encoder. Declared here so `update_model_arch` doesn't reject them when
# loading repos like `stabilityai/stable-audio-open-1.0` that ship the
# full HF config.
n_positions: int = 512
decoder_start_token_id: int = 0
output_past: bool = True
task_specific_params: dict | None = None
stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"),
(".qkv_proj", ".k", "k"),
(".qkv_proj", ".v", "v"),
])
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings, _is_final_layernorm])
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
def __post_init__(self):
super().__post_init__()
act_info = self.feed_forward_proj.split("-")
self.dense_act_fn: str = act_info[-1]
self.is_gated_act: bool = act_info[0] == "gated"
if self.feed_forward_proj == "gated-gelu":
self.dense_act_fn = "gelu_new"
self.tokenizer_kwargs = {
"truncation": True,
"max_length": self.text_len,
"add_special_tokens": True,
"return_attention_mask": True,
"return_tensors": "pt",
}
self.hidden_size = self.d_model
@dataclass
class T5LargeArchConfig(T5ArchConfig):
"""T5 Large architecture config with parameters for your specific model."""
d_model: int = 1024
d_kv: int = 128
d_ff: int = 65536
num_layers: int = 24
num_decoder_layers: int | None = 24
num_heads: int = 128
decoder_start_token_id: int = 0
n_positions: int = 512
task_specific_params: dict | None = None
@dataclass
class T5Config(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=T5ArchConfig)
prefix: str = "t5"
@dataclass
class T5LargeConfig(TextEncoderConfig):
"""T5 Large configuration for your specific model."""
arch_config: TextEncoderArchConfig = field(default_factory=T5LargeArchConfig)
prefix: str = "t5"
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Config for the T5-Gemma encoder used by daVinci-MagiHuman.
The reference pipeline uses `transformers.models.t5gemma.T5GemmaEncoderModel`
on `google/t5gemma-9b-9b-ul2`. That is a gated Google repository, so the
encoder weights are not bundled inside GAIR/daVinci-MagiHuman; they are
loaded from the T5-Gemma HF repo directly.
Encoder shape (verified from google/t5gemma-9b-9b-ul2/config.json):
layers=42, hidden=3584, heads=16, kv_heads=8, head_dim=256,
intermediate=14336, rope_theta=10000.0, max_pos=8192,
layer_types alternate sliding_attention / full_attention.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from v2._vendor.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class T5GemmaEncoderArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(default_factory=lambda: ["T5GemmaEncoderModel"])
hidden_size: int = 3584
num_hidden_layers: int = 42
num_attention_heads: int = 16
num_key_value_heads: int = 8
head_dim: int = 256
intermediate_size: int = 14336
max_position_embeddings: int = 8192
rope_theta: float = 10000.0
vocab_size: int = 256000
# MagiHuman fixes prompt embed length at 640 via pad_or_trim.
text_len: int = 640
pad_token_id: int = 0
eos_token_id: int = 1
# Path to the upstream gated repo. When set, the FastVideo loader will
# pull the encoder directly via `T5GemmaEncoderModel.from_pretrained`.
t5gemma_model_path: str = "google/t5gemma-9b-9b-ul2"
t5gemma_dtype: str = "bfloat16"
# The HF T5-Gemma encoder is lazy-loaded on first forward (see
# `fastvideo/models/encoders/t5gemma.py`), so no FastVideo-owned
# submodules exist at FSDP-apply time. An empty list makes
# `shard_model()` log a warning and return cleanly instead of raising
# "No layer modules were sharded" — sharding of the lazy HF model is
# the activation pipeline's responsibility.
_fsdp_shard_conditions: list = field(default_factory=list)
def __post_init__(self) -> None:
super().__post_init__()
# WHY: upstream `t5_gemma_model.py:25` tokenizes without
# padding/max_length, then `prompt_process.py` pad_or_trim-s the
# encoded states. Keep only tensor return here so
# MagiHumanLatentPreparationStage can pad/trim post-encode while
# preserving the real original prompt length.
self.tokenizer_kwargs.pop("truncation", None)
self.tokenizer_kwargs.pop("max_length", None)
self.tokenizer_kwargs.pop("padding", None)
@dataclass
class T5GemmaEncoderConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=T5GemmaEncoderArchConfig)
prefix: str = "t5gemma"
@@ -0,0 +1,4 @@
from v2._vendor.configs.models.upsamplers.hunyuan15 import SRTo720pUpsamplerConfig, SRTo1080pUpsamplerConfig
from v2._vendor.configs.models.upsamplers.base import UpsamplerConfig
__all__ = ["SRTo720pUpsamplerConfig", "SRTo1080pUpsamplerConfig", "UpsamplerConfig"]
@@ -0,0 +1,7 @@
from dataclasses import dataclass
from v2._vendor.configs.models.base import ModelConfig
@dataclass
class UpsamplerConfig(ModelConfig):
pass
@@ -0,0 +1,20 @@
from dataclasses import dataclass
from v2._vendor.configs.models.upsamplers.base import UpsamplerConfig
@dataclass
class SRTo720pUpsamplerConfig(UpsamplerConfig):
in_channels: int = 0
out_channels: int = 0
hidden_channels: int = 64
num_blocks: int = 6
global_residual: bool = False
@dataclass
class SRTo1080pUpsamplerConfig(UpsamplerConfig):
z_channels: int = 0
out_channels: int = 0
block_out_channels: tuple[int, ...] = (0, 0)
num_res_blocks: int = 2
is_residual: bool = False
@@ -0,0 +1,24 @@
from v2._vendor.configs.models.vaes.cosmosvae import CosmosVAEConfig
from v2._vendor.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from v2._vendor.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
from v2._vendor.configs.models.vaes.gen3cvae import Gen3CVAEConfig
from v2._vendor.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from v2._vendor.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from v2._vendor.configs.models.vaes.ltx2vae import LTX2VAEConfig
from v2._vendor.configs.models.vaes.oobleck import OobleckVAEArchConfig, OobleckVAEConfig
from v2._vendor.configs.models.vaes.flux2vae import Flux2VAEConfig
from v2._vendor.configs.models.vaes.wanvae import WanVAEConfig
__all__ = [
"GameCraftVAEConfig",
"HunyuanVAEConfig",
"WanVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Gen3CVAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
"OobleckVAEArchConfig",
"OobleckVAEConfig",
"Flux2VAEConfig",
]
@@ -0,0 +1,38 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class AutoencoderKLArchConfig(VAEArchConfig):
_name_or_path: str = ""
act_fn: str = "silu"
block_out_channels: tuple[int, ...] | list[int] = field(default_factory=list)
down_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
up_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
force_upcast: bool = True
in_channels: int = 3
latent_channels: int = 4
latents_mean: tuple[float, ...] | list[float] | None = None
latents_std: tuple[float, ...] | list[float] | None = None
layers_per_block: int = 1
mid_block_add_attention: bool = True
norm_num_groups: int = 32
out_channels: int = 3
sample_size: int = 32
scaling_factor: float | torch.Tensor = 0.18215
shift_factor: float | None = None
use_post_quant_conv: bool = True
use_quant_conv: bool = True
temporal_compression_ratio: int = 1
spatial_compression_ratio: int = 8
@dataclass
class AutoencoderKLVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=AutoencoderKLArchConfig)
+145
View File
@@ -0,0 +1,145 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import dataclasses
from dataclasses import dataclass, field
from typing import Any
import torch
from v2._vendor.configs.models.base import ArchConfig, ModelConfig
from v2._vendor.utils import StoreBoolean
@dataclass
class VAEArchConfig(ArchConfig):
scaling_factor: float | torch.Tensor = 0
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
@dataclass
class VAEConfig(ModelConfig):
arch_config: VAEArchConfig = field(default_factory=VAEArchConfig)
# FastVideoVAE-specific parameters
load_encoder: bool = True
load_decoder: bool = True
tile_sample_min_height: int = 256
tile_sample_min_width: int = 256
tile_sample_min_num_frames: int = 16
tile_sample_stride_height: int = 192
tile_sample_stride_width: int = 192
tile_sample_stride_num_frames: int = 12
blend_num_frames: int = 0
use_tiling: bool = True
use_temporal_tiling: bool = True
use_parallel_tiling: bool = True
# When True, latent preparation skips the schedule shift on frames
# whose temporal index is below the model's first-frame conditioning
# threshold. LTX-2 reads this in the latent prep stage.
use_temporal_scaling_frames: bool = True
def __post_init__(self):
self.blend_num_frames = self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames
@staticmethod
def add_cli_args(parser: Any, prefix: str = "vae-config") -> Any:
"""Add CLI arguments for VAEConfig fields"""
parser.add_argument(
f"--{prefix}.load-encoder",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.load_encoder",
default=VAEConfig.load_encoder,
help="Whether to load the VAE encoder",
)
parser.add_argument(
f"--{prefix}.load-decoder",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.load_decoder",
default=VAEConfig.load_decoder,
help="Whether to load the VAE decoder",
)
parser.add_argument(
f"--{prefix}.tile-sample-min-height",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_min_height",
default=VAEConfig.tile_sample_min_height,
help="Minimum height for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-min-width",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_min_width",
default=VAEConfig.tile_sample_min_width,
help="Minimum width for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-min-num-frames",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_min_num_frames",
default=VAEConfig.tile_sample_min_num_frames,
help="Minimum number of frames for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-stride-height",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_stride_height",
default=VAEConfig.tile_sample_stride_height,
help="Stride height for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-stride-width",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_stride_width",
default=VAEConfig.tile_sample_stride_width,
help="Stride width for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.tile-sample-stride-num-frames",
type=int,
dest=f"{prefix.replace('-', '_')}.tile_sample_stride_num_frames",
default=VAEConfig.tile_sample_stride_num_frames,
help="Stride number of frames for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.blend-num-frames",
type=int,
dest=f"{prefix.replace('-', '_')}.blend_num_frames",
default=VAEConfig.blend_num_frames,
help="Number of frames to blend for VAE tile sampling",
)
parser.add_argument(
f"--{prefix}.use-tiling",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_tiling",
default=VAEConfig.use_tiling,
help="Whether to use tiling for VAE",
)
parser.add_argument(
f"--{prefix}.use-temporal-tiling",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_temporal_tiling",
default=VAEConfig.use_temporal_tiling,
help="Whether to use temporal tiling for VAE",
)
parser.add_argument(
f"--{prefix}.use-parallel-tiling",
action=StoreBoolean,
dest=f"{prefix.replace('-', '_')}.use_parallel_tiling",
default=VAEConfig.use_parallel_tiling,
help="Whether to use parallel tiling for VAE",
)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
kwargs = {}
for attr in dataclasses.fields(cls):
value = getattr(args, attr.name, None)
if value is not None:
kwargs[attr.name] = value
return cls(**kwargs)
@@ -0,0 +1,216 @@
"""Cosmos 2.5 (Wan2.1-style) VAE config and checkpoint-key mapping."""
from __future__ import annotations
import re
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class Cosmos25VAEArchConfig(VAEArchConfig):
_name_or_path: str = ""
base_dim: int = 96
decoder_base_dim: int | None = None
z_dim: int = 16
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
is_residual: bool = False
in_channels: int = 3
out_channels: int = 3
patch_size: int | None = None
scale_factor_temporal: int = 4
scale_factor_spatial: int = 8
clip_output: bool = True
latents_mean: tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
# Simple 1:1 renames. More complex decoder remapping is handled by
# `map_official_key()`.
param_names_mapping: dict[str, str] = field(
default_factory=lambda: {
r"^conv1\.(.*)$": r"quant_conv.\1",
r"^conv2\.(.*)$": r"post_quant_conv.\1",
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
})
@staticmethod
def map_official_key(key: str) -> str | None:
"""Map a single official checkpoint key into FastVideo key space."""
def map_residual_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^residual\.0\.gamma$", sub):
return f"{prefix}.norm1.gamma"
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv1.{m.group(1)}"
if re.match(r"^residual\.3\.gamma$", sub):
return f"{prefix}.norm2.gamma"
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv2.{m.group(1)}"
m = re.match(r"^shortcut\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv_shortcut.{m.group(1)}"
return None
def map_attn_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^norm\.gamma$", sub):
return f"{prefix}.norm.gamma"
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
if m:
return f"{prefix}.to_qkv.{m.group(1)}"
m = re.match(r"^proj\.(weight|bias)$", sub)
if m:
return f"{prefix}.proj.{m.group(1)}"
return None
def map_resample_subkey(prefix: str, sub: str) -> str | None:
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
if m:
return f"{prefix}.resample.1.{m.group(1)}"
m = re.match(r"^time_conv\.(weight|bias)$", sub)
if m:
return f"{prefix}.time_conv.{m.group(1)}"
return None
m = re.match(r"^conv1\.(weight|bias)$", key)
if m:
return f"quant_conv.{m.group(1)}"
m = re.match(r"^conv2\.(weight|bias)$", key)
if m:
return f"post_quant_conv.{m.group(1)}"
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_in.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
if m:
return f"{m.group(1)}.norm_out.gamma"
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_out.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if m:
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1", m.group(2))
m = re.match(r"^encoder\.downsamples\.(\d+)\.(.*)$", key)
if m:
idx = int(m.group(1))
sub = m.group(2)
if sub.startswith("residual.") or sub.startswith("shortcut."):
return map_residual_subkey(f"encoder.down_blocks.{idx}", sub)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return map_resample_subkey(f"encoder.down_blocks.{idx}", sub)
return None
m = re.match(r"^decoder\.upsamples\.(\d+)\.(.*)$", key)
if m:
uidx = int(m.group(1))
sub = m.group(2)
if uidx in (0, 1, 2):
block_i, res_i = 0, uidx
elif uidx == 3:
block_i, res_i = 0, None
elif uidx in (4, 5, 6):
block_i, res_i = 1, uidx - 4
elif uidx == 7:
block_i, res_i = 1, None
elif uidx in (8, 9, 10):
block_i, res_i = 2, uidx - 8
elif uidx == 11:
block_i, res_i = 2, None
elif uidx in (12, 13, 14):
block_i, res_i = 3, uidx - 12
else:
return None
if res_i is None:
return map_resample_subkey(
f"decoder.up_blocks.{block_i}.upsamplers.0",
sub,
)
return map_residual_subkey(
f"decoder.up_blocks.{block_i}.resnets.{res_i}",
sub,
)
return None
temporal_compression_ratio: int = 4
spatial_compression_ratio: int = 8
def __post_init__(self):
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(self.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(1, self.z_dim, 1, 1, 1)
self.temporal_compression_ratio = self.scale_factor_temporal
self.spatial_compression_ratio = self.scale_factor_spatial
@dataclass
class Cosmos25VAEConfig(VAEConfig):
"""Cosmos2.5 VAE config."""
arch_config: Cosmos25VAEArchConfig = field(default_factory=Cosmos25VAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames) * 2
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
import torch
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class CosmosVAEArchConfig(VAEArchConfig):
_name_or_path: str = ""
base_dim: int = 96
z_dim: int = 16
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
decoder_base_dim: int | None = None
is_residual: bool = False
in_channels: int = 3
out_channels: int = 3
patch_size: int | None = None
scale_factor_temporal: int = 4
scale_factor_spatial: int = 8
clip_output: bool = True
latents_mean: tuple[float, ...] = (
-0.7571,
-0.7089,
-0.9113,
0.1075,
-0.1745,
0.9653,
-0.1517,
1.5508,
0.4134,
-0.0715,
0.5517,
-0.3632,
-0.1922,
-0.9497,
0.2503,
-0.2921,
)
latents_std: tuple[float, ...] = (
2.8184,
1.4541,
2.3275,
2.6558,
1.2196,
1.7708,
2.6052,
2.0743,
3.2687,
2.1526,
2.8652,
1.5579,
1.6382,
1.1253,
2.8251,
1.9160,
)
temporal_compression_ratio = 4
spatial_compression_ratio = 8
def __post_init__(self):
self.scaling_factor: torch.Tensor = 1.0 / torch.tensor(self.latents_std).view(1, self.z_dim, 1, 1, 1)
self.shift_factor: torch.Tensor = torch.tensor(self.latents_mean).view(1, self.z_dim, 1, 1, 1)
self.temporal_compression_ratio = self.scale_factor_temporal
self.spatial_compression_ratio = self.scale_factor_spatial
@dataclass
class CosmosVAEConfig(VAEConfig):
arch_config: CosmosVAEArchConfig = field(default_factory=CosmosVAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
def __post_init__(self):
self.blend_num_frames = (self.tile_sample_min_num_frames - self.tile_sample_stride_num_frames) * 2
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
# Copied and adapted from: https://github.com/sglang-ai/sglang
from dataclasses import dataclass, field
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class Flux2VAEArchConfig(VAEArchConfig):
"""Architecture configuration for Flux2 VAE model."""
# Flux2 VAE-specific architecture parameters
in_channels: int = 3
out_channels: int = 3
down_block_types: tuple[str, ...] = (
"DownEncoderBlock2D",
"DownEncoderBlock2D",
"DownEncoderBlock2D",
"AttnDownEncoderBlock2D",
)
up_block_types: tuple[str, ...] = (
"AttnUpDecoderBlock2D",
"UpDecoderBlock2D",
"UpDecoderBlock2D",
"UpDecoderBlock2D",
)
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
latent_channels: int = 16
norm_num_groups: int = 32
sample_size: int = 512
force_upcast: bool = False
use_quant_conv: bool = True
use_post_quant_conv: bool = True
mid_block_add_attention: bool = True
batch_norm_eps: float = 1e-5
batch_norm_momentum: float = 0.1
patch_size: tuple[int, int] = (1, 1)
# Latent scaling for decode: avoid division-by-zero; match Flux/Flux2 convention (e.g. 0.13025)
scaling_factor: float = 0.13025
# Spatial compression (for images, this is typically 8)
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 1 # Images don't have temporal dimension
@dataclass
class Flux2VAEConfig(VAEConfig):
"""Configuration for Flux2 VAE model."""
arch_config: Flux2VAEArchConfig = field(default_factory=Flux2VAEArchConfig)
# Flux2 is an image model, so disable temporal tiling
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
@@ -0,0 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
"""
GameCraft VAE config - matches official config.json from Hunyuan-GameCraft-1.0.
"""
from dataclasses import dataclass, field
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class GameCraftVAEArchConfig(VAEArchConfig):
"""Architecture config matching official AutoencoderKLCausal3D config.json."""
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: tuple[str, ...] = (
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
)
up_block_types: tuple[str, ...] = (
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
)
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
time_compression_ratio: int = 4 # alias for DecoderCausal3D
mid_block_add_attention: bool = True
mid_block_causal_attn: bool = True
sample_size: int = 256 # from config.json
sample_tsize: int = 64 # from config.json
def __post_init__(self):
self.spatial_compression_ratio = 2**(len(self.block_out_channels) - 1)
@dataclass
class GameCraftVAEConfig(VAEConfig):
"""Full config for GameCraft VAE."""
arch_config: VAEArchConfig = field(default_factory=GameCraftVAEArchConfig)
@@ -0,0 +1,14 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from v2._vendor.configs.models.vaes.cosmosvae import CosmosVAEConfig
@dataclass
class Gen3CVAEConfig(CosmosVAEConfig):
"""
GEN3C VAE config placeholder.
GEN3C uses tokenizer-backed VAE loading logic at runtime, but we keep a
model-specific config class so pipeline/model configs stay model-scoped.
"""
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class Hunyuan15VAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 32
block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024)
layers_per_block: int = 2
spatial_compression_ratio: int = 16
temporal_compression_ratio: int = 4
downsample_match_channel: bool = True
upsample_match_channel: bool = True
scaling_factor: float = 1.03682
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) - 1)
@dataclass
class Hunyuan15VAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=Hunyuan15VAEArchConfig)
@@ -0,0 +1,39 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from v2._vendor.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class HunyuanVAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
down_block_types: tuple[str, ...] = (
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
"HunyuanVideoDownBlock3D",
)
up_block_types: tuple[str, ...] = (
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
"HunyuanVideoUpBlock3D",
)
block_out_channels: tuple[int, ...] = (128, 256, 512, 512)
layers_per_block: int = 2
act_fn: str = "silu"
norm_num_groups: int = 32
scaling_factor: float = 0.476986
spatial_compression_ratio: int = 8
temporal_compression_ratio: int = 4
mid_block_add_attention: bool = True
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) - 1)
@dataclass
class HunyuanVAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=HunyuanVAEArchConfig)

Some files were not shown because too many files have changed in this diff Show More