Compare commits

..
Author SHA1 Message Date
SolitaryThinker c59f568d56 [docs] cosmos3: port_feedback.md (pitfalls/lessons) + native-reuse audit 2026-07-13 03:23:49 -07:00
SolitaryThinker fdea9e9898 [docs] cosmos3: full-omni parity summary (all components bit-exact) 2026-07-13 03:23:49 -07:00
SolitaryThinker 9b19098434 [feat] cosmos3 PR4: deepstack reasoner forward (image-conditioned reasoning)
Add the one new native piece for image-conditioned reasoning: a deepstack-capable
reasoner backbone forward. The other parts are already proven/reusable — the
Qwen3-VL vision_encoder (transformers, bit-exact vs framework) and the multimodal
prefill (masked-scatter of image features + get_rope_index, standard transformers
Qwen3-VL boilerplate, reusable like the tokenizer/encoder).

- `Cosmos3VFMTransformer.reason_forward(inputs_embeds, position_ids,
  deepstack_embeds, visual_pos_mask)`: runs the und (causal) layers over a single
  text(+image) sequence, injecting per-layer deepstack visual embeds at the visual
  positions for the first `len(deepstack_embeds)` layers (Qwen3-VL deepstack),
  then final `norm`. Mirrors the framework `reasoner_forward`.
- Parity (`test_cosmos3_reasoning_parity::test_deepstack_reasoner_forward_matches_framework`):
  native `reason_forward` vs framework `reasoner_forward` given identical
  inputs_embeds + positions + per-layer deepstack + visual mask — hidden states
  max abs diff = 0.0, mean = 0.0.

All five omni modalities are now bit-exact vs the framework (video/audio/action
gen, text reasoning, vision encoding, deepstack reasoner). Full suite: 150 passed.
2026-07-13 03:23:49 -07:00
SolitaryThinker 24e2e73454 [feat] cosmos3 PR4: vision_encoder parity (transformers Qwen3-VL == framework)
Image-conditioned reasoning needs the Qwen3-VL vision_encoder. The checkpoint's
vision_encoder is a standard transformers `Qwen3VLVisionModel` (the framework
ships its own copy of the same model). Like the Qwen2 tokenizer, FastVideo reuses
the transformers model (no diffusers) rather than re-porting a 27-layer ViT.

Parity (`test_cosmos3_vision_encoder_parity.py`): transformers `Qwen3VLVisionModel`
vs the framework `Qwen3VLVisionModel`, tiny config, weights copied across,
identical (hidden_states, grid_thw) forward — embeds max abs diff = 0.0, mean = 0.0
across grids. Also verified on the REAL 1.15 GB checkpoint: both strict-load and
produce identical [N, out_hidden] embeds (max=mean=0.0).

All five omni modalities now have bit-exact framework parity: video gen
(T2V/I2V/T2I), audio gen (t2vs), action gen, text reasoning, vision encoding.
Remaining: wire image-conditioned reasoning (deepstack multimodal prefill joining
the two already-proven components). Full suite: 149 passed.
2026-07-13 03:23:49 -07:00
SolitaryThinker e5bd10678f [feat] cosmos3 PR4: text reasoning (und pathway + lm_head), bit-exact
Add the reasoning (VLM text-generation) capability. The framework reasoner uses
only the und (causal) pathway weights (no _moe_gen) + embed_tokens / norm /
lm_head — all already present in the native DiT — so a text-only forward + lm_head
reproduces it with no new model code.

- `cosmos3_generate_reasoner_text(transformer, input_ids, max_new_tokens)`:
  greedy text reasoning via the und causal backbone + lm_head (text-only prefill,
  re-prefill per step; KV-cache fast path is a later optimization). Mirrors the
  framework `generate_reasoner_text` (text-only path).
- Parity (`test_cosmos3_reasoning_parity.py`, framework = oracle):
  * prefill logits (native text-only forward + lm_head vs framework
    `reasoner_forward` + lm_head, per-position): max abs diff = 0.0, mean = 0.0;
  * greedy generation vs framework `generate_reasoner_text(do_sample=False)`:
    token-for-token identical across seeds.

Text reasoning is fully supported + bit-exact. Image-conditioned reasoning
additionally needs the Qwen3-VL `vision_encoder` (a transformers ViT, tracked
separately). Full suite: 146 passed, 0 skipped.
2026-07-13 03:23:49 -07:00
SolitaryThinker 5c749a6911 [misc]: simplify cosmos3 action (fold action split into _split_flat_latent)
Collapse the byte-identical _split_action_latent into the generic
_split_flat_latent; all three modality specs now share one splitter.
Behavior-preserving; 139 tests green.
2026-07-13 03:23:47 -07:00
SolitaryThinker 5623c1ba1b [feat] cosmos3 PR3: action (domain-aware) DiT pathway + packing + CFG glue
Add the action (multi-embodiment world-model) modality. The Nano checkpoint
ships action weights (action_gen=true, 32 embodiment domains); this activates
them, bit-exact vs the framework.

- DiT forward: action encode/decode mirroring sound but domain-aware — pack each
  [T, action_dim] latent with a per-token embodiment domain id, `action_proj_in`
  (= framework `action2llm`, a `DomainAwareLinear`) + `action_modality_embed` +
  timestep scatter; decode via `action_proj_out` (= `llm2action`) on noisy hidden
  states -> `_unpack_action`. Mirrors `_encode_action`/`_decode_action`.
- Sound/action share the vision "full" split; sequence/flat order is
  [vision | action | sound] (matches the framework per-sample concat).
- Packing: `Cosmos3ActionItem` + action fields. Action tokens are `(T,)`-shaped,
  domain-tagged, with 3D-MRoPE at the vision temporal offset and
  `start_frame_offset=1` (framework `_pack_action_tokens`).
- CFG velocity / `Cosmos3DenoiseEngine` take optional `action_specs`
  (`Cosmos3ActionSpec`); the joint denoise machinery now handles action.

Parity (`test_cosmos3_action_parity.py`, framework model+pack = oracle, 3
embodiment domains):
- action packing fields + `position_ids` exact;
- `preds_vision` + `preds_action` (forward): max abs diff = 0.0, mean = 0.0;
- combined [vision|action] CFG velocity: max abs diff = 0.0, mean = 0.0.

Action is processed bit-exactly at the model + packing + denoise-glue level. A
full action2world real-weights run is gated on real robot-action input data (not
available here); the bit-exact parity is the correctness proof. Full suite: 139
passed, 0 skipped.
2026-07-13 03:23:46 -07:00
SolitaryThinker c72fc2de2c [misc]: simplify cosmos3 audio (fold sound split into _split_flat_latent)
Collapse the byte-identical _split_sound_latent into the generic
_split_flat_latent (duck-typed on numel/shape). Behavior-preserving;
130 tests green.
2026-07-13 03:23:45 -07:00
SolitaryThinker 41e2d3ee8b [docs] cosmos3 PR2 (audio/t2vs) complete; bit-exact across components 2026-07-13 03:23:43 -07:00
SolitaryThinker 1ec268cf80 [feat] cosmos3 PR2: t2vs pipeline (joint denoise + AVAE decode + AV mux)
Wire text-to-video+sound end-to-end, completing the audio PR.

- CFG velocity (`cosmos3_get_cfg_velocity`) + `Cosmos3DenoiseEngine` take an
  optional `sound_specs`: the denoise latent becomes a combined [vision | sound]
  flat vector; each step packs vision+sound, forwards the DiT, zeros the
  prediction on conditioning frames for both, and returns the concatenated
  velocity. The UniPC scheduler steps the joint vector. `Cosmos3SoundSpec` +
  `_split_sound_latent` added.
- `Cosmos3DenoisingStage` t2vs path (gated on `COSMOS3_T2VS`): size the sound
  latent from the video duration (framework `create_placeholder_audio` +
  `get_latent_num_samples`), append sound noise, joint-denoise, split, AVAE-decode
  the sound latent, and set `batch.extra["audio"]` / `["audio_sample_rate"]` so
  the generator muxes a stereo 48 kHz AAC track into the mp4. Sound AVAE is
  lazy-loaded from `<model_path>/sound_tokenizer` (`_get_sound_vae`).
- Parity (`test_cosmos3_sound_parity::test_t2vs_cfg_velocity_matches_framework`):
  the combined [vision|sound] sequential-CFG velocity vs a framework-DiT oracle —
  max abs diff = 0.0, mean abs diff = 0.0 across t2vs/i2vs cases.
- Example `basic_cosmos3_t2vs_new_api.py`.

Verified on B200 (real weights, 1280x704, 29f, 35 steps): coherent ocean/sunset
video WITH a real stereo 48 kHz audio track (mean -10.2 dB, peak -0.4 dB; not
silence). Full cosmos3 suite: 130 passed, 0 skipped.
2026-07-13 03:23:42 -07:00
SolitaryThinker 637d0bf943 [feat] cosmos3 PR2: DiT sound pathway + sound packing (bit-exact parity)
Activate the dormant sound MoT heads and pack the sound modality, completing
the t2vs DiT path (audio component 2).

DiT forward (`Cosmos3VFMTransformer`):
- Encode sound: pack each [C,T] latent -> [T,C], `audio_proj_in` (= framework
  `sound2llm`) + `audio_modality_embed`, scatter timestep embeds onto noisy
  frames, scatter into the (full) gen split. Mirrors framework `_encode_sound`.
- Decode sound: `audio_proj_out` (= `llm2sound`) on noisy sound hidden states ->
  `_unpack_sound` to per-sample [C,T] (framework `_decode_sound`).
- Reuses the general `_scatter_timestep_embeds` (works for (T,1,1) shapes).

Sound packing (`sequence_packing.py`): `Cosmos3SoundItem` + sound fields on
`Cosmos3SampleInputs`/`Cosmos3PackedSequence`/`to_dit_kwargs`. Sound shares the
vision "full" split (preserving the causal+full 2-split invariant), with
`(T,1,1)` token shapes, a `(T,1)` condition mask, mse-loss/timestep bookkeeping
per noisy frame, and 3D-MRoPE temporal positions starting at the vision temporal
offset (parallel to vision; `start_frame_offset=0`, tcf=1; does not advance the
offset). Mirrors framework `_pack_sound_tokens`.

Parity (`test_cosmos3_sound_parity.py`, framework model+pack = oracle):
- sound packing fields + `position_ids` exact vs framework `pack_input_sequence`
  (has_sound) across t2vs/i2vs cases.
- `preds_vision` AND `preds_sound`: max abs diff = 0.0, mean abs diff = 0.0
  (both native-pack and framework-pack inputs).
Tiny framework builder gains a `sound_gen` flag. Full suite: 127 passed.
2026-07-13 03:23:40 -07:00
SolitaryThinker 417e653604 [docs] cosmos3 PR2: AVAE sound decoder done; remaining audio components 2026-07-13 03:23:39 -07:00
SolitaryThinker 9de2ab7505 [feat] cosmos3 PR2: native AVAE sound decoder + bit-exact parity
First audio component: the Cosmos3 sound tokenizer's decode path (t2vs needs
only DECODE — generate sound latents, decode to waveform).

Finding: the shipped `sound_tokenizer` checkpoint is decoder-only (`decoder.*`,
the SpectrogramConvNeXt encoder is not exported) in diffusers AutoencoderOobleck
naming, but with SnakeBeta activations (alpha+beta, logscale) and
weight_g/weight_v weight-norm — i.e. exactly FastVideo's existing native
`OobleckVAE` decoder. So no from-scratch port: reuse it.

- `OobleckDecoderBlock`: add `output_padding=stride % 2` to the transpose conv.
  A provable no-op for even strides (Stable Audio: [2,4,4,8,8]); required for the
  odd stride in Cosmos3 (`[2,4,5,6,8]`), where the framework's
  `output_padding=stride%2` otherwise makes the decode 1 sample longer per
  odd-stride block. Without it, parity diverged (60 vs 59 samples).
- `Cosmos3SoundVAE` (`fastvideo/models/audio/cosmos3_avae.py`): decoder-only
  wrapper — `decode` (Oobleck decoder + clamp to [-1,1], matching AVAEModel),
  `get_latent_num_samples` (N // hop_size=1920), `from_pretrained` (reads the
  checkpoint config, strict-loads `decoder.*`). 48 kHz stereo, 473M params.
- Parity: `test_cosmos3_avae_parity.py` maps the framework OobleckDecoder
  (`avae_utils.models`, Sequential naming, the oracle) into the FastVideo decoder
  and asserts bit-exact decode across stride patterns incl. odd (5) — max abs
  diff 0.0. Verified the real 1.9GB checkpoint strict-loads and decodes
  [1,64,25] -> [1,2,48000] (1s @ 48kHz stereo).

Full cosmos3 suite: 121 passed, 0 skipped.
2026-07-13 03:23:38 -07:00
SolitaryThinker 7b4f757fc2 [docs] cosmos3 PR2: audio port plan (AVAE + DiT sound pathway + t2vs) 2026-07-13 03:23:37 -07:00
SolitaryThinker 35fe3db250 [docs] cosmos3: record T2I verification + resolution-based flow_shift 2026-07-13 03:23:36 -07:00
SolitaryThinker 2114d19ced [feat] cosmos3: resolution-based flow_shift + T2I check & example
Add and verify the text-to-image path, and fix the UniPC flow_shift selection
it exposed.

Bug: the stage chose flow_shift by task (`3.0 if is_t2i else 10.0`). The
framework selects it purely by the named resolution bucket
(`OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256->3.0, 480->5.0,
720/768->10.0); T2V/I2V/T2I share a shift at a given resolution. The task-based
rule only coincidentally matched (T2V@720, T2I@256). The canonical Cosmos3 T2I
is 960x960 (the "720" bucket -> 10.0), so `is_t2i->3.0` was wrong for real T2I.

- Replace with `_flow_shift_for_resolution(h, w)` mapping the longest side to the
  framework bucket (<=320:3.0, 640-832:5.0, 960-1360:10.0), applied to all tasks.
- Parity: `test_cosmos3_flow_shift_parity.py` checks the mapping against the
  framework's `{VIDEO,IMAGE}_RES_SIZE_INFO` x `_RESOLUTION_SHIFT_DEFAULTS` for
  every (resolution, aspect) in the 8B rows (20 cases).
- Update the call-graph mode-dispatch tests for resolution-based shift.

Also harden `_image_to_video_tensor`'s tensor branch to respect the FastVideo
[-1,1] convention for already-preprocessed conditioning tensors (the PIL path,
used by the real pipeline, stays framework-exact and bit-parity-tested).

Example: `examples/inference/basic/basic_cosmos3_t2i_new_api.py` (num_frames=1,
960x960). Verified on B200 (real weights, 35 steps): a coherent red-panda image
matching the prompt, flow_shift=10.0. Full cosmos3 suite: 118 passed, 0 skipped.
2026-07-13 03:23:34 -07:00
SolitaryThinker 895cef22a3 [docs] cosmos3: record I2V real-weights verification (feat/cosmos3-i2v) 2026-07-13 03:23:33 -07:00
SolitaryThinker 284c50fc96 [feat] cosmos3: I2V image conditioning (static-repeat) + parity + example
Wire up and verify the Cosmos3 image-to-video path. The denoise stage already
handled the I2V conditioning math (encode -> condition mask -> keep frame 0
clean -> velocity zeroed on condition frames, all matching the framework), but
the conditioning pixel clip was built wrong.

Bug: `_image_to_video_tensor` zero-filled every frame after frame 0. The
framework (`cosmos_framework.inference.vision.build_conditioned_video_batch`)
instead fills frame 0 with the image and **repeats the last conditioning frame**
for the rest of the clip (a static video) before VAE-encoding. Because the Wan
VAE is temporal (4x), latent frame 0 (the kept-clean condition frame) depends on
several pixel frames, so zero-filling produces a wrong conditioning latent.

Fix `_image_to_video_tensor` to be framework-faithful:
- aspect-preserving resize + center crop + uint8 quantization, then `/127.5 - 1`
  (mirrors `load_conditioning_image` / `_resize_and_center_crop`);
- static-repeat fill across all frames (mirrors `build_conditioned_video_batch`).

Parity: add `test_cosmos3_i2v_conditioning_parity.py` comparing the FastVideo
conditioning clip against the framework's `load_conditioning_image` +
repeat-fill across aspect/size/frame-count cases — bit-exact (max abs diff 0.0).

Example: `examples/inference/basic/basic_cosmos3_i2v_new_api.py` (new
VideoGenerator API, image via `InputConfig(image_path=...)`,
`assets/images/cyclist.jpg` default).

Verified on B200 (1280x704, 29f, 35 steps, real weights): frame 0 reproduces
the conditioning image and the clip generates coherent forward motion following
the prompt. Full cosmos3 suite: 98 passed, 0 skipped.
2026-07-13 03:23:31 -07:00
SolitaryThinker 94d9dda5cb [misc]: simplify cosmos3 port (reuse embed, dedup prompt, hoist)
Behavior-preserving cleanup (95 tests green):
- reuse fastvideo.layers.visual_embedding.timestep_embedding in the DiT
- single-home COSMOS3_VIDEO_NEGATIVE_PROMPT in the lightweight presets
  module; drop the unused duplicate from cosmos3_pipeline
- drop the write-only Cosmos3VisionSpec.clean_latent field + its kwarg
- hoist device = next(transformer.parameters()).device out of the per-pass
  _run closure in cosmos3_get_cfg_velocity
2026-07-13 03:23:30 -07:00
SolitaryThinker f661bbf104 [docs] cosmos3: record real-weights E2E acceptance + scheduler fix (I003) 2026-07-13 03:23:28 -07:00
SolitaryThinker a2a204c25d [bugfix] cosmos3: fix black-video UniPC config; wire E2E inference
The first real-weights Cosmos3-Nano T2V run produced an all-black video.
Root cause: the checkpoint's diffusers-style scheduler_config.json sets
use_karras_sigmas=true, and FastVideo's vendored UniPC checks karras
*before* use_flow_sigmas, so it built diffusion (beta) sigmas instead of
flow-matching sigmas -> scheduler.step diverged to NaN latents -> black
frames. The DiT/CFG velocity itself was clean.

The framework samples with FlowUniPCMultistepScheduler (pure flow:
shift + num_train_timesteps only). FastVideo's vendored UniPC is the same
algorithm, so the fix is config, not a port:

- Coerce the loaded scheduler to the flow setup in
  Cosmos3OmniDiffusersPipeline.initialize_pipeline (kill karras/exp/beta
  sigmas, force use_flow_sigmas + flow_prediction + final_sigmas_type=zero)
  and rebuild the runtime scheduler from it.
- Use FastVideo's native UniPC (not diffusers) in the pipeline and tests,
  honoring the "no diffusers at runtime" constraint.

Parity:
- Add test_cosmos3_scheduler_parity.py: FastVideo vendored UniPC (flow
  config) vs the framework FlowUniPCMultistepScheduler. timesteps bit-exact,
  sigmas within fp32 epsilon (~1e-8), full multi-step trajectory < ~1e-6
  across shift in {10 (video), 3 (t2i)} and steps in {4,10,35}.
- Fix test_cosmos3_denoise_cfg_parity to use the framework scheduler as the
  oracle (it previously compared diffusers-vs-diffusers, so the scheduler
  was never actually checked against the framework).

Integration fixes to run end-to-end on real weights:
- registry: map checkpoint DiT class Cosmos3OmniTransformer ->
  Cosmos3VFMTransformer.
- component_loader: route the "text_tokenizer" module to TokenizerLoader;
  filter scheduler config to the class's __init__ params (robust to
  diffusers schema drift like shift_terminal/sigma_min/sigma_max).
- DiT: materialize_non_persistent_buffers (recompute rotary_emb.inv_freq
  after meta-device FSDP load) + cast vision latents / timestep embeds to
  the compute dtype (no-ops in the fp32 parity tests).
- sequence_packing: move packed tensors to the model device in to_dit_kwargs.
- config: drop the unused identity text-preprocess (Cosmos3 has no text
  encoder; the DiT tokenizes in the denoise stage).

Adds examples/inference/basic/basic_cosmos3_new_api.py (new VideoGenerator
API). Verified: 1280x704, 29 frames, 35 steps on B200 -> coherent video
matching the prompt (no NaNs, real pixel variance). Full cosmos3 parity
suite green.
2026-07-13 03:23:27 -07:00
SolitaryThinker 4533447e58 [misc] cosmos3 PORT_STATUS: PR1 video core complete (framework parity) 2026-07-13 03:23:24 -07:00
SolitaryThinker c94d625844 [feat] cosmos3: native video pipeline + framework denoise parity 2026-07-13 03:23:23 -07:00
SolitaryThinker 7427b9d2d3 [feat] cosmos3: native MoT sequence-packing (video) + framework parity 2026-07-13 03:22:40 -07:00
SolitaryThinker 9dab81f180 [misc] cosmos3 PORT_STATUS: strict-load done; only pipeline remains 2026-07-13 03:22:39 -07:00
SolitaryThinker 840b43b4c3 [feat] cosmos3: strict-load verified (identity; needs_conversion=no) 2026-07-13 03:22:37 -07:00
SolitaryThinker ace34f866a [misc] cosmos3 PORT_STATUS: VAE component framework parity verified 2026-07-13 03:22:37 -07:00
SolitaryThinker b5f90a6b77 [feat] cosmos3: VAE config (Wan2.2 reuse) + framework parity test 2026-07-13 03:22:36 -07:00
SolitaryThinker ed8f118acd [misc] cosmos3 PORT_STATUS: DiT framework parity verified (both rope modes) 2026-07-13 03:22:36 -07:00
SolitaryThinker b379ea96b4 [test] cosmos3: DiT unified_3d_mrope framework parity (real checkpoint) 2026-07-13 03:22:34 -07:00
SolitaryThinker 6771f4ec2e [feat] cosmos3: native DiT + framework parity (Cosmos3VFMTransformer) 2026-07-13 03:22:33 -07:00
SolitaryThinker b48eafe515 [misc] cosmos3 PORT_STATUS: PR1 progress (arch config, framework parity ref) 2026-07-13 03:22:32 -07:00
SolitaryThinker 1d199050af [test] cosmos3: official framework DiT parity reference (CPU/SDPA) 2026-07-13 03:22:30 -07:00
SolitaryThinker 8b65f7ce6c [misc] cosmos3 PORT_STATUS: framework-reference pivot + Phase 1 findings 2026-07-13 03:22:30 -07:00
SolitaryThinker 2d2657821e [feat] cosmos3: arch config 1:1 with Cosmos3-Nano checkpoint 2026-07-13 03:22:29 -07:00
SolitaryThinker ded49ab1c8 [misc]: cosmos3 PORT_STATUS: post-rebase state + fv-cosmos3 env 2026-07-13 03:22:29 -07:00
SolitaryThinker 10f5372f42 [misc]: cosmos3 resume: PORT_STATUS + README for official diffusers ref 2026-07-13 03:22:27 -07:00
SolitaryThinker ed4125e7d8 [feat]: Cosmos3 checkpoint conversion script (cosmos3_convert.py)
scripts/checkpoint_conversion/cosmos3_convert.py wraps Cosmos3OmniDiffusersPipeline._remap_ckpt_key. convert_state_dict(src) returns (new_state, skipped, unmapped) for shard-by-shard remap; convert_checkpoint_dir reads *.safetensors shards and writes a consolidated FastVideo-format state dict (sharded index.json TODO when real weights land). smoke_test() exercises all 14 representative remap branches synthetically (no real weights needed) and returns nonzero exit on failure; runs in CI as a remap-drift guard. Usage: python scripts/checkpoint_conversion/cosmos3_convert.py --smoke-test.
2026-07-13 03:22:26 -07:00
SolitaryThinker 26e78f8fb1 [feat]: Cosmos3Config + registry wire-up for nvidia/Cosmos3-Nano
Cosmos3Config(PipelineConfig): reuses Cosmos25VAEConfig (Cosmos3 uses DistributedAutoencoderKLWan); text_encoder_configs=() since Cosmos3LanguageModel lives inside the DiT; flow_shift=1.0 default (T2V/I2V engine init; T2I uses 3.0 per-request via _set_flow_shift). Registry entry registers BEFORE the generic cosmos detector to win path precedence (same pattern as GEN3C). Resolves nvidia/Cosmos3-Nano -> Cosmos3Config without regressing Cosmos25/GEN3C/Cosmos.
2026-07-13 03:22:26 -07:00
SolitaryThinker 12f94fd6e7 [feat]: Cosmos3 pipeline (Cosmos3OmniDiffusersPipeline)
Pipeline class with: _remap_ckpt_key static method (verbatim port from pipeline_cosmos3.py:319-409, 14-rule checkpoint remap UND/GEN split + lm_head skip); _set_flow_shift method with lazy UniPCMultistepScheduler construction; diffuse() with sequential 3-mode CFG denoising loop + I2V velocity_mask + image_latent re-injection (ported from pipeline_cosmos3.py:883-1033); forward() with T2I/T2V/I2V mode dispatch + flow_shift selection (ported from pipeline_cosmos3.py:1037-1206); 4 helper-method stubs raise NotImplementedError. Dual inheritance (nn.Module, ComposedPipelineBase) matches upstream pattern.
2026-07-13 03:21:44 -07:00
SolitaryThinker facd035b97 [feat]: Cosmos3 DiT skeleton (Cosmos3VFMTransformer + Cosmos3LanguageModel)
Ports math (compute_mrope_position_ids_text/_vision, patchify, unpatchify) verbatim from vllm-omni transformer_cosmos3.py:113-177,1009-1036. Module tree (language_model.layers.*.self_attn.{q,k,v,o}_proj + gen_layers.*.cross_attention.* + vae2llm + llm2vae + time_embedder + norm_moe_gen) exposes the parameter names the checkpoint converter targets. All layer forward() raise NotImplementedError until weights publish; module instantiation + state-dict key tests pass.
2026-07-13 03:21:44 -07:00
SolitaryThinker c3c6c0cb2d [test]: Cosmos3 local-test parity scaffold (Tier A, 15 tests)
Adds 8 CPU-only parity test files + conftest stubs (StubScheduler, StubCosmos3VAE, StubCosmos3Transformer) under tests/local_tests/cosmos3/ mirroring the vllm-omni reference at tests/diffusion/models/cosmos3/conftest.py. All 15 tests skip until the FastVideo modules land (subsequent commits). Reference: vllm-omni PR #3454 @ 8536f5b1421f.
2026-07-13 03:21:44 -07:00
Mac Lee 1ea2517e22 [ci]: extend LoRA training CI timeout (#1589) 2026-07-13 02:47:28 -07:00
MookandSolitaryThinker 0c63528c59 [perf] Cache RoPE position-embedding tables across denoising steps (#1442)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-13 02:34:51 -07:00
b063f8ca41 [feat] Fix FLUX.1-dev port: native RoPE, parity tests, SSIM reference (#1321)
Co-authored-by: Ishan Vaish <ivaish@ucsd.edu>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-13 01:49:42 -07:00
William Lin e7fff0173a [bugfix]: benchmark_weight_loading_comparison.py — iterate safe_open via .keys() (#1378) 2026-07-12 22:50:45 -07:00
Shreejith SGandH1yori233 d82abc271e [feat] Add GLM-Image inference support (#1030)
Co-authored-by: H1yori233 <k1kong@ucsd.edu>
2026-07-12 22:42:42 -07:00
Guian FangandSolitaryThinker 970409962f [feat] Add AnyFlow any-step video distillation (pretrain + on-policy) (#1371)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-12 02:32:08 +00:00
Raghav K 055586703d [perf]: register a real backward for FA2 default + masked/varlen custom ops (training-under-compile) (#1388) 2026-07-12 00:45:27 +00:00
Mac Lee 5d89f86675 [ci] Stop forcing FA4 in model-load lanes (#1561) 2026-07-11 14:10:36 -07:00
Satyam Srivastava 19a51a1fe6 [ci] Trigger performance benchmarks for performance code changes (#1583) 2026-07-10 20:21:33 -07:00
William Lin d3232cea5a [ci]: gate the full-suite trigger on pre-commit and docs build (#1572) 2026-07-11 02:56:49 +00:00
Raghav KandSolitaryThinker 0c90c8c24d [bugfix] nvfp4: cast fp32 inputs to bf16 instead of asserting (#1488)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-10 21:39:17 +00:00
Mingjia HuoandClaude Fable 5 4c08ffce49 [feat] World model training using third person games (#1443)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
2026-07-10 05:09:07 +00:00
Atharv Ramesh af4a77553c [ci]: add SSIM reference bootstrap flow (#1522) (#1547) 2026-07-10 01:49:38 +00:00
alexzmsandSolitaryThinker c096fda1eb [docs] Add LTX-2.3 distilled inference run configs (t2v/i2v × 5+2/8+3 × resolutions) (#1568)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-09 18:53:12 +00:00
William Lin 8f47e85be0 [bugfix]: retry remote image downloads in load_image (#1570) 2026-07-09 07:33:24 -07:00
William Lin 90d3bd19eb [infra] Deliver per-job Buildkite env to Modal CI at runtime, not as image layers (#1569) 2026-07-09 06:33:37 -07:00
Mac LeeandSolitaryThinker afb4f7d3c5 [ci]: emit v2 performance result schema (#1551)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-09 06:11:15 +00:00
02e1143f22 [feat] Add Kandinsky-5 T2V/I2V pipeline support (#1471)
Co-authored-by: Aryan Kumar <aryan5v@users.noreply.github.com>
Co-authored-by: leffff <levnovitskiy@gmail.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 14:43:30 -07:00
KaredandSolitaryThinker e2f4d1a7b5 [feat]: add SwanLab tracker (#1461)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 08:35:38 +00:00
Kaiqin KongandSolitaryThinker f037351146 [feat] Add Clean-history Teacher Forcing and Causal Consistency Distillation (#1505)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 08:13:11 +00:00
Mac LeeandSolitaryThinker 1ee11e08dc [ci]: add performance fingerprint cohorts (#1546)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 07:24:00 +00:00
595f0ea60e [feat] Add DreamX-World 5B Cam and AR pipelines (#1538)
Co-authored-by: Suckl <Suckl@users.noreply.github.com>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-07 06:16:18 +00:00
William Lin d921832cd2 [misc]: reserve CPU/memory for timing-sensitive Modal CI lanes (#1566) 2026-07-06 22:13:53 -07:00
William Lin 629697629a [bugfix]: free CUDA memory between train-framework model tests (LongCat OOM on L40S) (#1565) 2026-07-06 21:27:30 -07:00
William Lin a25313beec [ci]: wire fastvideo/tests/ops/ into the unit-test lane (#1559) 2026-07-06 12:07:26 -07:00
William Lin dbde64385b [bugfix]: bump FA4 pin to the CuTe DSL 4.6 compatible rev (#1564) 2026-07-06 12:06:59 -07:00
William Lin 9d909f5f04 [test]: remove dead and duplicate tests (-489 lines) (#1556) 2026-07-05 15:53:40 -07:00
William Lin 76b0550c15 [ci]: run pre-commit on fork PRs without manual approval (#1555) 2026-07-05 14:18:16 -07:00
William Lin 384c1e9493 [misc]: update reseed-performance-baseline skill for the hf_store move (#1545 follow-up) (#1553) 2026-07-05 14:16:55 -07:00
William Lin b1dbcc93f6 [misc]: reformat fastvideo/performance to the repo yapf config (#1554) 2026-07-05 14:16:20 -07:00
William Lin b93833772e [ci]: guard against test directories no CI lane collects (#1552) 2026-07-05 14:07:38 -07:00
Mac Lee 30b523edd6 [ci] Normalize performance stage component metrics (#1475) (#1550) 2026-07-05 14:05:26 -07:00
Mac LeeandSolitaryThinker 6aab7f3832 [ci] cover Hunyuan 1.5 chat-list text preprocessing (#1518)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-05 12:05:25 -07:00
Mac Lee 9cd53fe5f8 [ci] Add metric-specific performance thresholds (#1545) 2026-07-05 12:04:33 -07:00
Mac Lee 6a32cf3a5e [ci]: expose LoRA extraction slash command (#1542) 2026-07-05 06:45:31 -07:00
William Lin 98be9b3da2 [bugfix]: address the three remaining #1447 review findings (#1549) 2026-07-05 06:43:27 -07:00
Mac Lee c53e85b767 [ci] Add v2 performance benchmark config identity fields (#1544) 2026-07-05 06:18:15 -07:00
William Lin 40a8bd2d3b [bugfix]: skip ThunderKittens kernels on aarch64 and document the kernel build matrix (#1548) 2026-07-05 06:13:11 -07:00
Mac LeeandSolitaryThinker 31aa115611 [bugfix]: preserve FSDP hooks for RMSNorm qk norms (#1513)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-05 04:15:49 +00:00
zainnhandSolitaryThinker 98ac10a528 [infra] Auto-rebuild CUDA images when docker/Dockerfile changes on main (#1526)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-04 16:25:46 -07:00
William Lin a5a6d171e5 [attn] Make FA4 explicit opt-in via FASTVIDEO_FA4 and delete the runtime fallback machinery (#1540) 2026-07-04 14:51:00 -07:00
704 changed files with 38507 additions and 72752 deletions
-207
View File
@@ -1,207 +0,0 @@
# v2 ← M\*: Architecture Gap-Analysis & Improvement Roadmap
**Status:** exploration, flagged for review. **Date:** 2026-06-19.
**Source paper:** *M\*: A Modular, Extensible, Serving System for Multimodal Models* (arXiv 2606.12688,
Stanford/UW/CMU; Jha, Sagan, Kamahori, …, Kasikci, S. Wang). It is a universal serving runtime for composite
multimodal models built on the **Walk Graph** abstraction (a model is a dataflow graph `G`; a request is a
*Walk* — a labeled subgraph — and the runtime executes walks). It beats vLLM-Omni (~20% lower T2I latency on
**BAGEL**, up to 2.64× on I2I), SGLang-Omni (2.7× TTS throughput on **Qwen3-Omni**), and native V-JEPA2
rollout (12.5×). It explicitly names **FastVideo's own** sparse/sliding-tile attention, xDiT/PipeFusion/USP,
Inferix, and FlashDrive as techniques integratable into the graph runtime.
**Method:** a 28-agent workflow — 6 parallel v2-subsystem maps → 10 M\*-dimension analyses, each
*adversarially verified against the actual v2 code* → synthesis + a completeness critic. The critic's
corrections and three P0 claims were then **spot-verified by hand** (file:line below). This doc folds those
corrections in; it is the corrected, authoritative synthesis.
---
## 1. Executive summary
v2 already implements the **harder half** of M\*'s thesis and in several axes **exceeds** it:
- v2's `Program` *is* M\*'s graph `G` (typed `ComponentNode`/`ModelLoopNode` + edges).
- v2's `shared_weight_components` *is* M\*'s cross-Walk node sharing — BAGEL/Cosmos3/LTX2 each bind two
`ModelLoopNode`s to **one resident transformer** (`instance.component()` returns the same live object). This
is the exact MoT serving property the omni cards in this repo already express.
- v2 adds three things M\* (serving-only) has **no equivalent for**: a required+validated per-loop **cost
model**, a non-negotiable **interleave bit-parity gate**, and an **integrated training plane** (RL→distill
flywheel driving the *same* serving Loop).
- The `extend/` plugin seam (interceptors/observers/registry with capability negotiation) is precisely the
hook M\*'s "extensible / integrate FastVideo-STA, xDiT, Inferix, FlashDrive" call-out asks for — **v2
already has the seam M\* only gestures at.**
What v2 lacks is M\*'s **declarative authoring layer above the substrate**, and — the key insight — *much of
that substrate is already authored but inert*: v2 has declared the metadata for "minimum components per
request" (`required_for`/`optional_for` on every omni card) and "branch as a cache axis" (`guidance_sig`,
`CacheKey`) but **never wired it to an executor**. The substrate is ~80% built and switched off.
**Highest-leverage cluster:** three small, parity-safe wires that turn on inert substrate and unblock the
BAGEL/Qwen-Omni/Cosmos3 latency wins M\* measured **on the exact models this repo already runs** — plus one
P1 that aligns v2 with the paper's headline "extensible" claim using a seam v2 already has.
### Verified P0 correctness findings (spot-checked by hand)
1. **Runner divergence (real bug).** `v2/runtime/engine.py:88` → `nodes = self.program.nodes`;
`v2/runtime/disaggregated.py:96` → `nodes = self.program.active_nodes(self.request)`. The inline and
disaggregated runners execute *different node sets*. ✅ confirmed.
2. **EOS is faked.** `v2/recipes/omni/ar_loop.py` docstring says "done on EOS/max_tokens"; `next()` (`:46-48`)
checks **only** `max_tokens`. M\*'s marquee `DynamicLoop` use case (EOS) is unimplemented in the loop that
serves the Qwen-Omni Thinker/Talker and Cosmos3 reasoner. ✅ confirmed.
3. **`required_for`/`optional_for` have zero runtime consumers** (grep outside `specs.py`/recipes/tests is
empty). The min-components metadata is declared on every card and never read. ✅ confirmed.
---
## 2. Dimension table (corrected)
| # | Dimension | v2 status | Gap | Priority | Effort | Payoff | Action |
|---|---|---|---|---|---|---|---|
| 1 | Min-components per request (`required_for` + `when_task`) | substrate built, **inert** | real, cheap | **P0** | S | Consume `required_for` in `active_nodes`; unify `engine.py:88` onto `active_nodes`; deliver via registry/card builder so all ~40 cards inherit it |
| 2 | Real EOS + declarative `DynamicLoop` | early-exit emergent; **EOS faked** | real | **P0** | S | `ARDecodeLoop` honors `eos_id` + `req.sampling.stop`; add `LoopSpec.dynamic_stop` + `register_loop_stop`. **Training-enabling** (world-model rollout horizon) |
| 3 | CFG/branch as label over one paged KV pool | absent (`PagedKVCache` is a counter) | real | **P1** | L | `(namespace,label)` paged store w/ one budget; reuse `guidance_sig` for hash (NOT `partition_field`); by-ref via existing `InProcKVConnector`. AR path only (diffusion has no KV) |
| 4 | `extend/` plugin seam → integrate FastVideo-STA / Inferix | **seam exists, unused for attn** | real (paper headline) | **P1** | M | Expose FastVideo sparse/sliding-tile attention + Inferix block-diffusion as `Interceptor`/`EngineKind` plugins — the paper's named integration targets, on this repo's own code |
| 5 | `ParitySpec.output_determinism` (C3 distributional) | C3 rung defined, **0 users** | real, dormant | **P1** | S | Add field; `compare_outputs` consults it. **Training-enabling** (SDE/FlowGRPO stochastic rollouts) |
| 6 | Registry-driven delivery of #1 | present, not leveraged | integration | **P1** | S | Express `when_task`/min-components through `WorkflowRegistry`/card builders, not 3 bespoke recipe patches |
| 7 | Serving conductor + pluggable data plane | conductor exists (`serving/http.py`); **single-process transport** | real | **P2** | L | v2 already has the step-scheduled worker surface; gap is ZeroMQ/Mooncake + direct worker→worker tensor routing (today `InProcKVConnector` only) |
| 8 | Fleet/Dynamo placement + replicas | **live** (`deploy/fleet.py`,`dynamo.py`) | partial | **P2** | M | Fleet-level placement/affinity/replica is real & ≥M\*; missing piece is only the intra-engine `(node,Walk)→rank` map decoupled from model code |
| 9 | Per-node TP / SP + cross-rank transport | axis vocab **exists** (`sp` incl.); not wired to runtime | partial | **P2** | XL | Wire declarative degrees into runtime; Wan/LTX are **SP-native** (TP is a no-op there); populate `parallel_plan_hash` on the serving cache path |
| 10 | Named Walks + per-model state machine | `Program`=G, sharing real; no Walk/SM | real | **P2** | M | Defer until a *re-entrant* phase graph (Thinker↔Talker, rollout) needs it; #1 captures the min-components win without it |
| 11 | Declarative `Parallel/Sequential/Loop` IR | imperative loop classes | real (authoring) | **P2** | M | Thin Section IR lowering to flat `Program`; scope to one AR recipe |
| 12 | Streaming `ChunkPolicy` + `StreamBuffer` | causal-chunk emit **already ships** (`wan_causal`); `EdgeKind.STREAM` inert | real | **P2** | L | Declarative `ChunkPolicy` vocab over the existing chunk mechanism; needs concurrent producer/consumer runner (= pipelined scheduling). Inferix integration point |
| 13 | Speculative deferred-termination; loop-spanning CUDA graphs; N+1 prefetch; attn double-buffer | absent / per-step capture (14 cards) | real | **P3** | L | Gate behind a real GPU executor; unobservable on CPU-toy CI; loop-span needs an `allows_interleaving=False` carve-out |
| — | Cost model + interleave/consistency parity | **exceeds M\*** | none | **guard** | — | Do not regress; keep `step_cost_model` mandatory + `bit_identical` default |
| — | Integrated training plane (flywheel, weight-sync) | **exceeds M\*** | none | **guard** | — | Protect train==serve loop identity with a toy fixture |
---
## 3. P0/P1 deep-dives (sequenced)
```
PR-1 (P0) min-components ──┐
PR-2 (P0) real EOS ─┼─► prereqs for honest "DynamicLoop" + min-component claims; both training-enabling
PR-3 (P1) output_determinism (independent)
PR-5 (P1) extend/ plugin: FastVideo-STA / Inferix as Interceptors (independent; highest paper-alignment)
PR-4 (P1) CFG-as-label paged pool ──► depends on PR-2 (AR loop is the only KV consumer)
```
PR-1, PR-2, PR-3, PR-5 are mutually independent; PR-4 depends on PR-2.
### PR-1 (P0) — Turn on the inert min-components substrate + fix runner divergence
- **Change.** Extend `Program.active_nodes(request)` (`v2/program/specs.py`) to also drop any node whose bound
`ComponentSpec.required_for` (`v2/card/specs.py:144`) excludes `request.task` (and isn't in `optional_for`).
**Fix the bug:** change `v2/runtime/engine.py:88` to `nodes = self.program.active_nodes(self.request)` so the
inline `ProgramRunner` matches `DisaggregatedRunner` (`disaggregated.py:96`). Deliver the `when_task` gating
through the **registry/card builder** (`recipes/__init__.py`, `program/workflow.py:WorkflowRegistry`) so all
~40 cards inherit it uniformly — not three bespoke `program.py` patches.
- **Why (this repo's models).** BAGEL T2I currently steps the AR-text loop and Cosmos3 t2v materializes the
reasoner even though the cards declare `transformer required_for={'reason','t2i'}`, `vae required_for={'t2i'}`.
On the GPU backend that is wasted resident-weight load + wasted steps on every single-modality request —
exactly M\*'s "execute the MINIMUM components per request," delivered by consuming existing metadata.
- **Risk/invariant.** Validate in `ModelCard.validate()` that every active node's `reads` are produced by an
active node for each declared `TaskType` (avoid dropping a producer). Pure node-id filtering ⇒ serial and
interleaved still walk the same filtered list ⇒ §9.3 interleave bit-parity holds by construction. CPU-toy clean.
### PR-2 (P0) — Real EOS + declarative `dynamic_stop` *(also training-enabling)*
- **Change.** In `v2/recipes/omni/ar_loop.py`, `advance()` reads the emitted token; if it equals the model
`eos_id` (toy backend exposes `EOS=0`) or matches `req.sampling.stop` (`params.py:21`, currently dead),
register termination; `next()` returns `Done()` on stop OR `max_tokens`. Add `StopRegistry` to `LoopState` +
`register_loop_stop(name)` to the `LoopContext` protocol (`contracts.py:204`) and to
`DisaggregatedRunner`'s `RuntimeLoopContext`. Add `LoopSpec.dynamic_stop: bool=False`, opt the AR cards in.
- **Why.** The docstring-vs-code lie sits in the loop serving Qwen-Omni Thinker/Talker and the Cosmos3 reasoner;
M\*'s second named `DynamicLoop` use case (world-model **rollout horizon**) is exactly what `self_forcing` RL
needs — so this is both a serving-credibility fix and a training enabler (raise its payoff accordingly).
- **Risk/invariant.** `dynamic_stop=False` is byte-identical back-compat. Must pass **all three** parity gates:
serial==interleaved AND disaggregated==inline. **Not** in this PR: speculative deferred-termination (unobservable
on CPU-toy, fights the interleave invariant — P3, gated on GPU executor).
### PR-3 (P1) — `ParitySpec.output_determinism` (close the dormant C3 hole) *(training-enabling)*
- **Change.** Add `output_determinism: str = "bit_identical"` to `ParitySpec` (`card/specs.py:88`); make
`compare_outputs` (`parity/interleave_gate.py:54`) consult it (`bit_identical` → today's exact check;
`distributional` → a moment/tolerance check — land a simple moment match first; a real KS test is new code).
- **Why.** `ConsistencyLevel.C3` is defined and used by zero recipes; an SDE/FlowGRPO stochastic rollout cannot
honestly declare its parity contract and would falsely fail the bit-identical gate. Additive; default unchanged.
### PR-5 (P1) — Expose FastVideo's own attention + Inferix as `extend/` plugins *(highest paper-alignment)*
- **Change.** Use the existing `extend/{interceptors,observers,registry}.py` seam (capability-negotiated, with
per-(request,branch) `plugin_state` that already passes the interleave gate) to register FastVideo's
sparse/sliding-tile attention and Inferix-style block-diffusion as `Interceptor`s / an `EngineKind` plugin.
- **Why.** M\*'s title is "Modular, **Extensible**" and it explicitly lists FastVideo-STA, xDiT/PipeFusion/USP,
Inferix, FlashDrive as integratable. v2 already has the seam M\* only describes — this is where v2 most
directly answers the paper, using this repo's own attention code. Low risk (the seam + capability negotiation
already exist and are tested).
### PR-4 (P1) — CFG/branch as a LABEL over one paged KV pool
- **Change.** Rewrite `PagedKVCache` (`cache/classes.py:155-172`) from a block *counter* into a real
`(namespace,label)->[block-handle]` store with **one shared `total_blocks` budget** (M\*'s single-pool
property). Reuse the existing-but-unpopulated `CacheKey.guidance_sig` (`keys.py:53`) for the hash. Thread the
label through `ar_loop.py` (alloc/append/get per `(request_id, branch)`; prefill once per shared-prefix label;
combine via `CFGPolicy.combine`). Wire `ResourceRequest.cache_blocks` (`contracts.py:64`, zero consumers) into
admission per (class,label).
- **Why.** The dossier-identified driver of M\*'s BAGEL win (3 CFG contexts as 3 labels over ONE pool vs dense
per-context). Targets AR_DECODE (BAGEL `generate_text`, omni Thinker); **correctly excludes diffusion**
(Wan/LTX are bidirectional, no KV — their CFG stays dense-but-batched).
- **Corrections to bake in.** Do **NOT** add `branch_label` to `CacheKey.partition_field()` (CFG branches share
embeddings; partitioning by branch is a semantic bug). Do **NOT** add a new by-ref type — reuse
`InProcKVConnector` + `TransferManifest.cache_key`. Wiring `cache_blocks` admission is greenfield ⇒ effort **L**.
CPU version proves label/sharing semantics; the real latency win needs a FlashInfer paged kernel (out of scope)
— **merge** with a future "real KVCacheEngine" effort rather than landing isolated.
---
## 4. What v2 already does ≥ M\* — do NOT regress
1. **Required+validated cost model** on every `LoopSpec` (13-kind `WorkUnitKind`) — typed, pre-GPU-validated.
2. **Interleave bit-parity as a hard gate** (`parity.interleave_required=True` on 40+ cards). M\* has no such
gate (its speculative scheduling deliberately wastes steps). Load-bearing invariant; every new primitive
must pass it.
3. **C0–C4 consistency ladder** wired into RL methods, with first-divergence tap reporting. No M\* equivalent.
4. **Integrated training plane** — DiffusionNFT/DMD2/self_forcing, RL→distill flywheel, `WeightSyncController`
hot weight-sync with drain-to-boundary + scoped cache invalidation, driving the **same** serving Loop.
M\* is serving-only. Protect with a toy fixture asserting `rollout_loop` drives the served Loop object.
5. **CPU-toy parity for the whole stack** — loops/CFG/caches/parity/RL run in CI without a GPU. Every new
primitive must ship a toy exercise (this is what makes all PRs above testable without H100s).
6. **Partition-not-flush cache invalidation** + four independent per-class pools.
7. **`extend/` plugin seam** with capability negotiation (a 4-step distilled card *rejects* a residual-skip
interceptor) — M\* describes extensibility; v2 has the mechanism.
8. **Dynamo citizenship** (`deploy/dynamo.py`: one `DeploymentCard`+cost model, two consumers) — beyond M\*'s
self-contained runtime.
---
## 5. Dropped / merged / deferred (and why)
- **DROP declarative `Parallel` as a CFG-execution win.** The runner walks nodes linearly (ignores
`Program.edges`), so `Parallel` lowers to sequential sugar and the CFG 3-pass braid is already one
co-scheduled `WorkPlan.run`; splitting it risks the interleave gate. Salvage only the no-op refactor
extracting `branch_forward` from `WanDenoiseLoop._velocity`. Reassign `Parallel` to the placement workstream.
- **MERGE the full Walk/state-machine layer** into "defer until a re-entrant phase graph needs it" (PR-1 gets the
min-components win with ~20 lines, no new abstraction). If built: the validator must check a walk's node-id
order is a *subsequence* of `program.nodes` (not just membership) or the runner can reorder and break parity.
- **MERGE `StreamBuffer`/`ChunkPolicy` into pipelined-scheduling.** Causal-chunk emit *already ships*
(`wan_causal/loop.py` per-chunk `StepResult.emit` + slab-KV); the gap is the declarative `ChunkPolicy` vocab
+ a concurrent producer/consumer runner. If built: keep all policies pure (per-request `StreamBuffer` history,
not shared edge state) and restrict the bit-identical claim to the token-only handoff.
- **MERGE CFG-fan-out exec + cross-rank transport + PD loop-splitting into a multi-GPU-runtime program.** These
need real collectives (`v2/distributed/` is a stub) and KV-by-reference (KV lives in `CacheManager`, not the
transferable `slots`). **Keep cheaply now:** the *declarative* halves — per-component degree, `(node,Walk)`
placement key with node-only fallback, `ReplicaSet` under `LocalFleet`, and populate `parallel_plan_hash` on
the **serving** cache path (it is already populated in `training/behavior.py:40` — the gap is serving-only).
- **DEFER** speculative deferred-termination, loop-spanning CUDA graphs, N+1 prefetch, attention-plan
double-buffer — all gated on a real GPU executor; benefit unobservable on CPU-toy CI. Keep the cheap
`EngineKind` tag (`STATELESS|KV_CACHE|DIFFUSION`) now. Correct the stale `cudagraph.py:51-52` docstring
(per-step capture ships in 14 cards, not just wan21).
- **RESCOPE per-node TP.** Wan/LTX use `ReplicatedLinear` + **sequence parallelism** (`sp`), not TP; the `sp`
axis already exists in `parallel/plan.py:AXIS_NAMES`. The work is wiring degrees into the runtime, not
inventing vocabulary; a `tp_size=2` "one-line activation" is a no-op for the shipped models.
---
## 6. The first integration test, if/when multi-GPU placement work starts
The **live Qwen-Omni 2-GPU bring-up** (Thinker on rank 0, Talker+Code2Wav on rank 1; see
`v2_debug_videos/vlm.md` Session 4) is the natural first validation target for any `(node,Walk)→rank`
placement work — it is the one place this repo already has real multi-rank composite-model execution.
---
## Anchor files for P0/P1
`v2/program/specs.py`, `v2/runtime/engine.py` (**line 88 fix**), `v2/runtime/disaggregated.py`,
`v2/recipes/omni/ar_loop.py`, `v2/loop/contracts.py`, `v2/card/specs.py`, `v2/cache/{classes.py,keys.py}`,
`v2/parity/interleave_gate.py`, `v2/extend/{interceptors,registry}.py`, `recipes/__init__.py` +
`v2/program/workflow.py` (registry-driven delivery).
@@ -69,7 +69,7 @@ approval, then upload reviewed accepted baseline records.
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
| `source_results` | Yes | One or more local paths or Buildkite artifact URLs for accepted shifted performance JSONs. Prefer normalized `normalized_perf_*.json` artifacts emitted by `compare_baseline.py`. Accept `source_result` as an alias only for a single JSON. |
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `PERF_MAX_REGRESSION` if set, otherwise `0.05` (5%). |
| `max_intra_batch_regression` | No | Maximum allowed regression of any source JSON against the source batch median. Default: `0.05` (5%). |
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
Hardcoded defaults:
@@ -148,8 +148,7 @@ For each metric with at least two non-null source values:
4. Stop if any source record regresses against the batch median by more than
`max_intra_batch_regression`.
Default `max_intra_batch_regression` to `PERF_MAX_REGRESSION` when set,
otherwise `0.05`. Print a table with per-source values, batch median, and
Default `max_intra_batch_regression` to `0.05`. Print a table with per-source values, batch median, and
worst intra-batch regression.
This check prevents uploading a mixed batch where one JSON is materially
@@ -183,7 +182,7 @@ present, that run is not a valid source for baseline reseeding.
### 2. Sync and back up existing HF records under /tmp
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
Use `fastvideo/performance/hf_store.py` helpers directly. Do **not** use
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
records, while this step must only fetch and back up existing history.
@@ -192,7 +191,7 @@ The sync command pattern is:
```bash
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
python -c 'from fastvideo.performance.hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
```
Then back up only the sanitized model directory under `/tmp`:
@@ -200,8 +199,8 @@ Then back up only the sanitized model directory under `/tmp`:
```bash
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
MODEL_SAFE=$(PYTHONPATH=fastvideo/tests/performance python - <<'PY'
from hf_store import sanitize
MODEL_SAFE=$(python - <<'PY'
from fastvideo.performance.hf_store import sanitize
print(sanitize("<model_id>"))
PY
)
@@ -235,7 +234,7 @@ first baseline seed. Continue, but report that baseline history was empty.
Load the last 5 successful records for the target:
```python
from hf_store import load_records_for_model
from fastvideo.performance.hf_store import load_records_for_model
records = load_records_for_model(
"/tmp/perf-tracking",
@@ -372,7 +371,7 @@ prepared records plus backup on disk.
Use the shared storage helper so the path and repo type match CI:
```python
from hf_store import upload_record
from fastvideo.performance.hf_store import upload_record
upload_record("<local_record_path>", record, strict=True)
```
@@ -460,7 +459,7 @@ directories created for this reseed. Never remove unrelated `/tmp` contents.
intentional baseline replacement.
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
median comparison, and persistence rules.
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
- `fastvideo/performance/hf_store.py` — HF sync, record loading,
`sanitize()`, and `upload_record()`.
- `fastvideo/tests/performance/test_inference_performance.py` — source result
JSON schema.
@@ -1,5 +1,9 @@
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 2,
"description": "Wan2.1 T2V 1.3B inference performance",
"model": {
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
+28 -1
View File
@@ -114,6 +114,17 @@ steps:
limit: 2
agents:
queue: "default"
- label: ":test_tube: LoRA Extraction Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "lora_extraction"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
retry:
automatic:
- exit_status: 128
limit: 3
- exit_status: -1
limit: 2
agents:
queue: "default"
- label: ":test_tube: Training Tests"
if: build.env("TEST_SCOPE") == "direct" && build.env("TEST_TYPE") == "training"
command: "timeout 90m .buildkite/scripts/pr_test.sh"
@@ -371,6 +382,21 @@ steps:
- TEST_TYPE=inference_lora
agents:
queue: "default"
- path:
- "scripts/lora_extraction/**"
- "fastvideo/tests/lora_extraction/**"
- "fastvideo/models/loader/**"
- "fastvideo/training/training_utils.py"
- "fastvideo/layers/lora/**"
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 90m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Extraction Tests"
env:
- TEST_TYPE=lora_extraction
agents:
queue: "default"
- path:
- "fastvideo/**"
- "pyproject.toml"
@@ -410,7 +436,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 25m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Training Tests"
env:
- TEST_TYPE=training_lora
@@ -455,6 +481,7 @@ steps:
- "fastvideo/layers/**"
- "fastvideo/worker/**"
- "fastvideo/entrypoints/**"
- "fastvideo/performance/**"
- "fastvideo/tests/performance/**"
- ".buildkite/performance-benchmarks/**"
- "pyproject.toml"
+23 -1
View File
@@ -80,6 +80,23 @@ MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUI
POST_RUN_HOOK=""
is_truthy() {
case "${1:-}" in
1|true|TRUE|yes|YES|on|ON) return 0 ;;
*) return 1 ;;
esac
}
ssim_bootstrap_args() {
local title="${PR_TITLE:-}"
local message="${BUILDKITE_MESSAGE:-}"
if is_truthy "${FASTVIDEO_SSIM_BOOTSTRAP_MODE:-}" \
|| [[ "$title" == *"[new-model]"* ]] \
|| [[ "$message" == *"[new-model]"* ]]; then
printf ' --bootstrap-mode'
fi
}
upload_performance_artifacts() {
SHORT_SHA=${BUILDKITE_COMMIT:0:7}
LOCAL_DIR="downloaded_reports"
@@ -172,7 +189,12 @@ case "$TEST_TYPE" in
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run $MODAL_SSIM_TEST_FILE::run_ssim_tests"
SSIM_BOOTSTRAP_ARGS=$(ssim_bootstrap_args)
if [ -n "$SSIM_BOOTSTRAP_ARGS" ]; then
log "SSIM bootstrap mode enabled for new-model reference draft generation"
fi
MODAL_COMMAND="$MODAL_ENV HF_API_KEY=$HF_API_KEY python3 -m modal run "
MODAL_COMMAND+="$MODAL_SSIM_TEST_FILE::run_ssim_tests$SSIM_BOOTSTRAP_ARGS"
;;
"training")
log "Running training tests..."
+133
View File
@@ -0,0 +1,133 @@
#!/usr/bin/env bash
# Gate the expensive Buildkite full suite on the cheap GitHub checks.
#
# Polls the workflow runs for the PR head commit and only exits 0 once the
# watched cheap workflows (pre-commit, docs build) have succeeded, so the
# 'ready' label cannot burn ~20 GPU lanes on a head that a cheap check has
# already doomed.
#
# Semantics:
# - watched run completed with a bad conclusion -> exit 1 (fail CLOSED:
# no full suite; the next push re-arms via the 'synchronize' trigger)
# - watched run cancelled -> still pending: the docs
# workflow's repo-global 'pages' concurrency group cancels runs superseded
# by unrelated pushes, so 'cancelled' is not a verdict on this PR
# - watched runs pending -> poll until done
# - docs run absent -> not applicable after a
# short grace period ('Deploy Documentation' is path-filtered on PRs)
# - pre-commit run absent -> keep polling: pre-commit
# is never path-filtered, so its absence is always anomalous
# - 'ready' label removed while waiting -> exit 1 (fail CLOSED:
# un-labeling is a deliberate maintainer action)
# - GitHub API unreachable or timeout -> exit 0 (fail OPEN,
# loud warning: never brick CI on a GitHub outage)
#
# Required env: PR_SHA (PR head commit), PR_NUMBER, GITHUB_REPOSITORY, GH_TOKEN.
set -euo pipefail
: "${PR_SHA:?PR_SHA (PR head commit) is required}"
: "${PR_NUMBER:?PR_NUMBER (pull request number) is required}"
: "${GITHUB_REPOSITORY:?GITHUB_REPOSITORY is required}"
# Workflow-level `name:` values that must be green before the full suite
# may start. "Deploy Documentation" is path-filtered on PRs, so its run may
# legitimately never exist; pre-commit always runs, so it must appear.
WATCHED_NAMES='["pre-commit", "Deploy Documentation"]'
WATCHED_REGEX='^(pre-commit|Deploy Documentation)$'
POLL_SECS="${POLL_SECS:-20}"
GRACE_SECS="${GRACE_SECS:-60}"
MAX_WAIT_SECS="${MAX_WAIT_SECS:-1500}"
# Bound each API call so a hung connection hits the 3-strike fail-open path
# instead of pinning the loop until the job timeout (which would fail closed
# on exactly the GitHub-outage case this script is meant to survive).
if command -v timeout >/dev/null 2>&1; then
gh_api() { timeout 30 gh api "$@"; }
else
gh_api() { gh api "$@"; } # macOS dev boxes; CI always has coreutils timeout
fi
# The workflow checked the label before starting the gate, but the wait can
# last ~25 min: re-check once before any exit 0 and fail closed if 'ready'
# was removed in the meantime. An API error here proceeds (the label was
# present when the gate started; never brick CI on an outage).
recheck_ready_label() {
local pr_json
if pr_json=$(gh_api "repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}" 2>/dev/null); then
if ! jq -e '[.labels[]?.name] | index("ready")' <<<"$pr_json" >/dev/null 2>&1; then
echo "::error::PR #${PR_NUMBER} no longer has the 'ready' label —" \
"NOT triggering the Buildkite full suite. Re-add the label to re-arm."
exit 1
fi
else
echo "::warning::Could not re-check the 'ready' label on PR #${PR_NUMBER}; proceeding (it was present when the gate started)."
fi
}
start=$(date +%s)
api_fails=0
missing=""
while true; do
elapsed=$(( $(date +%s) - start ))
if runs_json=$(gh_api "repos/${GITHUB_REPOSITORY}/actions/runs?head_sha=${PR_SHA}&per_page=100" 2>/dev/null) \
&& state=$(jq --arg re "$WATCHED_REGEX" '
[.workflow_runs[]? | select(.name // "" | test($re))]
| group_by(.name) | map(max_by(.id))
| map({name, status, conclusion})' <<<"$runs_json" 2>/dev/null); then
api_fails=0
echo "t+${elapsed}s watched checks: $(jq -c . <<<"$state")"
failed=$(jq -r '[.[] | select(.status == "completed"
and (.conclusion | IN("success", "skipped", "neutral", "cancelled") | not))]
| map(.name) | join(", ")' <<<"$state")
if [ -n "$failed" ]; then
echo "::error::Cheap check(s) failed on ${PR_SHA}: ${failed}." \
"NOT triggering the Buildkite full suite. Push a fix (the 'ready'" \
"label re-arms on every push), or re-run the failed check and then" \
"re-run this workflow."
exit 1
fi
# 'cancelled' counts as pending: wait for a re-run to reach a real verdict
# (bounded by MAX_WAIT, then the fail-open below).
pending=$(jq '[.[] | select(.status != "completed" or .conclusion == "cancelled")] | length' <<<"$state")
missing=$(jq -r --argjson watched "$WATCHED_NAMES" '($watched - map(.name)) | join(", ")' <<<"$state")
if [ "$pending" -eq 0 ]; then
if [ -z "$missing" ]; then
recheck_ready_label
echo "All watched cheap checks are green — full suite may proceed."
exit 0
fi
case "$missing" in
*pre-commit*)
echo "pre-commit run not found for ${PR_SHA} yet; waiting (pre-commit is never path-filtered, so its absence is anomalous)."
;;
*)
if [ "$elapsed" -ge "$GRACE_SECS" ]; then
recheck_ready_label
echo "::warning::Watched run(s) never appeared for ${PR_SHA}: ${missing} (path-filtered, likely not applicable). Proceeding on the checks that did run."
exit 0
fi
echo "Waiting up to ${GRACE_SECS}s grace for path-filtered run(s) to appear: ${missing}."
;;
esac
fi
else
api_fails=$(( api_fails + 1 ))
echo "::warning::GitHub API error querying workflow runs for ${PR_SHA} (attempt ${api_fails}/3)."
if [ "$api_fails" -ge 3 ]; then
recheck_ready_label
echo "::warning::FAILING OPEN: cannot query GitHub check status — triggering the full suite WITHOUT the cheap-check gate."
exit 0
fi
fi
if [ "$elapsed" -ge "$MAX_WAIT_SECS" ]; then
recheck_ready_label
echo "::warning::FAILING OPEN: watched checks still pending after $(( MAX_WAIT_SECS / 60 )) min${missing:+ (never appeared: ${missing})} — triggering the full suite anyway."
exit 0
fi
sleep "$POLL_SECS"
done
+122
View File
@@ -0,0 +1,122 @@
#!/usr/bin/env bash
# Self-test for gate_full_suite.sh using a mocked `gh`. No network, runs on
# any dev box: bash .github/scripts/test_gate_full_suite.sh
set -u
here=$(cd "$(dirname "$0")" && pwd)
tmp=$(mktemp -d)
trap 'rm -rf "$tmp"' EXIT
# Mock gh. Asserts the exact endpoint (including head_sha) it is called
# with — an endpoint typo in the gate script fails the test rather than
# silently serving canned data. On the runs endpoint it serves
# $MOCK_DIR/response_<call#>.json, sticking on the highest existing file,
# and exits 1 if none exist (simulates a GitHub API outage). On the pulls
# endpoint it serves $MOCK_DIR/pr.json, defaulting to a 'ready'-labeled PR.
cat > "$tmp/gh" <<'EOF'
#!/usr/bin/env bash
if [ "${1:-}" != "api" ]; then
echo "unexpected gh invocation: $*" >> "$MOCK_DIR/endpoint_error"
exit 2
fi
case "${2:-}" in
"repos/o/r/actions/runs?head_sha=deadbeef&per_page=100")
n=$(( $(cat "$MOCK_DIR/count" 2>/dev/null || echo 0) + 1 ))
echo "$n" > "$MOCK_DIR/count"
while [ "$n" -gt 0 ]; do
if [ -f "$MOCK_DIR/response_$n.json" ]; then
cat "$MOCK_DIR/response_$n.json"
exit 0
fi
n=$(( n - 1 ))
done
echo "api outage" >&2
exit 1
;;
"repos/o/r/pulls/42")
if [ -f "$MOCK_DIR/pr.json" ]; then
cat "$MOCK_DIR/pr.json"
else
echo '{"labels": [{"name": "ready"}]}'
fi
;;
*)
echo "unexpected gh endpoint: $2" >> "$MOCK_DIR/endpoint_error"
exit 2
;;
esac
EOF
chmod +x "$tmp/gh"
PC_OK='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "success"}'
PC_BAD='{"name": "pre-commit", "id": 1, "status": "completed", "conclusion": "failure"}'
PC_PENDING='{"name": "pre-commit", "id": 1, "status": "in_progress", "conclusion": null}'
DOCS_OK='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "success"}'
DOCS_BAD='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "failure"}'
DOCS_CANCELLED='{"name": "Deploy Documentation", "id": 2, "status": "completed", "conclusion": "cancelled"}'
OTHER='{"name": "Trigger Full Suite", "id": 3, "status": "in_progress", "conclusion": null}'
NULL_NAME='{"name": null, "id": 4, "status": "completed", "conclusion": "failure"}'
PC_OK_RERUN='{"name": "pre-commit", "id": 5, "status": "completed", "conclusion": "success"}'
fails=0
want_log="" # optional: expect() also greps out.log for this regex, then resets
pr_json="" # optional: served for the pulls (label re-check) endpoint, then resets
raw_body="" # optional: serve responses verbatim instead of wrapping in workflow_runs
expect() { # <name> <expected-exit> <response json>...
local name=$1 want=$2 dir i=1
shift 2
dir=$(mktemp -d "$tmp/test_XXXXXX")
for body in "$@"; do
if [ -n "$raw_body" ]; then
printf '%s' "$body" > "$dir/response_$i.json"
else
printf '{"workflow_runs": [%s]}' "$body" > "$dir/response_$i.json"
fi
i=$(( i + 1 ))
done
[ -n "$pr_json" ] && printf '%s' "$pr_json" > "$dir/pr.json"
( export PATH="$tmp:$PATH" MOCK_DIR="$dir" PR_SHA=deadbeef PR_NUMBER=42 \
GITHUB_REPOSITORY=o/r POLL_SECS=0 GRACE_SECS=1 MAX_WAIT_SECS=3
bash "$here/gate_full_suite.sh" > "$dir/out.log" 2>&1 )
local rc=$?
if [ "$rc" -ne "$want" ]; then
echo "FAIL: $name (exit $rc, want $want)"
cat "$dir/out.log"
fails=1
elif [ -f "$dir/endpoint_error" ]; then
echo "FAIL: $name (mock gh got an unexpected call)"
cat "$dir/endpoint_error"
fails=1
elif [ -n "$want_log" ] && ! grep -Eq "$want_log" "$dir/out.log"; then
echo "FAIL: $name (log does not match: $want_log)"
cat "$dir/out.log"
fails=1
else
echo "ok: $name"
fi
want_log="" pr_json="" raw_body=""
}
expect "both green -> proceed" 0 "$PC_OK, $DOCS_OK, $OTHER, $NULL_NAME"
expect "docs build failed -> blocked" 1 "$PC_OK, $DOCS_BAD"
expect "pre-commit failed -> blocked" 1 "$PC_BAD"
expect "pending then green -> proceed" 0 "$PC_PENDING" "$PC_OK, $DOCS_OK"
want_log="never appeared.*Deploy Documentation"
expect "docs run absent (path-filtered) -> proceed after grace" 0 "$PC_OK"
expect "API outage -> fail open" 0
want_log="FAILING OPEN"
expect "pending past MAX_WAIT -> fail open" 0 "$PC_PENDING"
want_log="FAILING OPEN"
expect "unrelated runs only -> no grace, fail open at MAX_WAIT" 0 "$OTHER"
expect "cancelled docs then green -> proceed" 0 \
"$PC_OK, $DOCS_CANCELLED" "$PC_OK, $DOCS_OK"
want_log="FAILING OPEN"
expect "cancelled docs forever -> fail open at MAX_WAIT" 0 "$PC_OK, $DOCS_CANCELLED"
want_log="FAILING OPEN"
expect "pre-commit absent -> no grace, fail open at MAX_WAIT" 0 "$DOCS_OK"
expect "duplicate run names -> latest wins" 0 "$PC_BAD, $PC_OK_RERUN, $DOCS_OK"
raw_body=1
expect "garbage response body -> fail open" 0 "this is not json"
pr_json='{"labels": [{"name": "other"}]}'
expect "ready label removed mid-gate -> blocked" 1 "$PC_OK, $DOCS_OK"
exit "$fails"
+22 -2
View File
@@ -1,7 +1,11 @@
name: pre-commit
on:
pull_request:
# pull_request_target instead of pull_request: the workflow definition and
# the hook config are always taken from the BASE branch, so fork /
# first-time-contributor PRs run immediately without a maintainer clicking
# "Approve and run". The PR head is checked out as data only.
pull_request_target:
branches: [main]
workflow_call:
inputs:
@@ -15,12 +19,25 @@ permissions:
jobs:
pre-commit:
if: github.event_name == 'workflow_call' || github.event.pull_request.draft != true
if: github.event.pull_request.draft != true
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
with:
ref: ${{ inputs.ref || '' }}
# For PR events, lint the PR head — but keep the hook definitions from
# the base branch so an untrusted PR cannot alter what gets executed.
- name: Save trusted hook config
if: github.event_name == 'pull_request_target'
run: cp .pre-commit-config.yaml "$RUNNER_TEMP/trusted-pre-commit-config.yaml"
- uses: actions/checkout@v4
if: github.event_name == 'pull_request_target'
with:
ref: ${{ github.event.pull_request.head.sha }}
persist-credentials: false
- name: Restore trusted hook config
if: github.event_name == 'pull_request_target'
run: cp "$RUNNER_TEMP/trusted-pre-commit-config.yaml" .pre-commit-config.yaml
- uses: actions/setup-python@v5
with:
python-version: "3.12"
@@ -30,3 +47,6 @@ jobs:
- uses: pre-commit/action@v3.0.1
with:
extra_args: --all-files --hook-stage manual
# After pre-commit so a self-test failure cannot mask lint failures.
- name: Full-suite gate self-test
run: bash .github/scripts/test_gate_full_suite.sh
+19 -11
View File
@@ -52,6 +52,7 @@ jobs:
core.setOutput('pr_sha', pr.head.sha);
core.setOutput('pr_branch', pr.head.ref);
core.setOutput('pr_number', String(prNumber));
core.setOutput('pr_title', pr.title);
- name: Trigger Full Suite
if: steps.perm.outputs.has_write == 'true'
@@ -60,6 +61,7 @@ jobs:
PR_SHA: ${{ steps.label.outputs.pr_sha }}
PR_BRANCH: ${{ steps.label.outputs.pr_branch }}
PR_NUMBER: ${{ steps.label.outputs.pr_number }}
PR_TITLE: ${{ steps.label.outputs.pr_title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
@@ -71,6 +73,7 @@ jobs:
--arg commit "$PR_SHA" \
--arg branch "$PR_BRANCH" \
--arg message "Full Suite for PR #${PR_NUMBER} (via /merge)" \
--arg pr_title "$PR_TITLE" \
--argjson pr_id "$PR_NUMBER" \
'{
commit: $commit,
@@ -80,11 +83,12 @@ jobs:
pull_request_id: $pr_id,
pull_request_base_branch: "main",
env: {
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring)
}
}')"
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring),
PR_TITLE: $pr_title
}
}')"
parse-command:
if: >-
@@ -125,7 +129,7 @@ jobs:
set -euo pipefail
TEST_NAME=$(echo "$COMMENT" | grep -oP '(?<=/test\s)\S+' | head -1 || true)
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
VALID="encoder vae transformer kernel unit dreamverse ssim training lora-inference lora-training lora-extraction distillation self-forcing vsa vmoba performance api train-framework eval full fastcheck pre-commit"
if [ -z "$TEST_NAME" ] || ! echo "$VALID" | grep -qw "$TEST_NAME"; then
echo "Unknown test: '$TEST_NAME'. Valid: $VALID"
exit 1
@@ -136,6 +140,7 @@ jobs:
[kernel]=kernel_tests [unit]=unit_test [dreamverse]=dreamverse_app
[ssim]=ssim [training]=training
[lora-inference]=inference_lora [lora-training]=training_lora
[lora-extraction]=lora_extraction
[distillation]=distillation_dmd [self-forcing]=self_forcing
[vsa]=training_vsa [vmoba]=inference_vmoba
[performance]=performance [api]=api_server
@@ -240,6 +245,7 @@ jobs:
TEST_SCOPE: ${{ needs.parse-command.outputs.test_scope }}
FULL_SUITE: ${{ needs.parse-command.outputs.full_suite }}
TEST_TYPE: ${{ needs.parse-command.outputs.test_type }}
PR_TITLE: ${{ github.event.issue.title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
@@ -256,6 +262,7 @@ jobs:
--arg full_suite "$FULL_SUITE" \
--arg test_type "$TEST_TYPE" \
--arg pr_number "$PR_NUMBER" \
--arg pr_title "$PR_TITLE" \
'{
commit: $commit,
branch: $branch,
@@ -265,8 +272,9 @@ jobs:
pull_request_base_branch: "main",
env: {
TEST_SCOPE: $test_scope,
FULL_SUITE: $full_suite,
TEST_TYPE: $test_type,
PR_NUMBER: $pr_number
}
}')"
FULL_SUITE: $full_suite,
TEST_TYPE: $test_type,
PR_NUMBER: $pr_number,
PR_TITLE: $pr_title
}
}')"
+21 -1
View File
@@ -7,6 +7,7 @@ on:
permissions:
contents: read
pull-requests: read
actions: read
concurrency:
group: full-suite-${{ github.event.pull_request.number }}
@@ -18,6 +19,8 @@ jobs:
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|| github.event.action == 'synchronize'
runs-on: ubuntu-latest
# Gate below may wait for cheap checks (up to MAX_WAIT_SECS = 25 min).
timeout-minutes: 35
steps:
- name: Check ready label
id: check
@@ -49,6 +52,20 @@ jobs:
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
done
# Checks out the BASE branch (default for pull_request_target), so PR
# authors cannot tamper with the gate script.
- name: Checkout gate script
if: steps.check.outputs.has_ready == 'true'
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
- name: Wait for pre-commit and docs build
if: steps.check.outputs.has_ready == 'true'
env:
GH_TOKEN: ${{ github.token }}
PR_SHA: ${{ github.event.pull_request.head.sha }}
PR_NUMBER: ${{ github.event.pull_request.number }}
run: bash .github/scripts/gate_full_suite.sh
- name: Trigger Buildkite Full Suite
if: steps.check.outputs.has_ready == 'true'
env:
@@ -56,6 +73,7 @@ jobs:
PR_SHA: ${{ github.event.pull_request.head.sha }}
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
PR_NUMBER: ${{ github.event.pull_request.number }}
PR_TITLE: ${{ github.event.pull_request.title }}
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
run: |
@@ -67,6 +85,7 @@ jobs:
--arg commit "$PR_SHA" \
--arg branch "$PR_BRANCH" \
--arg message "Full Suite for PR #${PR_NUMBER}" \
--arg pr_title "$PR_TITLE" \
--argjson pr_id "$PR_NUMBER" \
'{
commit: $commit,
@@ -78,6 +97,7 @@ jobs:
env: {
TEST_SCOPE: "full",
FULL_SUITE: "true",
PR_NUMBER: ($pr_id | tostring)
PR_NUMBER: ($pr_id | tostring),
PR_TITLE: $pr_title
}
}')"
+30 -5
View File
@@ -13,12 +13,33 @@ on:
required: false
default: false
type: boolean
# Auto-rebuild the CUDA images when their Dockerfile changes on main. The CUDA
# matrix is the only lane that builds from docker/Dockerfile, so a path-scoped
# push trigger is a sufficient change detector on its own -- no separate
# detect-changes/paths-filter job is needed now that there is a single
# in-scope Dockerfile. Dreamverse (apps/dreamverse/docker/Dockerfile) and the
# rocm Dockerfile stay manual-dispatch only.
push:
branches: [main]
paths:
- 'docker/Dockerfile'
permissions:
contents: read
packages: write
# One static group, no cancellation: every run of this workflow writes the same
# mutable registry tags (latest, py3.12-latest, ...), so runs must serialize —
# concurrent push/dispatch runs would race on those tags, and cancelling a run
# mid-publish can strand the cu126/cu130 tag families at different commits. An
# in-flight superseded build wastes its runner time, but its tags are then
# overwritten by the newer queued run. GitHub keeps a single pending run per
# group: the newest queued run replaces any older queued one.
concurrency:
group: infra-build-image
cancel-in-progress: false
jobs:
# CUDA matrix: Python 3.12 x {12.6.3, 13.0.0} x {amd64, arm64}. Each architecture
# builds natively and pushes only by digest; publish-cuda-manifests is the sole
@@ -28,7 +49,11 @@ jobs:
# aliases; 13.0.0/cu130 is published under explicit versioned tags. Flash-attn
# 2.8.3 comes from the architecture-specific prebuilt releases.
build-cuda-images:
if: ${{ github.event.inputs.build_cuda_matrix == 'true' }}
# Runs on a manual dispatch when build_cuda_matrix is set, or automatically
# on a push that changed docker/Dockerfile (inputs are null on push). The
# repository guard keeps fork syncs from auto-building; manual dispatch
# still works in forks.
if: ${{ (github.event_name == 'push' && github.repository == 'hao-ai-lab/FastVideo') || github.event.inputs.build_cuda_matrix == 'true' }}
strategy:
fail-fast: false
matrix:
@@ -75,10 +100,10 @@ jobs:
secrets: inherit
publish-cuda-manifests:
# !cancelled(): a failed sibling build leg must not skip the manifests for a
# CUDA lane whose own digests all exist; the digest-count check below fails
# the incomplete lane loudly instead.
if: ${{ !cancelled() && github.event.inputs.build_cuda_matrix == 'true' }}
# !cancelled(): publish lanes whose digests exist even if a sibling build
# leg failed (the digest-count check fails incomplete lanes); it also
# bypasses skipped-needs propagation, hence the explicit skipped check.
if: ${{ !cancelled() && needs.build-cuda-images.result != 'skipped' }}
needs: build-cuda-images
runs-on: ubuntu-latest
permissions:
+8 -2
View File
@@ -37,6 +37,11 @@ logs/
official_weights/
converted_weights/
# Cosmos3 local parity assets (symlinked from main worktree)
/official_weights/
/converted_weights/
/cosmos-framework
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
**/.cache/**
@@ -72,8 +77,7 @@ docs/distillation/examples/
# Python pickle files
*.pkl
# Reference videos
!fastvideo/tests/ssim/reference_videos/**/*.mp4
# Reference videos (negations must come after the catch-all on line below)
# Static images
!docs/assets/images/**/*.png
@@ -127,6 +131,8 @@ apps/dreamverse/web/.env.production.local
.sisyphus/
openspec/
fastvideo/tests/ssim/reference_videos/**
!fastvideo/tests/ssim/reference_videos/**/*.mp4
!fastvideo/tests/ssim/reference_videos/**/*.png
# Editor logs and local Python version pins (accidentally committed)
*.nvimlog
-2
View File
@@ -10,8 +10,6 @@ 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/.*|
+3
View File
@@ -84,9 +84,12 @@ RUN source /opt/venv/bin/activate \
FFMPEG_NATIVE_CXX=/usr/bin/g++ \
bash /opt/FastVideo/apps/dreamverse/scripts/install_native_ffmpeg.sh
# FASTVIDEO_FA4: FA4 (flash_attn.cute) is opt-in; this image installs it via
# the dreamverse extra and is validated with it, so enable it here.
ENV FASTVIDEO_DREAMVERSE_HOME=/var/lib/dreamverse \
STREAM_MODE=av_fmp4 \
FASTVIDEO_ENABLE_PROMPT_SAFETY=0 \
FASTVIDEO_FA4=1 \
HF_HOME=/root/.cache/huggingface
RUN mkdir -p /var/lib/dreamverse
+8 -3
View File
@@ -12,13 +12,17 @@ Defaults:
- `HF_REPO_ID=FastVideo/performance-tracking`
- `PERFORMANCE_TRACKING_ROOT=/tmp/fastvideo-perf-dashboard`
- `PERF_MAX_REGRESSION=0.05`
Records can include source metadata:
Records can include source metadata and rolling-baseline policy context:
- `run_source`: `pr`, `local`, `scheduled_main`, or `unknown`
- `baseline_eligible`: only successful scheduled-main records should be true
- Buildkite metadata such as branch, PR number, build URL, build ID, and job ID
- `regression_thresholds`: per-metric rolling-baseline percent and absolute
floors used for recomputed status context
Dashboard/API metric payloads expose `threshold_exceeded` for raw threshold
crossings; `regressed` remains the gated CI-failure signal.
Set one of `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` if the
configured dataset repo requires authenticated access:
@@ -89,7 +93,8 @@ Trend charts show metric-specific axes and exact point details on hover/focus:
- PR number, branch, and Buildkite URL when present
The latest status table uses the stored JSON `success` value. Recomputed
baseline context is shown separately and does not override stored status.
baseline context applies each metric's percent and absolute regression floors
and does not override stored status.
## API
@@ -1,6 +1,7 @@
import { useEffect, useMemo, useState } from "react";
import { fetchSummary, fetchTrends, refreshData, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
import { fetchSummary, fetchTrends, refreshData } from "./api";
import type { CohortValue, RunSource, SummaryResponse, TrendGroup, TrendPoint } from "./api";
const METRIC_KEYS = ["latency", "throughput", "memory", "text_encoder_time_s", "dit_time_s", "vae_decode_time_s"];
const RUN_SOURCES: Array<{ value: "" | RunSource; label: string }> = [
@@ -109,6 +110,61 @@ function metricLabel(metricKey: string) {
return METRIC_DEFINITIONS[metricKey]?.label ?? metricKey;
}
type CohortFields = {
model_id: string;
gpu_type: string;
workload_id: CohortValue;
variant_id: CohortValue;
benchmark_version: CohortValue;
recipe_fingerprint: CohortValue;
hardware_profile_id: CohortValue;
software_profile_id: CohortValue;
};
function cohortValue(value: CohortValue) {
if (value === null || value === undefined || value === "") {
return "legacy";
}
return String(value);
}
function shortCohortValue(value: CohortValue) {
const text = cohortValue(value);
if (text === "legacy" || text.length <= 14) {
return text;
}
return text.slice(0, 12);
}
function cohortKey(cohort: CohortFields) {
return [
cohort.model_id,
cohort.gpu_type,
cohortValue(cohort.workload_id),
cohortValue(cohort.variant_id),
cohortValue(cohort.benchmark_version),
cohortValue(cohort.recipe_fingerprint),
cohortValue(cohort.hardware_profile_id),
cohortValue(cohort.software_profile_id)
].join("|");
}
function cohortTitle(cohort: CohortFields) {
const workload = cohortValue(cohort.workload_id);
const variant = cohortValue(cohort.variant_id);
const version = cohortValue(cohort.benchmark_version);
const versionLabel = version === "legacy" ? version : `v${version}`;
return `${workload} / ${variant} / ${versionLabel}`;
}
function cohortDetail(cohort: CohortFields) {
return [
`recipe ${shortCohortValue(cohort.recipe_fingerprint)}`,
shortCohortValue(cohort.hardware_profile_id),
shortCohortValue(cohort.software_profile_id)
].join(" | ");
}
function formatMetricValue(metricKey: string, value: number | null | undefined, tooltip = false) {
const definition = METRIC_DEFINITIONS[metricKey];
if (!definition) {
@@ -171,7 +227,9 @@ function TrendChart({ group, metricKey }: { group: TrendGroup; metricKey: string
top: `${(activePoint.y / height) * 100}%`
}
: undefined;
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}`;
const ariaLabel = `${metricLabel(metricKey)} trend for ${group.model_id} on ${group.gpu_type}, ${cohortTitle(
group
)}`;
return (
<div className="chart-shell">
@@ -419,7 +477,7 @@ export default function App() {
<section className="panel">
<div className="panel-header">
<h2>Latest Status</h2>
<span>{latestRows.length} model/GPU groups</span>
<span>{latestRows.length} comparison cohorts</span>
</div>
{latestRows.length === 0 ? (
<div className="empty">No records match the selected filters.</div>
@@ -432,6 +490,7 @@ export default function App() {
<th>Recomputed</th>
<th>Model</th>
<th>GPU</th>
<th>Cohort</th>
<th>Commit</th>
<th>Source</th>
<th>Baseline</th>
@@ -440,11 +499,13 @@ export default function App() {
<th>Throughput</th>
<th>Memory</th>
<th>Worst</th>
<th>Exceeded</th>
<th>Failing</th>
</tr>
</thead>
<tbody>
{latestRows.map((row) => (
<tr key={`${row.model_id}-${row.gpu_type}`}>
<tr key={cohortKey(row)}>
<td>
<span className={`badge ${row.status}`}>{row.status}</span>
</td>
@@ -455,6 +516,12 @@ export default function App() {
</td>
<td>{row.model_id}</td>
<td>{row.gpu_type}</td>
<td>
<div className="cohort-cell">
<strong>{cohortTitle(row)}</strong>
<span>{cohortDetail(row)}</span>
</div>
</td>
<td>{shortSha(row.commit_sha)}</td>
<td>
<span className={`source-badge source-${row.run_source}`}>{runSourceLabel(row.run_source)}</span>
@@ -465,6 +532,12 @@ export default function App() {
<td>{formatNumber(row.metrics.throughput?.current, 3)}</td>
<td>{formatNumber(row.metrics.memory?.current, 1)}</td>
<td>{formatNumber(row.worst_regression_pct, 1)}%</td>
<td>
{row.threshold_exceeded_metrics.length
? row.threshold_exceeded_metrics.join(", ")
: "none"}
</td>
<td>{row.failing_metrics.length ? row.failing_metrics.join(", ") : "none"}</td>
</tr>
))}
</tbody>
@@ -487,11 +560,13 @@ export default function App() {
) : (
trends.map((group) =>
METRIC_KEYS.map((metricKey) => (
<article className="trend-card" key={`${group.model_id}-${group.gpu_type}-${metricKey}`}>
<article className="trend-card" key={`${cohortKey(group)}-${metricKey}`}>
<div>
<h3>{metricLabel(metricKey)}</h3>
<p>
{group.model_id} | {group.gpu_type}
<span>{cohortTitle(group)}</span>
<span>{cohortDetail(group)}</span>
</p>
</div>
<TrendChart group={group} metricKey={metricKey} />
+22 -4
View File
@@ -2,11 +2,28 @@ export type MetricValue = {
current: number | null;
baseline: number | null;
regression_pct: number | null;
absolute_delta: number | null;
threshold_percent: number;
threshold_absolute: number;
gated: boolean;
threshold_exceeded: boolean;
regressed: boolean;
label: string;
lower_is_better: boolean;
precision: number;
};
export type CohortValue = string | number | null;
export type ComparisonCohort = {
workload_id: CohortValue;
variant_id: CohortValue;
benchmark_version: CohortValue;
recipe_fingerprint: CohortValue;
hardware_profile_id: CohortValue;
software_profile_id: CohortValue;
};
export type SummaryRow = {
model_id: string;
gpu_type: string;
@@ -15,7 +32,8 @@ export type SummaryRow = {
success: boolean;
baseline_n: number;
worst_regression_pct: number | null;
regression_threshold_pct: number;
threshold_exceeded_metrics: string[];
failing_metrics: string[];
computed_regression_status: "pass" | "fail";
status: "pass" | "fail";
run_source: RunSource;
@@ -27,7 +45,7 @@ export type SummaryRow = {
build_id: string;
job_id: string;
metrics: Record<string, MetricValue>;
};
} & ComparisonCohort;
export type RunSource = "pr" | "local" | "scheduled_main" | "unknown";
@@ -61,13 +79,13 @@ export type TrendPoint = {
build_id: string;
job_id: string;
metrics: Record<string, number | null>;
};
} & ComparisonCohort;
export type TrendGroup = {
model_id: string;
gpu_type: string;
points: TrendPoint[];
};
} & ComparisonCohort;
export type TrendsResponse = {
groups: TrendGroup[];
@@ -149,6 +149,11 @@ h3 {
font-size: 0.82rem;
}
.trend-card p {
display: grid;
gap: 2px;
}
.stat strong {
display: block;
margin-top: 8px;
@@ -186,7 +191,7 @@ h3 {
table {
width: 100%;
min-width: 1120px;
min-width: 1260px;
border-collapse: collapse;
}
@@ -209,6 +214,25 @@ td {
font-size: 0.9rem;
}
.cohort-cell {
display: grid;
gap: 2px;
}
.cohort-cell strong,
.trend-card p span {
color: #1b2836;
font-size: 0.78rem;
font-weight: 700;
}
.cohort-cell span,
.trend-card p span + span {
color: #607080;
font-family: ui-monospace, SFMono-Regular, Menlo, Monaco, Consolas, "Liberation Mono", monospace;
font-size: 0.72rem;
}
.badge {
display: inline-flex;
align-items: center;
@@ -0,0 +1,22 @@
{
"alpha_yaw": 0.08734091699186919,
"alpha_pitch": 0.08169667696275307,
"alpha_turn": 5.724587470723463e-17,
"beta_fwd": 0.02842768078408099,
"beta_strafe": 0.022531015077067108,
"focal_length": 457.0,
"frame_shape": [
352,
640
],
"calibrated_from": [
"1_wasd_only",
"camera",
"camera4hold_alpha1",
"fully_random",
"wasdonly_alpha1",
"wasd4holdrandview_simple_1key1mouse1"
],
"residual_rms": 15.890399609478676,
"n_equations": 4125232
}
-56
View File
@@ -1,56 +0,0 @@
# 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.
+6 -4
View File
@@ -68,7 +68,7 @@ ARG FLASH_ATTN_WHEEL_RELEASE_ARM64=https://github.com/mjun0812/flash-attention-p
# cutlass-4.4 `cute.core.ThrMma` API, which crashes on the cutlass-dsl 4.5 that
# flashinfer/quack pull in. After the wheel install we overlay this cutlass-4.5-safe
# upstream cute (flash-attn-4) so the image runs FA4 instead of the FA2 fallback.
ARG FA4_CUTE_REF=940cd9680f3315f2f06b43ab5bea2c2cf2d96806
ARG FA4_CUTE_REF=82d6441eec5d4dfec120153db2c0145ae855a083
# Provided automatically by BuildKit/buildx (e.g. "amd64" / "arm64") and used to
# select the prebuilt flash-attn wheel. Empty under a plain `docker build` without
@@ -161,7 +161,7 @@ RUN --mount=type=cache,target=/opt/uv/cache \
uv pip install flash-attn==${FLASH_ATTN_VERSION} --no-build-isolation; \
fi
# Overlay the cutlass-4.5-safe upstream FA4 cute (FA4_CUTE_REF) over the
# Overlay the CuTe-DSL-4.6-compatible upstream FA4 cute (FA4_CUTE_REF) over the
# wheel/source one so the image runs FA4, not the FA2 fallback. This pulls the FA4
# runtime stack (cutlass-dsl, quack-kernels, apache-tvm-ffi, torch-c-dlpack-ext) --
# the same deps the [dreamverse] extra already installs in CI; the installed torch
@@ -170,12 +170,14 @@ RUN --mount=type=cache,target=/opt/uv/cache \
# rmtree clears the wheel's stale cute files first to avoid an install conflict.
# Then verify both survive so a broken overlay fails the build instead of shipping
# an FA2-less image. x86 only: the FA4 stack (quack-kernels etc.) is unvalidated on
# arm64 / GB10 (sm_121), so there we skip the overlay and FA4 falls back to FA2.
# arm64 / GB10 (sm_121), so there we skip the overlay; FA4 is opt-in
# (FASTVIDEO_FA4=1) and errors if set without the overlay, so leave it unset on
# arm64 and the image runs FA3/FA2 as usual.
RUN --mount=type=cache,target=/opt/uv/cache \
source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
if [ "${TARGETARCH}" = "arm64" ]; then \
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; FA4 falls back to FA2)"; \
echo "Skipping FA4 cute overlay on arm64 (FA4 stack unvalidated there; do not set FASTVIDEO_FA4)"; \
else \
python -c "import glob, shutil; [shutil.rmtree(d, ignore_errors=True) for d in glob.glob('/opt/venv/lib/python*/site-packages/flash_attn/cute')]" && \
uv pip install "flash-attn-4 @ git+https://github.com/Dao-AILab/flash-attention.git@${FA4_CUTE_REF}#subdirectory=flash_attn/cute" && \
+9
View File
@@ -99,10 +99,18 @@ status.
Full Suite is also path-filtered. It validates broader behavior before Mergify
can merge a PR.
A `ready`-labeled PR does not hit Buildkite immediately:
`ci-trigger-full-suite.yml` first runs `.github/scripts/gate_full_suite.sh`,
which waits for the cheap Tier-1 checks (pre-commit, docs build) on the PR
head. A red cheap check blocks the suite (fail closed; the next push re-arms
it), while a GitHub outage or a >25 min wait lets it run anyway (fail open).
`/test full` bypasses the gate.
| Buildkite label | `TEST_TYPE` | Main watched paths |
|---|---|---|
| SSIM Tests | `ssim` | `fastvideo/**/*.py`, `pyproject.toml`, `docker/Dockerfile` |
| LoRA Inference Tests | `inference_lora` | LoRA tests, loader, transformer tests, pipelines, LoRA layers |
| LoRA Extraction Tests | `lora_extraction` | LoRA extraction scripts/tests, loader, training utilities, LoRA layers |
| Training Tests | `training` | `fastvideo/**`, `pyproject.toml`, `docker/Dockerfile` |
| Distillation DMD Tests | `distillation_dmd` | `fastvideo/training/*distillation_pipeline.py` |
| Self-Forcing Tests | `self_forcing` | self-forcing distillation pipeline and tests |
@@ -144,6 +152,7 @@ Valid direct test names:
| `/test training` | `training` |
| `/test lora-inference` | `inference_lora` |
| `/test lora-training` | `training_lora` |
| `/test lora-extraction` | `lora_extraction` |
| `/test distillation` | `distillation_dmd` |
| `/test self-forcing` | `self_forcing` |
| `/test vsa` | `training_vsa` |
+234 -46
View File
@@ -72,14 +72,18 @@ fastvideo/tests/performance/
│ writes Markdown summary + (optionally) uploads new records
├── dashboard.py
│ └── builds time-series Plotly HTML from HF history
└── hf_store.py # shared HF I/O + DataFrame helpers
fastvideo/performance/
├── hf_store.py # shared HF I/O + DataFrame helpers
└── metric_policy.py # shared rolling-baseline threshold policy
```
The HF dataset (`FastVideo/performance-tracking` by default) holds one
normalized JSON per `(model_id, gpu_type, run)` tuple. The rolling baseline is
the median of the last 5 successful, baseline-eligible records for that
model+GPU. PR and local records are visible in the dashboard but are not
baseline eligible.
normalized JSON per run. For v2 records, the rolling baseline is the median of
the last 5 successful, baseline-eligible records in the same comparison cohort:
`model_id`, `gpu_type`, `workload_id`, `variant_id`, `benchmark_version`,
`recipe_fingerprint`, `hardware_profile_id`, and `software_profile_id`. PR and
local records are visible in the dashboard but are not baseline eligible.
## Planned Coverage
@@ -92,25 +96,28 @@ and recipe changes instead of treating all records for a model as equivalent.
## Metrics
Each benchmark records six metrics:
Each benchmark records six metrics. The rolling-baseline comparator also has a
per-metric policy with direction, percent threshold, absolute threshold, and a
`gated` flag.
| Metric | Raw key | Normalized key | Direction |
|---|---|---|---|
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better |
| Video throughput | `throughput_fps` | `throughput` | Higher is better |
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better |
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better |
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better |
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better |
| Metric | Raw key | Normalized key | Direction | Default rolling policy |
|---|---|---|---|---|
| End-to-end generation latency | `avg_generation_time_s` | `latency` | Lower is better | 8% and 0.5 s |
| Video throughput | `throughput_fps` | `throughput` | Higher is better | 8% and 0.05 FPS |
| Peak GPU memory | `max_peak_memory_mb` | `memory` | Lower is better | 5% and 256 MB |
| Text encoder time | `text_encoder_time_s` | `text_encoder_time_s` | Lower is better | 5% and 0.25 s |
| DiT denoising time | `dit_time_s` | `dit_time_s` | Lower is better | 5% and 0.25 s |
| VAE decode time | `vae_decode_time_s` | `vae_decode_time_s` | Lower is better | 5% and 0.25 s |
`test_inference_performance.py` temporarily sets `FASTVIDEO_STAGE_LOGGING=1`
while it runs so pipeline stage execution times are available in
`generate_video(...).logging_info`. Stage logs use pipeline-unique keys such as
`prompt_encoding_stage` so duplicate stage classes do not collide. For
`PipelineStage` entries, the extractor maps the `stage_class` field:
`TextEncodingStage` maps to `text_encoder_time_s`, `DenoisingStage` and
`DmdDenoisingStage` map to `dit_time_s`, and `DecodingStage` maps to
`vae_decode_time_s`, with a fallback for older logs that used the class name as
`PipelineStage` entries, shared component stage bases emit a stable
`component_metric`: text encoding stages map to `text_encoder_time_s`,
denoising stages and subclasses map to `dit_time_s`, and decoding stages map to
`vae_decode_time_s`. The extractor falls back to known `stage_class` names for
older logs that do not include `component_metric` or that used the class name as
the stage key. Generator-side timings such as `PostDecodeFrameProcessStage`,
`VideoSaveStage`, and `AudioMuxStage` are intentionally ignored. If a pipeline
does not report one of the mapped stages, that component metric is stored as
@@ -152,13 +159,29 @@ unrealistic memory growth, and optionally large component-specific slowdowns
even when the rolling baseline is empty. They are hand-set with generous
headroom and almost never need touching.
### Rolling baseline (per `(model_id, gpu_type)`)
### Rolling baseline (per comparison cohort)
`compare_baseline.py` loads the last 5 successful, baseline-eligible records
for the same `(model_id, gpu_type)` from the HF dataset, computes the median
for each available metric, and fails if the current run regresses by more than
`PERF_MAX_REGRESSION` (default 5%). For latency, memory, and component times,
higher values are regressions. For throughput, lower values are regressions.
for the same comparison cohort from the HF dataset, computes the median for
each available metric, and evaluates the current run with the metric's
rolling regression policy. For v2 records, that cohort is `model_id`,
`gpu_type`, `workload_id`, `variant_id`, `benchmark_version`,
`recipe_fingerprint`, `hardware_profile_id`, and `software_profile_id`. For
latency, memory, and component times, higher values are regressions. For
throughput, lower values are regressions.
A metric exceeds its rolling threshold when both of these are true:
```text
percent_delta > threshold_percent
absolute_delta > threshold_absolute
```
Gated metrics fail CI when that threshold crossing happens. Set `gated: false`
for metrics that should remain visible in reports and the dashboard without
failing CI. Dashboard/API payloads expose `threshold_exceeded` separately from
`regressed`, where `regressed` means a gated CI failure. Missing or `null`
metrics are skipped.
This is the **drift detector** — it catches sub-threshold regressions that
slowly add up. Only scheduled-main successful records are baseline eligible.
@@ -172,6 +195,51 @@ agent skill to advance the rolling median.
## Schemas
### Benchmark config (`.buildkite/performance-benchmarks/tests/*.json`)
Benchmark configs without `config_schema_version` are treated as legacy v1
configs and remain loadable. New or migrated configs should use
`config_schema_version: 2` and include explicit comparable identity fields:
```jsonc
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"config_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 2
}
```
`benchmark_id` is still required in this phase because raw artifact names,
generated-video directories, normalized record paths, and the current rolling
baseline comparator still depend on it. The v2 identity fields are config
metadata that make the measured workload explicit:
| Field | Purpose |
|---|---|
| `workload_id` | Stable benchmark family, such as `wan-t2v`. |
| `variant_id` | Intentional recipe family, including model size and parallelism config, such as `1.3b-sp2`. |
| `benchmark_version` | Version of the measurement protocol and comparison policy. |
If a config declares `config_schema_version: 2`, loading fails clearly when any
required v2 identity field is missing. If v2 identity or metadata fields are
added without `config_schema_version: 2`, loading also fails so partial
migrations do not silently run as v1 configs. Optional v2 metadata fields
reserved for follow-up work, such as `metric_threshold_policy` and
`quality_metadata`, must be JSON objects when present. (`recipe` is emitted
by the harness and is not config-declarable.)
Recipe fingerprinting, hardware/software profile IDs, exact-identity
comparison, and dashboard cohort grouping land with this change: v2 records
compare only within their identity cohort, and a record that opens a NEW
cohort is marked `baseline_status: "initialized_new_cohort"` (regression
gating starts once that cohort accumulates history). Legacy v1 configs still
run and are normalized for reporting, but their records skip rolling-baseline
comparison entirely (`baseline_status: "skipped_missing_identity"`, never
baseline eligible); only static thresholds gate them. Metric-specific
threshold policies and promoted baselines remain separate follow-ups.
### Raw record (`results/perf_*.json`)
Written by `test_inference_performance.py`. One file per benchmark run.
@@ -179,6 +247,10 @@ Written by `test_inference_performance.py`. One file per benchmark run.
```jsonc
{
"benchmark_id": "wan-t2v-1.3b-2gpu",
"result_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 2,
"model_short_name": "Wan2.1-T2V-1.3B-Diffusers",
"device": "NVIDIA L40S",
"num_gpus": 2,
@@ -196,12 +268,62 @@ Written by `test_inference_performance.py`. One file per benchmark run.
"max_dit_time_s": 10.0,
"max_vae_decode_time_s": 10.0
},
"regression_thresholds": {
"latency": {
"threshold_percent": 0.10,
"threshold_absolute": 1.0,
"gated": true
}
},
"commit": "<full sha>",
"run_source": "pr",
"branch": "feature/perf-change",
"pr_number": "1234",
"test_scope": "direct",
"build_url": "https://buildkite.example/build",
"build_id": "<buildkite-build-id>",
"job_id": "<buildkite-job-id>",
"timestamp": "2026-05-08T22:00:00+00:00",
"quality_metadata": { "quality_status": "canonical" },
"text_encoder_time_s": 2.141,
"dit_time_s": 8.437,
"vae_decode_time_s": 3.208
"vae_decode_time_s": 3.208,
"recipe": {
"recipe_schema_version": 1,
"benchmark": {
"benchmark_id": "wan-t2v-1.3b-2gpu",
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 2
},
"model": { "model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" },
"init_kwargs": { "num_gpus": 2, "sp_size": 2, "tp_size": 1 },
"generation_kwargs": { "height": 480, "width": 832, "num_frames": 45 },
"inputs": { "prompt_count": 1, "prompt_sha256": ["<measured-prompt-sha256>"] },
"attention": { "requested_backend": "FLASH_ATTN", "resolved_backend": "FLASH_ATTN" }
},
"recipe_fingerprint": "<sha256>",
"hardware_profile": {
"device_type": "cuda",
"gpu_count": 2,
"gpus": [{ "name": "NVIDIA L40S", "memory_gb": 48, "compute_capability": "8.9" }],
"interconnect": "none_or_partial"
},
"hardware_profile_id": "hw-<sha256-prefix>",
"software_profile": {
"python": "3.12",
"pytorch": "2.12",
"cuda": "13.0",
"packages": {
"fastvideo_kernel": "0.3.2",
"flashinfer": "0.2.11",
"nvidia_cutlass_dsl": "4.5.0",
"triton": "3.4.1"
}
},
"software_profile_id": "sw-<sha256-prefix>",
"environment_metadata": { "env": { "IMAGE_VERSION": "py3.12-cuda13.0.0" } },
"environment_fingerprint": "env-<sha256-prefix>"
}
```
@@ -213,6 +335,10 @@ result, used as the rolling-baseline source of truth.
```jsonc
{
"model_id": "wan-t2v-1.3b-2gpu",
"result_schema_version": 2,
"workload_id": "wan-t2v",
"variant_id": "1.3b-sp2",
"benchmark_version": 2,
"timestamp": "2026-05-08T22:00:00+00:00",
"commit_sha": "<full sha>",
"gpu_type": "NVIDIA L40S",
@@ -222,34 +348,69 @@ result, used as the rolling-baseline source of truth.
"text_encoder_time_s": 2.141,
"dit_time_s": 8.437,
"vae_decode_time_s": 3.208,
"regression_thresholds": {
"latency": {
"threshold_percent": 0.08,
"threshold_absolute": 0.5,
"gated": true
}
},
"recipe_fingerprint": "<sha256>",
"hardware_profile_id": "hw-<sha256-prefix>",
"software_profile_id": "sw-<sha256-prefix>",
"environment_fingerprint": "env-<sha256-prefix>",
"run_source": "pr",
"branch": "feature/perf-change",
"pr_number": "1234",
"test_scope": "direct",
"build_url": "https://buildkite.example/build",
"build_id": "<buildkite-build-id>",
"job_id": "<buildkite-job-id>",
"quality_metadata": { "quality_status": "canonical" },
"success": true
}
```
### Compatibility with legacy records
Older records in the HF dataset may not have component timing fields. The
comparator ignores missing or `null` metrics when computing a median, and the
dashboard lists skipped plots for metric series that have no non-null values.
Records missing both `run_source` and `baseline_eligible` are treated as legacy
successful main/full-suite uploads and remain eligible for rolling baselines.
Older records in the HF dataset may not have `result_schema_version`,
component timing fields, or v2 identity/profile fields. Records without
`result_schema_version` are treated as v1. The comparator ignores missing or
`null` metrics when computing a median, and the dashboard lists skipped plots
for metric series that have no non-null values. Records missing both
`run_source` and `baseline_eligible` are treated as legacy successful
main/full-suite uploads and remain eligible for rolling baselines.
Current `perf_*.json` artifacts that lack the v2 comparison identity are
normalized for reporting but skip rolling-baseline comparison and are not marked
baseline eligible.
New records compare only against the same `model_id`, `gpu_type`,
`workload_id`, `variant_id`, `benchmark_version`, `recipe_fingerprint`,
`hardware_profile_id`, and `software_profile_id` cohort.
`environment_metadata` and `environment_fingerprint` are audit data and are not
part of the comparison key.
The recipe prompt digests describe the prompts actually measured by the
benchmark run; extra configured prompts are ignored unless the benchmark runner
executes them.
Software profile package cohorts keep exact versions for relevant
attention/kernel packages, including FastVideo kernels, FlashAttention,
FlashInfer, Cutlass DSL, SageAttention, Triton, and xFormers when installed.
## Environment variable reference
| Variable | Default | Used by | Purpose |
|---|---|---|---|
| `PERF_MAX_REGRESSION` | `0.05` | `compare_baseline.py` | Per-metric regression fraction that fails the build. |
| `PERFORMANCE_TRACKING_ROOT` | `/tmp/perf-tracking` | `compare_baseline.py`, `dashboard.py` | Local directory the HF dataset is synced to. |
| `PERF_REPORTS_DIR` | `/root/data/perf_reports` | `compare_baseline.py`, `dashboard.py` | Where the Markdown summary and Plotly HTML get written for Buildkite to pick up. |
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `hf_store.py` | HF dataset repo holding rolling-baseline records. |
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `hf_store.py` | Required for upload or private dataset reads. |
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
| `HF_REPO_ID` | `FastVideo/performance-tracking` | `fastvideo/performance/hf_store.py` | HF dataset repo holding rolling-baseline records. |
| `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, `HF_TOKEN` | unset | `fastvideo/performance/hf_store.py` | Required for upload or private dataset reads. |
| `PERF_RUN_SOURCE` | inferred | `compare_baseline.py`, `test_inference_performance.py` | Source metadata for uploaded records: `pr`, `local`, `scheduled_main`, or `unknown`. |
| `PERF_UPLOAD_POLICY` | `never` | `compare_baseline.py` | Upload policy: `never`, `pass`, or `always`. |
| `PERF_PYTEST_RC` | unset | `compare_baseline.py` | Static-threshold pytest exit code, used so scheduled-main failures can be uploaded with `success=false`. |
| `TEST_SCOPE` | unset | `compare_baseline.py` | CI context used to infer scheduled-main runs together with `BUILDKITE_BRANCH=main`. |
| `BUILDKITE_BRANCH`, `BUILDKITE_COMMIT`, `BUILDKITE_PULL_REQUEST` | unset | `compare_baseline.py`, `test_inference_performance.py` | CI metadata stamped into records. |
| `DASHBOARD_DAYS` | `30` | `dashboard.py` | Lookback window for the Plotly trend pages. |
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
| `PERFORMANCE_TRACKING_SYNC_REUSE_TTL_SECONDS` | `3600` | `fastvideo/performance/hf_store.py` | Freshness window for reusing an existing HF sync when requested by dashboard consumers. |
| `FASTVIDEO_STAGE_LOGGING` | set by the pytest test | `test_inference_performance.py` | Enables pipeline stage timing capture for component metrics during benchmark runs. |
## CI integration
@@ -260,17 +421,22 @@ point is `fastvideo/tests/modal/pr_test.py:run_performance_tests` and the
Buildkite artifact upload is in
`.buildkite/scripts/pr_test.sh:upload_performance_artifacts`.
Each performance build runs pytest first. If that fixed-threshold phase fails,
`compare_baseline.py` is skipped, so Markdown summaries and normalized JSON
artifacts are not emitted. The dashboard still runs best-effort for
observability. When pytest passes, the rolling-baseline phase emits:
Each performance build runs pytest first. PR and direct runs only continue to
`compare_baseline.py` when that fixed-threshold phase passes; if pytest fails,
Markdown summaries and normalized JSON artifacts are not emitted. Scheduled
main runs set `PERF_UPLOAD_POLICY=always`, so they still run
`compare_baseline.py` (with `PERF_PYTEST_RC` set) after a fixed-threshold
failure. Those failed scheduled main runs emit summaries and normalized
records, upload records with `success=false`, and are excluded from future
rolling baselines. The dashboard still runs best-effort for observability.
When the rolling-baseline phase runs, it emits:
* **Markdown summary** — appended to `$GITHUB_STEP_SUMMARY` when that variable
is set, and written as `perf_<sha>_<ts>.md` for Buildkite upload. Contains a
per-benchmark row with current vs. baseline values for latency, throughput,
memory, text encoder time, DiT time, and VAE decode time.
* **Plotly dashboard** — `dashboard_<sha>_<ts>.html` showing time-series for
each metric grouped by `(model_id, gpu_type)`.
each metric grouped by comparison cohort.
* **Normalized records** — `normalized_perf_*.json`, one per benchmark.
Useful as input to the
[`reseed-performance-baseline`](https://github.com/hao-ai-lab/FastVideo/blob/main/.agents/skills/reseed-performance-baseline/SKILL.md)
@@ -279,11 +445,16 @@ observability. When pytest passes, the rolling-baseline phase emits:
## Adding a new benchmark
1. Drop a new JSON config into
`.buildkite/performance-benchmarks/tests/<name>.json`. Required keys:
`.buildkite/performance-benchmarks/tests/<name>.json`. New configs should
use v2 identity fields:
```json
{
"benchmark_id": "<unique-id>",
"config_schema_version": 2,
"workload_id": "<stable-workload-id>",
"variant_id": "<variant, e.g. 1.3b-sp2>",
"benchmark_version": 1,
"model": { "model_path": "...", "model_short_name": "..." },
"init_kwargs": { "num_gpus": 1, ... },
"generation_kwargs": { "num_frames": 45, ... },
@@ -299,9 +470,19 @@ observability. When pytest passes, the rolling-baseline phase emits:
"max_vae_decode_time_s": 10.0
},
"default": { "max_generation_time_s": 120.0, "max_peak_memory_mb": 30000.0 }
},
"regression_thresholds": {
"latency": { "threshold_percent": 0.10, "threshold_absolute": 1.0, "gated": true }
}
}
```
}
```
Legacy v1 configs without `config_schema_version` still load, but should not
gain v2 identity or metadata fields until they are migrated to
`config_schema_version: 2`. For v2 configs, `workload_id`, `variant_id`,
and `benchmark_version` are part of the comparison key; benchmark runs
fail if any of these identity fields are missing.
2. The pytest test auto-discovers all configs — no test code needed. CI
picks it up on the next `/test performance` run.
@@ -320,10 +501,17 @@ observability. When pytest passes, the rolling-baseline phase emits:
a useful fixed gate. The rolling baseline will still track component times
when static component thresholds are omitted.
6. Omit `regression_thresholds` to use the default rolling-baseline policy, or
include only benchmark-specific deviations. Tune these independently from
the fixed thresholds when a metric is noisy or should be informational. The
fixed `thresholds` block is an absolute pytest ceiling. The
`regression_thresholds` block controls rolling-baseline comparisons against
recent scheduled-main records.
## Troubleshooting
**"No baseline for ... Initializing"** — first run for this `(model_id,
gpu_type)`. Run will pass and (if persisting) seed the first record.
**"No baseline for ... Initializing"** — first run for this comparison cohort.
Run will pass and (if persisting) seed the first record.
**Persistent failure right after a torch / kernel / image upgrade** —
genuine regression *or* baseline drift. Compare the failing normalized record
@@ -336,5 +524,5 @@ pipelines that did not report a mapped component stage.
**Component timing is `null`** — the generated result did not include a mapped
stage in `logging_info.stages`. Check that the pipeline emits stage logging
and that the stage name is listed in `STAGE_METRIC_MAP` in
`test_inference_performance.py`.
and that the stage emits `component_metric` or is covered by the legacy
`STAGE_METRIC_MAP` fallback in `test_inference_performance.py`.
+24
View File
@@ -180,6 +180,30 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
--device-folder L40S_reference_videos
```
### SSIM Bootstrap Mode
Normal SSIM runs are strict: if a reference video or latent is missing, the
test fails. For new-model PRs, CI can run SSIM in bootstrap mode so missing
references are uploaded as draft artifacts for review instead of immediately
blocking on a missing canonical reference.
Buildkite enables SSIM bootstrap mode when either condition is true:
- the PR title or Buildkite message contains `[new-model]`;
- `FASTVIDEO_SSIM_BOOTSTRAP_MODE=1` is set for the Buildkite job.
Bootstrap mode passes `--ssim-bootstrap-mode` to pytest. When a generated
artifact is available, the test uploads it under the `drafts/...` namespace in
the SSIM reference repo and marks that case as expected-failed. After reviewing
the draft, promote it into the canonical reference layout:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py promote-draft \
--quality-tier default \
--device-folder L40S_reference_videos \
--model-id <model_id>
```
## CI Integration
FastVideo CI tests are orchestrated by Buildkite and run on Modal GPU
@@ -191,6 +191,9 @@ surfaces:
sources:
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
color_correction_strength:
sources:
- fastvideo.configs.pipelines.dreamx_world.DreamXWorld5BARPipelineConfig
default_camera_rotation:
sources:
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
@@ -455,6 +458,8 @@ surfaces:
guidance_scale: request.sampling.guidance_scale
guidance_scale_2: request.sampling.guidance_scale_2
guidance_rescale: request.sampling.guidance_rescale
use_embedded_guidance: request.sampling.use_embedded_guidance
true_cfg_scale: request.sampling.true_cfg_scale
boundary_ratio: request.sampling.boundary_ratio
sigmas: request.sampling.sigmas
enable_teacache: request.runtime.enable_teacache
+116
View File
@@ -0,0 +1,116 @@
# 🌊 AnyFlow Any-Step Video Distillation
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
```
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
```
for any ``t > r``.
## 📊 Model Overview
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
## ⚙️ Inference
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
## 🧠 Algorithm
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
### Stage 1 — Pretrain (flow-map central-difference)
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
- the remainder is free.
The student forward at ``(t, r)`` is trained against the central-difference target
```
target = (eps - x_0) - (t - r) · dF/dt
```
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
### Stage 2 — On-policy DMD
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
## 🚀 Training Scripts
### Stage 1 — pretrain
```bash
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
```
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
- Global batch size: 32 (8 GPUs × 4 per-GPU)
- Learning rate: 5e-5
- Flow shift: 5.0
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
- ``weight_type``: ``beta08``
- ``fuse_guidance_scale``: 3.0
- Training steps: 6000
### Stage 2 — on-policy
```bash
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
```
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
**Key configuration**:
- Global batch size: 8 (8 GPUs × 1 per-GPU)
- Learning rate: 2e-6
- Flow shift: 5.0
- ``student_sample_steps``: 4
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
- ``real_score_guidance_scale``: 3.0
- ``generator_update_interval``: 5 (DMD2 alternation)
- Training steps: 4000
## 🔌 Loading published AnyFlow checkpoints
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
## 🧭 Note on ``fuse_guidance_scale``
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
```
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
```
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
+17
View File
@@ -74,6 +74,23 @@ uv pip install ninja
python setup.py install
```
### Flash Attention 4 (opt-in)
FastVideo never auto-selects FlashAttention-4 (`flash_attn.cute`) just because it
is installed: its CuTeDSL kernels JIT-compile per shape family and can fail at
runtime on some GPU/shape combinations. To use FA4, install the pinned
`flash-attn-4` build (see the `flash-attn-4` source in `pyproject.toml`) and set:
```bash
export FASTVIDEO_FA4=1
```
On GPUs below sm90 a capability gate routes to FlashAttention-2 the calls FA4
cannot serve there: grad-enabled (training) attention (FA4's backward requires
sm90+) and GQA attention (FA4's `pack_gqa` fails to JIT-compile below sm90).
On sm90+ both run on FA4. If FA4 is unusable while `FASTVIDEO_FA4=1` is set,
FastVideo fails loudly instead of silently falling back.
### FP4 Flash Attention 4 (Blackwell only)
**`FLASH_ATTN`** with **`--nvfp4_fa4`**
+2
View File
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
+91
View File
@@ -0,0 +1,91 @@
# Training Trackers
FastVideo can send training metrics and validation media to Weights & Biases
or SwanLab. Tracking runs only on global rank 0, and local tracker files are
stored under `<output_dir>/tracker`.
## Supported Trackers
| Value | Backend | Installation |
|-------|---------|--------------|
| `wandb` | Weights & Biases | Included with FastVideo |
| `swanlab` | SwanLab | Install the optional `swanlab` dependency |
| `none` | Disable external tracking | No additional package |
You can enable more than one backend, for example `trackers: [wandb, swanlab]`.
Metrics and validation media are converted to the artifact type required by
each backend.
## Install SwanLab
For a published FastVideo installation, install the SwanLab extra:
```bash
uv pip install "fastvideo[swanlab]"
```
For an editable source checkout, include the same extra during installation:
```bash
uv pip install -e ".[swanlab]"
```
If FastVideo is already installed, you can install the compatible SDK directly:
```bash
uv pip install "swanlab>=0.6.7"
```
Authenticate once before starting a training run:
```bash
swanlab login
```
See the [SwanLab login documentation](https://docs.swanlab.cn/en/api/cli-swanlab-login.html)
for non-interactive and self-hosted setups.
## Configure Tracking
Select SwanLab in the YAML config used by the modular training framework:
```yaml
training:
checkpoint:
output_dir: outputs/my_run
tracker:
trackers: [swanlab]
project_name: my_project
run_name: my_run
```
To log to both supported services:
```yaml
training:
tracker:
trackers: [wandb, swanlab]
project_name: my_project
run_name: my_run
```
An empty or omitted `trackers` list selects W&B when `project_name` is set.
Use an explicit `none` entry to disable external tracking:
```yaml
training:
tracker:
trackers: [none]
```
## Validation Videos
SwanLab currently accepts GIF video artifacts. FastVideo converts validation
MP4 files and in-memory video arrays to GIF automatically before logging them.
For video files, FastVideo uses the sampling frame rate supplied by the caller,
or the source file's frame rate when no value is supplied. In-memory arrays use
the frame rate supplied by the caller. Both forms fall back to 16 FPS when no
frame rate is available.
For details about configuring validation callbacks, see
[Training Infrastructure](train_infra.md#callbacks-pluggable-hooks).
+49
View File
@@ -161,6 +161,21 @@ training:
decay_interval_steps: 0
```
`training.data.data_path` can also mix multiple preprocessed datasets by using a mapping from dataset path to repeat count:
```yaml
training:
data:
data_path:
data/zeldam2-clean: 1
data/multi3d_games: 2
```
The repeat count duplicates that dataset's parquet file list before shuffling/sampling, so the example above trains with roughly twice as much `multi3d_games` exposure as `zeldam2-clean`. Paths are just suggested locations; use any local path that contains a FastVideo preprocessed parquet dataset.
See [Training Trackers](trackers.md) to configure Weights & Biases or SwanLab,
including SwanLab installation and authentication.
### `callbacks` — Pluggable hooks
Callbacks run at specific points in the training loop (before/after optimizer
@@ -323,6 +338,40 @@ Self-Forcing inherits all DMD2 parameters, plus:
| `enable_gradient_in_rollout` | `true` | Enable backprop through rollout |
| `start_gradient_frame` | `0` | Frame index where gradients begin |
### Streaming Long Tuning
`StreamingLongTuningMethod` extends Self-Forcing for LongLive-style rollouts. It
keeps a streaming state, generates overlapping chunks, and trains only the new
frames while preserving context from earlier chunks.
For the MatrixGame2/Zelda world-model example, self-forcing and long tuning are
separate runs: first train or load the 1k-step self-forcing checkpoint using
`examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml`,
then run
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
from that checkpoint for the 3k-step streaming long-tuning stage.
```yaml
method:
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
streaming_chunk_size: 9
streaming_max_length: 39
streaming_fixed_overlap_latents: 3
streaming_reencode_overlap_anchor: true
streaming_anchor_inject_k: 1
streaming_require_full_blocks: true
multi_phased_distill_schedule:
- stage: streaming_long
start_step: 0
end_step: 3000
num_latent_t: 39
streaming_training: true
```
See
`examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml`
for a complete MatrixGame2/Zelda configuration.
---
## Callbacks
@@ -1,94 +0,0 @@
# v2 porting status — fastvideo models → the v2 (recipe, runtime) substrate
Goal: every model in fastvideo's registry resolves through the **v2 `VideoGenerator`** / `Engine`
(typed `fastvideo.api` configs + the real torch backend) to a recipe that can construct and run it.
**Scope: ALL fastvideo models (achieved).** v2 now resolves **63/64** of fastvideo's registered HF ids
by exact id (PRIMARY), plus the architecture fallback for local/unregistered checkpoints. The single
remaining id — `FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers` — is **environment-blocked**: its VSA
(Sparse-Linear Attention) kernels require `nvcc` (not built in this bring-up). It arch-resolves to the
base Wan card but needs the VSA kernel build to run faithfully.
Dispatch is **architecture-driven** (`v2/registry.py`): exact HF id → short-name → architecture
inference from the checkpoint (pipeline / transformer / VAE class names + `z_dim`, `transformer_2`,
`spatial_upsampler`). Adding a model is one `_BUCKET_C` row (HF ids → builders + transformer class).
## The porting mechanism — self-contained recipe packages
Every net-new arch is a **self-contained recipe package** (`v2/recipes/<arch>/` = `card.py` `loop.py`
`program.py` [+ `sampler.py`] + an optional `v2/platform/backends/torch_<arch>.py` adapter). The card
declares its torch adapter via **`ComponentSpec.adapter="module:Class"`** (the `_explicit_adapter` seam in
`torch_backend.py`) instead of editing the shared `_make_dit`/`_make_vae`/`_make_text_encoder` dispatch —
so a port adds **only new files**, never touching shared code, and parallel ports never conflict. New
samplers/loops live in-package. Registration is one row in `v2/registry.py:_BUCKET_C`.
## Working today (GPU-verified, real video/audio) — committed on `v2`
| Official example(s) | Model | v2 card |
|---|---|---|
| `basic.py`, `basic_mps.py`, `basic_ray.py` | `Wan-AI/Wan2.1-T2V-1.3B-Diffusers` | wan21 |
| `basic_self_forcing_causal.py` | `wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers` | wan_causal |
| `basic_ltx2_distilled.py` | `FastVideo/LTX2-Distilled-Diffusers` (2-stage + spatial upsampler) | ltx2 |
| `basic_wan2_2_ti2v.py` | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | wan2.2-ti2v |
| `basic_wan2_2.py` | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` (MoE + CPU expert offload) | wan2.2-a14b |
| `basic_ltx2.py` | `Davids048/LTX2-Base-Diffusers` | ltx2 base |
| `basic_ltx2_3_distilled.py` | `FastVideo/LTX-2.3-Distilled-Diffusers` (joint T2VS, video+audio) | ltx2.3-distilled |
Plus the **Wan2.1 i2v cluster** (Fun-1.3B-InP GPU-verified; I2V-14B-480P/720P + Wan2.2-I2V-A14B MoE reuse
the i2v card) — CLIP image-encoder + first-frame `[mask|cond]` → 36ch DiT.
## GPU bring-up results (real weights on H100 NVL, single-GPU, TORCH_SDPA)
**20 models generate real video/audio on GPU** — the 7 above + **13 of the newly-ported** archs, each run
end-to-end through the real `VideoGenerator` (resolve → stamp → CUDA load → generate). The rest are blocked
by a **fastvideo-shared-code / missing-kernel / HF-access** wall, NOT a v2 recipe bug (the v2 recipes are
faithful — e.g. cosmos25's DiT+VAE produced finite output; only its Qwen2.5-VL encoder hit a library
incompat). All ports also resolve + run end-to-end on the CPU toy backend (`test_bucket_c_ports.py`).
| GPU status | Models |
|---|---|
| ✅ **Verified** (real GPU output) | stable_audio (audio), matrixgame2, matrixgame3, gen3c, wan_fun_control, lucy_edit, hunyuangamecraft, hunyuan_video, hunyuan_video15, longcat (13.58B), sfwan22 (2×14B MoE, expert offload), lingbotworld (2×14B, offload), fastwan (TI2V-5B-FullAttn DMD) |
| 🚫 fastvideo/env-blocked | **cosmos25** (DiT+VAE ran; Qwen2.5-VL encoder → transformers 5.12.1 incompat in fastvideo); **kandinsky5** (fastvideo registry registers a bare `PipelineConfig`); **hyworld** (fastvideo DiT hardcodes `flash_attn`, not built); **turbowan** 1.3B/i2v + **fastwan** VSA-variants (SLA/VSA sparse-attn params + Triton kernels need nvcc) |
| 🚫 access-blocked (HF-gated) | cosmos2, flux2, sd35 (no HF token in this env) |
To unblock the env-blocked: build `fastvideo-kernel` (SLA/VSA Triton, needs nvcc); pin a fastvideo-compatible
`transformers` for the Qwen2.5-VL encoder; add a Kandinsky5 `PipelineConfig` + an SDPA fallback in the
hyworld DiT (all fastvideo-side / environment, not v2 recipe work).
## Newly ported (recipe details)
Each resolves through the registry AND runs end-to-end on the CPU toy backend via the public `Engine`
path (the `v2/tests/test_bucket_c_ports.py` regression guard), emitting the correct modality artifact.
**15 net-new architectures** (each a new `TorchComponent` adapter + recipe):
- **cosmos2** (Cosmos-Predict2-2B-Video2World) — EDM-Karras denoiser; new `CosmosDenoiseLoop` +
`build_karras_sigmas` (the reference port). **cosmos25** (Cosmos-Predict2.5 2B/14B) — flow-match,
per-frame plain-sigma timestep, Reason1/Qwen2.5-VL encoder. **gen3c** (GEN3C) — EDM + 82ch pose-buffer.
- **hunyuan_video** (+FastHunyuan) — reuses WanDenoiseLoop, dual LLaMA+CLIP encoders, Hunyuan VAE.
**hunyuan_video15** (480p/720p). **hunyuangamecraft**, **hyworld** — interactive (camera/action).
- **longcat** (T2V/I2V/VC). **kandinsky5** (5.0 T2V Lite).
- **sd35** (MMDiT, image, triple-encoder). **flux2** (dev/klein, MMDiT image). **stable_audio** (audio).
- **lingbotworld** (camera/Plucker), **matrixgame2**, **matrixgame3** — interactive world models.
**5 Wan-family variants** (reuse the Wan/Causal arch, new in-package sampler/loop/conditioning):
- **turbowan** — rCM few-step (faithful RCMScheduler port), 1.3B/14B T2V + I2V-A14B MoE.
- **lucy_edit** — v2v editor (video-VAE-encode node → 96ch DiT input). **wan_fun_control** — control input.
- **sfwan22** — Self-Forcing Wan2.2-A14B causal + MoE (i2v + t2v). **fastwan** — DMD 3-step (TI2V-5B-FullAttn
loadable; VSA-trained variants + non-strict `to_gate_compress` load are BRINGUP).
BRINGUP scope per port (documented in each package): GPU load/run; for interactive/world-model archs the
action/camera/memory conditioning needs a request-API extension (the t2v/degenerate path is what
CPU-verifies); video2world/i2v frame-replace conditioning is threaded but inert without conditioning inputs.
## Environment
v2 bring-up runs **single-GPU, resident, on the `TORCH_SDPA` backend** (no fastvideo-kernel / VSA / FP4).
The box has been rescheduled across hosts/arches/python versions mid-session; rebuild the venv for the
current arch when that happens: `uv venv --python 3.12 .venv`; comment out `fastvideo-kernel` in
`pyproject.toml`; `uv pip install -e ".[dev]"`. Source `/home/scratch.willlin_ent/.bringup_env`
(`HF_HOME=./.cache` on scratch, `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`). v2 CPU mini: 240 passed, 2 skipped.
## How to add a model to the v2 substrate
1. `v2/recipes/<arch>/` — card (declare adapters via `ComponentSpec.adapter`; per-model `SamplingDefaults`),
loop (reuse `WanDenoiseLoop`/`chunk_rollout` or a new in-package loop+sampler), program.
2. `v2/platform/backends/torch_<arch>.py` — a `TorchComponent` subclass (only the forward semantics) if the
arch is genuinely new; reuse `WanDiT`/`LTX2DiT`/`WanVAE`/`T5Encoder` via `load_id` when it isn't.
3. One row in `v2/registry.py:_BUCKET_C` (HF ids → builders; `transformer_cls` for the arch fallback, or
`""` for explicit-id-only capability variants of an existing arch).
4. CPU-verify: it resolves + runs on the toy backend (auto-covered by `test_bucket_c_ports.py`). Then GPU
bring-up (`stamp_*_checkpoints` → real weights) per BRINGUP notes.
@@ -0,0 +1,81 @@
import os
import time
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OffloadConfig,
OutputConfig,
SamplingConfig,
)
# NVIDIA Cosmos3-Nano omni world model — image-to-video (I2V) path through
# FastVideo's native Cosmos3 pipeline. The input image conditions latent frame 0
# (kept clean during denoising); the rest of the clip is generated to follow it.
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
# ``official_weights/cosmos3``) to skip the Hugging Face download.
OUTPUT_PATH = "video_samples_cosmos3_i2v"
def main():
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
image_path = os.environ.get("COSMOS3_IMAGE_PATH", "assets/images/cyclist.jpg")
generator_config = GeneratorConfig(
model_path=model_name,
engine=EngineConfig(
num_gpus=1,
use_fsdp_inference=False,
offload=OffloadConfig(
text_encoder=True,
pin_cpu_memory=True,
dit=False,
vae=False,
),
),
)
load_start_time = time.perf_counter()
generator = VideoGenerator.from_config(generator_config)
load_time = time.perf_counter() - load_start_time
prompt = (
"A mountain biker rides forward along the sunlit forest trail, wheels "
"kicking up dust as trees and dappled light sweep past, smooth cinematic "
"tracking shot from behind."
)
request = GenerationRequest(
prompt=prompt,
inputs=InputConfig(image_path=image_path),
sampling=SamplingConfig(
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
# overridable via env for quick smoke runs.
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
guidance_scale=6.0,
fps=24,
seed=1024,
),
output=OutputConfig(
output_path=OUTPUT_PATH,
save_video=True,
return_frames=False,
),
)
start_time = time.perf_counter()
result = generator.generate(request)
gen_time = time.perf_counter() - start_time
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate video: {gen_time} seconds")
print(f"Output written to: {result.video_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,77 @@
import os
import time
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
SamplingConfig,
)
# NVIDIA Cosmos3-Nano omni world model — this example exercises the
# text-to-video (T2V) path through FastVideo's native Cosmos3 pipeline.
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
# ``official_weights/cosmos3``) to skip the Hugging Face download.
OUTPUT_PATH = "video_samples_cosmos3"
def main():
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
generator_config = GeneratorConfig(
model_path=model_name,
engine=EngineConfig(
num_gpus=1,
use_fsdp_inference=False,
offload=OffloadConfig(
text_encoder=True,
pin_cpu_memory=True,
dit=False,
vae=False,
),
),
)
load_start_time = time.perf_counter()
generator = VideoGenerator.from_config(generator_config)
load_time = time.perf_counter() - load_start_time
prompt = (
"A golden retriever puppy runs across a sunlit meadow toward the camera, "
"ears flopping and wildflowers swaying in the breeze. Shallow depth of "
"field, warm afternoon light, smooth cinematic tracking shot."
)
request = GenerationRequest(
prompt=prompt,
sampling=SamplingConfig(
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
# overridable via env for quick smoke runs.
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
guidance_scale=6.0,
fps=24,
seed=1024,
),
output=OutputConfig(
output_path=OUTPUT_PATH,
save_video=True,
return_frames=False,
),
)
start_time = time.perf_counter()
result = generator.generate(request)
gen_time = time.perf_counter() - start_time
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate video: {gen_time} seconds")
print(f"Output written to: {result.video_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,78 @@
import os
import time
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
SamplingConfig,
)
# NVIDIA Cosmos3-Nano omni world model — text-to-image (T2I) path through
# FastVideo's native Cosmos3 pipeline. T2I is the single-frame case
# (num_frames=1); the canonical Cosmos3 T2I resolution is 960x960 (the model's
# "720" bucket, UniPC flow_shift=10.0). Point COSMOS3_MODEL_PATH at a local
# diffusers checkpoint (e.g. ``official_weights/cosmos3``) to skip the download.
OUTPUT_PATH = "video_samples_cosmos3_t2i"
def main():
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
generator_config = GeneratorConfig(
model_path=model_name,
engine=EngineConfig(
num_gpus=1,
use_fsdp_inference=False,
offload=OffloadConfig(
text_encoder=True,
pin_cpu_memory=True,
dit=False,
vae=False,
),
),
)
load_start_time = time.perf_counter()
generator = VideoGenerator.from_config(generator_config)
load_time = time.perf_counter() - load_start_time
prompt = (
"A photograph of a red panda sitting on a mossy log in a misty bamboo "
"forest, soft golden morning light filtering through the leaves, shallow "
"depth of field, crisp fur detail, serene atmosphere."
)
request = GenerationRequest(
prompt=prompt,
sampling=SamplingConfig(
# T2I is single-frame; canonical Cosmos3 T2I is 960x960. Overridable
# via env for quick smoke runs.
num_frames=1,
height=int(os.environ.get("COSMOS3_HEIGHT", "960")),
width=int(os.environ.get("COSMOS3_WIDTH", "960")),
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
guidance_scale=6.0,
fps=24,
seed=1024,
),
output=OutputConfig(
output_path=OUTPUT_PATH,
save_video=True,
return_frames=False,
),
)
start_time = time.perf_counter()
result = generator.generate(request)
gen_time = time.perf_counter() - start_time
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate image: {gen_time} seconds")
print(f"Output written to: {result.video_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,67 @@
import os
import time
# t2vs (text -> video + sound). The Cosmos3 denoise stage generates a joint
# [vision | sound] latent and AVAE-decodes the sound to a waveform muxed into the
# mp4. The joint-sound path is gated on COSMOS3_T2VS (set here for the example).
os.environ.setdefault("COSMOS3_T2VS", "1")
from fastvideo import VideoGenerator # noqa: E402
from fastvideo.api import ( # noqa: E402
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
SamplingConfig,
)
OUTPUT_PATH = "video_samples_cosmos3_t2vs"
def main():
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
generator_config = GeneratorConfig(
model_path=model_name,
engine=EngineConfig(
num_gpus=1,
use_fsdp_inference=False,
offload=OffloadConfig(text_encoder=True, pin_cpu_memory=True, dit=False, vae=False),
),
)
load_start = time.perf_counter()
generator = VideoGenerator.from_config(generator_config)
load_time = time.perf_counter() - load_start
prompt = (
"Ocean waves crash against a rocky shore at sunset, white foam spraying "
"into the air as seagulls wheel overhead. Golden light, cinematic wide "
"shot, the rhythmic roar of the surf."
)
request = GenerationRequest(
prompt=prompt,
sampling=SamplingConfig(
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
guidance_scale=6.0,
fps=24,
seed=1024,
),
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True, return_frames=False),
)
start = time.perf_counter()
result = generator.generate(request)
gen_time = time.perf_counter() - start
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate video+sound: {gen_time} seconds")
print(f"Output written to: {result.video_path}")
if __name__ == "__main__":
main()
@@ -0,0 +1,64 @@
import os
from fastvideo import VideoGenerator
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
def _env_int(name: str, default: int) -> int:
return int(os.getenv(name, str(default)))
def _env_float(name: str, default: float) -> float:
return float(os.getenv(name, str(default)))
def main():
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
generator = VideoGenerator.from_pretrained(
model_name,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
override_pipeline_cls_name="DreamXWorldPipeline",
)
prompt = os.getenv(
"DREAMX_WORLD_PROMPT",
"A cinematic first-person drive through a futuristic coastal city at "
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
)
image_path = os.getenv(
"DREAMX_WORLD_IMAGE_PATH",
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
)
kwargs = {
"output_path": OUTPUT_PATH,
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
"action_speed_list": [
float(value)
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
],
}
if image_path:
kwargs["image_path"] = image_path
try:
generator.generate_video(prompt, **kwargs)
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+140
View File
@@ -0,0 +1,140 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import contextlib
import os
import re
DEFAULT_PROMPTS = [
"a photo of a cat",
(
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"
),
]
def _safe_filename(text: str, max_len: int = 100) -> str:
"""Make a stable, filesystem-friendly filename base."""
s = text[:max_len].strip()
s = s.replace(os.sep, "_")
if os.altsep:
s = s.replace(os.altsep, "_")
s = re.sub(r"\s+", " ", s)
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
s = s.strip(" .")
return s or "prompt"
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
if not os.path.isdir(out_dir):
return
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
for fn in os.listdir(out_dir):
if pattern.match(fn):
with contextlib.suppress(FileNotFoundError):
os.remove(os.path.join(out_dir, fn))
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
)
p.add_argument(
"--model-path",
default="official_weights/FLUX.1-dev",
help="Local Diffusers checkpoint dir or HF repo id.",
)
p.add_argument(
"--out-dir",
"--outdir",
default="outputs/flux_dev/samples",
help="Directory for saved PNG outputs.",
)
p.add_argument(
"--prompt",
action="append",
default=None,
help="Prompt. Repeat for multiple images.",
)
p.add_argument(
"--backend",
default=None,
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
)
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
p.add_argument("--height", type=int, default=1024, help="Output height.")
p.add_argument("--width", type=int, default=1024, help="Output width.")
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
return p.parse_args()
def main() -> None:
args = parse_args()
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
if args.backend:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
from fastvideo import VideoGenerator
os.makedirs(args.out_dir, exist_ok=True)
init_kwargs = {
"num_gpus": args.num_gpus,
"workload_type": "t2i",
"sp_size": 1,
"tp_size": 1,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": False,
"vae_cpu_offload": False,
"image_encoder_cpu_offload": False,
"pin_cpu_memory": False,
"use_fsdp_inference": False,
}
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
**init_kwargs,
)
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = (
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
)
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
generation_kwargs = {
"output_path": output_path,
"height": args.height,
"width": args.width,
"num_frames": 1,
"fps": 1,
"num_inference_steps": args.steps,
"guidance_scale": args.guidance,
"use_embedded_guidance": True,
"true_cfg_scale": 1.0,
"seed": seed,
"save_video": True,
}
generator.generate_video(prompt, **generation_kwargs)
print(f"[flux] done. outputs written to: {args.out_dir}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+107
View File
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""Run GLM-Image text-to-image generation through FastVideo.
User story:
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
text-to-image generation command, saved as a PNG."
"""
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
parser.add_argument(
"--model-path",
default="zai-org/GLM-Image",
help="HF id or local diffusers-format GLM-Image weights directory.",
)
parser.add_argument(
"--output",
default="image_output/landscape.png",
help="Output PNG path.",
)
parser.add_argument(
"--prompt",
default=("A beautiful landscape photography with rolling hills, "
"a winding river, and a vibrant sunset in the background. "
"Warm golden light, photorealistic style."),
help="Text prompt.",
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--guidance-scale", type=float, default=1.5)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
# pipeline class come from the model's registered defaults — don't override.
generator_config = GeneratorConfig(
model_path=args.model_path,
trust_remote_code=True,
engine=EngineConfig(
num_gpus=args.num_gpus,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
save_video=False,
return_frames=True,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
frames = result.frames
if frames is not None and len(frames):
Image.fromarray(frames[0]).save(output)
print(f"Saved image to {output}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,37 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_i2v"
IMAGE_PATH = "assets/girl.png"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-I2V-Pro-distilled-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-I2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-I2V-Lite-5s-Diffusers"
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = (
"A woman stands up and walks away"
)
_ = generator.generate_video(
prompt,
image_path=IMAGE_PATH,
output_path=OUTPUT_PATH,
save_video=True,
height=1024,
width=1024,
num_frames=121,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,37 @@
from fastvideo import VideoGenerator
OUTPUT_PATH = "video_samples_kandinsky5_t2v"
def main():
generator = VideoGenerator.from_pretrained(
"kandinskylab/Kandinsky-5.0-T2V-Lite-sft-5s-Diffusers",
# "kandinskylab/Kandinsky-5.0-T2V-Pro-sft-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-T2V-Lite-distilled16steps-5s-Diffusers"
# "kandinskylab/Kandinsky-5.0-T2V-Pro-distilled-5s-Diffusers"
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True,height=512, width=768, num_frames=121)
prompt2 = (
"A majestic lion strides across the golden savanna, its powerful frame "
"glistening under the warm afternoon sun. The tall grass ripples gently in "
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
"cinematic.")
_ = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True, height=512, width=768, num_frames=121)
if __name__ == "__main__":
main()
+120
View File
@@ -0,0 +1,120 @@
# SPDX-License-Identifier: Apache-2.0
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
User story:
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
want a minimal edit command (text + image -> edited image), saved as a PNG."
GLM-Image is a single unified pipeline: passing a condition image switches it
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
write pass), so the generator config is identical to `basic_glm_image.py` — the
`inputs.pil_image` on the request is what selects the edit mode.
"""
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
parser.add_argument(
"--model-path",
default="zai-org/GLM-Image",
help="HF id or local diffusers-format GLM-Image weights directory.",
)
parser.add_argument(
"--image",
default="assets/images/couple.jpg",
help="Condition image to edit.",
)
parser.add_argument(
"--output",
default="image_output/edited.png",
help="Output PNG path.",
)
parser.add_argument(
"--prompt",
default="Change the background to a snowy mountain landscape at golden hour.",
help="Edit instruction.",
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--guidance-scale", type=float, default=1.5)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
condition = Image.open(args.image).convert("RGB")
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
# pipeline class come from the model's registered defaults — don't override.
# The pipeline is registered as t2i; passing inputs.pil_image below switches
# it to the edit path.
generator_config = GeneratorConfig(
model_path=args.model_path,
trust_remote_code=True,
engine=EngineConfig(
num_gpus=args.num_gpus,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
inputs=InputConfig(pil_image=condition),
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
save_video=False,
return_frames=True,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
frames = result.frames
if frames is not None and len(frames):
Image.fromarray(frames[0]).save(output)
print(f"Saved image to {output}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
-36
View File
@@ -1,36 +0,0 @@
"""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
@@ -1,29 +0,0 @@
"""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()
@@ -1,36 +0,0 @@
"""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()
@@ -1,87 +0,0 @@
"""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()
@@ -1,30 +0,0 @@
"""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()
@@ -1,40 +0,0 @@
"""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()
@@ -1,38 +0,0 @@
"""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()
+69
View File
@@ -0,0 +1,69 @@
# LTX-2.3 distilled inference configs
Ready-to-run `fastvideo generate` run configs for the LTX-2.3
distilled model (`FastVideo/LTX-2.3-Distilled-Diffusers`), covering both
workloads (t2v / i2v), both two-stage step schedules (`5+2`, `8+3` = denoise
+ refine), and four resolutions.
```bash
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
```
Each config is self-contained (no preset registry needed): the two-stage
refine is wired via `generator.pipeline.preset_overrides.refine`, and the
base sampling knobs live under `request.sampling`. The refine upsampler
auto-resolves from the model's `spatial_upscaler`.
## Configs
| workload | schedule | resolution (HxW) | file |
|---|---|---|---|
| t2v | 5+2 | 1280x832 | `t2v_5s2_1280x832.yaml` |
| t2v | 5+2 | 1024x1536 | `t2v_5s2_1024x1536.yaml` |
| t2v | 5+2 | 768x1280 | `t2v_5s2_768x1280.yaml` |
| t2v | 5+2 | 512x768 | `t2v_5s2_512x768.yaml` |
| t2v | 8+3 | 1280x832 | `t2v_8s3_1280x832.yaml` |
| t2v | 8+3 | 1024x1536 | `t2v_8s3_1024x1536.yaml` |
| t2v | 8+3 | 768x1280 | `t2v_8s3_768x1280.yaml` |
| t2v | 8+3 | 512x768 | `t2v_8s3_512x768.yaml` |
| i2v | 5+2 | 1280x832 | `i2v_5s2_1280x832.yaml` |
| i2v | 5+2 | 1024x1536 | `i2v_5s2_1024x1536.yaml` |
| i2v | 5+2 | 768x1280 | `i2v_5s2_768x1280.yaml` |
| i2v | 5+2 | 512x768 | `i2v_5s2_512x768.yaml` |
| i2v | 8+3 | 1280x832 | `i2v_8s3_1280x832.yaml` |
| i2v | 8+3 | 1024x1536 | `i2v_8s3_1024x1536.yaml` |
| i2v | 8+3 | 768x1280 | `i2v_8s3_768x1280.yaml` |
| i2v | 8+3 | 512x768 | `i2v_8s3_512x768.yaml` |
## Overriding without editing a file
Dotted overrides (prefixes `generator.` / `request.`) let you tweak any field:
```bash
# swap prompt
fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml \
--request.prompt "a red fox running through fresh snow"
# change output path / gpu count
fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml \
--request.output.output_path outputs/preview.mp4 \
--generator.engine.num_gpus 4
```
## i2v
The `i2v_*` configs take a first-frame image via
`request.extensions.ltx2_images` (`[[path, frame_offset, weight]]`). Edit the
path in the file, or override it:
```bash
fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml \
--request.extensions.ltx2_images '[["/data/portrait.jpg", 0, 1.0]]'
```
## Schedules
`5+2` is the fast preview schedule; `8+3` is the higher-quality distilled
recipe. Refine (`preset_overrides.refine.num_inference_steps`) only accepts 2
or 3 steps.
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1024x1536.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_1024x1536.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_1280x832.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_1280x832.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_512x768.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_512x768.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 5+2 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_5s2_768x1280.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_5s2_768x1280.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1024x1536.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_1024x1536.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_1280x832.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_1280x832.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_512x768.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_512x768.mp4
save_video: true
@@ -0,0 +1,38 @@
# LTX-2.3 distilled i2v — 8+3 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/i2v_8s3_768x1280.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: i2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
The subject slowly turns toward the camera with a soft, natural expression, hair and clothing swaying gently, shallow depth of field, subtle cinematic motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
# i2v conditioning — replace the path with your own first-frame image.
extensions:
ltx2_images:
- ["/path/to/your/first_frame.jpg", 0, 1.0]
ltx2_image_crf: 0.0
output:
output_path: outputs/ltx2_3_i2v_8s3_768x1280.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1024x1536.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_1024x1536.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_1280x832.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_1280x832.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_512x768.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_512x768.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 5+2 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_5s2_768x1280.yaml
#
# Stage 1 denoises for 5 steps at half resolution; the latents are then
# spatially upsampled and refined for 2 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 2
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 5
output:
output_path: outputs/ltx2_3_t2v_5s2_768x1280.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 1024x1536.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1024x1536.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1024
width: 1536
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_1024x1536.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 1280x832.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_1280x832.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 1280
width: 832
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_1280x832.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 512x768.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_512x768.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 512
width: 768
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_512x768.mp4
save_video: true
@@ -0,0 +1,33 @@
# LTX-2.3 distilled t2v — 8+3 two-stage at 768x1280.
#
# Run:
# fastvideo generate --config examples/inference/ltx2_3/t2v_8s3_768x1280.yaml
#
# Stage 1 denoises for 8 steps at half resolution; the latents are then
# spatially upsampled and refined for 3 steps (refine only supports 2 or 3).
# The refine upsampler auto-resolves from the model's `spatial_upscaler`.
generator:
model_path: FastVideo/LTX-2.3-Distilled-Diffusers
engine:
num_gpus: 1
pipeline:
workload_type: t2v
preset_overrides:
refine:
enabled: true
num_inference_steps: 3
guidance_scale: 1.0
add_noise: true
request:
prompt: >-
A cinematic drone shot flying over dramatic coastal cliffs at sunrise, golden light spilling across the water, gentle waves breaking on the rocks below, ultra-detailed, smooth camera motion.
sampling:
height: 768
width: 1280
num_frames: 121
fps: 24
guidance_scale: 1.0
num_inference_steps: 8
output:
output_path: outputs/ltx2_3_t2v_8s3_768x1280.mp4
save_video: true
@@ -0,0 +1,78 @@
# Causal Consistency Distillation: Wan 2.1 T2V 1.3B Causal
#
# ODE-data-free distillation. A frozen teacher takes a single CFG Euler step;
# the student matches an EMA copy of itself at the next timestep, all under
# clean-history teacher forcing.
#
# All three roles initialize from the SAME checkpoint (the teacher-forcing
# AR-diffusion model). Point init_from at that checkpoint for a real run.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
ema:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
method:
_target_: fastvideo.train.methods.consistency_model.causal_cd.CausalConsistencyDistillationMethod
discrete_cd_N: 48
guidance_scale: 3.0
ema_decay: 0.99
ema_start_step: 200
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_cd
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_cd_shift5
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
pipeline:
flow_shift: 5
@@ -0,0 +1,105 @@
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
#
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
# rollout from pure noise. Teacher provides the real score, critic
# learns the fake score; both inherited from DMD2Method.
#
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
# or with the NVIDIA-released checkpoint
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
# the paper weights (the delta_embedder rename is handled by the
# param_names_mapping in WanVideoArchConfig).
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: <PATH_TO_PRETRAIN_CKPT>
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 3.0
dmd_denoising_steps: [999, 937, 833, 624]
warp_denoising_step: false
# AnyFlow rollout knobs.
student_sample_steps: 4
use_mean_velocity: true
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
# Critic optimizer (DMD2 inherited).
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
attn_kind: vsa
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/preprocessed
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_anyflow_onpolicy
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: anyflow-wan
run_name: wan2.1_t2v_anyflow_onpolicy
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
pipeline:
flow_shift: 5.0
dit_config:
r_embedder: true
r_embedder_fusion: gated
r_embedder_gate_value: 0.25
r_embedder_deltatime_type: r
@@ -0,0 +1,83 @@
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
#
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
# can be sampled at arbitrary NFE in the on-policy stage.
#
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
diffusion_ratio: 0.5
consistency_ratio: 0.25
epsilon: 5 # finite-difference step in absolute train-timestep units
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
fuse_guidance_scale: 3.0
# shift is taken from pipeline.flow_shift below.
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/preprocessed
dataloader_num_workers: 4
train_batch_size: 4
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 6000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_anyflow_pretrain
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: anyflow-wan
run_name: wan2.1_t2v_anyflow_pretrain
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
pipeline:
flow_shift: 5.0
dit_config:
# Enable AnyFlow dual-timestep conditioning. The student loads from
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
# get initialized identically to time_embedder via deep-copy in
# WanTimeTextImageEmbedding.__init__.
r_embedder: true
r_embedder_fusion: gated
r_embedder_gate_value: 0.25
r_embedder_deltatime_type: r
@@ -100,9 +100,6 @@ callbacks:
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 81
# Validation/inference uses standard CFG in both clean and Self-Forcing,
# so this directly matches Self-Forcing guidance_scale=3.0.
guidance_scale: 3.0
pipeline:
flow_shift: 5
@@ -0,0 +1,73 @@
# DFSFT (Diffusion-Forcing SFT), frame-wise: Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model with a block size of 1 frame
# - Training: each frame gets its own independent noise level (frame-wise
# diffusion forcing), versus the chunk-wise variant that shares one noise
# level across num_frames_per_block frames.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
num_frames_per_block: 1
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 1
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_dfsft_framewise
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_dfsft_framewise_shift5_gauss_weight
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 5
@@ -0,0 +1,72 @@
# TFSFT (Teacher-Forcing SFT): Wan 2.1 T2V 1.3B Causal
#
# - Student: trainable causal Wan model
# - Training: inhomogeneous timesteps per chunk, but the causal transformer
# denoises the current block while attending to *clean* history (clean_x),
# not its own noisy rollout.
models:
student:
_target_: fastvideo.train.models.wan.WanCausalModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.tfsft.TeacherForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/Wan-Syn_77x448x832_600k
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 18
num_height: 448
num_width: 832
num_frames: 69
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_causal_tfsft
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 3
tracker:
project_name: distillation_wan_r
run_name: wan2.1_causal_tfsft_shift5_gauss_weight
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.wan.wan_causal_pipeline.WanCausalPipeline
dataset_file: examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/validation_4.json
every_steps: 50
sampling_steps: [40]
guidance_scale: 6.0
num_frames: 69
pipeline:
flow_shift: 5
+96 -9
View File
@@ -1,32 +1,119 @@
# World-Model: Matrix-Game 2.0 I2V
Three training scenarios for the Matrix-Game 2.0 I2V world model on the
new YAML-driven trainer (`fastvideo/train/entrypoint/train.py`).
Training scenarios for the Matrix-Game 2.0 I2V world model on Solaris (Minecraft)
data and Zelda data, using the new YAML-driven trainer
(`fastvideo/train/entrypoint/train.py`).
## Solaris Configs
| Config | Method | Student | Notes |
|---|---|---|---|
| `finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
| `dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
| `self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
| `solaris/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Multi-step SFT from `mg_bidirectional_Solaris`. |
| `solaris/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Diffusion-Forcing SFT with chunkwise timesteps. |
| `solaris/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Matrix-Game 2.0 DMD/Self-Forcing distillation; teacher = bidirectional, critic = bidirectional. |
## Zelda Configs
| Config | Method | Student | Notes |
|---|---|---|---|
| `zelda/finetune_i2v.yaml` | `FineTuneMethod` | `MatrixGame2Model` (bidirectional) | Zelda bidirectional I2V finetuning from `FastVideo/Matrix-Game-2.0-Base-Diffusers`. Uses 33-frame clips and Zelda validation with action overlays. |
| `zelda/dfsft_causal_i2v.yaml` | `DiffusionForcingSFTMethod` | `MatrixGame2CausalModel` | Zelda causal Diffusion-Forcing SFT from `mignonjia/mg_bidirectional_zelda`. Uses the same Zelda data, resolution, optimizer, and validation defaults as the Zelda finetune config. |
| `zelda/self_forcing_causal_i2v.yaml` | `SelfForcingMethod` | `MatrixGame2CausalModel` | Zelda DMD/Self-Forcing distillation; student init = `mignonjia/mg_causal_zelda`, teacher = bidirectional (`mignonjia/mg_bidirectional_zelda`), critic = bidirectional. |
| `zelda/streaming_long_tuning_causal_i2v.yaml` | `StreamingLongTuningMethod` | `MatrixGame2CausalModel` | LongLive-style streaming long tuning from the 1k-step Zelda self-forcing checkpoint. |
Zelda world-model distillation is a two-run workflow: first run
`zelda/self_forcing_causal_i2v.yaml` to train or load the 1k-step
self-forcing checkpoint (`mignonjia/mg_sf_distilled_zelda_1k_steps`), then run
`zelda/streaming_long_tuning_causal_i2v.yaml` for the 3k-step streaming
long-tuning stage. The long-tuning YAML starts from that 1k-step checkpoint; it
does not run the short self-forcing stage inside the same config.
## Zelda Training Data
The Zelda training configs use `data/zeldam2-clean` as a suggested local path.
Download the dataset from Hugging Face before running those configs:
```bash
python scripts/huggingface/download_hf.py \
--repo_id mignonjia/zeldam2-clean \
--local_dir data/zeldam2-clean \
--repo_type dataset
```
You can store the dataset elsewhere; update `training.data.data_path` in the
YAML to point at that location.
## Multi3D Training Data
`zelda/finetune_i2v.yaml` includes an optional, commented-out Multi3D entry.
Enable it only when you want to mix Zelda with multi-game data from
`data/multi3d_games`. You can store this dataset anywhere; before enabling it,
update the matching commented `training.data.data_path` key in the YAML to the
correct location.
To mix datasets in a training YAML, set `training.data.data_path` to a
path-to-repeat-count mapping. For example, `zelda/finetune_i2v.yaml` can use
`data/zeldam2-clean: 1` and `# data/multi3d_games: 10`; uncommenting the
Multi3D entry repeats the multi-game parquet list ten times before training
samples are shuffled.
## World Model Validation Data
The Zelda validation configs expect a small public validation bundle under
`data/zelda_validation_data`.
Download it from Hugging Face before running the Zelda scenarios:
```bash
python scripts/huggingface/download_hf.py \
--repo_id mignonjia/zelda_validation_data \
--local_dir data/zelda_validation_data \
--repo_type dataset
```
The bundle contains `validation_zelda.json`, `images/`, and `actions/`.
The Zelda configs point
`callbacks.validation.dataset_file` at
`data/zelda_validation_data/validation_zelda.json`.
## Usage
### Solaris
```bash
bash examples/train/run.sh \
examples/train/scenario/worldmodel/finetune_i2v.yaml
examples/train/scenario/worldmodel/solaris/finetune_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/self_forcing_causal_i2v.yaml
examples/train/scenario/worldmodel/solaris/self_forcing_causal_i2v.yaml
```
### Zelda
```bash
# Finetuning / DFSFT
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/finetune_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/dfsft_causal_i2v.yaml
# Distillation / long tuning
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/self_forcing_causal_i2v.yaml
bash examples/train/run.sh \
examples/train/scenario/worldmodel/zelda/streaming_long_tuning_causal_i2v.yaml
```
Override any field on the command line:
```bash
bash examples/train/run.sh \
examples/train/scenario/worldmodel/dfsft_causal_i2v.yaml \
examples/train/scenario/worldmodel/solaris/dfsft_causal_i2v.yaml \
--training.distributed.num_gpus 8 \
--training.optimizer.learning_rate 1e-5
```
@@ -97,4 +97,4 @@ callbacks:
guidance_scale: 6.0
pipeline:
flow_shift: 5
flow_shift: 5
@@ -0,0 +1,94 @@
# Diffusion-Forcing SFT: Zelda world model I2V Causal
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
init_from: mignonjia/mg_bidirectional_zelda
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.dfsft.DiffusionForcingSFTMethod
chunk_size: 3
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path:
data/zeldam2-clean: 1
dataloader_num_workers: 1
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 9
num_height: 480
num_width: 832
num_frames: 33
optimizer:
learning_rate: 2.0e-5
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 60000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_causal_dfsft
training_state_checkpointing_steps: 5000
checkpoints_total_limit: 3
tracker:
entity: hapo-exp
project_name: mg_1.3b_zelda
run_name: zelda_causal_dfsft
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 200
sampling_steps: [40]
sampling_timesteps: [1000, 975, 950, 925, 900, 875, 850, 825, 800, 775,
750, 725, 700, 675, 650, 625, 600, 575, 550, 525,
500, 475, 450, 425, 400, 375, 350, 325, 300, 275,
250, 225, 200, 175, 150, 125, 100, 75, 50, 25]
num_frames: 33
overlay_actions: true
guidance_scale: 6.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
@@ -0,0 +1,88 @@
# Matrix-Game 2.0 Zelda + multi-game I2V finetune.
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: FastVideo/Matrix-Game-2.0-Base-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path:
data/zeldam2-clean: 1
# data/multi3d_games: 10
dataloader_num_workers: 1
train_batch_size: 1
training_cfg_rate: 0.0 # unused for MatrixGame2 I2V; no text_embedding CFG dropout
seed: 42
num_latent_t: 9
num_height: 480
num_width: 832
num_frames: 33
optimizer:
learning_rate: 2.0e-5
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 60000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/matrixgame_finetune/checkpoints/zelda_with_mg_init
training_state_checkpointing_steps: 5000
checkpoints_total_limit: 3
tracker:
project_name: mg_1.3b_zelda
run_name: zelda_with_mg_init
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
validation:
_target_: fastvideo.train.callbacks.validation.ValidationCallback
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_i2v_pipeline.MatrixGame2I2VPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 200
sampling_steps: [40]
num_frames: 33
overlay_actions: true
guidance_scale: 6.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
@@ -0,0 +1,120 @@
# Self-forcing distillation: Zelda world model I2V Causal
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
init_from: mignonjia/mg_causal_zelda
trainable: true
teacher:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.self_forcing.SelfForcingMethod
rollout_mode: simulate
generator_update_interval: 5
dmd_denoising_steps: [1000, 750, 500, 250]
warp_denoising_step: true
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
chunk_size: 3
student_sample_type: sde
same_step_across_blocks: true
last_step_only: false
context_noise: 0.0
enable_gradient_in_rollout: true
start_gradient_frame: 0
# Critic optimizer
fake_score_learning_rate: 3.0e-7
fake_score_betas: [0.9, 0.95]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/zeldam2-clean
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1001
num_latent_t: 9
num_height: 480
num_width: 832
num_frames: 33
optimizer:
learning_rate: 3.0e-6
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/zelda_causal_self_forcing
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 1
tracker:
project_name: wangame_sf
run_name: mg2_self_forcing_9_latents
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 100
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 153
overlay_actions: true
keyboard_value_scale: 1.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
@@ -0,0 +1,140 @@
# MatrixGame2 I2V LongLive-style streaming distillation: Zelda world model I2V Causal
# Student init from self forcing after 1k steps
models:
student:
_target_: fastvideo.train.models.matrixgame2.matrixgame2_causal.MatrixGame2CausalModel
init_from: mignonjia/mg_sf_distilled_zelda_1k_steps
trainable: true
# transformer_override_safetensor: outputs/matrixgame_dmd/checkpoints/mg_zelda_sf_m2/checkpoint-1000_weight_only/ema/generator_ema.safetensors
teacher:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.matrixgame2.matrixgame2.MatrixGame2Model
init_from: mignonjia/mg_bidirectional_zelda
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.streaming_long_tuning.StreamingLongTuningMethod
rollout_mode: simulate
generator_update_interval: 5
dmd_denoising_steps: [1000, 750, 500, 250]
warp_denoising_step: true
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
chunk_size: 3
student_sample_type: sde
same_step_across_blocks: true
last_step_only: false
context_noise: 0.0
enable_gradient_in_rollout: true
start_gradient_frame: 0
streaming_training: true
streaming_chunk_size: 9
streaming_max_length: 39
streaming_fixed_overlap_latents: 3
streaming_reencode_overlap_anchor: true
streaming_anchor_inject_k: 1
streaming_require_full_blocks: true
multi_phased_distill_schedule:
- stage: streaming_long
start_step: 0
end_step: 3000
num_latent_t: 39
streaming_training: true
streaming_chunk_size: 9
streaming_max_length: 39
streaming_fixed_overlap_latents: 3
# Critic optimizer
fake_score_learning_rate: 3.0e-7
fake_score_betas: [0.9, 0.95]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/zeldam2-clean
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1001
num_latent_t: 39
num_height: 480
num_width: 832
num_frames: 153
optimizer:
learning_rate: 3.0e-6
betas: [0.9, 0.95]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 3000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/zelda_causal_long_tuning
training_state_checkpointing_steps: 1000
checkpoints_total_limit: 2
tracker:
project_name: wangame_sf
run_name: mg2_39only_streaming_long
model:
enable_gradient_checkpointing_type: full
dit_precision: fp32
callbacks:
ema:
decay: 0.99
start_iter: 200
grad_clip:
max_grad_norm: 1.0
validation:
pipeline_target: fastvideo.pipelines.basic.matrixgame2.matrixgame2_causal_dmd_pipeline.MatrixGame2CausalDMDPipeline
dataset_file: data/zelda_validation_data/validation_zelda.json
every_steps: 100
sampling_steps: [4]
sampling_timesteps: [1000, 750, 500, 250]
num_frames: 153
overlay_actions: true
keyboard_value_scale: 1.0
metrics:
enabled: true
names:
- vbench.imaging_quality
- vbench.aesthetic_quality
- vbench.temporal_flickering
- vbench.motion_smoothness
- vbench.subject_consistency
- vbench.background_consistency
- vbench.dynamic_degree
- optical_flow.synthetic_optical_flow
calibration_path: assets/eval/worldmodel_synthetic_flow_calibration.json
skip_missing_deps: true
strict: false
unload_after_validation: true
pipeline:
flow_shift: 5
dit_config:
local_attn_size: 6
sink_size: 1
+31 -5
View File
@@ -12,6 +12,9 @@ set(_FASTVIDEO_USER_CUDA_ARCH "${CMAKE_CUDA_ARCHITECTURES}")
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
endif()
if(NOT GPU_BACKEND)
set(GPU_BACKEND "CUDA")
endif()
if(GPU_BACKEND STREQUAL "ROCM")
enable_language(HIP)
@@ -50,7 +53,16 @@ if(NOT GPU_BACKEND STREQUAL "ROCM")
if(_FASTVIDEO_USER_CUDA_ARCH)
# Caller pinned -DCMAKE_CUDA_ARCHITECTURES (which torch ignores); translate it
# to the TORCH_CUDA_ARCH_LIST spelling: "121" -> "12.1", "90a" -> "9.0a".
# Only numeric spellings translate; keywords like "native"/"all" would
# otherwise be mangled into nonsense ("nativ.e").
foreach(_fv_arch IN LISTS _FASTVIDEO_USER_CUDA_ARCH)
if(NOT _fv_arch MATCHES "^[0-9]+[af]?$")
message(FATAL_ERROR
"fastvideo-kernel: CMAKE_CUDA_ARCHITECTURES='${_fv_arch}' is not "
"supported. Use a numeric arch (e.g. 90a, 121), set "
"TORCH_CUDA_ARCH_LIST directly (e.g. 9.0a), or unset both to "
"auto-detect from the visible GPU.")
endif()
string(REGEX MATCH "[af]$" _fv_suffix "${_fv_arch}")
string(REGEX REPLACE "[af]$" "" _fv_num "${_fv_arch}")
string(REGEX REPLACE "(.)$" ".\\1" _fv_num "${_fv_num}") # dot before the last digit
@@ -173,6 +185,14 @@ else()
endif()
endif()
# ThunderKittens headers don't compile on aarch64 hosts: plain char is unsigned
# there, and tk's base_types.cuh brace-initializes signed-char vector members
# from char (narrowing error). Skip TK until upstream is aarch64-clean.
if(ENABLE_TK_KERNELS AND CMAKE_SYSTEM_PROCESSOR MATCHES "^(aarch64|arm64)$")
message(STATUS "ThunderKittens kernels: forced OFF on ${CMAKE_SYSTEM_PROCESSOR} (tk headers are not aarch64-clean)")
set(ENABLE_TK_KERNELS OFF)
endif()
if(ENABLE_TK_KERNELS)
message(STATUS "ThunderKittens kernels: ENABLED")
else()
@@ -263,11 +283,6 @@ set(CUDA_FLAGS
"--expt-relaxed-constexpr"
"-Xcompiler=-fno-strict-aliasing"
"-Xcompiler=-fPIC"
# ARM/aarch64 defaults `char` to unsigned, but ThunderKittens headers assume the
# x86 signed-char behavior (else base_types.cuh hits "narrowing conversion from
# char to signed char"). Force signed char so TK compiles on Grace Hopper; this
# is a no-op on x86_64, where char is already signed.
"-Xcompiler=-fsigned-char"
"-DTORCH_COMPILE"
"-Xnvlink=--verbose"
"-Xptxas=--verbose"
@@ -426,3 +441,14 @@ if(ENABLE_ATTN_QAT_INFER)
install(TARGETS fp4attn_cuda LIBRARY DESTINATION .)
install(TARGETS fp4quant_cuda LIBRARY DESTINATION .)
endif()
# One-look answer to "what is this build producing?" — kept last so it is the
# final thing configure prints. The per-kernel matrix lives in README.md.
message(STATUS "============== fastvideo-kernel build summary ==============")
message(STATUS "host / backend: ${CMAKE_SYSTEM_PROCESSOR} / ${GPU_BACKEND}")
message(STATUS "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}")
message(STATUS "fastvideo_kernel_ops: ON (turbodiffusion int8-gemm/quant/rmsnorm/layernorm, all listed archs)")
message(STATUS " + TK sta/block_sparse (sm_90a only): ${ENABLE_TK_KERNELS}")
message(STATUS "fp4attn/fp4quant (sm_120a only, CUDA >= 12.8): ${ENABLE_ATTN_QAT_INFER}")
message(STATUS "Triton fallbacks ship in python/fastvideo_kernel/triton_kernels regardless.")
message(STATUS "============================================================")
+36
View File
@@ -2,6 +2,42 @@
CUDA kernels for FastVideo video generation.
## Kernel inventory
Compiled CUDA extensions (CMake, see the build summary printed at the end of every configure):
| Extension | Kernels | Sources | GPU arch | Build gate |
|---|---|---|---|---|
| `fastvideo_kernel._C.fastvideo_kernel_ops` | TurboDiffusion INT8 GEMM, quant, RMSNorm, LayerNorm | `csrc/turbodiffusion/` | every arch in `TORCH_CUDA_ARCH_LIST` | always built |
| same extension, optional part | ThunderKittens sliding-tile attention (`sta_fwd`) and VSA block-sparse (`block_sparse_fwd/bwd`) | `csrc/attention/*_h100.cu` | Hopper `sm_90a` only | `FASTVIDEO_KERNEL_BUILD_TK` (AUTO = ON iff `9.0a` is in the arch list; always OFF on aarch64 hosts — TK headers don't compile there) |
| `fp4attn_cuda`, `fp4quant_cuda` | FP4 attention + quantization ("attn_qat_infer", modified SageAttention3) | `attn_qat_infer/` | consumer Blackwell `sm_120a` only, CUDA ≥ 12.8 | `FASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER` (AUTO = ON iff `12.0a` is in the arch list) |
Runtime-JIT kernels (no build step, ship in every wheel/image):
| Kernels | Where | Used when |
|---|---|---|
| Triton: STA, VSA block-sparse, SLA, fused compress+topk, FP4 QAT training, quant/norm utils | `python/fastvideo_kernel/triton_kernels/` | automatic fallback when the matching C++ op is absent (`ops.py`, `turbodiffusion_ops.py`) |
| FA4 CuTe-DSL block-sparse (VSA-256 fastpath on `sm_100`) | `block_sparse_attn_cute_fwd.py` | optional `flash_attn.cute` dependency, see below |
| VMoBA `moba_attn_varlen` | `vmoba.py` | wraps flash-attn varlen |
## What gets built where, and when
| Surface | Trigger | Leg | `TORCH_CUDA_ARCH_LIST` | TK | FP4 |
|---|---|---|---|---|---|
| PyPI wheels (`.github/workflows/publish-kernel.yml`) | version bump in `fastvideo-kernel/pyproject.toml` on main, or manual dispatch | x86_64 cu126 | `9.0a` | ON | — (CUDA < 12.8) |
| | | x86_64 cu130 | `9.0a;12.0a` | ON | ON |
| | | aarch64 cu130 | `10.0a;12.0a` | — | ON |
| Docker images `ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev` (`.github/workflows/infra-build-image.yml`) | `docker/Dockerfile` changes on main, or manual dispatch | amd64 cuda12.6.3 + cuda13.0.0 | `9.0a` | ON | — |
| | | arm64 cuda12.6.3 (GH200) | `9.0a` | — (aarch64) | — |
| | | arm64 cuda13.0.0 (GB10 / DGX Spark) | `12.1` | — | — |
| Local `./build.sh` | manual | probes the visible GPU via torch | detected | ON iff sm_90 (non-aarch64 host) | ON iff sm_120 |
Notes:
- No Docker image ships the FP4 kernels; only the x86_64/aarch64 cu130 wheels do.
- On arm64 images (GH200 included) STA/VSA run on the Triton fallbacks, since TK never builds on aarch64.
- Kernel tests run on Buildkite GPU CI for PRs touching `fastvideo-kernel/**` (see `.buildkite/pipeline.yml`).
## Installation
### Standard Installation (Local Development)
+17
View File
@@ -46,6 +46,23 @@ fi
if git rev-parse --git-dir >/dev/null 2>&1; then
git submodule update --init --recursive include/cutlass include/tk
fi
# Fail fast with a clear message if the headers are still missing (e.g. a
# Docker context that excluded .git AND the submodule contents) instead of
# dying later in a wall of nvcc include errors. CUTLASS is consumed by the
# always-built turbodiffusion sources, so it is a hard error; ThunderKittens
# only feeds the TK-gated Hopper kernels, so a missing tree just warns (the
# TK gate resolves later, and non-SM90/ROCm builds never touch it).
if [ ! -d include/cutlass/include ]; then
echo "ERROR: include/cutlass/include is missing. Outside a git checkout the" >&2
echo " CUTLASS sources must already be present (run" >&2
echo " 'git submodule update --init --recursive include/cutlass include/tk'" >&2
echo " in the source checkout, or include them in the build context)." >&2
exit 1
fi
if [ ! -d include/tk/include ]; then
echo "WARNING: include/tk/include is missing; ThunderKittens (Hopper sm_90a)" >&2
echo " kernels cannot be built. Fine for non-SM90/ROCm targets." >&2
fi
# Install build dependencies
uv pip install scikit-build-core cmake ninja
+27
View File
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from fastvideo.api.sampling_param import SamplingParam
@dataclass
class FluxSamplingParam(SamplingParam):
prompt: str | None = "a photo of a cat"
negative_prompt: str = ""
num_videos_per_prompt: int = 1
seed: int = 0
num_frames: int = 1
height: int = 1024
width: int = 1024
fps: int = 1
num_inference_steps: int = 28
guidance_scale: float = 3.5
use_embedded_guidance: bool = True
true_cfg_scale: float = 1.0
+16
View File
@@ -90,6 +90,10 @@ class SamplingParam:
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
use_embedded_guidance: bool = False
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
true_cfg_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
sigmas: list[float] | None = None
@@ -325,6 +329,18 @@ class SamplingParam:
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--use-embedded-guidance",
action="store_true",
default=SamplingParam.use_embedded_guidance,
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
)
parser.add_argument(
"--true-cfg-scale",
type=float,
default=SamplingParam.true_cfg_scale,
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
)
parser.add_argument(
"--boundary-ratio",
type=float,
+1
View File
@@ -150,6 +150,7 @@ class SamplingConfig:
guidance_scale_2: float | None = None
guidance_rescale: float = 0.0
true_cfg_scale: float | None = None
use_embedded_guidance: bool | None = None
boundary_ratio: float | None = None
sigmas: list[float] | None = None
+21 -104
View File
@@ -5,99 +5,10 @@ import torch
import torch.nn.functional as F
from dataclasses import dataclass
try:
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
fa_version = "4"
except ImportError:
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 no longer have a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "3"
except ImportError:
from flash_attn import flash_attn_func as flash_attn_2_func
flash_attn_func = flash_attn_2_func
fa_version = "2"
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
# already a registered torch.library custom op, so dynamo treats it as a
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
# breaks the graph at the call site (observed: wanvideo.py self-attn,
# once per layer every step), which fragments the compiled region and
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
# opaque-but-traceable node. The kernel still runs eager inside the op
# (correct — flash-attn must run eager); only dynamo's treatment of the
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
if fa_version in ("2", "3"):
_fa_default = flash_attn_func
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
# call shape used by FlashAttentionImpl.forward's default branch
# (see `flash_attn_func_compilable(...)` call site below). The
# masked/no-pad and varlen / cross-attn paths use different
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
# are intentionally out of scope for this PR — wrapping them is a
# natural follow-up. The wrapper's signature is the contract: any
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
# return_attn_probs, ...) raises TypeError at the call site, so
# silent loss of kwargs is not a failure mode.
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
del softmax_scale, causal
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
# same dtype/device as q (head dim taken from v).
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Autograd carve-out. The custom op above registers a forward + fake
# kernel but NO backward (register_autograd), so it is opaque to
# autograd. Inference runs under no_grad / inference_mode and routes
# through the traceable custom op — that is the torch.compile win, and
# the only path this PR claims. Training backprops through attention,
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
# (itself an autograd.Function, so backward is correct) at the cost of a
# dynamo graph break on the training path — i.e. pre-PR behavior, no
# regression. Full autograd parity for the custom op (mirroring the FP4
# cute template) is a tracked follow-up.
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
elif fa_version == "4":
# FA4 path: `flash_attn_func` is already a torch.library custom op
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
# passthrough is enough — no extra registration needed.
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
else:
# Defensive: the probe above only ever sets fa_version to "2", "3",
# or "4"; an unexpected value means an import/probe regression and
# we want a loud error at import, not a silent NameError later.
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
f"'2', '3', or '4' from the import probe above.")
from fastvideo.attention.utils.flash_attn_default import (
fa_version,
flash_attn_func_compilable,
)
from fastvideo.attention.backends.abstract import (
AttentionBackend,
@@ -108,8 +19,6 @@ from fastvideo.attention.backends.abstract import (
from fastvideo.logger import init_logger
logger = init_logger(__name__)
_WARNED_NON_FA_DTYPE = False
logger.info("Using FlashAttention-%s backend", fa_version)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
@@ -271,12 +180,8 @@ class FlashAttentionImpl(AttentionImpl):
# SP). Cast through bf16 and restore, matching TORCH_SDPA's tolerance.
orig_dtype = query.dtype
if orig_dtype not in (torch.float16, torch.bfloat16):
global _WARNED_NON_FA_DTYPE
if not _WARNED_NON_FA_DTYPE:
_WARNED_NON_FA_DTYPE = True
logger.warning(
"FLASH_ATTN received %s inputs; casting to bfloat16 for the "
"kernel and restoring on output.", orig_dtype)
logger.warning_once(f"FLASH_ATTN received {orig_dtype} inputs; casting to "
f"bfloat16 for the kernel and restoring on output.")
query = query.to(torch.bfloat16)
key = key.to(torch.bfloat16)
value = value.to(torch.bfloat16)
@@ -293,9 +198,17 @@ class FlashAttentionImpl(AttentionImpl):
attn_metadata: FlashAttnMetadata,
):
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
# Route through the *_compilable wrappers so dynamo sees one
# traceable node for each masked entry point (the unpad/pad
# bookkeeping runs eager inside the custom op). On FA2 these
# wrappers go through ops with full register_autograd, so
# training also backprops through the op (no graph break on
# the training path); on FA3/FA4 they carve out to the
# autograd.Function for grad-enabled calls — see
# fastvideo/attention/utils/flash_attn_no_pad.py.
from fastvideo.attention.utils.flash_attn_no_pad import (
flash_attn_no_pad,
flash_attn_varlen_qk_no_pad,
flash_attn_no_pad_compilable as flash_attn_no_pad,
flash_attn_varlen_qk_no_pad_compilable as flash_attn_varlen_qk_no_pad,
)
attn_mask = attn_metadata.attn_mask
@@ -322,7 +235,11 @@ class FlashAttentionImpl(AttentionImpl):
)
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)
key_padding_mask = _key_padding_mask_from_attn_mask(attn_mask, attn_mask.shape[-1]).to(device=query.device)
if key_padding_mask.shape[-1] > qkv.shape[1]:
raise ValueError("Invalid key padding mask length for FLASH_ATTN: "
f"expected at most {qkv.shape[1]}, got {key_padding_mask.shape[-1]}")
attn_mask_padded = F.pad(key_padding_mask, (qkv.shape[1] - key_padding_mask.shape[-1], 0), value=True)
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)
+147
View File
@@ -0,0 +1,147 @@
# SPDX-License-Identifier: Apache-2.0
"""NABLA block-sparse flex-attention backend (Kandinsky5 "nabla" checkpoints).
The block mask is data-dependent: nablaT_v2 mean-pools 64-token blocks of the
fractal-ordered sequence, thresholds the softmaxed block map, and ORs it with a
precomputed spatio-temporal-window (STA) mask carried on the attention
metadata. The mask spans the full sequence, so this backend does not support
sequence parallelism — use it via LocalAttention only.
"""
import math
from dataclasses import dataclass
from typing import Any
import torch
try:
from torch.nn.attention.flex_attention import BlockMask, flex_attention
flex_attention = torch.compile(flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
CAN_USE_FLEX_ATTN = True
except ImportError:
CAN_USE_FLEX_ATTN = False
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
def nablaT_v2(
q: torch.Tensor,
k: torch.Tensor,
sta: torch.Tensor,
thr: float = 0.9,
) -> "BlockMask":
q = q.transpose(1, 2).contiguous()
k = k.transpose(1, 2).contiguous()
# Map estimation
B, h, S, D = q.shape
s1 = S // 64
qa = q.reshape(B, h, s1, 64, D).mean(-2)
ka = k.reshape(B, h, s1, 64, D).mean(-2).transpose(-2, -1)
map = qa @ ka
map = torch.softmax(map / math.sqrt(D), dim=-1)
# Map binarization
vals, inds = map.sort(-1)
cvals = vals.cumsum_(-1)
mask = (cvals >= 1 - thr).int()
mask = mask.gather(-1, inds.argsort(-1))
mask = torch.logical_or(mask, sta)
# BlockMask creation
kv_nb = mask.sum(-1).to(torch.int32)
kv_inds = mask.argsort(dim=-1, descending=True).to(torch.int32)
return BlockMask.from_kv_blocks(torch.zeros_like(kv_nb), kv_inds, kv_nb, kv_inds, BLOCK_SIZE=64, mask_mod=None)
class NablaAttentionBackend(AttentionBackend):
@staticmethod
def get_name() -> str:
return "NABLA_ATTN"
@staticmethod
def get_impl_cls() -> type["NablaAttentionImpl"]:
return NablaAttentionImpl
@staticmethod
def get_metadata_cls() -> type["NablaAttentionMetadata"]:
return NablaAttentionMetadata
@staticmethod
def get_builder_cls() -> type["NablaAttentionMetadataBuilder"]:
return NablaAttentionMetadataBuilder
@dataclass
class NablaAttentionMetadata(AttentionMetadata):
# Block-level STA window mask [1, 1, S/64, S/64], precomputed once per run.
sta_mask: torch.Tensor = None # type: ignore[assignment]
# Cumulative-probability threshold for block-map binarization.
P: float = 0.9
visual_shape: tuple[int, int, int] = (0, 0, 0)
class NablaAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build(
self,
current_timestep: int,
sta_mask: torch.Tensor,
P: float,
visual_shape: tuple[int, int, int],
**kwargs: Any,
) -> NablaAttentionMetadata:
return NablaAttentionMetadata(
current_timestep=current_timestep,
sta_mask=sta_mask,
P=P,
visual_shape=visual_shape,
)
class NablaAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
if not CAN_USE_FLEX_ATTN:
raise RuntimeError("NABLA attention requires torch.nn.attention.flex_attention, "
"which is unavailable in this PyTorch build.")
if causal:
raise ValueError("NABLA attention does not support causal masking.")
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: NablaAttentionMetadata,
) -> torch.Tensor:
# q/k/v: [B, S, heads, head_dim], fractal-ordered by the model; S % 64 == 0.
block_mask = nablaT_v2(query, key, attn_metadata.sta_mask, thr=attn_metadata.P)
return flex_attention(
query=query.transpose(1, 2),
key=key.transpose(1, 2),
value=value.transpose(1, 2),
block_mask=block_mask,
).transpose(1, 2)
+45 -1
View File
@@ -2,6 +2,7 @@
import torch
from dataclasses import dataclass
from torch.nn import functional as F
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
AttentionBackend, AttentionImpl, AttentionMetadata, AttentionMetadataBuilder)
from fastvideo.logger import init_logger
@@ -19,7 +20,7 @@ class SDPABackend(AttentionBackend):
@staticmethod
def get_name() -> str:
return "SDPA"
return "TORCH_SDPA"
@staticmethod
def get_impl_cls() -> type["SDPAImpl"]:
@@ -49,9 +50,51 @@ class SDPAMetadataBuilder(AttentionMetadataBuilder):
current_timestep: int,
attn_mask: torch.Tensor,
) -> SDPAMetadata:
# Store the mask exactly as passed. The metadata is cross-backend:
# call sites (HYWorld, HunyuanVideo15) build SDPAMetadata while the
# layer's selector may pick FLASH_ATTN, and the shared convention for
# padding masks is the tokenizer-style 2D [batch, key_len]. Any
# reshaping for torch.sdpa happens inside the SDPA impl
# (_normalize_attn_mask_for_sdpa).
return SDPAMetadata(current_timestep=current_timestep, attn_mask=attn_mask)
def _normalize_attn_mask_for_sdpa(
attn_mask: torch.Tensor | None,
query: torch.Tensor,
key: torch.Tensor,
) -> torch.Tensor | None:
if attn_mask is None:
return None
attn_mask = attn_mask.to(device=query.device)
# F.scaled_dot_product_attention only accepts bool or float masks;
# tokenizers commonly produce int64 0/1 padding masks.
if attn_mask.dtype != torch.bool and not attn_mask.dtype.is_floating_point:
attn_mask = attn_mask != 0
key_len = key.shape[-2]
if attn_mask.shape[-1] > key_len:
raise ValueError("Invalid attention mask length for SDPA: "
f"expected at most {key_len}, got {attn_mask.shape[-1]}")
if attn_mask.shape[-1] < key_len:
# Front-pad as "attend": double-stream layouts (HYWorld) prepend
# non-text tokens the tokenizer mask does not cover.
valid_value = True if attn_mask.dtype == torch.bool else 0.0
attn_mask = F.pad(attn_mask, (key_len - attn_mask.shape[-1], 0), value=valid_value)
if attn_mask.dim() == 2:
# In-tree producers pass 2D [batch, key_len] padding masks; lift to a
# broadcastable [batch, 1, 1, key_len] here so torch.sdpa does not
# reinterpret 2D as its documented [query_len, key_len] broadcast.
return attn_mask[:, None, None, :]
if attn_mask.dim() == 3:
return attn_mask[:, None, :, :]
if attn_mask.dim() == 4:
return attn_mask
raise ValueError(f"Unsupported attention mask shape for SDPA: {attn_mask.shape}")
class SDPAImpl(AttentionImpl):
def __init__(
@@ -82,6 +125,7 @@ class SDPAImpl(AttentionImpl):
attn_mask = attn_metadata.attn_mask if (attn_metadata is not None
and hasattr(attn_metadata, "attn_mask")) else None
attn_mask = _normalize_attn_mask_for_sdpa(attn_mask, query, key)
attn_kwargs = {
"attn_mask": attn_mask,
"dropout_p": self.dropout,
+5 -1
View File
@@ -252,6 +252,7 @@ class LocalAttention(nn.Module):
causal: bool = False,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
default_backend: AttentionBackendEnum | None = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
@@ -262,7 +263,10 @@ class LocalAttention(nn.Module):
num_kv_heads = num_heads
dtype = get_compute_dtype()
attn_backend = get_attn_backend(head_size, dtype, supported_attention_backends=supported_attention_backends)
attn_backend = get_attn_backend(head_size,
dtype,
supported_attention_backends=supported_attention_backends,
default_backend=default_backend)
impl_cls = attn_backend.get_impl_cls()
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
+9 -1
View File
@@ -84,8 +84,9 @@ def get_attn_backend(
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
default_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends)
return _cached_get_attn_backend(head_size, dtype, supported_attention_backends, default_backend)
@cache
@@ -94,6 +95,7 @@ def _cached_get_attn_backend(
dtype: torch.dtype,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
default_backend: AttentionBackendEnum | None = None,
) -> type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
@@ -112,6 +114,12 @@ def _cached_get_attn_backend(
if backend_by_env_var is not None:
selected_backend = backend_name_to_enum(backend_by_env_var)
# Layer-level default (e.g. a checkpoint that requires a specific sparse
# backend). Lower precedence than the global force and the env var, so
# users can still override it.
if selected_backend is None and default_backend is not None:
selected_backend = default_backend
# get device-specific attn_backend
from fastvideo.platforms import current_platform
+66 -73
View File
@@ -4,10 +4,9 @@ import functools
from collections.abc import Callable
import torch
from flash_attn import flash_attn_func as _flash_attn_2_func
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
from fastvideo.logger import init_logger
from fastvideo.platforms import current_platform
logger = init_logger(__name__)
@@ -16,7 +15,8 @@ if torch.cuda.is_available():
from flash_attn.cute.interface import _flash_attn_bwd, _flash_attn_fwd
except ImportError:
# flash_attn.cute (FA4) is simply not installed -- expected on builds
# without it; callers fall back to FA3/FA2 quietly.
# without it; callers handle the ImportError (the FASTVIDEO_FA4 gate in
# flash_attn.py raises, the FP4 probe treats FA4 as unavailable).
raise
except Exception as e:
# flash_attn.cute IS installed but failed to import -- almost always an
@@ -24,23 +24,57 @@ if torch.cuda.is_available():
# 'cutlass.cute.core' has no attribute 'ThrMma'" (an AttributeError, not
# ImportError). This is fixable by pinning a compatible
# nvidia-cutlass-dsl, so warn loudly, then re-raise as ImportError so
# callers fall back to FA3/FA2 instead of crashing worker init.
# callers can handle it uniformly.
logger.warning(
"flash_attn.cute (FA4) is installed but failed to import (%r); "
"falling back to FA3/FA2. This is usually an nvidia-cutlass-dsl "
"version mismatch -- pin a compatible nvidia-cutlass-dsl to "
"restore FA4.", e)
"flash_attn.cute (FA4) is installed but failed to import (%r). "
"This is usually an nvidia-cutlass-dsl version mismatch -- pin a "
"compatible nvidia-cutlass-dsl to restore FA4.", e)
raise ImportError(f"flash_attn.cute (FA4) import failed: {e!r}") from e
else:
# This error will be caught in flash_attn.py or flash_attn_no_pad.py
raise ImportError("flash_attn.cute is only available on CUDA devices; this error must be handled internally")
try:
# FA2 serves the calls FA4 cute cannot on pre-sm90 GPUs (backward, GQA).
# Optional so FA4-only installs can still import this module.
from flash_attn import flash_attn_func as _flash_attn_2_func
from flash_attn import flash_attn_varlen_func as _flash_attn_2_varlen_func
except ImportError:
_flash_attn_2_func = None
_flash_attn_2_varlen_func = None
def _check_dropout(dropout_p: float) -> None:
if dropout_p != 0.0:
raise NotImplementedError(f"flash_attn.cute does not support dropout (got dropout_p={dropout_p})")
@functools.cache
def _sm90_or_newer() -> bool:
return current_platform.has_device_capability(90)
def _use_fa2(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> bool:
if _sm90_or_newer():
return False
# Pre-sm90 FA4 cute limitations, both served by FA2 (deterministic
# capability gate, not a runtime fallback):
# * the backward asserts sm90+ (L40S/sm_89 dies on its arch check);
# * GQA fails CuTeDSL JIT in pack_gqa ("ValueError: Operation creation
# failed", observed on sm_89 with HunyuanGameCraft/LTX2).
if q.shape[-2] != k.shape[-2]:
return True
return torch.is_grad_enabled() and any(t.requires_grad for t in (q, k, v))
def _fa2_or_raise(fa2_func: Callable | None) -> Callable:
if fa2_func is None:
raise RuntimeError("this attention call cannot run on FA4 cute below sm90 (its backward and "
"GQA support require sm90+) and flash-attn 2, which serves it there, is "
"not installed.")
return fa2_func
@torch.library.custom_op(
"fastvideo::_flash_attn_cute_forward",
mutates_args=(),
@@ -243,70 +277,6 @@ torch.library.register_autograd(
)
# FA4's CuTeDSL kernels JIT-compile per shape family, and some configurations
# fail MLIR op creation at runtime even though the import succeeded (observed:
# GQA models on sm_89 dying in pack_gqa with "ValueError: Operation creation
# failed"). Degrade to FA2 once, process-wide, instead of crashing inference.
class _FA4Policy:
"""Per-call gate for the FA4 cute fast path, with FA2 as the fallback.
FA4 is skipped when:
* a previous call failed at runtime -- CuTeDSL JIT compilation is
shape-dependent, so the first failure disables FA4 for the rest of
the process instead of retrying a broken JIT on every call; or
* the call needs autograd -- FA4's backward asserts sm90+ (L40S/sm_89
dies on its arch check) and is unvalidated for training in this repo
(its lse is not even allocated through our inference-shaped custom
op), so training keeps the pre-FA4 behavior: FA2 on every device.
"""
def __init__(self) -> None:
self.broken = False
def use_fa4(self, *tensors: torch.Tensor) -> bool:
if self.broken:
return False
return not (torch.is_grad_enabled() and any(t.requires_grad for t in tensors))
def mark_broken(self, error: Exception) -> None:
if not self.broken:
self.broken = True
logger.warning(
"flash_attn.cute (FA4) failed at runtime (%r); falling back "
"to FA2 for the rest of this process.", error)
_FA4 = _FA4Policy()
def _with_fa2_fallback(fa2_func: Callable) -> Callable:
"""Pair an FA4 cute wrapper with its FA2 twin of the same signature.
The decorated body runs only when ``_FA4`` allows it; otherwise (or after
the first FA4 runtime failure) the call is served by ``fa2_func``.
``NotImplementedError`` is a contract error (e.g. dropout), not a JIT
failure, so it propagates without disabling FA4.
"""
def decorator(fa4_func: Callable) -> Callable:
@functools.wraps(fa4_func)
def wrapper(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *args, **kwargs) -> torch.Tensor:
if _FA4.use_fa4(q, k, v):
try:
return fa4_func(q, k, v, *args, **kwargs)
except NotImplementedError:
raise
except Exception as e: # CuTeDSL compile errors surface as ValueError
_FA4.mark_broken(e)
return fa2_func(q, k, v, *args, **kwargs)
return wrapper
return decorator
@_with_fa2_fallback(_flash_attn_2_func)
def flash_attn_func(
q: torch.Tensor,
k: torch.Tensor,
@@ -317,6 +287,16 @@ def flash_attn_func(
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
if _use_fa2(q, k, v):
return _fa2_or_raise(_flash_attn_2_func)(
q,
k,
v,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
_check_dropout(dropout_p)
out, _ = torch.ops.fastvideo._flash_attn_cute_forward(q, k, v, softmax_scale, causal, deterministic)
return out
@@ -392,7 +372,6 @@ def flash_attn_fp4_func(
return torch.ops.fastvideo._flash_attn_cute_fp4_forward(q, k, v, sfq, sfk, softmax_scale, causal)
@_with_fa2_fallback(_flash_attn_2_varlen_func)
def flash_attn_varlen_func(
q: torch.Tensor,
k: torch.Tensor,
@@ -407,6 +386,20 @@ def flash_attn_varlen_func(
deterministic: bool = False,
) -> torch.Tensor:
"""Only returns the output, not the lse."""
if _use_fa2(q, k, v):
return _fa2_or_raise(_flash_attn_2_varlen_func)(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
_check_dropout(dropout_p)
out, _ = torch.ops.fastvideo._flash_attn_cute_varlen_forward(
q,
@@ -0,0 +1,251 @@
# SPDX-License-Identifier: Apache-2.0
"""torch.compile-traceable wrapper for the FA2/FA3/FA4 default attention path.
The FA4/cute path (`fa_version == "4"`) is already a registered
`torch.library.custom_op` in `fastvideo.attention.utils.flash_attn_cute`, so
dynamo treats it as a graph node. The external FA2/FA3 ``flash_attn_func`` is
NOT — dynamo breaks the graph at the call site (observed: wanvideo.py
self-attn, once per layer every step), which fragments the compiled region
and blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom op
(mirrors the FP4 `flash_attn_cute` template) so it becomes an
opaque-but-traceable node. The kernel still runs eager inside the op
(correct — flash-attn must run eager); only dynamo's treatment of the
boundary changes, so numerics are unchanged (SSIM-gated).
Autograd: FA2 has full ``register_autograd`` parity — the custom op's
backward calls flash_attn's ``_flash_attn_backward`` directly, so training
backprops *through* the op (no graph break on the training path either).
FA3 currently keeps the no-backward + carve-out pattern from PR #1373
because FA3's private backward signature wants validation on a real Hopper
box (gated on Kuan-Hao's Modal FA3 setup PR). Once that lands the FA3 path
can mirror FA2.
Lives in `attention/utils/` (sibling of `flash_attn_cute.py` and
`flash_attn_no_pad.py`) so it can be imported by any backend that wants the
traceable FA default call without pulling in backend dispatch logic. The
backend (`attention/backends/flash_attn.py`) just imports
`flash_attn_func_compilable` and `fa_version` from here.
"""
import importlib.util
import torch
from fastvideo import envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Pick the same backend the rest of FastVideo picked for `flash_attn_func`
# (FA4/cute → FA3 → FA2). Mirror the precedence used in
# `attention/utils/flash_attn_no_pad.py` so the two probes always agree.
#
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1: its CuTeDSL
# kernels JIT-compile per shape family and can fail at runtime on some
# arch/shape combinations, so it is never auto-selected just because it is
# installed.
if envs.FASTVIDEO_FA4:
try:
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
except ImportError as e:
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
fa_version = "4"
else:
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 no longer has a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "3"
except ImportError:
from flash_attn import flash_attn_func as flash_attn_2_func
flash_attn_func = flash_attn_2_func
fa_version = "2"
try:
if importlib.util.find_spec("flash_attn.cute") is not None:
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
"set FASTVIDEO_FA4=1 to use it for inference.")
except ImportError:
pass
if fa_version == "2":
# Scope: this op covers exactly the q/k/v + softmax_scale + causal call
# shape used by FlashAttentionImpl.forward's default branch (see
# `flash_attn_func_compilable(...)` call site in
# `attention/backends/flash_attn.py`). The masked/no-pad and varlen /
# cross-attn paths use different entry points
# (`flash_attn_no_pad`, `flash_attn_varlen_*`) which live in
# `attention/utils/flash_attn_no_pad.py`. The wrapper's signature is the
# contract: any extra kwarg (dropout_p, window_size, alibi_slopes,
# deterministic, return_attn_probs, ...) raises TypeError at the call
# site, so silent loss of kwargs is not a failure mode.
from flash_attn.flash_attn_interface import _flash_attn_backward as _fa2_backward
_fa_default = flash_attn_func
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
# `return_attn_probs=True` asks FA2 to also return softmax_lse +
# S_dmask. We need softmax_lse to feed the backward; S_dmask is the
# dropout mask (always None here since dropout_p is fixed at 0).
out, softmax_lse, _ = _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal, return_attn_probs=True)
return out, softmax_lse
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del softmax_scale, causal
# FA2 default path: out = [batch, seqlen_q, nheads, head_dim_v],
# softmax_lse = [batch, nheads, seqlen_q], fp32 regardless of q dtype.
b, sq, hq = q.shape[0], q.shape[1], q.shape[2]
out = q.new_empty(b, sq, hq, v.shape[-1])
lse = q.new_empty(b, hq, sq, dtype=torch.float32)
return out, lse
def _flash_attn_default_setup_context(ctx, inputs, output):
q, k, v, softmax_scale, causal = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse)
# `lse` is an auxiliary output we save to feed FA2's backward; nobody
# should differentiate through it. Mark it non-differentiable so
# autograd errors loudly if a caller wires it into a loss, rather
# than silently producing zero/None grads through the `del grad_lse`
# in our backward.
ctx.mark_non_differentiable(lse)
# FA2's *forward* substitutes `1 / sqrt(head_dim)` for `softmax_scale=None`
# internally; FA2's *backward* (`_flash_attn_backward`) demands a concrete
# float in its C++ schema and rejects None at the binding boundary. Resolve
# the default here so the value saved on ctx (and passed to backward) is
# always a real float — matches what FA2's own autograd.Function does.
if softmax_scale is None:
softmax_scale = q.shape[-1]**-0.5
ctx.softmax_scale = softmax_scale
ctx.causal = causal
def _flash_attn_default_backward(ctx, grad_out, grad_lse):
# We only differentiate `out`; softmax_lse is saved-for-backward, not
# a real differentiable output. (Mirrors the FP4 cute template.)
del grad_lse
q, k, v, out, lse = ctx.saved_tensors
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
# FA2's `_flash_attn_backward` writes into dq/dk/dv in place. The
# extra kwargs (window_size_*, softcap, alibi_slopes, deterministic,
# rng_state) are pinned to the same defaults the forward wrapper
# uses — flash-attn==2.8.1 (the version FastVideo pins) requires
# all of them explicitly. `rng_state=None` is correct for our
# `dropout_p=0` configuration.
_fa2_backward(
grad_out,
q,
k,
v,
out,
lse,
dq,
dk,
dv,
dropout_p=0.0,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
alibi_slopes=None,
deterministic=False,
rng_state=None,
)
return dq, dk, dv, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_default_forward",
_flash_attn_default_backward,
setup_context=_flash_attn_default_setup_context,
)
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Backward is registered: autograd flows through the op (training
# path is also traceable; no carve-out needed). Public API matches
# `flash_attn_func` — returns just `out`; we drop the saved-for-
# backward `lse` here so callers see the original single-tensor
# contract.
out, _ = torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
return out
elif fa_version == "3":
# FA3 path: same forward+fake custom op as the original PR #1373, with
# the autograd carve-out kept. The full backward (mirroring the FA2 leg
# above) wants a Hopper box for grad-check validation, which we don't
# have until Kuan-Hao's Modal FA3 setup PR lands. Until then this keeps
# inference traceable + training correct (via the original
# autograd.Function path + a pre-PR-style graph break on training).
_fa_default = flash_attn_func
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
del softmax_scale, causal
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Autograd carve-out. The custom op above registers a forward + fake
# kernel but NO backward (register_autograd), so it is opaque to
# autograd. Inference runs under no_grad / inference_mode and routes
# through the traceable custom op — that is the torch.compile win, and
# the only path this PR claims. Training backprops through attention,
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
# (itself an autograd.Function, so backward is correct) at the cost of a
# dynamo graph break on the training path — i.e. pre-PR behavior, no
# regression. Full autograd parity for the custom op (mirroring the FP4
# cute template) is a tracked follow-up.
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
elif fa_version == "4":
# FA4 path: `flash_attn_func` is already a torch.library custom op
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
# passthrough is enough — no extra registration needed.
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
else:
# Defensive: the probe above only ever sets fa_version to "2", "3",
# or "4"; an unexpected value means an import/probe regression and
# we want a loud error at import, not a silent NameError later.
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
f"'2', '3', or '4' from the import probe above.")
+503 -14
View File
@@ -21,27 +21,46 @@ from einops import rearrange
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
from fastvideo import envs
def _resolve_flash_attn_varlen_func() -> Any:
try:
from fastvideo.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
return flash_attn_varlen_func_cute
except ImportError:
def _resolve_flash_attn_varlen_func() -> tuple[Any, str]:
if envs.FASTVIDEO_FA4:
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
# fail loudly here rather than fall through to FA3/FA2. RuntimeError,
# not ImportError: importers like bsa_attn.py treat ImportError as
# "flash-attn not installed" and silently degrade to reference kernels.
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
from fastvideo.attention.utils.flash_attn_cute import (
flash_attn_varlen_func as flash_attn_varlen_func_cute, )
except ImportError as e:
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
return flash_attn_varlen_func_interface
except ImportError:
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
return flash_attn_varlen_func_cute, "4"
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
return flash_attn_varlen_func_flash
return flash_attn_varlen_func_interface, "3"
except ImportError:
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
return flash_attn_varlen_func_flash, "2"
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
flash_attn_varlen_func_impl, _FA_VARLEN_VERSION = _resolve_flash_attn_varlen_func()
# FA2-only: the private varlen backward we register against the custom ops
# below. FA3 / FA4 have different private signatures and validation paths
# (Hopper / Blackwell boxes) — those legs keep the autograd carve-out
# pattern from PR #1373 until their setup PRs land.
if _FA_VARLEN_VERSION == "2":
from flash_attn.flash_attn_interface import (
_flash_attn_varlen_backward as _fa2_varlen_backward, )
def flash_attn_no_pad(
@@ -180,3 +199,473 @@ def flash_attn_varlen_qk_no_pad(
h=nheads,
)
return output
# ---------------------------------------------------------------------------
# torch.compile traceability + register_autograd parity for the masked /
# varlen attention paths.
#
# Wraps the two entry points `FlashAttentionImpl.forward` calls
# (`flash_attn_no_pad`, `flash_attn_varlen_qk_no_pad`) as
# `torch.library.custom_op`s so dynamo sees one traceable node — the
# internal unpad / pad bookkeeping (data-dependent `nnz` shapes) runs
# eager inside the op, and the op's outputs are the statically-shaped
# padded tensors. This mirrors the FA2 default-path wrapper in
# `fastvideo/attention/backends/flash_attn.py`.
#
# Autograd: on FA2 we register a real backward (`register_autograd`)
# that calls FA2's `_flash_attn_varlen_backward` on the unpadded form
# — re-unpadding the saved padded tensors using the saved mask. The
# `softmax_lse` from the varlen forward is naturally unpadded
# (`[nheads, total_q]`); we pad it to `[batch, nheads, seqlen]` on
# the way out (statically shaped) and re-unpad in backward. So
# training backprops *through* the op (no graph break on the training
# path either).
#
# FA3 / FA4 keep the autograd carve-out pattern from PR #1373: the
# custom op has forward + fake only, and `*_compilable` falls back to
# the original autograd.Function for grad-enabled calls. Those legs
# are gated on Hopper-class / Blackwell-class boxes for backward
# validation and ship as separate follow-ups.
if _FA_VARLEN_VERSION == "2":
# ---------- masked self-attention: flash_attn_no_pad (FA2) ----------
@torch.library.custom_op(
"fastvideo::_flash_attn_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_no_pad_forward(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
b, s, _three, h, d = qkv.shape
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
out_unpad, lse_unpad, _ = flash_attn_varlen_qkvpacked_func(x_unpad,
cu_seqlens,
max_s,
dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
return_attn_probs=True)
# Pad out: [nnz, h, d] -> [b, s, h, d]
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), indices, b, s),
"b s (h d) -> b s h d",
h=h)
# Pad lse: FA2 varlen returns [nheads, total_q]. Transpose to [total_q,
# nheads], pad to [b, s, nheads], permute to [b, nheads, s] — statically
# shaped so register_fake matches.
lse_padded = pad_input(lse_unpad.t().contiguous(), indices, b, s).permute(0, 2, 1).contiguous()
return out_padded, lse_padded
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
def _flash_attn_no_pad_forward_fake(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
b, s, _three, h, d = qkv.shape
out = qkv.new_empty(b, s, h, d)
lse = qkv.new_empty(b, h, s, dtype=torch.float32)
return out, lse
def _flash_attn_no_pad_setup_context(ctx, inputs, output):
qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic = inputs
out, lse = output
ctx.save_for_backward(qkv, out, lse, key_padding_mask)
# Auxiliary output, not differentiable — see default-path note.
ctx.mark_non_differentiable(lse)
# FA2's varlen backward requires a concrete float for softmax_scale.
if softmax_scale is None:
softmax_scale = qkv.shape[-1]**-0.5 # head_dim from qkv's last dim
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.dropout_p = dropout_p
ctx.deterministic = deterministic
def _flash_attn_no_pad_backward(ctx, grad_out, grad_lse):
# lse is saved-for-backward, not differentiated.
del grad_lse
qkv, out_padded, lse_padded, key_padding_mask = ctx.saved_tensors
b, s, _three, h, d = qkv.shape
# One `unpad_input` call (on qkv) gives us indices + cu_seqlens + max_s;
# reuse those for out / dout / lse below via direct indexing instead
# of redundant `unpad_input` calls (each of which would re-run
# `nonzero` + `cumsum` + a `.max().item()` GPU→CPU sync).
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
q_unpad, k_unpad, v_unpad = (t.contiguous() for t in x_unpad.unbind(dim=1))
# Direct-index variants reuse `indices` (computed above).
out_unpad = out_padded.flatten(0, 1)[indices].view(-1, h, d).contiguous()
dout_unpad = grad_out.flatten(0, 1)[indices].view(-1, h, d).contiguous()
# lse_padded [b, h, s] -> [b, s, h] -> [nnz, h] -> [h, nnz].
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[indices].t().contiguous()
dq_unpad = torch.empty_like(q_unpad)
dk_unpad = torch.empty_like(k_unpad)
dv_unpad = torch.empty_like(v_unpad)
_fa2_varlen_backward(
dout_unpad,
q_unpad,
k_unpad,
v_unpad,
out_unpad,
lse_unpad,
dq_unpad,
dk_unpad,
dv_unpad,
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_s,
max_seqlen_k=max_s,
dropout_p=ctx.dropout_p,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
alibi_slopes=None,
deterministic=ctx.deterministic,
rng_state=None,
)
# Re-pad each grad and stack into dqkv.
def _repad(dt_unpad: torch.Tensor) -> torch.Tensor:
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, b, s)
return rearrange(padded, "b s (h d) -> b s h d", h=h)
dqkv = torch.stack([_repad(dq_unpad), _repad(dk_unpad), _repad(dv_unpad)], dim=2)
# 6 inputs total: qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic.
return dqkv, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_no_pad_forward",
_flash_attn_no_pad_backward,
setup_context=_flash_attn_no_pad_setup_context,
)
# ---------- cross-attention: flash_attn_varlen_qk_no_pad (FA2) ----------
@torch.library.custom_op(
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_varlen_qk_no_pad_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
b, sq, h, d = query.shape
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
query_padding_mask)
k_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
key_padding_mask)
v_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h)
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h)
v_unpad = rearrange(v_unpad, "nnz (h d) -> nnz h d", h=h)
out_unpad, lse_unpad, _ = flash_attn_varlen_func_impl(q_unpad,
k_unpad,
v_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
return_attn_probs=True)
# Pad out: [nnz_q, h, d] -> [b, sq, h, d]
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), q_indices, b, sq),
"b s (h d) -> b s h d",
h=h)
# Pad lse: [h, nnz_q] -> [b, h, sq]
lse_padded = pad_input(lse_unpad.t().contiguous(), q_indices, b, sq).permute(0, 2, 1).contiguous()
return out_padded, lse_padded
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
def _flash_attn_varlen_qk_no_pad_forward_fake(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del key, query_padding_mask, key_padding_mask
del causal, dropout_p, softmax_scale, deterministic
b, sq, h, _ = query.shape
# `out`'s head_dim comes from value (d_v), matching the real forward's
# out_padded ([b, sq, h, d_v]); it can differ from query's d_q.
out = query.new_empty(b, sq, h, value.shape[-1])
lse = query.new_empty(b, h, sq, dtype=torch.float32)
return out, lse
def _flash_attn_varlen_qk_no_pad_setup_context(ctx, inputs, output):
(query, key, value, query_padding_mask, key_padding_mask, causal, dropout_p, softmax_scale,
deterministic) = inputs
out, lse = output
ctx.save_for_backward(query, key, value, out, lse, query_padding_mask, key_padding_mask)
# Auxiliary output, not differentiable — see default-path note.
ctx.mark_non_differentiable(lse)
if softmax_scale is None:
softmax_scale = query.shape[-1]**-0.5
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.dropout_p = dropout_p
ctx.deterministic = deterministic
def _flash_attn_varlen_qk_no_pad_backward(ctx, grad_out, grad_lse):
del grad_lse
(query, key, value, out_padded, lse_padded, query_padding_mask, key_padding_mask) = ctx.saved_tensors
b, sq, h, d = query.shape
sk = key.shape[1]
# One `unpad_input` call per distinct mask; reuse the returned
# indices via direct indexing for everything else that shares
# the same mask (v with k_mask; out/dout/lse with q_mask; the
# final repad of dk/dv also reuses k_indices). Avoids ~4
# redundant `unpad_input` calls + their GPU→CPU `.max().item()`
# syncs.
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
query_padding_mask)
k_unpad, k_indices, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
key_padding_mask)
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
v_unpad = value.flatten(0, 1)[k_indices].view(-1, h, d).contiguous()
# out / dout / lse follow q's shape, so index with q_indices.
out_unpad = out_padded.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
dout_unpad = grad_out.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[q_indices].t().contiguous()
dq_unpad = torch.empty_like(q_unpad)
dk_unpad = torch.empty_like(k_unpad)
dv_unpad = torch.empty_like(v_unpad)
_fa2_varlen_backward(
dout_unpad,
q_unpad,
k_unpad,
v_unpad,
out_unpad,
lse_unpad,
dq_unpad,
dk_unpad,
dv_unpad,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=ctx.dropout_p,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
alibi_slopes=None,
deterministic=ctx.deterministic,
rng_state=None,
)
# k_indices is already available from the unpad_input above —
# no need to recompute it for the dk/dv repad.
def _repad(dt_unpad: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int) -> torch.Tensor:
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, batch, seqlen)
return rearrange(padded, "b s (h d) -> b s h d", h=h)
dq_padded = _repad(dq_unpad, q_indices, b, sq)
dk_padded = _repad(dk_unpad, k_indices, b, sk)
dv_padded = _repad(dv_unpad, k_indices, b, sk)
# 9 inputs total: query, key, value, q_mask, k_mask, causal, dropout_p,
# softmax_scale, deterministic.
return dq_padded, dk_padded, dv_padded, None, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
_flash_attn_varlen_qk_no_pad_backward,
setup_context=_flash_attn_varlen_qk_no_pad_setup_context,
)
# ---------- public dispatchers (FA2: autograd flows through the op) -----
def flash_attn_no_pad_compilable(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
"""dynamo-traceable wrapper around ``flash_attn_no_pad`` (registered op,
full register_autograd on FA2 — both inference and training go through
the op, no graph break on either)."""
out, _ = torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
deterministic)
return out
def flash_attn_varlen_qk_no_pad_compilable(query,
key,
value,
query_padding_mask,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
"""dynamo-traceable wrapper around ``flash_attn_varlen_qk_no_pad`` (registered
op, full register_autograd on FA2)."""
out, _ = torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
key_padding_mask, causal, dropout_p,
softmax_scale, deterministic)
return out
else:
# ---------- FA3 / FA4: carve-out (forward+fake only, no real backward) ---
# Same pattern as the parked varlen-extension and the FA3 default leg in
# `fastvideo/attention/backends/flash_attn.py`. Real backward for these
# versions is a follow-up gated on Hopper / Blackwell box validation.
@torch.library.custom_op(
"fastvideo::_flash_attn_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_no_pad_forward(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
return flash_attn_no_pad( # type: ignore[no-untyped-call]
qkv,
key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
def _flash_attn_no_pad_forward_fake(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
b, s, _three, h, d = qkv.shape
return qkv.new_empty(b, s, h, d)
@torch.library.custom_op(
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_varlen_qk_no_pad_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
return flash_attn_varlen_qk_no_pad( # type: ignore[no-untyped-call]
query,
key,
value,
query_padding_mask=query_padding_mask,
key_padding_mask=key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
def _flash_attn_varlen_qk_no_pad_forward_fake(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
del key, query_padding_mask, key_padding_mask
del causal, dropout_p, softmax_scale, deterministic
b, sq, h, _ = query.shape
# `out`'s head_dim comes from value (d_v), matching the real forward's
# output ([b, sq, h, d_v]); it can differ from query's d_q.
return query.new_empty(b, sq, h, value.shape[-1])
def flash_attn_no_pad_compilable(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
if torch.is_grad_enabled() and qkv.requires_grad:
return flash_attn_no_pad(qkv,
key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
return torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
deterministic)
def flash_attn_varlen_qk_no_pad_compilable(query,
key,
value,
query_padding_mask,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
if torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad):
return flash_attn_varlen_qk_no_pad(query,
key,
value,
query_padding_mask=query_padding_mask,
key_padding_mask=key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
return torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
key_padding_mask, causal, dropout_p,
softmax_scale, deterministic)
+7 -3
View File
@@ -1,6 +1,9 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
from fastvideo.configs.models.dits.flux import FluxDiTConfig
from fastvideo.configs.models.dits.flux_2 import Flux2Config
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
@@ -13,7 +16,8 @@ from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig"
]
+119
View File
@@ -0,0 +1,119 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 VFM Transformer FastVideo dataclass configs.
Architecture is 1:1 with the published ``nvidia/Cosmos3-Nano`` checkpoint
(``transformer/config.json``; class ``Cosmos3OmniTransformer`` / framework
``Cosmos3VFMNetwork``). Field values match that config so the FastVideo native
DiT builds a parameter tree matching the checkpoint's state-dict surface
(814 tensors / 44 patterns, validated 2026-06-06).
Reference of record: ``cosmos-framework`` (NVIDIA). The checkpoint is a single
``layers`` ModuleList of dual-pathway (understanding/text + generation/vision)
decoder blocks; per layer: ``self_attn`` with und (``to_{q,k,v}``/``to_out``)
and gen (``add_{q,k,v}_proj``/``to_add_out``) projections + QK-norms, plus
``mlp`` (und) and ``mlp_moe_gen`` (gen), and four RMSNorms. Top level adds
``embed_tokens``/``norm``/``norm_moe_gen``/``lm_head``/``proj_in``/``proj_out``/
``time_embedder`` and dormant ``action_*``/``audio_*`` heads. The checkpoint
remap lives in ``scripts/checkpoint_conversion/cosmos3_convert.py``.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_cosmos3_transformer_block(name: str, module) -> bool:
"""FSDP shard boundary: the dual-pathway decoder blocks ``layers.{i}``."""
del module
parts = name.split(".")
return "layers" in parts and parts[-1].isdigit()
@dataclass
class Cosmos3ArchConfig(DiTArchConfig):
"""Architecture config for the Cosmos3 omni DiT (Cosmos3-Nano).
1:1 with ``transformer/config.json``. The action/sound heads ship in the
checkpoint, so they are constructed for strict-load parity even though the
PR1 video path (T2V/I2V/T2I) leaves them dormant.
"""
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_cosmos3_transformer_block])
# Conversion is owned by scripts/checkpoint_conversion/cosmos3_convert.py;
# the native module tree is the source of truth for parameter names.
param_names_mapping: dict = field(default_factory=dict)
# ---- Backbone (Qwen3-VL-text) ----
hidden_size: int = 4096
num_hidden_layers: int = 36
num_attention_heads: int = 32
num_key_value_heads: int = 8 # GQA (4 query groups)
head_dim: int = 128
intermediate_size: int = 12288
hidden_act: str = "silu"
vocab_size: int = 151936
rms_norm_eps: float = 1e-6
attention_bias: bool = False
qk_norm_for_diffusion: bool = True
qk_norm_for_text: bool = True
use_moe: bool = True # dual-pathway weights; sparse routing unused
joint_attn_implementation: str = "two_way"
freeze_und: bool = False
# ---- Position embedding (unified 3D MRoPE) ----
position_embedding_type: str = "unified_3d_mrope"
rope_theta: float = 5_000_000.0
max_position_embeddings: int = 262144
mrope_section: list[int] = field(default_factory=lambda: [24, 20, 20])
mrope_interleaved: bool = True
unified_3d_mrope_reset_spatial_ids: bool = True
temporal_modality_margin: int = 15000 # unified_3d_mrope_temporal_modality_margin
# ---- VAE / patch geometry ----
latent_patch_size: int = 2
latent_channel: int = 48
patch_latent_dim: int = 192 # latent_patch_size**2 * latent_channel
# ---- Diffusion conditioning ----
timestep_scale: float = 0.001
# ---- Temporal / FPS modulation ----
base_fps: float = 24.0
temporal_compression_factor: int = 4
enable_fps_modulation: bool = True
video_temporal_causal: bool = False
# ---- Action generation head (dormant in PR1 video path) ----
action_gen: bool = True
action_dim: int = 64
max_action_dim: int = 64
num_embodiment_domains: int = 32
# ---- Sound generation head (dormant in PR1 video path) ----
sound_gen: bool = True
sound_dim: int = 64
sound_latent_fps: float = 25.0
temporal_compression_factor_sound: int = 1
# ---- BaseDiT bookkeeping ----
in_channels: int = 48
out_channels: int = 48
def __post_init__(self) -> None:
super().__post_init__()
# Video DiT contract: latent channels == VAE z_dim.
self.num_channels_latents = self.latent_channel
if not self.out_channels:
self.out_channels = self.in_channels
# Derived: patchify packs latent_patch_size**2 spatial patches * channels.
self.patch_latent_dim = self.latent_patch_size**2 * self.latent_channel
@dataclass
class Cosmos3VideoConfig(DiTConfig):
"""Pipeline-level Cosmos3 DiT config (T2V / I2V / T2I share this surface)."""
arch_config: DiTArchConfig = field(default_factory=Cosmos3ArchConfig)
prefix: str = "Cosmos3"
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
@dataclass
class DreamXWorldArchConfig(WanVideoArchConfig):
"""DreamX-World DiT config with camera PRoPE control fields."""
add_control_adapter: bool = True
cam_method: str | None = "prope"
attn_compress: int = 1
cam_self_attn_layers: tuple[int, ...] | None = None
@dataclass
class DreamXWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=DreamXWorldArchConfig)
prefix: str = "Wan"
@dataclass
class DreamXWorldARArchConfig(DreamXWorldArchConfig):
"""DreamX-World-5B autoregressive causal DiT config."""
model_type: str = "ti2v"
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
text_dim: int = 4096
freq_dim: int = 256
attn_compress: int = 4
cam_self_attn_layers: tuple[int, ...] | None = tuple(range(30))
local_attn_size: int = 12
sink_size: int = 3
num_frames_per_block: int = 3
rope_cache_policy: str = "block_relativistic"
# The official AR checkpoint (AMAP-ML/DreamX-World ``model.safetensors``)
# already uses FastVideo's native key names and the converter copies the
# tensors verbatim, so every rule is an identity. The rules enumerate the
# full state-dict surface of ``DreamXWorldARTransformer3DModel`` (norm1 /
# norm2 / head.norm are affine-free and have no parameters).
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.\1",
r"^text_embedding\.([02])\.(.*)$": r"text_embedding.\1.\2",
r"^time_embedding\.([02])\.(.*)$": r"time_embedding.\1.\2",
r"^time_projection\.1\.(.*)$": r"time_projection.1.\1",
r"^blocks\.(\d+)\.self_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.self_attn.\2.\3",
r"^blocks\.(\d+)\.self_attn\.norm_(q|k)\.weight$": r"blocks.\1.self_attn.norm_\2.weight",
r"^blocks\.(\d+)\.cross_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.cross_attn.\2.\3",
r"^blocks\.(\d+)\.cross_attn\.norm_(q|k)\.weight$": r"blocks.\1.cross_attn.norm_\2.weight",
r"^blocks\.(\d+)\.cam_self_attn\.(q_proj|k_proj|v_proj|out_proj)\.(.*)$": r"blocks.\1.cam_self_attn.\2.\3",
r"^blocks\.(\d+)\.cam_self_attn\.norm_(q|k)\.weight$": r"blocks.\1.cam_self_attn.norm_\2.weight",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm3.\2",
r"^blocks\.(\d+)\.ffn\.([02])\.(.*)$": r"blocks.\1.ffn.\2.\3",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.modulation",
r"^head\.head\.(.*)$": r"head.head.\1",
r"^head\.modulation$": r"head.modulation",
})
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
@dataclass
class DreamXWorldARConfig(DreamXWorldConfig):
arch_config: DiTArchConfig = field(default_factory=DreamXWorldARArchConfig)
prefix: str = "Wan"
+27
View File
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class FluxTransformer2DArchConfig(DiTArchConfig):
patch_size: int = 1
in_channels: int = 64
out_channels: int | None = None
num_layers: int = 19
num_single_layers: int = 38
attention_head_dim: int = 128
num_attention_heads: int = 24
joint_attention_dim: int = 4096
pooled_projection_dim: int = 768
guidance_embeds: bool = True
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
@dataclass
class FluxDiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=FluxTransformer2DArchConfig)
prefix: str = "flux"
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class GlmImageDiTArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
hidden_size: int = 4096
num_attention_heads: int = 32
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
num_layers: int = 30
text_embed_dim: int = 1472
time_embed_dim: int = 512
condition_dim: int = 256
prior_vq_quantizer_codebook_size: int = 16384
patch_size: int = 2
max_height: int = 2048
max_width: int = 2048
qk_norm: str = "layer_norm"
eps: float = 1e-5
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
})
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
def __post_init__(self):
super().__post_init__()
self.num_channels_latents = self.out_channels
@dataclass
class GlmImageDiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
prefix: str = "GlmImage"
+14 -4
View File
@@ -2,14 +2,24 @@
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def _is_kandinsky5_transformer_block(n: str, m) -> bool:
return ("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
@dataclass
class Kandinsky5ArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [
lambda n, m:
("text_transformer_blocks" in n or "visual_transformer_blocks" in n) and n.split(".")[-1].isdigit()
])
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_kandinsky5_transformer_block])
# NABLA block-sparse attention for attention_type="nabla" checkpoints, plus
# the dense backends every DiT supports.
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.NABLA_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
# Native FastVideo implementation uses the same parameter names as diffusers
# except FFN internals: Diffusers FFN uses `in_layer/out_layer`, while
+14
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Literal
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -19,6 +20,11 @@ class WanVideoArchConfig(DiTArchConfig):
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
# AnyFlow dual-timestep checkpoints expose delta_embedder weights with the
# same internal layout as time_embedder. The regex is harmless on plain
# Wan checkpoints (no delta_embedder keys to match).
r"^condition_embedder\.delta_embedder\.linear_1\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_in.\1",
r"^condition_embedder\.delta_embedder\.linear_2\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
@@ -86,6 +92,14 @@ class WanVideoArchConfig(DiTArchConfig):
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
rope_cache_policy: str = "absolute"
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
# the legacy single-timestep forward (no delta_embedder allocated, no
# extra computation on the embedder forward path).
r_embedder: bool = False
r_embedder_fusion: Literal["additive", "gated"] = "additive"
r_embedder_gate_value: float = 0.25
r_embedder_deltatime_type: Literal["r", "t-r"] = "r"
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels

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