Compare commits

...
Author SHA1 Message Date
SolitaryThinker 8b30a222a3 [docs] cosmos3: port_feedback.md (pitfalls/lessons) + native-reuse audit 2026-07-20 12:19:46 -07:00
SolitaryThinker 624a638cb3 [docs] cosmos3: full-omni parity summary (all components bit-exact) 2026-07-20 12:19:46 -07:00
SolitaryThinker 1241ee0600 [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-20 12:19:46 -07:00
SolitaryThinker cbdd5e2c70 [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-20 12:19:46 -07:00
SolitaryThinker 8d1c8fa01a [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-20 12:19:46 -07:00
SolitaryThinker 78f9ce3c37 [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-20 12:19:46 -07:00
SolitaryThinker 54def005c4 [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-20 12:19:46 -07:00
SolitaryThinker 8309acae69 [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-20 12:19:46 -07:00
SolitaryThinker 7926926ad1 [docs] cosmos3 PR2 (audio/t2vs) complete; bit-exact across components 2026-07-20 12:19:46 -07:00
SolitaryThinker eb9d439bab [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-20 12:19:46 -07:00
SolitaryThinker d6affff3a1 [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-20 12:19:46 -07:00
SolitaryThinker 1bd77cb466 [docs] cosmos3 PR2: AVAE sound decoder done; remaining audio components 2026-07-20 12:19:46 -07:00
SolitaryThinker 35900c32ba [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-20 12:19:46 -07:00
SolitaryThinker 2945890df4 [docs] cosmos3 PR2: audio port plan (AVAE + DiT sound pathway + t2vs) 2026-07-20 12:19:46 -07:00
SolitaryThinker 4bd8457072 [docs] cosmos3: record T2I verification + resolution-based flow_shift 2026-07-20 12:19:46 -07:00
SolitaryThinker f13d017fc1 [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-20 12:19:45 -07:00
SolitaryThinker 5fe5b3b12f [docs] cosmos3: record I2V real-weights verification (feat/cosmos3-i2v) 2026-07-20 12:19:45 -07:00
SolitaryThinker ec61d3d0bf [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-20 12:19:45 -07:00
SolitaryThinker c665dd83c5 [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-20 12:19:45 -07:00
SolitaryThinker b4ffc8e22e [docs] cosmos3: record real-weights E2E acceptance + scheduler fix (I003) 2026-07-20 12:19:45 -07:00
SolitaryThinker 455cd83bed [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-20 12:19:45 -07:00
SolitaryThinker 4747d51265 [misc] cosmos3 PORT_STATUS: PR1 video core complete (framework parity) 2026-07-20 12:19:45 -07:00
SolitaryThinker 7ae57e6fbd [feat] cosmos3: native video pipeline + framework denoise parity 2026-07-20 12:19:45 -07:00
SolitaryThinker 8931f1b5d7 [feat] cosmos3: native MoT sequence-packing (video) + framework parity 2026-07-20 12:19:17 -07:00
SolitaryThinker 42eb451b50 [misc] cosmos3 PORT_STATUS: strict-load done; only pipeline remains 2026-07-20 12:19:17 -07:00
SolitaryThinker f52e88e397 [feat] cosmos3: strict-load verified (identity; needs_conversion=no) 2026-07-20 12:19:17 -07:00
SolitaryThinker 5cdbf06a23 [misc] cosmos3 PORT_STATUS: VAE component framework parity verified 2026-07-20 12:19:17 -07:00
SolitaryThinker 9e799d5bf4 [feat] cosmos3: VAE config (Wan2.2 reuse) + framework parity test 2026-07-20 12:19:17 -07:00
SolitaryThinker 468c4d505e [misc] cosmos3 PORT_STATUS: DiT framework parity verified (both rope modes) 2026-07-20 12:19:17 -07:00
SolitaryThinker 99e8c404b6 [test] cosmos3: DiT unified_3d_mrope framework parity (real checkpoint) 2026-07-20 12:19:17 -07:00
SolitaryThinker 4e6a6a3d17 [feat] cosmos3: native DiT + framework parity (Cosmos3VFMTransformer) 2026-07-20 12:19:17 -07:00
SolitaryThinker 5784e0fcb5 [misc] cosmos3 PORT_STATUS: PR1 progress (arch config, framework parity ref) 2026-07-20 12:19:17 -07:00
SolitaryThinker 576a006d06 [test] cosmos3: official framework DiT parity reference (CPU/SDPA) 2026-07-20 12:19:17 -07:00
SolitaryThinker 423f8d7eff [misc] cosmos3 PORT_STATUS: framework-reference pivot + Phase 1 findings 2026-07-20 12:19:17 -07:00
SolitaryThinker a91482ed14 [feat] cosmos3: arch config 1:1 with Cosmos3-Nano checkpoint 2026-07-20 12:19:17 -07:00
SolitaryThinker f0c838a78d [misc]: cosmos3 PORT_STATUS: post-rebase state + fv-cosmos3 env 2026-07-20 12:19:17 -07:00
SolitaryThinker 13c76b710b [misc]: cosmos3 resume: PORT_STATUS + README for official diffusers ref 2026-07-20 12:19:17 -07:00
SolitaryThinker 4bd3eff562 [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-20 12:19:17 -07:00
SolitaryThinker 041e21821a [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-20 12:19:17 -07:00
SolitaryThinker d254b7455d [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-20 12:18:53 -07:00
SolitaryThinker b0f1b8ad96 [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-20 12:18:53 -07:00
SolitaryThinker 5d778c94a1 [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-20 12:18:53 -07:00
William Lin 65f3b946b9 [feat] Add LTX-2 and LTX-2.3 fine-tuning to the modular trainer (#1624) 2026-07-20 11:58:39 -07:00
Zhang Peiyuan 191fcbf46c [feat] add qat docs (#1621) 2026-07-19 18:37:47 -07:00
Junda Su 755a4e4470 [new-model] Add LingBot-Video Dense and MoE/refiner T2V inference (#1595) 2026-07-18 20:18:16 -07:00
9709b7513b [feat] Port NVFP4 QAT/QAD to modular train framework (#1619)
Co-authored-by: Peiyuan Zhang <email>

Co-authored-by: Peiyuan <a>
2026-07-18 15:01:09 -07:00
William Lin 32cd603515 Revert docs trusted-branch-only workflow (#1618) 2026-07-17 00:22:56 -07:00
Junda Su d4bdd3621a [new-model] Port LingBot-World-v2 (#1579) 2026-07-16 19:51:53 -07:00
William Lin e2f8322842 [ci]: skip unused Buildkite submodule checkout (#1614) 2026-07-16 18:37:57 -07:00
6966f9e0bc [fix] Z-Image (#1236) draft port: rebase + strict-load contract + bf16 encoder parity + PORT_STATUS (#1339)
Co-authored-by: Mrinaal Dogra <mdogra@ucsd.edu>
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-16 18:18:25 -07:00
201 changed files with 25577 additions and 387 deletions
+2
View File
@@ -1,6 +1,8 @@
env:
IMAGE_VERSION: "py3.12-latest"
BUILDKITE_CLEAN_CHECKOUT: true
# Buildkite only launches Modal; remote jobs initialize their own submodules.
BUILDKITE_GIT_SUBMODULES: false
notify:
- github_commit_status:
+8 -14
View File
@@ -11,9 +11,7 @@ on:
- 'requirements-mkdocs.txt'
- 'scripts/check_docs_links.py'
- '.github/workflows/infra-docs.yml'
# Run the trusted base-branch workflow so fork PRs can be skipped without
# waiting for maintainer approval.
pull_request_target:
pull_request:
branches: [ main ]
paths:
- 'docs/**'
@@ -26,19 +24,21 @@ on:
permissions:
contents: read
pages: write
id-token: write
concurrency:
group: "pages"
cancel-in-progress: false
jobs:
build:
# MkDocs executes repository code; only trusted same-repository PRs run it.
if: github.event_name == 'push' || github.event.pull_request.head.repo.full_name == github.repository
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v4
with:
ref: ${{ github.event.pull_request.head.sha || github.sha }}
fetch-depth: 0
persist-credentials: false
- name: Setup Python
uses: actions/setup-python@v5
@@ -52,7 +52,6 @@ jobs:
run: uv pip install --system -r requirements-mkdocs.txt
- name: Setup Pages
if: github.event_name == 'push'
uses: actions/configure-pages@v4
- name: Build documentation
@@ -62,22 +61,17 @@ jobs:
run: python scripts/check_docs_links.py
- name: Upload artifact
if: github.event_name == 'push'
uses: actions/upload-pages-artifact@v3
with:
path: ./site
deploy:
permissions:
pages: write
id-token: write
concurrency: pages
environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}
runs-on: ubuntu-latest
needs: build
if: github.event_name == 'push'
if: github.ref == 'refs/heads/main'
steps:
- name: Deploy to GitHub Pages
id: deployment
+7
View File
@@ -6,6 +6,7 @@ results/
wandb/
*.ipynb
*.jpg
!examples/dataset/lingbotworld2/image.jpg
*.safetensors
*.mp4
*.png
@@ -34,9 +35,15 @@ env
*.log
weights/
logs/
/Z-Image/
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/**
+1 -1
View File
@@ -9,7 +9,7 @@
**FastVideo is a unified post-training and real-time inference framework for accelerated video generation.**
## NEWS
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), check out the [Blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/06/23`: Release FastWan-QAD: 5s of Video generated in 1.8s E2E. See the [FastWan-QAD models](https://huggingface.co/FastVideo/FastWan-QAD-FP8-1.3B), [Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/), and [blog](https://haoailab.com/blogs/fastwan-qad/).
- `2026/03/17`: Release demo: Into the Dreamverse: Vibe Directing in FastVideo, check out the [Blog](https://haoailab.com/blogs/dreamverse/).
- `2026/03/13`: Release demo: Create a 5s 1080p Video in 4.5s with FastVideo on a Single GPU, check out the [Blog](https://haoailab.com/blogs/fastvideo_realtime_1080p/).
- `2025/11/19`: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py).
+3 -1
View File
@@ -8,13 +8,15 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
* **[Sliding Tile Attention (STA)](sta/index.md)**: STA kernel support is kept in
`fastvideo-kernel`; full FastVideo STA pipeline workflow is archived in
`sta_do_not_delete`.
* **[Attn-QAT Training](../training/attn_qat.md)**: Runtime-JIT Triton forward
and backward kernels for role-local quantization-aware training.
* **Backend development guide**: See the developer guide at
[Attention Backend Development](../contributing/attention_backend.md).
## General Build Instructions
These instructions apply to building the `fastvideo-kernel` package from
source, which includes both STA and VSA kernels.
source, which includes STA, VSA, and Attn-QAT kernels.
### Prerequisites
+3 -4
View File
@@ -255,10 +255,9 @@ If you add a new CI test category:
### Documentation
`.github/workflows/infra-docs.yml` builds documentation for same-repository PRs
that touch `docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow
itself. Fork PRs skip this executable build instead of waiting for maintainer
approval. On pushes to `main`, it also deploys the built site to GitHub Pages.
`.github/workflows/infra-docs.yml` builds documentation for PRs that touch
`docs/**`, `mkdocs.yml`, `requirements-mkdocs.txt`, or the workflow itself. On
pushes to `main`, it also deploys the built site to GitHub Pages.
The docs job:
@@ -333,15 +333,25 @@ surfaces:
use_distill:
sources: [fastvideo.configs.pipelines.longcat.LongCatT2V480PConfig, fastvideo.configs.pipelines.longcat.LongCatT2V704PConfig]
scheduler_arch:
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
text_encoder_archs:
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
tokenizer_archs:
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
transformer_arch:
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
vae_arch:
sources: [fastvideo.configs.pipelines.sd35.SD35Config]
sources:
- fastvideo.configs.pipelines.sd35.SD35Config
- fastvideo.configs.pipelines.zimage.ZImagePipelineConfig
expand_timesteps:
sources:
- fastvideo.configs.pipelines.wan.FastWan2_2_TI2V_5B_Config
@@ -421,6 +431,8 @@ surfaces:
frame_receptive_field: "MagiHuman internal data-proxy receptive-field setting."
image_conditioning: "MagiHuman preset variant marker for reference-image conditioning."
ref_audio_offset: "MagiHuman internal data-proxy audio alignment offset."
scheduler_sigma_min: "Z-Image scheduler parity invariant; not part of the public typed inference API."
scheduler_use_reference_discrete_timesteps: "Z-Image scheduler parity invariant; not part of the public typed inference API."
sr_local_attn_layers: "MagiHuman SR internal sparse-attention layer selection."
text_offset: "MagiHuman internal data-proxy text alignment offset."
vae_stride: "MagiHuman internal VAE/data-proxy stride setting."
@@ -438,6 +450,7 @@ surfaces:
grid_sizes: request.inputs.grid_sizes
pose: request.inputs.pose
c2ws_plucker_emb: request.inputs.c2ws_plucker_emb
action_path: request.inputs.action_path
refine_from: request.inputs.refine_from
stage1_video: request.inputs.stage1_video
prompt: request.prompt
@@ -447,6 +460,7 @@ surfaces:
output_video_name: request.output.output_video_name
num_videos_per_prompt: request.sampling.num_videos_per_prompt
seed: request.sampling.seed
max_sequence_length: request.sampling.max_sequence_length
num_frames: request.sampling.num_frames
height: request.sampling.height
width: request.sampling.width
@@ -456,7 +470,10 @@ surfaces:
num_inference_steps: request.sampling.num_inference_steps
num_inference_steps_sr: request.sampling.num_inference_steps_sr
guidance_scale: request.sampling.guidance_scale
batch_cfg: request.sampling.batch_cfg
guidance_scale_2: request.sampling.guidance_scale_2
cfg_normalization: request.sampling.cfg_normalization
cfg_truncation: request.sampling.cfg_truncation
guidance_rescale: request.sampling.guidance_rescale
use_embedded_guidance: request.sampling.use_embedded_guidance
true_cfg_scale: request.sampling.true_cfg_scale
@@ -508,7 +525,6 @@ surfaces:
internal_only:
data_type: "Derived from the request shape and not a public input."
latents: "Pre-generated diffusion latents supplied by parity/debug harnesses; not a public input."
max_sequence_length: "Model-specific text-encoder sequence cap; not part of the public typed inference API."
sampling_param_extensions: {}
+6
View File
@@ -2,6 +2,12 @@
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computation, enabling much faster video generation.
!!! tip "Attn-QAT DMD2 workflow"
The modular trainer also provides a Wan2.1 MixKit recipe that first
fine-tunes with fake-quantized attention, then distills the student to
timesteps `[1000, 757, 522]` while teacher and critic remain on Flash
Attention. See [Attn-QAT Training](../training/attn_qat.md).
## 📊 Model Overview
We provide two distilled models:
+147
View File
@@ -0,0 +1,147 @@
# Attn-QAT Training
Attn-QAT simulates low-bit attention during training while keeping the rest of
the training method unchanged. In the modular `fastvideo/train` framework it is
a per-role model option, not a separate training method: supervised fine-tuning
and DMD2 still own their losses and optimizer cadence.
This guide covers the QAD Wan2.1-T2V-1.3B MixKit workflow:
1. run a 4,000-step supervised Attn-QAT fine-tune;
2. export the stage-1 DCP checkpoint to Diffusers format; and
3. distill the student to three denoising steps with DMD2.
The ready-to-run configs and wrappers are in
`examples/train/scenario/qad_wan2_1_mixkit/`.
## Role-local attention backends
A DMD2 run owns three independent model roles. Configure the attention backend
on each role so fake quantization is applied only to the student:
```yaml
models:
student:
attention_backend: ATTN_QAT_TRAIN
teacher:
attention_backend: FLASH_ATTN
critic:
attention_backend: FLASH_ATTN
```
The override is active only while that role's transformer is constructed, then
the previous process-wide backend is restored. This lets student, teacher, and
critic use different implementations in one process. Invalid role-level names
fail during configuration instead of silently selecting another backend.
See [Training Infrastructure](train_infra.md) for the complete model-role
configuration reference.
## Prerequisites
- Install FastVideo and make the `fastvideo-kernel` Python package importable.
`ATTN_QAT_TRAIN` intentionally fails instead of falling back to dense
attention when its kernel cannot be loaded.
- Prepare the precomputed MixKit VAE latents and text embeddings.
- Run the commands below from the repository root. The supplied recipe expects
four GPUs by default; set `NUM_GPUS` to override it.
Download the published preprocessed dataset:
```bash
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
```
## Stage 1: supervised Attn-QAT fine-tuning
The stage-1 config uses `ATTN_QAT_TRAIN` on the student, sequence parallelism
across four GPUs, FP32 master weights, and 4,000 optimizer steps:
```bash
NUM_GPUS=4 \
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
```
Pass a dataset directory as the first positional argument when it differs from
the default:
```bash
NUM_GPUS=4 \
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh \
/path/to/combined_parquet_dataset
```
The wrapper calls `examples/train/run.sh`; the YAML file remains the source of
truth for optimizer, validation, checkpointing, and distributed settings.
## Export the stage-1 checkpoint
Modular training checkpoints use Distributed Checkpoint (DCP) format. Export
the student before using it to initialize stage 2:
```bash
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh \
checkpoints/wan_t2v_qat_finetune/checkpoint-4000 \
checkpoints/wan_t2v_qat_finetune/diffusers
```
Both arguments are optional; the command above shows their defaults.
## Stage 2: three-step DMD2 distillation
Stage 2 loads the exported student weights, keeps Attn-QAT on the student, and
uses Flash Attention for the teacher and critic:
```bash
NUM_GPUS=4 \
bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh \
data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset \
checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
```
The migrated recipe preserves these behaviors:
| Behavior | Modular configuration |
|---|---|
| Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` |
| Teacher and critic full-precision attention | Role-local `FLASH_ATTN` |
| Generator update every five critic steps | `method.generator_update_interval: 5` |
| Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` |
| Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` |
| Legacy guidance `cond + 2(cond - uncond)` | Standard CFG scale `3.0` |
| Stage handoff | DCP checkpoint to Diffusers export to student override weights |
The timestep ratios apply to randomly sampled teacher and critic score
timesteps; `dmd_denoising_steps` separately controls the student rollout. See
[DMD Distillation](../distillation/dmd.md) for general DMD concepts.
## Architecture-specific Triton routing
The training kernel is runtime-JIT-compiled Triton code and selects its route on
every call. It supports different query and key/value sequence lengths for
cross-attention; key and value must have the same sequence length.
| Hardware/configuration | Route |
|---|---|
| SM100, validated non-causal BF16 QAT configuration with head dimension 128 | Large-tile forward and split 64x64 backward; optimized backward requires a 16-aligned KV length |
| SM120, including RTX 5090 | Previous forward tiling with joined quantized/STE P@V operations and a shallower backward pipeline for long sequences |
| Unsupported configurations | Previous Triton implementation |
Warp specialization is disabled automatically on SM100 and SM120 because the
Triton 3.7 NVWS compiler pass aborts for this kernel on Blackwell. No user
setting is required.
The available tuning and comparison controls are:
| Environment variable | Default | Effect |
|---|---|---|
| `FASTVIDEO_ATTN_QAT_FWD_MODE` | `fast` | Selects `fast`, `balanced`, or `reference` forward tiling on the SM100 optimized route |
| `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` | `0` | Set to `1` to recompute reference-order softmax statistics and keep `dV` bitwise-compatible on the SM100 optimized route |
| `FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED` | `1` | Set to `0` to force the previous SM100 forward and backward for comparison |
| `FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV` | `1` | Set to `0` to compare SM120 against the split P@V path |
The first invocation JIT-compiles the selected configuration; later calls reuse
the Triton cache. To measure the production shape, run
`python benchmarks/benchmark_attn_qat_train.py` from `fastvideo-kernel/`.
For import and backend-selection failures, see [Debugging](../utilities/debugging.md).
+13
View File
@@ -62,6 +62,18 @@ bash examples/training/finetune/wan_t2v_1.3B/crush_smol/finetune_t2v.sh
- Gradient checkpointing: `--enable_gradient_checkpointing_type "full"`
- Memory scaling: Increase `--sp_size` or reduce `--num_latent_t` to fit in memory
## Attention Quantization-Aware Training
Attn-QAT fine-tunes a model while simulating low-bit attention in the forward
and backward passes. The modular trainer can select the backend per model role,
so a later DMD2 stage can keep fake quantization on the student while the
teacher and critic use Flash Attention.
The ready-to-run Wan2.1 MixKit workflow includes supervised fine-tuning,
checkpoint export, and three-step DMD2 distillation:
**→ [Follow the Attn-QAT training guide](attn_qat.md)**
## LoRA Finetuning
LoRA (Low-Rank Adaptation) trains lightweight adapters while keeping the base model frozen. This significantly reduces memory usage and training time.
@@ -166,6 +178,7 @@ Ready-to-run training scripts are available for multiple models:
| Wan2.1 I2V 14B | I2V | `examples/training/finetune/wan_i2v_14B_480p/crush_smol/` |
| Wan2.1-Fun 1.3B InP | I2V | `examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/` |
| Wan2.1 VSA | T2V/I2V | `examples/training/finetune/Wan2.1-VSA/Wan-Syn-Data/` |
| Wan2.1 T2V 1.3B Attn-QAT | QAT SFT + DMD2 | `examples/train/scenario/qad_wan2_1_mixkit/` |
Each example includes:
+6 -1
View File
@@ -43,6 +43,9 @@ Ready-to-run examples with preprocessing scripts, training launchers, and valida
**→ [Browse all training examples](examples/examples_training_index.md)**
For the complete two-stage Wan2.1 MixKit quantization-aware workflow, see
**[Attn-QAT Training](attn_qat.md)**.
Each example includes:
- `download_dataset.sh` — download sample data
@@ -59,9 +62,11 @@ FastVideo supports several training approaches:
| **Full finetune** | Adapt entire model to a new domain or style |
| **LoRA finetune** | Lightweight adaptation with frozen base weights |
| **VSA finetune** | Finetune with Variable Sparse Attention for efficiency |
| **Attn-QAT** | Train with fake-quantized attention, optionally followed by DMD2 distillation |
## Next Steps
1. **Get started**: Pick an example from the [training examples index](examples/examples_training_index.md)
2. **Prepare data**: Follow [data preprocessing](data_preprocess.md) for your own dataset
3. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
3. **Train with quantized attention**: Follow the [Attn-QAT two-stage recipe](attn_qat.md)
4. **Run inference**: After training, see [inference examples](../inference/examples/examples_inference_index.md)
+3
View File
@@ -81,6 +81,7 @@ Common model parameters:
| `disable_custom_init_weights` | `false` | Skip custom weight initialization (use for teacher/critic) |
| `flow_shift` | `3.0` | Timestep shifting factor |
| `enable_gradient_checkpointing_type` | `null` | Gradient checkpointing (`"full"` or `null`) |
| `attention_backend` | `null` | Optional role-local backend for Wan models (for example `ATTN_QAT_TRAIN`); overrides the process default only while this role's transformer is built |
Which roles are needed depends on the training method:
@@ -298,6 +299,8 @@ method:
| `dmd_denoising_steps` | *(required)* | Timestep schedule for student rollout |
| `generator_update_interval` | `1` | Update student every N critic steps |
| `real_score_guidance_scale` | `1.0` | CFG scale for teacher predictions |
| `min_timestep_ratio` | `0.0` | Lower bound for randomly sampled teacher/critic score timesteps |
| `max_timestep_ratio` | `1.0` | Upper bound for randomly sampled teacher/critic score timesteps |
| `fake_score_learning_rate` | *(required)* | Critic optimizer learning rate |
| `fake_score_betas` | *(required)* | Critic optimizer Adam betas |
| `fake_score_lr_scheduler` | *(required)* | Critic LR scheduler type |
+4 -1
View File
@@ -78,7 +78,10 @@ If forcing a backend fails, verify optional dependencies are installed:
- `SAGE_ATTN_THREE`: upstream `sageattn3` package
- `ATTN_QAT_INFER`: `fastvideo-kernel` checkout/source install that exposes
`attn_qat_infer`
- `ATTN_QAT_TRAIN`: `fastvideo-kernel` install exposing `fastvideo_kernel`
- `ATTN_QAT_TRAIN`: `fastvideo-kernel`; its runtime-JIT Triton implementation
selects an optimized route on SM100, joins the quantized and STE P@V paths on
SM120, and retains the previous route for unsupported configurations. See
[Attn-QAT Training](../training/attn_qat.md) for architecture controls.
As a fallback, use:
+17
View File
@@ -0,0 +1,17 @@
# LingBot World 2 Example Dataset
These files were copied unchanged from the LingBot World 2 repository for the
FastVideo causal-fast inference example.
- Repository: `https://github.com/Robbyant/lingbot-world-v2.git`
- Source commit: `94f43115de8d4a4f9f282126528c300a0b232c5f`
- Source directory: `examples/03`
## Files
- `image.jpg`: source image for image-to-video generation. SHA-256:
`6ee3dacfef32cfef504dd698adb8a660cf15f686535c52fed4903fef27c0edd0`
- `poses.npy`: camera-to-world trajectory matrices. SHA-256:
`bd0a23a696e184b0b43e7767eb432bfe644690560fe327fa96961affc941c404`
- `intrinsics.npy`: camera intrinsic parameters. SHA-256:
`821fca6cf957ae8fbb1181307f02479efb1705e04c9e05734cd02fb43462e082`
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.
Binary file not shown.
@@ -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,52 @@
"""Generate a five-second Dense LingBot-Video clip with the official defaults."""
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
def parse_args() -> argparse.Namespace:
"""Parse the converted checkpoint and output paths for the sample."""
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--model-path",
type=Path,
required=True,
help="Path to a converted Dense LingBot-Video checkpoint.",
)
parser.add_argument(
"--output-path",
type=Path,
default=Path("outputs/lingbot-video/dense-t2v"),
help="Directory for the generated video.",
)
return parser.parse_args()
def main() -> None:
"""Load the converted Dense checkpoint and generate the default T2V sample."""
args = parse_args()
generator = VideoGenerator.from_pretrained(
str(args.model_path),
num_gpus=1,
use_fsdp_inference=False,
text_encoder_cpu_offload=True,
vae_cpu_offload=False,
pin_cpu_memory=True,
)
try:
generator.generate({
"prompt": "A red fox runs through fresh snow at sunrise.",
"output": {
"output_path": str(args.output_path),
"save_video": True,
"return_frames": False,
},
})
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,52 @@
# SPDX-License-Identifier: Apache-2.0
"""Run LingBot World 2 14B causal-fast I2V generation with FastVideo."""
import os
from pathlib import Path
from fastvideo import VideoGenerator
REPO_ROOT = Path(__file__).resolve().parents[3]
DATASET_DIR = REPO_ROOT / "examples" / "dataset" / "lingbotworld2"
OUTPUT_PATH = REPO_ROOT / "outputs" / "lingbotworld2_causal_fast.mp4"
def main() -> None:
"""Load the native FastVideo LingBot World 2 causal-fast pipeline and generate one video."""
generator = VideoGenerator.from_pretrained(
os.environ["LINGBOTWORLD2_MODEL_PATH"],
num_gpus=8,
sp_size=8,
hsdp_shard_dim=8,
use_fsdp_inference=True,
dit_layerwise_offload=False,
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
pin_cpu_memory=True,
override_pipeline_cls_name="LingBotWorld2CausalFastPipeline",
)
try:
generator.generate_video(
"A serene lakeside scene with a lone tree standing in calm water, surrounded by distant snow-capped mountains under a bright blue sky with drifting white clouds; gentle ripples reflect the tree and sky, creating a tranquil, meditative atmosphere.",
image_path=str(DATASET_DIR / "image.jpg"),
action_path=str(DATASET_DIR),
output_path=str(OUTPUT_PATH),
save_video=True,
height=480,
width=832,
num_frames=65,
num_inference_steps=4,
guidance_scale=1.0,
negative_prompt="",
fps=16,
seed=42,
)
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+96
View File
@@ -0,0 +1,96 @@
# SPDX-License-Identifier: Apache-2.0
"""Run Z-Image-Turbo text-to-image generation through FastVideo.
User story:
"I want the official Z-Image-Turbo defaults and a deterministic PNG from
a local or Hugging Face checkpoint."
"""
import argparse
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
DEFAULT_PROMPT = (
"Young Chinese woman in red Hanfu, intricate embroidery. Impeccable makeup, red floral forehead pattern. "
"Elaborate high bun, golden phoenix headdress, red flowers, beads. Holds round folding fan with lady, trees, bird. "
"Neon lightning-bolt lamp (⚡️), bright yellow glow, above extended left palm. Soft-lit outdoor night background, "
"silhouetted tiered pagoda (西安大雁塔), blurred colorful distant lights."
)
DEFAULT_REVISION = "f332072aa78be7aecdf3ee76d5c247082da564a6"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run Z-Image-Turbo text-to-image generation.")
parser.add_argument("--model-path", default="Tongyi-MAI/Z-Image-Turbo")
parser.add_argument("--revision", default=DEFAULT_REVISION)
parser.add_argument("--output", default="outputs/zimage/zimage_turbo.png")
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
parser.add_argument("--negative-prompt", default="")
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=8)
parser.add_argument("--guidance-scale", type=float, default=0.0)
parser.add_argument("--max-sequence-length", type=int, default=512)
parser.add_argument("--cfg-normalization", action=argparse.BooleanOptionalAction, default=False)
parser.add_argument("--cfg-truncation", type=float, default=1.0)
parser.add_argument("--seed", type=int, default=42)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
revision=args.revision,
engine=EngineConfig(
num_gpus=1,
parallelism=ParallelismConfig(tp_size=1, sp_size=1),
use_fsdp_inference=False,
),
# The model registry selects the native zimage_turbo preset.
pipeline=PipelineSelection(workload_type="t2i"),
))
try:
generator.generate(
GenerationRequest(
prompt=args.prompt,
negative_prompt=args.negative_prompt,
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
max_sequence_length=args.max_sequence_length,
cfg_normalization=args.cfg_normalization,
cfg_truncation=args.cfg_truncation,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output),
save_video=True,
return_frames=False,
),
))
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+4
View File
@@ -91,3 +91,7 @@ examples/train/
```
See `configs/README.md` and `scenario/README.md` for details.
The featured QAD Wan2.1 MixKit scenario runs Attn-QAT supervised fine-tuning,
checkpoint export, and three-step DMD2. See the
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
+3
View File
@@ -25,6 +25,7 @@ models:
disable_custom_init_weights: false # default: false
flow_shift: 3.0 # default: 3.0
enable_gradient_checkpointing_type: null # default: null (falls back to training.model)
attention_backend: null # default: null (global/default); role-local when set
teacher:
_target_: fastvideo.train.models.wan.WanModel
@@ -52,6 +53,8 @@ method:
rollout_mode: simulate # required: "simulate" or "data_latent"
generator_update_interval: 5 # default: 1
dmd_denoising_steps: [1000, 750, 500, 250] # SDE timestep schedule
min_timestep_ratio: 0.0 # score-model timestep lower bound
max_timestep_ratio: 1.0 # score-model timestep upper bound
# Critic optimizer (all required — no fallback)
fake_score_learning_rate: 8.0e-6
@@ -0,0 +1,95 @@
# LTX-2.3 T2V overfitting test config.
#
# Overfits on a single short video (480x832, 81 frames @ 24fps) to
# verify the LTX-2 training plugin works end-to-end on an LTX-2.3
# checkpoint (gated attention, cross-attention AdaLN, 4096-d
# post-connector text embeddings, no in-DiT caption projection).
# Uses the distilled checkpoint (validation is 8 sampling steps).
#
# Preprocess data first (writes data/ltx2_3_overfit_preprocessed):
# CUDA_VISIBLE_DEVICES=0 \
# LTX2_OVERFIT_MODEL=FastVideo/LTX-2.3-Distilled-Diffusers \
# LTX2_OVERFIT_OUTPUT_DIR=data/ltx2_3_overfit_preprocessed \
# python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
#
# Run:
# NUM_GPUS=4 bash examples/train/run.sh examples/train/configs/overfit_ltx2_3_t2v.yaml
models:
student:
_target_: fastvideo.train.models.ltx2.LTX2Model
init_from: FastVideo/LTX-2.3-Distilled-Diffusers
trainable: true
enable_gradient_checkpointing_type: full
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/ltx2_3_overfit_preprocessed
dataloader_num_workers: 0
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
training_cfg_rate: 0.0
seed: 42
num_latent_t: 11 # (81 - 1) / 8 + 1
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: 300
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/ltx2_3_overfit
# A full training-state checkpoint is ~150GB for the 13B trainable
# video branch; disable saves for the overfit smoke run.
training_state_checkpointing_steps: 0
checkpoints_total_limit: 1
tracker:
trackers: [wandb]
project_name: fastvideo_ltx2
run_name: ltx2_3_overfit
model:
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
# so the default noise-minus-clean target reproduces the official
# unweighted masked-MSE (mask is all-ones for plain T2V).
precondition_outputs: false
enable_gradient_checkpointing_type: full
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.ltx2.ltx2_pipeline.LTX2Pipeline
dataset_file: data/ltx2_3_overfit_preprocessed/validation_prompts.json
every_steps: 50
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Required so the LTX2T2VConfig pipeline config is resolved from
# init_from (without a `pipeline:` key the loader falls back to a
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
pipeline: {}
@@ -0,0 +1,90 @@
# LTX-2 T2V overfitting test config.
#
# Overfits on a single short video (480x832, 81 frames @ 24fps) to
# verify the LTX-2 training plugin works end-to-end. Uses the
# distilled checkpoint (validation is 8 sampling steps, single pass).
#
# Preprocess data first (writes data/ltx2_overfit_preprocessed):
# CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
#
# Run:
# NUM_GPUS=4 bash examples/train/run.sh examples/train/configs/overfit_ltx2_t2v.yaml
models:
student:
_target_: fastvideo.train.models.ltx2.LTX2Model
init_from: FastVideo/LTX2-Distilled-Diffusers
trainable: true
enable_gradient_checkpointing_type: full
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/ltx2_overfit_preprocessed
dataloader_num_workers: 0
train_batch_size: 1
# LTX2Model requires 0.0: CFG dropout would zero post-connector
# embeddings, which is not the model's unconditional input.
training_cfg_rate: 0.0
seed: 42
num_latent_t: 11 # (81 - 1) / 8 + 1
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: 300
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/ltx2_overfit
# A full training-state checkpoint is ~150GB for the 13B trainable
# video branch; disable saves for the overfit smoke run.
training_state_checkpointing_steps: 0
checkpoints_total_limit: 1
tracker:
trackers: [wandb]
project_name: fastvideo_ltx2
run_name: ltx2_overfit
model:
# LTX2Model.predict_noise converts the DiT's x0 output to velocity,
# so the default noise-minus-clean target reproduces the official
# unweighted masked-MSE (mask is all-ones for plain T2V).
precondition_outputs: false
enable_gradient_checkpointing_type: full
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.ltx2.ltx2_pipeline.LTX2Pipeline
dataset_file: data/ltx2_overfit_preprocessed/validation_prompts.json
every_steps: 50
sampling_steps: [8]
guidance_scale: 1.0
num_frames: 81
# Required so the LTX2T2VConfig pipeline config is resolved from
# init_from (without a `pipeline:` key the loader falls back to a
# generic PipelineConfig and the LTX-2 DiT cannot be constructed).
pipeline: {}
+5 -2
View File
@@ -5,9 +5,12 @@ all configs, scripts, and data needed to run a complete workflow.
```
scenario/
└── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
├── ode_init_self_forcing_wan_causal/ # KD → export → Self-Forcing
└── qad_wan2_1_mixkit/ # Attn-QAT SFT → export → DMD2
```
See the `usage.md` inside each scenario for step-by-step instructions.
See the `usage.md` inside each scenario for step-by-step instructions. The QAD
workflow is also documented in the website-visible
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
For single-step configs, see `examples/train/configs/`.
+15
View File
@@ -0,0 +1,15 @@
#!/usr/bin/env bash
# Export a modular-trainer DCP checkpoint for stage-2 initialization.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
cd "${REPO_ROOT}"
CHECKPOINT_DIR=${1:-checkpoints/wan_t2v_qat_finetune/checkpoint-4000}
OUTPUT_DIR=${2:-checkpoints/wan_t2v_qat_finetune/diffusers}
python -m fastvideo.train.entrypoint.dcp_to_diffusers \
--role student \
--checkpoint "${CHECKPOINT_DIR}" \
--output-dir "${OUTPUT_DIR}"
+20
View File
@@ -0,0 +1,20 @@
#!/usr/bin/env bash
# Run the modular Attn-QAT finetune recipe.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
cd "${REPO_ROOT}"
DATA_DIR=${1:-data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset}
NUM_GPUS=${NUM_GPUS:-4}
export NUM_GPUS
export FASTVIDEO_ATTN_QAT_FWD_EXACT_M=${FASTVIDEO_ATTN_QAT_FWD_EXACT_M:-0}
bash "${REPO_ROOT}/examples/train/run.sh" \
"${SCRIPT_DIR}/stage1_attn_qat_finetune.yaml" \
--training.data.data_path "${DATA_DIR}" \
--training.distributed.num_gpus "${NUM_GPUS}" \
--training.distributed.sp_size "${NUM_GPUS}" \
--training.distributed.hsdp_replicate_dim 1 \
--training.distributed.hsdp_shard_dim "${NUM_GPUS}"
+27
View File
@@ -0,0 +1,27 @@
#!/usr/bin/env bash
# Run modular DMD2 with Attn-QAT on the student only.
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../../.." && pwd)"
cd "${REPO_ROOT}"
DATA_DIR=${1:-data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset}
INIT_WEIGHTS=${2:-checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors}
NUM_GPUS=${NUM_GPUS:-4}
export NUM_GPUS
if [[ ! -f "${INIT_WEIGHTS}" ]]; then
echo "Missing exported stage-1 weights: ${INIT_WEIGHTS}" >&2
echo "Run export_stage1.sh before stage 2." >&2
exit 1
fi
bash "${REPO_ROOT}/examples/train/run.sh" \
"${SCRIPT_DIR}/stage2_attn_qat_dmd.yaml" \
--models.student.transformer_override_safetensor "${INIT_WEIGHTS}" \
--training.data.data_path "${DATA_DIR}" \
--training.distributed.num_gpus "${NUM_GPUS}" \
--training.distributed.sp_size 1 \
--training.distributed.hsdp_replicate_dim "${NUM_GPUS}" \
--training.distributed.hsdp_shard_dim 1
@@ -0,0 +1,72 @@
# QAD stage 1: Attn-QAT finetune of Wan2.1-T2V-1.3B on MixKit.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
# Role-local: does not change the backend used by other models loaded in
# this process.
attention_backend: ATTN_QAT_TRAIN
method:
_target_: fastvideo.train.methods.fine_tuning.finetune.FineTuneMethod
training:
distributed:
num_gpus: 4
sp_size: 4
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
dataloader_num_workers: 1
train_batch_size: 1
training_cfg_rate: 0.1
seed: 1000
num_latent_t: 20
num_height: 480
num_width: 832
num_frames: 77
optimizer:
learning_rate: 1.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: checkpoints/wan_t2v_qat_finetune
training_state_checkpointing_steps: 500
checkpoints_total_limit: 0
tracker:
project_name: wan_t2v_qat_finetune
run_name: wan_t2v_qat_finetune
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.wan.wan_pipeline.WanPipeline
dataset_file: examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json
every_steps: 50
sampling_steps: [50]
guidance_scale: 5.0
pipeline:
flow_shift: 1
@@ -0,0 +1,101 @@
# QAD stage 2: distill the Attn-QAT student to three sampling steps.
#
# Export the stage-1 DCP checkpoint first, then point
# models.student.transformer_override_safetensor at the exported weight file.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
transformer_override_safetensor: checkpoints/wan_t2v_qat_finetune/diffusers/transformer/model.safetensors
trainable: true
attention_backend: ATTN_QAT_TRAIN
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
attention_backend: FLASH_ATTN
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
attention_backend: FLASH_ATTN
method:
_target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method
rollout_mode: data_latent
generator_update_interval: 5
dmd_denoising_steps: [1000, 757, 522]
min_timestep_ratio: 0.02
max_timestep_ratio: 0.98
# The modular DMD method uses standard CFG:
# uncond + scale * (cond - uncond). This is equivalent to the legacy
# recipe's cond + 2.0 * (cond - uncond).
real_score_guidance_scale: 3.0
# The legacy recipe inherited these values from its global optimizer.
fake_score_learning_rate: 2.0e-6
fake_score_betas: [0.9, 0.999]
fake_score_lr_scheduler: constant
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 4
hsdp_shard_dim: 1
data:
data_path: data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 20
num_height: 480
num_width: 832
num_frames: 77
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: 2000
gradient_accumulation_steps: 1
checkpoint:
output_dir: checkpoints/wan_t2v_distill_dmd_qat
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
tracker:
project_name: wan_t2v_distill_dmd_qat
run_name: wan_t2v_distill_dmd_qat
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.wan.wan_dmd_pipeline.WanDMDPipeline
dataset_file: examples/training/finetune/wan_t2v_1.3B/crush_smol/validation.json
every_steps: 200
sampling_steps: [3]
sampling_timesteps: [1000, 757, 522]
guidance_scale: 6.0
pipeline:
flow_shift: 3
@@ -0,0 +1,26 @@
# QAD Wan2.1 MixKit Attn-QAT
This scenario runs entirely on the modular `fastvideo/train` stack. The
student's attention backend is configured per role, so DMD2 can keep the
teacher and critic on Flash Attention while the student uses the fake-quantized
Attn-QAT kernel.
From the repository root:
```bash
# 1. Download the preprocessed MixKit data.
bash examples/training/finetune/wan_t2v_1.3B/mixkit/download_mixkit_data.sh
# 2. Stage 1: Attn-QAT supervised finetune.
NUM_GPUS=4 bash examples/train/scenario/qad_wan2_1_mixkit/run_stage1.sh
# 3. Export the stage-1 DCP checkpoint to a Diffusers weight file.
bash examples/train/scenario/qad_wan2_1_mixkit/export_stage1.sh
# 4. Stage 2: three-step Attn-QAT DMD2 distillation.
NUM_GPUS=4 bash examples/train/scenario/qad_wan2_1_mixkit/run_stage2.sh
```
The two YAML configs are also directly runnable through `examples/train/run.sh`.
The wrapper scripts only provide dataset/checkpoint paths and derive distributed
dimensions from `NUM_GPUS`.
@@ -52,7 +52,7 @@ for the full parameter reference.
## Train (QAT finetune)
With the data in place, run the quantization-aware finetune. The 4-bit attention
path is **config-driven** — selected purely by an env var, no monkey-patching:
path is **config-driven** and selected by an environment variable:
```bash
bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
@@ -60,10 +60,19 @@ bash examples/training/finetune/wan_t2v_1.3B/mixkit/finetune_qat.sh
NUM_GPUS=4 bash .../mixkit/finetune_qat.sh data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/
```
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` routes attention through the
fake-quantized Triton kernel (straight-through estimator), so the DiT learns to
absorb FP4 attention error. This kernel is Triton, so it runs on both `sm_100`
(B200/GB200) and `sm_120` (RTX 5090).
`FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN` keeps the fake-quantized Triton
forward and backward (straight-through estimator). Both kernels ship in
`fastvideo-kernel`: SM100 automatically uses the optimized Triton path for the
production non-causal, head-dimension-128 configuration, while SM120 GPUs such
as RTX 5090 join the quantized and STE P@V operations and use a shallower
backward pipeline for long sequences. Set
`FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV=0` to compare against the split P@V path.
The script defaults
`FASTVIDEO_ATTN_QAT_FWD_EXACT_M=0` for throughput; set it to `1` to reproduce
the previous forward softmax statistic and bitwise-compatible `dV`.
For the website-visible modular SFT-to-DMD2 workflow, see the
[Attn-QAT training guide](https://haoailab.com/FastVideo/training/attn_qat/).
## Train stage 2 (QAT DMD distillation to 3 steps)
@@ -71,8 +80,8 @@ Distill the QAT-finetuned generator down to **3 sampling steps**. Only the
generator is quantized (Attn-QAT); the teacher (`real_score`) and critic
(`fake_score`) stay full precision. This is enforced in the loader
(`component_loader.py`, via the `_loading_teacher_critic_model` flag), so the
same global `ATTN_QAT_TRAIN` env reaches **only** the generator — no per-model
flags or monkey-patching.
same global `ATTN_QAT_TRAIN` env reaches **only** the generator, with no
per-model flags.
```bash
# generator init = the stage-1 finetune checkpoint
@@ -2,20 +2,23 @@
# QAD recipe — quantization-aware finetune of Wan2.1-T2V-1.3B with fake-quant
# (Attn-QAT) attention.
#
# The 4-bit attention path is selected purely by env var (config-driven, no
# monkey-patching): FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN routes attention
# through the fake-quantized Triton kernel (straight-through estimator), so the
# DiT learns to absorb FP4 attention error instead of fighting it.
# The 4-bit attention path is selected by env var:
# FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN keeps the fake-quantized Triton
# forward and backward, so the DiT learns to absorb FP4 attention error instead
# of fighting it. SM100 selects the optimized kernels; RTX 5090 keeps the
# previous Triton implementation.
#
# Data: run download_mixkit_data.sh first (preprocessed Parquet).
#
# Verified end-to-end on Blackwell (GB200/sm_100): the ATTN_QAT_TRAIN backend is
# selected (not a fallback), forward+backward run, loss/grad are healthy, and
# validation generates videos. The kernel is Triton so it runs on sm_100 and
# sm_120 alike (the FP4 inference kernel, by contrast, is sm_120-only).
# validation generates videos.
set -euo pipefail
REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../../../.." && pwd)"
export PYTHONPATH="${REPO_ROOT}/fastvideo-kernel/python${PYTHONPATH:+:${PYTHONPATH}}"
export FASTVIDEO_ATTENTION_BACKEND=ATTN_QAT_TRAIN # <-- enables Attn-QAT training
export FASTVIDEO_ATTN_QAT_FWD_EXACT_M=${FASTVIDEO_ATTN_QAT_FWD_EXACT_M:-0}
export WANDB_MODE=${WANDB_MODE:-online}
export TOKENIZERS_PARALLELISM=false
@@ -23,6 +26,8 @@ MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR=${1:-"data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset/"}
VALIDATION_FILE="$(dirname "$0")/../crush_smol/validation.json"
NUM_GPUS=${NUM_GPUS:-4}
MAX_TRAIN_STEPS=${MAX_TRAIN_STEPS:-4000}
VALIDATION_SAMPLING_STEPS=${VALIDATION_SAMPLING_STEPS:-50}
torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
fastvideo/training/wan_training_pipeline.py \
@@ -30,15 +35,15 @@ torchrun --nnodes 1 --nproc_per_node "${NUM_GPUS}" \
--hsdp_replicate_dim 1 --hsdp_shard_dim "${NUM_GPUS}" \
--model_path "${MODEL_PATH}" --pretrained_model_name_or_path "${MODEL_PATH}" \
--data_path "${DATA_DIR}" --dataloader_num_workers 1 \
--max_train_steps 2000 --train_batch_size 1 --train_sp_batch_size 1 \
--max_train_steps "${MAX_TRAIN_STEPS}" --train_batch_size 1 --train_sp_batch_size 1 \
--gradient_accumulation_steps 1 \
--num_latent_t 20 --num_height 480 --num_width 832 --num_frames 77 \
--enable_gradient_checkpointing_type full \
--log_validation --validation_dataset_file "${VALIDATION_FILE}" \
--validation_steps 200 --validation_sampling_steps 50 --validation_guidance_scale 3.0 \
--learning_rate 5e-5 --mixed_precision bf16 --weight_decay 1e-4 --max_grad_norm 1.0 \
--validation_steps 50 --validation_sampling_steps "${VALIDATION_SAMPLING_STEPS}" --validation_guidance_scale 5.0 \
--learning_rate 1e-6 --mixed_precision bf16 --weight_decay 0.01 --max_grad_norm 1.0 \
--weight_only_checkpointing_steps 500 --training_state_checkpointing_steps 500 \
--tracker_project_name wan_t2v_qat_finetune --output_dir checkpoints/wan_t2v_qat_finetune \
--inference_mode False --training_cfg_rate 0.1 --not_apply_cfg_solver \
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 \
--dit_precision fp32 --num_euler_timesteps 50 --ema_start_step 0 --flow_shift 1 \
--multi_phased_distill_schedule "4000-1"
+27
View File
@@ -108,6 +108,33 @@ out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
## Benchmark
### Attn-QAT training
The default shape matches one sequence-parallel rank of the 4-GPU
Wan2.1-T2V-1.3B MixKit recipe (`B=1, H=3, L=31200, D=128`):
```bash
cd fastvideo-kernel
python benchmarks/benchmark_attn_qat_train.py
```
The benchmark reports both conventional attention FLOPs and the extra matrix
multiplications executed by the QAT straight-through path. Override
`--peak-tflops` when running on a GPU other than RTX 5090.
The QAT kernel is entirely Triton and routes by architecture at runtime. SM100
uses a large-tile forward and split 64x64 backward for the production
non-causal, head-dimension-128 configuration with a 16-aligned KV length. SM120
(including RTX 5090) keeps the previous tiling but joins the quantized and STE
P@V operations and uses a shallower backward pipeline for long sequences. Set
`FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV=0` to compare against the split P@V path.
Unsupported configurations retain the previous implementation. Set
`FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED=0` to benchmark that previous path on SM100. Forward tuning is available through
`FASTVIDEO_ATTN_QAT_FWD_MODE=fast|balanced|reference`; exact reference-order
softmax statistics are controlled by `FASTVIDEO_ATTN_QAT_FWD_EXACT_M` and are
disabled by default for maximum throughput. Set it to `1` for reference-order
statistics and bitwise-compatible `dV`.
### VSA (block-sparse) TFLOPs
After building/installing `fastvideo-kernel`, run:
@@ -0,0 +1,127 @@
#!/usr/bin/env python3
"""Benchmark the Attn-QAT training kernel on a single GPU.
Defaults model one rank of the 4-GPU Wan2.1-T2V-1.3B MixKit recipe:
``B=1, H=12/4, L=20*30*52, D=128``.
"""
from __future__ import annotations
import argparse
import statistics
import time
from collections.abc import Callable
import torch
from fastvideo_kernel.triton_kernels.attn_qat_train import attention
RTX_5090_DENSE_BF16_TFLOPS = 209.5
def _qat_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
consumer_blackwell = torch.cuda.get_device_capability()[0] == 12
return attention(
q,
k,
v,
False,
q.shape[-1]**-0.5,
True, # use_qat_qkv_backward
False, # smooth_k
not consumer_blackwell, # warp_specialize
True, # IS_QAT
False, # two_level_quant_P
True, # fake_quant_P
True, # use_high_prec_o
False, # smooth_q
False, # use_global_sf_P
False, # use_global_sf_QKV
)
def _measure_ms(fn: Callable[[], object], warmup: int, repeat: int) -> tuple[float, float, float]:
for _ in range(warmup):
fn()
torch.cuda.synchronize()
samples = []
for _ in range(repeat):
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
fn()
end.record()
end.synchronize()
samples.append(start.elapsed_time(end))
return statistics.median(samples), min(samples), max(samples)
def _format_result(
label: str,
timing_ms: tuple[float, float, float],
algorithmic_flops: int,
executed_matmul_flops: int,
peak_tflops: float,
) -> str:
median_ms, min_ms, max_ms = timing_ms
algorithmic_tflops = algorithmic_flops / (median_ms * 1e9)
executed_tflops = executed_matmul_flops / (median_ms * 1e9)
return (
f"{label}: {median_ms:.3f} ms (min={min_ms:.3f}, max={max_ms:.3f}), "
f"algorithmic={algorithmic_tflops:.2f} TFLOPS/{100 * algorithmic_tflops / peak_tflops:.2f}% MFU, "
f"executed_matmul={executed_tflops:.2f} TFLOPS/{100 * executed_tflops / peak_tflops:.2f}% MFU"
)
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--batch-size", type=int, default=1)
parser.add_argument("--heads", type=int, default=3, help="Heads per SP rank; Wan 1.3B has 12 total.")
parser.add_argument("--query-length", type=int, default=31_200)
parser.add_argument("--kv-length", type=int, help="Defaults to --query-length.")
parser.add_argument("--head-dim", type=int, default=128)
parser.add_argument("--warmup", type=int, default=3)
parser.add_argument("--repeat", type=int, default=10)
parser.add_argument(
"--peak-tflops",
type=float,
default=RTX_5090_DENSE_BF16_TFLOPS,
help="Dense BF16 Tensor TFLOPS with FP32 accumulation; default is RTX 5090 boost-clock peak.",
)
args = parser.parse_args()
kv_length = args.kv_length or args.query_length
torch.manual_seed(0)
q_shape = (args.batch_size, args.heads, args.query_length, args.head_dim)
kv_shape = (args.batch_size, args.heads, kv_length, args.head_dim)
q = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
v = torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16, requires_grad=True)
grad_out = torch.randn_like(q)
compile_start = time.perf_counter()
output = _qat_attention(q, k, v)
torch.cuda.synchronize()
compile_seconds = time.perf_counter() - compile_start
forward_ms = _measure_ms(lambda: _qat_attention(q, k, v), args.warmup, args.repeat)
backward_ms = _measure_ms(
lambda: torch.autograd.grad(output, (q, k, v), grad_out, retain_graph=True),
args.warmup,
args.repeat,
)
base_flops = args.batch_size * args.heads * args.query_length * kv_length * args.head_dim
# Conventional attention FLOPs are 4*base forward and 10*base backward.
# QAT additionally computes the STE high-precision P@V path in forward and
# the quantized-P dV path in backward, for 6*base and 14*base matmul FLOPs.
print(f"device: {torch.cuda.get_device_name()}")
print(f"q: {q_shape}; k/v: {kv_shape}; compile+first-forward: {compile_seconds:.3f} s")
print(_format_result("forward", forward_ms, 4 * base_flops, 6 * base_flops, args.peak_tflops))
print(_format_result("backward", backward_ms, 10 * base_flops, 14 * base_flops, args.peak_tflops))
if __name__ == "__main__":
main()
@@ -20,14 +20,68 @@ def supports_host_descriptor():
return is_cuda() and torch.cuda.get_device_capability()[0] >= 9
def is_sm100(device=None):
return is_cuda() and torch.cuda.get_device_capability(device) == (10, 0)
def is_blackwell():
return is_cuda() and torch.cuda.get_device_capability()[0] == 10
def is_consumer_blackwell():
return is_cuda() and torch.cuda.get_device_capability()[0] == 12
def is_hopper():
return is_cuda() and torch.cuda.get_device_capability()[0] == 9
def _sm100_optimization_enabled():
return os.environ.get("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "1") != "0"
def _sm100_exact_m_enabled():
return os.environ.get("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "0") != "0"
def _consumer_blackwell_join_qat_pv_enabled():
"""Return whether SM120 uses the joined quantized/STE P@V path."""
return os.environ.get("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "1") != "0"
def _use_sm100_optimized_qat(
device,
head_dim: int,
causal: bool,
is_qat: bool,
fake_quant_p: bool,
two_level_quant_p: bool,
use_global_sf_p: bool,
) -> bool:
"""Return whether this call matches the validated SM100 fast path."""
return (
_sm100_optimization_enabled()
and is_sm100(device)
and head_dim == 128
and not causal
and is_qat
and fake_quant_p
and not two_level_quant_p
and not use_global_sf_p
)
def _select_sm100_forward_config(n_ctx_q: int, n_ctx_kv: int, mode: str):
n_ctx = max(n_ctx_q, n_ctx_kv)
if mode == "reference":
return 32, 32, 4, 4 if n_ctx >= 16_384 else 5
if n_ctx <= 2_048:
return 32, 32, 4, 5
if mode == "balanced":
return 64, 32, 4, 4
return 128, 128, 8, 3
@triton.jit
def _mul_alpha(acc, alpha, BM: tl.constexpr, BN: tl.constexpr):
acc0, acc1 = acc.reshape([BM, 2, BN // 2]).permute(0, 2, 1).split()
@@ -47,7 +101,8 @@ def _attn_fwd_inner(acc, high_prec_acc, l_i, m_i, q, q_valid,
IS_QAT: tl.constexpr,
fake_quant_P: tl.constexpr = True,
two_level_quant_P: tl.constexpr = False,
use_global_sf_P: tl.constexpr = True):
use_global_sf_P: tl.constexpr = True,
JOIN_QAT_PV: tl.constexpr = False):
# range of values handled by this stage (kv blocks)
if STAGE == 1:
lo, hi = 0, start_m * BLOCK_M
@@ -113,10 +168,19 @@ def _attn_fwd_inner(acc, high_prec_acc, l_i, m_i, q, q_valid,
v = desc_v.load([offsetv_y, 0])
v = tl.where(kv_valid[:, None], v, 0.0)
p = p.to(dtype)
# note that this non transposed v for FP8 is only supported on Blackwell
acc = tl.dot(p, v.to(dtype), acc)
if IS_QAT:
high_prec_acc = tl.dot(high_prec_p, v, high_prec_acc)
# Keep the quantized and STE paths in one tensor-core operation. They
# share V, so joining along M avoids issuing two small, independent
# dot operations for every KV tile.
if IS_QAT and JOIN_QAT_PV:
joined_p = tl.join(p, high_prec_p).permute(2, 0, 1).reshape([2 * BLOCK_M, BLOCK_N])
joined_acc = tl.join(acc, high_prec_acc).permute(2, 0, 1).reshape([2 * BLOCK_M, HEAD_DIM])
joined_acc = tl.dot(joined_p, v.to(dtype), joined_acc)
acc, high_prec_acc = joined_acc.reshape([2, BLOCK_M, HEAD_DIM]).permute(1, 2, 0).split()
else:
# note that this non transposed v for FP8 is only supported on Blackwell
acc = tl.dot(p, v.to(dtype), acc)
if IS_QAT:
high_prec_acc = tl.dot(high_prec_p, v, high_prec_acc)
# update m_i and l_i
# place this at the end of the loop to reduce register pressure
l_i = l_i * alpha + l_ij
@@ -204,6 +268,7 @@ def _attn_fwd(sm_scale, M,
fake_quant_P: tl.constexpr = True,
two_level_quant_P: tl.constexpr = False,
use_global_sf_P: tl.constexpr = True,
JOIN_QAT_PV: tl.constexpr = False,
):
dtype = tl.float8e5 if FP8_OUTPUT else tl.bfloat16
tl.static_assert(BLOCK_N <= HEAD_DIM)
@@ -268,7 +333,7 @@ def _attn_fwd(sm_scale, M,
offset_y_kv, dtype, start_m, qk_scale,
BLOCK_M, HEAD_DIM, BLOCK_N,
4 - STAGE, offs_m, offs_n, N_CTX_KV,
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
)
# stage 2: on-band
if STAGE & 2:
@@ -278,7 +343,7 @@ def _attn_fwd(sm_scale, M,
offset_y_kv, dtype, start_m, qk_scale,
BLOCK_M, HEAD_DIM, BLOCK_N,
2, offs_m, offs_n, N_CTX_KV,
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P
warp_specialize, IS_HOPPER, IS_QAT, fake_quant_P, two_level_quant_P, use_global_sf_P, JOIN_QAT_PV
)
# epilogue
m_i += tl.math.log2(l_i)
@@ -292,6 +357,52 @@ def _attn_fwd(sm_scale, M,
desc_high_prec_o.store([off_hz, start_m * BLOCK_M, 0], high_prec_acc[None, :, :])
@triton.jit
def _attn_fwd_exact_m(
desc_q,
desc_k,
M,
sm_scale,
N_CTX_Q,
N_CTX_KV,
HEAD_DIM: tl.constexpr,
BLOCK_M: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""Reproduce the legacy 32x32 forward softmax statistic exactly."""
start_m = tl.program_id(0) * BLOCK_M
off_hz = tl.program_id(1)
offs_m = start_m + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)
q_valid = offs_m < N_CTX_Q
q_base = off_hz * N_CTX_Q
kv_base = off_hz * N_CTX_KV
q = desc_q.load([q_base + start_m, 0])
q = tl.where(q_valid[:, None], q, 0.0)
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
l_i = tl.full([BLOCK_M], 1.0, tl.float32)
qk_scale = sm_scale * 1.44269504
for start_n in tl.range(0, N_CTX_KV, BLOCK_N):
start_n = tl.multiple_of(start_n, BLOCK_N)
kv_valid = start_n + offs_n < N_CTX_KV
k = desc_k.load([kv_base + start_n, 0])
k = tl.where(kv_valid[:, None], k, 0.0)
qk = tl.dot(q, tl.trans(k))
qk = tl.where(kv_valid[None, :], qk, -1.0e6)
m_ij = tl.maximum(m_i, tl.max(qk, axis=1) * qk_scale)
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
l_ij = tl.sum(p.to(tl.bfloat16), axis=1)
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
m_i = m_ij
m_i += tl.math.log2(l_i)
tl.store(M + off_hz * N_CTX_Q + offs_m, m_i, mask=q_valid)
@triton.jit
def _attn_bwd_preprocess(O, DO,
Delta,
@@ -835,10 +946,33 @@ class _attention(torch.autograd.Function):
assert HEAD_DIM_Q == HEAD_DIM_K and HEAD_DIM_K == HEAD_DIM_V
assert HEAD_DIM_K in {16, 32, 64, 128, 256}
# Triton 3.7's NVWS pass aborts for this kernel on Blackwell. Keep the
# architecture guard next to the kernel so direct callers and the
# FastVideo backend follow the same supported path on sm_100/sm_120.
consumer_blackwell = is_consumer_blackwell()
blackwell = is_blackwell() or consumer_blackwell
warp_specialize = warp_specialize and not blackwell
# Support different sequence lengths for q and k/v (needed for cross attention)
N_CTX_Q = q.shape[2] # Query sequence length
N_CTX_KV = k.shape[2] # Key/Value sequence length (may differ from query)
assert k.shape[2] == v.shape[2], "k and v must have the same sequence length"
sm100_optimized = (
q.dtype == torch.bfloat16
and k.dtype == q.dtype
and v.dtype == q.dtype
and k.device == q.device
and v.device == q.device
and _use_sm100_optimized_qat(
q.device,
HEAD_DIM_K,
causal,
IS_QAT,
fake_quant_P,
two_level_quant_P,
use_global_sf_P,
)
)
# smoothing k from SageAttn
ctx.k_mean = None
@@ -924,7 +1058,20 @@ class _attention(torch.autograd.Function):
else:
extra_kern_args["maxnreg"] = 80
BLOCK_M, BLOCK_N = 32, 32
qkv_block_m, qkv_block_n = 32, 32
fwd_block_m, fwd_block_n = 32, 32
fwd_num_warps, fwd_num_stages = 4, 2
fwd_mode = "legacy"
if sm100_optimized:
fwd_mode = os.environ.get("FASTVIDEO_ATTN_QAT_FWD_MODE", "fast").lower()
if fwd_mode not in {"fast", "balanced", "reference"}:
raise ValueError(
f"FASTVIDEO_ATTN_QAT_FWD_MODE={fwd_mode!r} "
"(want fast|balanced|reference)"
)
fwd_block_m, fwd_block_n, fwd_num_warps, fwd_num_stages = _select_sm100_forward_config(
N_CTX_Q, N_CTX_KV, fwd_mode
)
if IS_QAT:
fake_q = torch.empty_like(q)
fake_k = torch.empty_like(k)
@@ -942,8 +1089,8 @@ class _attention(torch.autograd.Function):
desc_v = fake_v
H = q.shape[1]
grid_1 = (triton.cdiv(q.shape[2], BLOCK_M), q.shape[0] * q.shape[1], 1)
grid_2 = (triton.cdiv(k.shape[2], BLOCK_N), q.shape[0] * q.shape[1], 1)
grid_1 = (triton.cdiv(q.shape[2], qkv_block_m), q.shape[0] * q.shape[1], 1)
grid_2 = (triton.cdiv(k.shape[2], qkv_block_n), q.shape[0] * q.shape[1], 1)
fake_quantize_q[grid_1](
q, fake_q,
@@ -952,7 +1099,7 @@ class _attention(torch.autograd.Function):
fake_q.stride(0), fake_q.stride(1),
fake_q.stride(2), fake_q.stride(3),
H, N_CTX_Q,
BLOCK_M=BLOCK_M, HEAD_DIM=HEAD_DIM_K,
BLOCK_M=qkv_block_m, HEAD_DIM=HEAD_DIM_K,
use_global_sf=use_global_sf_QKV,
)
fake_quantize_kv[grid_2](
@@ -962,14 +1109,14 @@ class _attention(torch.autograd.Function):
fake_k.stride(0), fake_k.stride(1),
fake_k.stride(2), fake_k.stride(3),
H, N_CTX_KV,
BLOCK_N=BLOCK_N, HEAD_DIM=HEAD_DIM_K,
BLOCK_N=qkv_block_n, HEAD_DIM=HEAD_DIM_K,
use_global_sf=use_global_sf_QKV,
)
# Apply pre-hook to set block shapes on tensor descriptors
_host_descriptor_pre_hook({
"BLOCK_M": BLOCK_M,
"BLOCK_N": BLOCK_N,
"BLOCK_M": fwd_block_m,
"BLOCK_N": fwd_block_n,
"HEAD_DIM": HEAD_DIM_K,
"desc_q": desc_q,
"desc_k": desc_k,
@@ -986,7 +1133,7 @@ class _attention(torch.autograd.Function):
N_CTX_Q=N_CTX_Q,
N_CTX_KV=N_CTX_KV,
HEAD_DIM=HEAD_DIM_K,
BLOCK_M=BLOCK_M, BLOCK_N=BLOCK_N,
BLOCK_M=fwd_block_m, BLOCK_N=fwd_block_n,
FP8_OUTPUT=q.dtype == torch.float8_e5m2,
STAGE=stage,
warp_specialize=warp_specialize,
@@ -995,10 +1142,40 @@ class _attention(torch.autograd.Function):
fake_quant_P=fake_quant_P,
two_level_quant_P=two_level_quant_P,
use_global_sf_P=use_global_sf_P,
num_warps=4,
num_stages=2,
JOIN_QAT_PV=(consumer_blackwell and _consumer_blackwell_join_qat_pv_enabled()),
num_warps=fwd_num_warps,
num_stages=fwd_num_stages,
**extra_kern_args
)
exact_m = _sm100_exact_m_enabled()
if (
sm100_optimized
and fwd_mode != "reference"
and exact_m
and (fwd_block_m != 32 or fwd_block_n != 32)
):
# The large forward tile changes a legal reduction order. Restore
# the legacy statistic so dV remains bitwise-compatible while the
# two output paths retain the faster large-tile PV computation.
assert isinstance(desc_q, TensorDescriptor)
assert isinstance(desc_k, TensorDescriptor)
desc_q.block_shape = [32, HEAD_DIM_K]
desc_k.block_shape = [32, HEAD_DIM_K]
stats_grid = (triton.cdiv(N_CTX_Q, 32), q.shape[0] * q.shape[1])
_attn_fwd_exact_m[stats_grid](
desc_q,
desc_k,
M,
sm_scale,
N_CTX_Q,
N_CTX_KV,
HEAD_DIM=HEAD_DIM_K,
BLOCK_M=32,
BLOCK_N=32,
num_warps=8,
num_stages=4,
)
o_for_bwd = high_prec_o if IS_QAT and use_high_prec_o else o
if IS_QAT:
@@ -1018,6 +1195,7 @@ class _attention(torch.autograd.Function):
ctx.smooth_q = smooth_q
ctx.use_global_sf_P = use_global_sf_P
ctx.warp_specialize = warp_specialize
ctx.sm100_optimized = sm100_optimized
return o
@staticmethod
@@ -1032,7 +1210,10 @@ class _attention(torch.autograd.Function):
N_CTX_KV = k.shape[2]
assert k.shape[2] == v.shape[2], "k and v must have the same sequence length"
PRE_BLOCK = 128
NUM_STAGES = 3
# Long video sequences are occupancy-bound on consumer Blackwell: a
# third software-pipeline stage consumes shared memory without hiding
# additional latency. Shorter sequences retain the deeper pipeline.
NUM_STAGES = 2 if is_consumer_blackwell() and max(N_CTX_Q, N_CTX_KV) >= 8192 else 3
NUM_WARPS = 4
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 32, 32, 32
if not ctx.use_qat_qkv_backward:
@@ -1057,7 +1238,85 @@ class _attention(torch.autograd.Function):
# _, q_m = triton_group_mean(q)
q_m = q_m.repeat_interleave(q.shape[2] // q_m.shape[2], dim=2) # B,H,L,D
if N_CTX_Q == N_CTX_KV:
sm100_optimized_backward = (
getattr(ctx, "sm100_optimized", False)
and ctx.use_qat_qkv_backward
and not ctx.smooth_k
and not ctx.smooth_q
and N_CTX_KV % 16 == 0
)
if sm100_optimized_backward:
# Keeping dQ and dK/dV in separate programs allows 64x64 tiles
# without carrying all three fp32 accumulators at once. On SM100
# this is substantially faster than the legacy 32x32 combined
# self-attention program with the same math and BF16 parity bounds.
block_m, block_n = 64, 64
grid_dq = ((N_CTX_Q + block_m - 1) // block_m, 1, BATCH * N_HEAD)
_attn_bwd_dq_cross[grid_dq](
q,
arg_k,
v,
ctx.sm_scale,
do,
dq,
M,
delta,
q.stride(0),
k.stride(0),
q.stride(1),
k.stride(1),
q.stride(2),
k.stride(2),
q.stride(3),
k.stride(3),
N_HEAD,
N_CTX_Q,
N_CTX_KV,
ctx.k_mean,
BLOCK_M2=block_m,
BLOCK_N2=block_n,
HEAD_DIM=ctx.HEAD_DIM,
SMOOTH_K=False,
warp_specialize=False,
num_warps=8,
num_stages=2,
)
grid_dkdv = ((N_CTX_KV + block_n - 1) // block_n, 1, BATCH * N_HEAD)
_attn_bwd_dkdv_cross[grid_dkdv](
q,
arg_k,
v,
ctx.sm_scale,
do,
dk,
dv,
M,
delta,
q_m,
q.stride(0),
k.stride(0),
q.stride(1),
k.stride(1),
q.stride(2),
k.stride(2),
q.stride(3),
k.stride(3),
N_HEAD,
N_CTX_Q,
N_CTX_KV,
BLOCK_M1=block_m,
BLOCK_N1=block_n,
HEAD_DIM=ctx.HEAD_DIM,
IS_QAT=True,
two_level_quant_P=False,
fake_quant_P=True,
SMOOTH_Q=False,
use_global_sf_P=False,
warp_specialize=False,
num_warps=8,
num_stages=3,
)
elif N_CTX_Q == N_CTX_KV:
# Use existing kernel for self-attention (same sequence lengths)
grid = ((N_CTX_KV + BLOCK_N1 - 1) // BLOCK_N1, 1, BATCH * N_HEAD)
_attn_bwd[grid](
@@ -1074,10 +1333,10 @@ class _attention(torch.autograd.Function):
IS_QAT=ctx.IS_QAT,
SMOOTH_K=ctx.smooth_k,
two_level_quant_P=ctx.two_level_quant_P,
fake_quant_P=ctx.fake_quant_P,
SMOOTH_Q=ctx.smooth_q,
use_global_sf_P=ctx.use_global_sf_P,
warp_specialize=ctx.warp_specialize,
fake_quant_P=ctx.fake_quant_P,
SMOOTH_Q=ctx.smooth_q,
use_global_sf_P=ctx.use_global_sf_P,
warp_specialize=ctx.warp_specialize,
num_warps=NUM_WARPS,
num_stages=NUM_STAGES
)
@@ -0,0 +1,190 @@
# SPDX-License-Identifier: Apache-2.0
import math
import pytest
import torch
from fastvideo_kernel.triton_kernels import attn_qat_train as kernel
def _production_route_kwargs():
return {
"device": torch.device("cuda"),
"head_dim": 128,
"causal": False,
"is_qat": True,
"fake_quant_p": True,
"two_level_quant_p": False,
"use_global_sf_p": False,
}
def test_sm100_production_configuration_uses_optimized_route(monkeypatch):
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", raising=False)
assert kernel._use_sm100_optimized_qat(**_production_route_kwargs())
@pytest.mark.parametrize(
("override", "value"),
[
("head_dim", 64),
("causal", True),
("is_qat", False),
("fake_quant_p", False),
("two_level_quant_p", True),
("use_global_sf_p", True),
],
)
def test_unsupported_configuration_keeps_legacy_route(monkeypatch, override, value):
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
kwargs = _production_route_kwargs()
kwargs[override] = value
assert not kernel._use_sm100_optimized_qat(**kwargs)
def test_non_sm100_and_debug_switch_keep_legacy_route(monkeypatch):
kwargs = _production_route_kwargs()
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: False)
assert not kernel._use_sm100_optimized_qat(**kwargs)
monkeypatch.setattr(kernel, "is_sm100", lambda device=None: True)
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "0")
assert not kernel._use_sm100_optimized_qat(**kwargs)
def test_exact_m_is_opt_in(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", raising=False)
assert not kernel._sm100_exact_m_enabled()
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "1")
assert kernel._sm100_exact_m_enabled()
def test_sm120_joined_pv_is_enabled_by_default_and_can_be_disabled(monkeypatch):
monkeypatch.delenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", raising=False)
assert kernel._consumer_blackwell_join_qat_pv_enabled()
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "0")
assert not kernel._consumer_blackwell_join_qat_pv_enabled()
@pytest.mark.parametrize(
("n_ctx", "mode", "expected"),
[
(2_048, "fast", (32, 32, 4, 5)),
(4_096, "fast", (128, 128, 8, 3)),
(4_096, "balanced", (64, 32, 4, 4)),
(4_096, "reference", (32, 32, 4, 5)),
(31_200, "reference", (32, 32, 4, 4)),
],
)
def test_sm100_forward_config_selection(n_ctx, mode, expected):
assert kernel._select_sm100_forward_config(n_ctx, n_ctx, mode) == expected
@pytest.mark.skipif(
not torch.cuda.is_available() or torch.cuda.get_device_capability() != (10, 0),
reason="SM100 parity test",
)
@pytest.mark.parametrize(("q_length", "kv_length"), [(2_112, 2_112), (2_112, 2_080)])
def test_sm100_optimized_forward_backward_matches_legacy(monkeypatch, q_length, kv_length):
torch.manual_seed(7)
q_shape = (1, 1, q_length, 128)
kv_shape = (1, 1, kv_length, 128)
inputs = [
torch.randn(q_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
]
grad_out = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
flags = (
True, # use_qat_qkv_backward
False, # smooth_k
True, # warp_specialize (disabled internally on Blackwell)
True, # IS_QAT
False, # two_level_quant_P
True, # fake_quant_P
True, # use_high_prec_o
False, # smooth_q
False, # use_global_sf_P
False, # use_global_sf_QKV
)
def run(optimized: bool):
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM100_OPTIMIZED", "1" if optimized else "0")
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_MODE", "fast")
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_FWD_EXACT_M", "1")
q, k, v = [tensor.clone().requires_grad_(True) for tensor in inputs]
output = kernel.attention(
q,
k,
v,
False,
1.0 / math.sqrt(q_shape[-1]),
*flags,
)
output.backward(grad_out)
return output.detach(), q.grad, k.grad, v.grad
legacy = run(False)
optimized = run(True)
assert (optimized[0].float() - legacy[0].float()).abs().max().item() <= 1e-2
assert (optimized[1].float() - legacy[1].float()).abs().max().item() <= 4e-3
assert (optimized[2].float() - legacy[2].float()).abs().max().item() <= 4e-3
assert torch.equal(optimized[3], legacy[3])
@pytest.mark.skipif(
not torch.cuda.is_available() or torch.cuda.get_device_capability()[0] != 12,
reason="SM120 parity test",
)
@pytest.mark.parametrize(("q_length", "kv_length"), [(2_112, 2_112), (2_112, 2_080)])
def test_sm120_joined_pv_forward_backward_matches_split_path(monkeypatch, q_length, kv_length):
torch.manual_seed(11)
q_shape = (1, 1, q_length, 128)
kv_shape = (1, 1, kv_length, 128)
inputs = [
torch.randn(q_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
torch.randn(kv_shape, device="cuda", dtype=torch.bfloat16),
]
grad_out = torch.randn(q_shape, device="cuda", dtype=torch.bfloat16)
flags = (
True, # use_qat_qkv_backward
False, # smooth_k
True, # warp_specialize (disabled internally on Blackwell)
True, # IS_QAT
False, # two_level_quant_P
True, # fake_quant_P
True, # use_high_prec_o
False, # smooth_q
False, # use_global_sf_P
False, # use_global_sf_QKV
)
def run(joined_pv: bool):
monkeypatch.setenv("FASTVIDEO_ATTN_QAT_SM120_JOIN_QAT_PV", "1" if joined_pv else "0")
q, k, v = [tensor.clone().requires_grad_(True) for tensor in inputs]
output = kernel.attention(
q,
k,
v,
False,
1.0 / math.sqrt(q_shape[-1]),
*flags,
)
output.backward(grad_out)
return output.detach(), q.grad, k.grad, v.grad
split = run(False)
joined = run(True)
assert torch.equal(joined[0], split[0])
assert torch.equal(joined[1], split[1])
assert torch.equal(joined[2], split[2])
assert torch.equal(joined[3], split[3])
+2
View File
@@ -304,6 +304,8 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
for key in _LTX2_REFINE_FLAT_KEYS:
if key in refine:
kwargs[f"ltx2_refine_{key}"] = refine[key]
if "enabled" in refine:
kwargs["refine_enabled"] = refine["enabled"]
kwargs.update(preset_overrides)
kwargs.update(deepcopy(normalized.pipeline.experimental))
return FastVideoArgs.from_kwargs(**kwargs)
+26 -1
View File
@@ -51,8 +51,9 @@ class SamplingParam:
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
# Camera control inputs (LingBotWorld)
# Camera control inputs (LingBotWorld and LingBotWorld2)
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
action_path: str | None = None # Directory containing poses.npy and intrinsics.npy
# Refine inputs (LongCat 480p->720p upscaling)
# Path-based refine (load stage1 video from disk, e.g. MP4)
@@ -89,7 +90,13 @@ class SamplingParam:
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
batch_cfg: bool = False
guidance_scale_2: float | None = None
# Z-Image CFG controls. ``cfg_normalization=True`` caps the guided
# prediction norm at the positive-prediction norm; ``cfg_truncation``
# disables CFG above the normalized-noise threshold.
cfg_normalization: bool = False
cfg_truncation: float | None = 1.0
# 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).
@@ -323,6 +330,24 @@ class SamplingParam:
default=SamplingParam.guidance_scale,
help="Classifier-free guidance scale",
)
parser.add_argument(
"--cfg-normalization",
action=StoreBoolean,
default=SamplingParam.cfg_normalization,
help="Cap Z-Image CFG prediction norm to the positive-prediction norm",
)
parser.add_argument(
"--cfg-truncation",
type=float,
default=SamplingParam.cfg_truncation,
help="Disable Z-Image CFG above this normalized-noise threshold",
)
parser.add_argument(
"--batch-cfg",
action=StoreBoolean,
default=SamplingParam.batch_cfg,
help="Evaluate conditional and unconditional CFG branches in one batch",
)
parser.add_argument(
"--guidance-rescale",
type=float,
+5
View File
@@ -130,6 +130,7 @@ class InputConfig:
keyboard_cond: Any | None = None
grid_sizes: Any | None = None
c2ws_plucker_emb: Any | None = None
action_path: str | None = None
refine_from: str | None = None
stage1_video: Any | None = None
@@ -138,6 +139,7 @@ class InputConfig:
class SamplingConfig:
num_videos_per_prompt: int = 1
seed: int = 1024
max_sequence_length: int | None = None
num_frames: int = 125
height: int = 720
width: int = 1280
@@ -147,7 +149,10 @@ class SamplingConfig:
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
batch_cfg: bool = False
guidance_scale_2: float | None = None
cfg_normalization: bool = False
cfg_truncation: float | None = 1.0
guidance_rescale: float = 0.0
true_cfg_scale: float | None = None
use_embedded_guidance: bool | None = None
+18 -9
View File
@@ -18,14 +18,14 @@ from fastvideo.logger import init_logger
logger = init_logger(__name__)
_project_root = Path(__file__).resolve().parent.parent.parent.parent
_kernel_root = _project_root / "fastvideo-kernel"
_kernel_python_root = _kernel_root / "python"
_kernel_python_root = _project_root / "fastvideo-kernel" / "python"
_attn_qat_train_attention: Callable[..., torch.Tensor] | None = None
_attn_qat_train_import_attempted = False
_attn_qat_train_import_error: ImportError | None = None
def _ensure_kernel_paths() -> None:
for path in (_project_root, _kernel_root, _kernel_python_root):
for path in (_project_root, _kernel_python_root):
path_str = str(path)
if path_str not in sys.path:
sys.path.insert(0, path_str)
@@ -34,6 +34,7 @@ def _ensure_kernel_paths() -> None:
def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
global _attn_qat_train_attention
global _attn_qat_train_import_attempted
global _attn_qat_train_import_error
if _attn_qat_train_import_attempted:
return _attn_qat_train_attention
@@ -42,8 +43,11 @@ def _get_attn_qat_train_attention() -> Callable[..., torch.Tensor] | None:
_ensure_kernel_paths()
try:
_attn_qat_train_attention = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train").attention
except ImportError:
triton_qat = importlib.import_module("fastvideo_kernel.triton_kernels.attn_qat_train")
_attn_qat_train_attention = triton_qat.attention
logger.info("ATTN_QAT_TRAIN loaded FastVideo's architecture-optimized Triton kernel")
except ImportError as exc:
_attn_qat_train_import_error = exc
_attn_qat_train_attention = None
return _attn_qat_train_attention
@@ -60,8 +64,9 @@ def attn_qat_train(q_BLHD: torch.Tensor,
sm_scale: float | None = None) -> torch.Tensor:
attention = _get_attn_qat_train_attention()
if attention is None:
raise ImportError("fastvideo_kernel.triton_kernels.attn_qat_train is not available. "
"Please ensure the FastVideo kernel package is installed.")
detail = f" Original import error: {_attn_qat_train_import_error}" if _attn_qat_train_import_error else ""
raise ImportError("ATTN_QAT_TRAIN requires FastVideo's fastvideo-kernel package. Install it or make "
f"fastvideo-kernel/python importable.{detail}")
q_BHLD = q_BLHD.permute(0, 2, 1, 3).contiguous()
k_BHLD = k_BLHD.permute(0, 2, 1, 3).contiguous()
@@ -69,7 +74,11 @@ def attn_qat_train(q_BLHD: torch.Tensor,
use_qat_qkv_backward = True
smooth_k = False
warp_specialize = True
# Triton 3.7's NVWS pass aborts while compiling this kernel on Blackwell.
# The kernel has a supported non-warp-specialized path, so use it on both
# datacenter (sm_100) and consumer (sm_120) Blackwell GPUs.
capability_major = torch.cuda.get_device_capability()[0]
warp_specialize = capability_major not in (10, 12)
is_qat = True
two_level_quant_p_sage3 = False
fake_quant_p_bwd = True
@@ -106,7 +115,7 @@ class AttnQatTrainBackend(AttentionBackend):
@staticmethod
def get_supported_head_sizes() -> list[int]:
return [64, 96, 128, 160, 192, 224, 256]
return [128]
@staticmethod
def get_name() -> str:
+26
View File
@@ -32,6 +32,27 @@ def backend_name_to_enum(backend_name: str) -> AttentionBackendEnum | None:
None
def coerce_attn_backend(attn_backend: AttentionBackendEnum | str | None, ) -> AttentionBackendEnum | None:
"""Normalize an explicit backend selection.
Environment-variable parsing remains permissive via
:func:`backend_name_to_enum`, but typed/config-driven call sites should
fail fast on typos instead of silently falling back to another backend.
"""
if attn_backend is None or isinstance(attn_backend, AttentionBackendEnum):
return attn_backend
if not isinstance(attn_backend, str) or not attn_backend.strip():
raise ValueError("attention backend must be a non-empty string, "
f"an AttentionBackendEnum, or None; got {attn_backend!r}")
backend_name = attn_backend.strip().upper()
backend = backend_name_to_enum(backend_name)
if backend is None:
raise ValueError(f"Unknown attention backend {attn_backend!r}. "
f"Expected one of {sorted(AttentionBackendEnum.__members__)}")
return backend
def get_env_variable_attn_backend() -> AttentionBackendEnum | None:
'''
Get the backend override specified by the FastVideo attention
@@ -69,6 +90,11 @@ def global_force_attn_backend(attn_backend: AttentionBackendEnum | None) -> None
'''
global forced_attn_backend
forced_attn_backend = attn_backend
# Backend selection is cached by tensor shape/dtype, while the global
# override is intentionally not part of that cache key. Invalidate cached
# resolutions whenever the override changes so independently constructed
# role models can bind different attention implementations.
_cached_get_attn_backend.cache_clear()
def get_global_forced_attn_backend() -> AttentionBackendEnum | None:
+5 -1
View File
@@ -12,12 +12,16 @@ from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.magi_human import MagiHumanVideoConfig
from fastvideo.configs.models.dits.stable_audio import StableAudioConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
from fastvideo.configs.models.dits.lingbotworld2 import LingBotWorld2CausalFastVideoConfig
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig"
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
"ZImageDiTConfig"
]
+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,70 @@
# SPDX-License-Identifier: Apache-2.0
"""Architecture configuration for LingBot-Video Dense and MoE DiTs."""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def _is_lingbot_video_block(name: str, module: object) -> bool:
"""Select top-level transformer blocks for FSDP and compilation."""
del module
parts = name.split(".")
return len(parts) == 2 and parts[0] == "blocks" and parts[1].isdigit()
@dataclass
class LingBotVideoArchConfig(DiTArchConfig):
"""One-to-one representation of the released transformer config JSON."""
_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_lingbot_video_block])
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
)
param_names_mapping: dict = field(default_factory=lambda: {r"^(.*)$": r"\1"})
patch_size: tuple[int, int, int] = (1, 2, 2)
in_channels: int = 16
out_channels: int = 16
hidden_size: int = 2048
num_attention_heads: int = 16
depth: int = 24
intermediate_size: int = 6144
text_dim: int = 2560
freq_dim: int = 256
norm_eps: float = 1e-6
rope_theta: float = 256.0
axes_dims: tuple[int, int, int] = (32, 48, 48)
axes_lens: tuple[int, int, int] = (8192, 1024, 1024)
qkv_bias: bool = False
out_bias: bool = True
patch_embed_bias: bool = True
timestep_mlp_bias: bool = True
num_experts: int = 0
num_experts_per_tok: int = 8
moe_intermediate_size: int = 512
decoder_sparse_step: int = 1
mlp_only_layers: tuple[int, ...] = ()
n_shared_experts: int | None = None
score_func: str = "sigmoid"
norm_topk_prob: bool = True
n_group: int | None = None
topk_group: int | None = None
routed_scaling_factor: float = 1.0
def __post_init__(self) -> None:
"""Populate FastVideo loader fields from the released architecture."""
super().__post_init__()
self.num_channels_latents = self.in_channels
self.attention_head_dim = self.hidden_size // self.num_attention_heads
@dataclass
class LingBotVideoConfig(DiTConfig):
"""FastVideo component configuration for LingBot-Video transformers."""
arch_config: DiTArchConfig = field(default_factory=LingBotVideoArchConfig)
@@ -0,0 +1,55 @@
# 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 "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class LingBotWorld2CausalFastArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
model_type: str = "i2v"
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
in_dim: int = 36
dim: int = 5120
ffn_dim: int = 13824
freq_dim: int = 256
text_dim: int = 4096
out_dim: int = 16
num_heads: int = 40
num_layers: int = 40
qk_norm: bool = True
cross_attn_norm: bool = True
eps: float = 1e-6
local_attn_size: int = 18
sink_size: int = 6
chunk_size: int = 4
sample_shift: float = 10.0
num_train_timesteps: int = 1000
timesteps_index: tuple[int, int, int, int] = (0, 250, 500, 750)
max_area: int = 480 * 832
def __post_init__(self):
super().__post_init__()
self.hidden_size = self.dim
self.num_attention_heads = self.num_heads
self.attention_head_dim = self.dim // self.num_heads
self.in_channels = self.in_dim
self.out_channels = self.out_dim
self.num_channels_latents = self.out_dim
@dataclass
class LingBotWorld2CausalFastVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=LingBotWorld2CausalFastArchConfig)
prefix: str = "Wan"
+60
View File
@@ -0,0 +1,60 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum
def is_zimage_block(name: str, module) -> bool:
parts = name.split(".")
return len(parts) >= 2 and parts[-2] in {"noise_refiner", "context_refiner", "layers"} and parts[-1].isdigit()
@dataclass
class ZImageDiTArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_zimage_block])
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (AttentionBackendEnum.TORCH_SDPA, )
all_patch_size: tuple[int, ...] = (2, )
all_f_patch_size: tuple[int, ...] = (1, )
in_channels: int = 16
dim: int = 3840
n_layers: int = 30
n_refiner_layers: int = 2
n_heads: int = 30
n_kv_heads: int = 30
norm_eps: float = 1e-5
qk_norm: bool = True
cap_feat_dim: int = 2560
rope_theta: float = 256.0
t_scale: float = 1000.0
axes_dims: tuple[int, ...] = (32, 48, 48)
axes_lens: tuple[int, ...] = (1536, 512, 512)
adaln_embed_dim: int = 256
frequency_embedding_size: int = 256
timestep_mid_size: int = 1024
max_period: int = 10000
seq_multi_of: int = 32
def __post_init__(self) -> None:
super().__post_init__()
if len(self.all_patch_size) != len(self.all_f_patch_size):
raise ValueError("all_patch_size and all_f_patch_size must have equal length")
if self.dim % self.n_heads:
raise ValueError("dim must be divisible by n_heads")
if self.dim // self.n_heads != sum(self.axes_dims):
raise ValueError("attention head dimension must equal sum(axes_dims)")
if len(self.axes_dims) != len(self.axes_lens) or any(dim % 2 for dim in self.axes_dims):
raise ValueError("RoPE axes require matching lengths and even dimensions")
self.hidden_size = self.dim
self.num_attention_heads = self.n_heads
self.num_channels_latents = self.in_channels
self.out_channels = self.in_channels
@dataclass
class ZImageDiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=ZImageDiTArchConfig)
prefix: str = "ZImage"
@@ -2,6 +2,7 @@ from fastvideo.configs.models.encoders.base import (BaseEncoderOutput, EncoderCo
TextEncoderConfig)
from fastvideo.configs.models.encoders.clip import (CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5ArchConfig, LingBotWorld2UMT5Config
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
@@ -9,6 +10,7 @@ from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
from fastvideo.configs.models.encoders.mistral3 import Mistral3TextConfig
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
from fastvideo.configs.models.encoders.stable_audio_conditioner import (StableAudioConditionerArchConfig,
StableAudioConditionerConfig)
from fastvideo.configs.models.encoders.t5gemma import T5GemmaEncoderConfig
@@ -17,5 +19,6 @@ __all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig", "BaseEncoderOutput", "CLIPTextConfig",
"CLIPVisionConfig", "WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig", "Qwen2_5_VLConfig",
"Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig", "SiglipVisionConfig", "StableAudioConditionerArchConfig",
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig"
"StableAudioConditionerConfig", "T5GemmaEncoderConfig", "Qwen3TextConfig", "Mistral3TextConfig",
"LingBotWorld2UMT5ArchConfig", "LingBotWorld2UMT5Config", "LingBotVideoQwen3VLTextConfig"
]
@@ -83,6 +83,7 @@ class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
is_chat_model: bool = False
treat_empty_as_dot: bool = False
chat_template_enable_thinking: bool = field(default=False, kw_only=True)
@dataclass
@@ -0,0 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
"""Qwen3-VL text-only encoder configuration used by LingBot-Video."""
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextArchConfig, Qwen3TextConfig
@dataclass
class LingBotVideoQwen3VLTextArchConfig(Qwen3TextArchConfig):
"""Exact Qwen3-VL language-model architecture released with LingBot-Video."""
architectures: list[str] = field(default_factory=lambda: ["LingBotVideoQwen3VLTextModel"])
vocab_size: int = 151936
hidden_size: int = 2560
intermediate_size: int = 9728
num_hidden_layers: int = 36
num_attention_heads: int = 32
num_key_value_heads: int = 8
max_position_embeddings: int = 262144
rms_norm_eps: float = 1e-6
rope_theta: float = 5000000.0
rope_scaling: dict | None = None
mrope_interleaved: bool = True
mrope_section: tuple[int, int, int] = (24, 20, 20)
bos_token_id: int = 151643
eos_token_id: int = 151645
pad_token_id: int = 151643
text_len: int = 37698
output_hidden_states: bool = True
require_processor: bool = True
def __post_init__(self) -> None:
"""Match the official processor call used by LingBotVideoPipeline."""
self.tokenizer_kwargs = {
"truncation": True,
"max_length": self.text_len,
"padding": "longest",
"return_tensors": "pt",
}
@dataclass
class LingBotVideoQwen3VLTextConfig(Qwen3TextConfig):
"""FastVideo loader config for the LingBot-Video text-only Qwen3-VL path."""
arch_config: TextEncoderArchConfig = field(default_factory=LingBotVideoQwen3VLTextArchConfig)
prefix: str = "language_model"
is_chat_model: bool = False
@@ -0,0 +1,40 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class LingBotWorld2UMT5ArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(default_factory=lambda: ["LingBotWorld2T5EncoderModel"])
vocab_size: int = 256384
dim: int = 4096
dim_attn: int = 4096
dim_ffn: int = 10240
num_heads: int = 64
num_layers: int = 24
num_buckets: int = 32
text_len: int = 512
hidden_size: int = 4096
dropout: float = 0.1
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs = {
"padding": "max_length",
"truncation": True,
"max_length": self.text_len,
"add_special_tokens": True,
"return_attention_mask": True,
"return_tensors": "pt",
}
@dataclass
class LingBotWorld2UMT5Config(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(default_factory=LingBotWorld2UMT5ArchConfig)
prefix: str = "text_encoder"
@@ -1,5 +1,6 @@
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.cosmos3vae import Cosmos3VAEConfig
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
@@ -16,6 +17,7 @@ __all__ = [
"WanVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Cosmos3VAEConfig",
"Gen3CVAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
+277
View File
@@ -0,0 +1,277 @@
"""Cosmos3 (Wan2.2-TI2V-5B) VAE config and checkpoint-key mapping.
The Cosmos3 checkpoint VAE is literally ``Wan-AI/Wan2.2-TI2V-5B-Diffusers``
(diffusers ``AutoencoderKLWan``), so this config locks the Wan2.2 geometry:
residual down/up blocks, ``patch_size=2``, ``z_dim=48``, ``base_dim=160``,
``decoder_base_dim=256``, and ``scale_factor_spatial=16``. The 48-dim
``latents_mean``/``latents_std`` are taken verbatim from the Cosmos3
checkpoint's ``vae/config.json`` (identical to the canonical Wan2.2-TI2V-5B
statistics).
Mirrors the :class:`Cosmos25VAEArchConfig` pattern. ``param_names_mapping`` /
``map_official_key`` translate the *official* Wan2.2 VAE state-dict keys
(nested-residual naming, e.g. ``encoder.downsamples.{b}.downsamples.{j}`` and
``decoder.upsamples.{b}.upsamples.{j}``) into FastVideo's ``AutoencoderKLWan``
key space. The standard diffusers checkpoint already ships native FastVideo
keys, so these helpers exist for parity tooling and official ``.pth`` loading.
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig, WanVAEConfig
@dataclass
class Cosmos3VAEArchConfig(WanVAEArchConfig):
# Wan2.2-TI2V-5B geometry (differs from the Wan2.1 WanVAEArchConfig
# defaults: residual blocks, patch_size=2, z_dim=48, base_dim=160,
# decoder_base_dim=256, scale_factor_spatial=16, 12 patch channels).
_name_or_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
base_dim: int = 160
decoder_base_dim: int | None = 256
z_dim: int = 48
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
is_residual: bool = True
in_channels: int = 12
out_channels: int = 12
patch_size: int | None = 2
scale_factor_temporal: int = 4
scale_factor_spatial: int = 16
clip_output: bool = False
# 48-dim statistics copied verbatim from the Cosmos3 checkpoint
# (official_weights/cosmos3/vae/config.json).
latents_mean: tuple[float, ...] = (
-0.2289,
-0.0052,
-0.1323,
-0.2339,
-0.2799,
0.0174,
0.1838,
0.1557,
-0.1382,
0.0542,
0.2813,
0.0891,
0.157,
-0.0098,
0.0375,
-0.1825,
-0.2246,
-0.1207,
-0.0698,
0.5109,
0.2665,
-0.2108,
-0.2158,
0.2502,
-0.2055,
-0.0322,
0.1109,
0.1567,
-0.0729,
0.0899,
-0.2799,
-0.123,
-0.0313,
-0.1649,
0.0117,
0.0723,
-0.2839,
-0.2083,
-0.052,
0.3748,
0.0152,
0.1957,
0.1433,
-0.2944,
0.3573,
-0.0548,
-0.1681,
-0.0667,
)
latents_std: tuple[float, ...] = (
0.4765,
1.0364,
0.4514,
1.1677,
0.5313,
0.499,
0.4818,
0.5013,
0.8158,
1.0344,
0.5894,
1.0901,
0.6885,
0.6165,
0.8454,
0.4978,
0.5759,
0.3523,
0.7135,
0.6804,
0.5833,
1.4146,
0.8986,
0.5659,
0.7069,
0.5338,
0.4889,
0.4917,
0.4069,
0.4999,
0.6866,
0.4093,
0.5709,
0.6065,
0.6415,
0.4944,
0.5726,
1.2042,
0.5458,
1.6887,
0.3971,
1.06,
0.3943,
0.5537,
0.5444,
0.4089,
0.7468,
0.7744,
)
# Simple 1:1 renames. The nested-residual block remapping (encoder
# downsamples / decoder upsamples / middle / head) is handled by
# ``map_official_key()``.
param_names_mapping: dict[str, str] = field(
default_factory=lambda: {
r"^conv1\.(.*)$": r"quant_conv.\1",
r"^conv2\.(.*)$": r"post_quant_conv.\1",
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
})
@staticmethod
def map_official_key(key: str) -> str | None:
"""Map a single official Wan2.2 VAE key into FastVideo key space.
Handles the residual (Wan2.2) module layout where each down/up block
is a nested ``Sequential`` (``downsamples.{b}.downsamples.{j}`` /
``upsamples.{b}.upsamples.{j}``) rather than the flat Wan2.1 indexing.
Returns ``None`` for keys with no FastVideo counterpart.
"""
def map_residual_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^residual\.0\.gamma$", sub):
return f"{prefix}.norm1.gamma"
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv1.{m.group(1)}"
if re.match(r"^residual\.3\.gamma$", sub):
return f"{prefix}.norm2.gamma"
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv2.{m.group(1)}"
m = re.match(r"^shortcut\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv_shortcut.{m.group(1)}"
return None
def map_attn_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^norm\.gamma$", sub):
return f"{prefix}.norm.gamma"
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
if m:
return f"{prefix}.to_qkv.{m.group(1)}"
m = re.match(r"^proj\.(weight|bias)$", sub)
if m:
return f"{prefix}.proj.{m.group(1)}"
return None
def map_resample_subkey(prefix: str, sub: str) -> str | None:
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
if m:
return f"{prefix}.resample.1.{m.group(1)}"
m = re.match(r"^time_conv\.(weight|bias)$", sub)
if m:
return f"{prefix}.time_conv.{m.group(1)}"
return None
m = re.match(r"^conv1\.(weight|bias)$", key)
if m:
return f"quant_conv.{m.group(1)}"
m = re.match(r"^conv2\.(weight|bias)$", key)
if m:
return f"post_quant_conv.{m.group(1)}"
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_in.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
if m:
return f"{m.group(1)}.norm_out.gamma"
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_out.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if m:
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1", m.group(2))
# Encoder: downsamples.{block}.downsamples.{j}.* (nested residual layout)
m = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
if m:
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return map_resample_subkey(f"encoder.down_blocks.{block_i}.downsampler", sub)
return map_residual_subkey(f"encoder.down_blocks.{block_i}.resnets.{res_i}", sub)
# Decoder: upsamples.{block}.upsamples.{j}.* (nested residual layout)
m = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
if m:
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return map_resample_subkey(f"decoder.up_blocks.{block_i}.upsampler", sub)
return map_residual_subkey(f"decoder.up_blocks.{block_i}.resnets.{res_i}", sub)
return None
# ``__post_init__`` (scaling_factor / shift_factor / compression ratios) is
# inherited unchanged from ``WanVAEArchConfig``.
@dataclass
class Cosmos3VAEConfig(WanVAEConfig):
"""Cosmos3 VAE config (reuses FastVideo's Wan2.2 ``AutoencoderKLWan``).
Subclasses :class:`WanVAEConfig` so the model reads the same runtime flags
(``use_feature_cache``, ``use_light_vae``, tiling) and only swaps in the
Cosmos3 = Wan2.2 ``arch_config``.
"""
arch_config: Cosmos3VAEArchConfig = field(default_factory=Cosmos3VAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
# ``__post_init__`` (blend_num_frames) is inherited from ``WanVAEConfig``.
+4 -1
View File
@@ -7,6 +7,8 @@ from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyua
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5I2VConfig, Kandinsky5T2VConfig
from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V480PConfig
from fastvideo.configs.pipelines.lingbot_video import LingBotVideoT2VConfig
from fastvideo.configs.pipelines.matrixgame2 import MatrixGame2I2V480PConfig
from fastvideo.configs.pipelines.matrixgame3 import MatrixGame3I2V720PConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
@@ -19,5 +21,6 @@ __all__ = [
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
"Kandinsky5I2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"Kandinsky5I2VConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig", "MatrixGame2I2V480PConfig",
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
]
+69
View File
@@ -0,0 +1,69 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 pipeline configuration.
Reference of record: the official ``cosmos-framework`` / ``nvidia/Cosmos3-Nano``
checkpoint (``model_index.json``). Cosmos3 is structurally different from
Cosmos 2.5:
- Dual-pathway (UND + GEN) DiT lives entirely inside ``Cosmos3VFMTransformer``
(``Cosmos3VideoConfig``).
- No separate text encoder — the Qwen3-VL-text backbone is inside the DiT, so
``text_encoder_configs`` is the empty tuple. The Qwen2 tokenizer is loaded as
the ``text_tokenizer`` checkpoint module by the component loader.
- VAE is Wan2.2 ``AutoencoderKLWan`` (z_dim=48, scale_factor_spatial=16),
configured by ``Cosmos3VAEConfig`` (the checkpoint's exact latents_mean/std).
- Scheduler is FastVideo-native ``UniPCMultistepScheduler`` configured for
pure flow matching (flow_prediction, use_flow_sigmas), equivalent to the
framework's ``FlowUniPCMultistepScheduler``. The checkpoint's diffusers-style
scheduler config (karras/sigma_min/max) is coerced to the flow setup in
``Cosmos3OmniDiffusersPipeline.initialize_pipeline``.
- T2I default ``flow_shift`` is 3.0 (set per-request by ``_set_flow_shift``);
T2V/I2V use the engine-init default of 1.0 baked into this config.
"""
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits.cosmos3 import (Cosmos3ArchConfig, Cosmos3VideoConfig)
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.vaes import Cosmos3VAEConfig # Wan2.2 AutoencoderKLWan
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class Cosmos3Config(PipelineConfig):
"""Configuration for the Cosmos3 video generation pipeline (T2V/I2V/T2I).
Wires the framework-parity-verified Cosmos3 components: the native
``Cosmos3VideoConfig`` DiT, the Wan2.2 ``Cosmos3VAEConfig`` VAE, the Qwen2
tokenizer (loaded as ``text_tokenizer``), and the UniPC scheduler.
"""
dit_config: DiTConfig = field(default_factory=lambda: Cosmos3VideoConfig(arch_config=Cosmos3ArchConfig()))
vae_config: VAEConfig = field(default_factory=Cosmos3VAEConfig)
# No separate text encoder: the Qwen3-VL-text backbone lives inside the DiT
# and the pipeline tokenizes in Cosmos3DenoisingStage, so all three
# text-encoder lists are empty (the generic text-encode stage is not used).
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=tuple)
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=tuple)
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(default_factory=tuple)
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
embedded_cfg_scale: float = 0.0
# T2V/I2V engine-init flow_shift (framework text2video/image2video default);
# T2I overrides to 3.0 per request via Cosmos3DenoisingStage._set_flow_shift.
flow_shift: float = 10.0
vae_tiling: bool = False
vae_sp: bool = False
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Dense LingBot-Video T2V pipeline configuration."""
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.lingbot_video import LingBotVideoQwen3VLTextConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
PROMPT_CROP_START = 140
PROMPT_TEMPLATE = ("<|im_start|>system\nGiven a user input that may include a text prompt alone, "
"a text prompt with an image reference, or a text prompt with a video reference "
'or a video reference alone, generate an "Enhanced prompt" that provides detailed '
"visual descriptions suitable for video generation. Evaluate the level of detail "
"in the user's input: if it is simple, enrich it by adding specifics about colors, "
"shapes, sizes, textures, lighting, motion dynamics, camera movement, temporal "
"progression, and spatial relationships to create vivid, concrete, and temporally "
"coherent scenes to create vivid and concrete scenes. Please generate only the "
"enhanced description for the prompt below and avoid including any additional "
"commentary or evaluations:<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n"
"<|im_start|>assistant\n")
def preprocess_lingbot_video_prompt(prompt: str) -> str:
"""Apply the released T2V system/user/assistant prompt template."""
return PROMPT_TEMPLATE.format(prompt)
def postprocess_lingbot_video_text(
outputs: BaseEncoderOutput,
attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Select the final hidden state, crop the template, and trim batch-one padding."""
if outputs.hidden_states is None:
raise ValueError("LingBot-Video requires text-encoder hidden states")
prompt_embeds = outputs.hidden_states[-1][:, PROMPT_CROP_START:]
prompt_mask = attention_mask[:, PROMPT_CROP_START:]
if prompt_embeds.shape[0] == 1:
true_length = int(prompt_mask[0].sum().item())
prompt_embeds = prompt_embeds[:, :true_length]
prompt_mask = prompt_mask[:, :true_length]
return prompt_embeds, prompt_mask
@dataclass
class LingBotVideoT2VConfig(PipelineConfig):
"""Released Dense T2V component wiring and numerical precision policy."""
dit_config: DiTConfig = field(default_factory=LingBotVideoConfig)
vae_config: VAEConfig = field(default_factory=WanVAEConfig)
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (LingBotVideoQwen3VLTextConfig(), ))
preprocess_text_funcs: tuple[Callable, ...] = field(default_factory=lambda: (preprocess_lingbot_video_prompt, ))
postprocess_text_funcs: tuple[Callable, ...] = field(default_factory=lambda: (postprocess_lingbot_video_text, ))
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
dit_precision: str = "bf16"
vae_precision: str = "fp32"
vae_decode_precision: str | None = "fp32"
vae_tiling: bool = False
vae_sp: bool = False
flow_shift: float | None = 3.0
embedded_cfg_scale: float | None = None
scheduler_step_in_fp32: bool = True
def __post_init__(self) -> None:
"""Load only the VAE decoder for the T2V workload."""
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@@ -0,0 +1,45 @@
# SPDX-License-Identifier: Apache-2.0
import html
from dataclasses import dataclass, field
import ftfy
import torch
from fastvideo.configs.models import DiTConfig
from fastvideo.configs.models.dits.lingbotworld2 import LingBotWorld2CausalFastVideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, LingBotWorld2UMT5Config
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.wan import Wan2_2_I2V_A14B_Config
def lingbotworld2_whitespace_preprocess(prompt: str) -> str:
text = ftfy.fix_text(prompt)
text = html.unescape(html.unescape(text))
return " ".join(text.strip().split())
def lingbotworld2_t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
assert outputs.last_hidden_state is not None
return outputs.last_hidden_state
@dataclass
class LingBotWorld2CausalFastI2V480PConfig(Wan2_2_I2V_A14B_Config):
dit_config: DiTConfig = field(default_factory=LingBotWorld2CausalFastVideoConfig)
vae_config: WanVAEConfig = field(default_factory=WanVAEConfig)
text_encoder_configs: tuple = field(default_factory=lambda: (LingBotWorld2UMT5Config(), ))
preprocess_text_funcs: tuple = field(default_factory=lambda: (lingbotworld2_whitespace_preprocess, ))
postprocess_text_funcs: tuple = field(default_factory=lambda: (lingbotworld2_t5_postprocess_text, ))
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
flow_shift: float | None = 10.0
boundary_ratio: float | None = 0.947
vae_precision: str = "fp32"
vae_decode_precision: str | None = "fp32"
vae_tiling: bool = False
vae_sp: bool = False
dit_precision: str = "bf16"
is_causal: bool = False
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
+52
View File
@@ -0,0 +1,52 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import EncoderConfig
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.qwen3 import Qwen3TextConfig
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def _zimage_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
if outputs.hidden_states is None:
raise RuntimeError("Z-Image requires Qwen3 hidden states")
return outputs.hidden_states[-2]
@dataclass
class ZImagePipelineConfig(PipelineConfig):
"""Configuration for the native Z-Image text-to-image pipeline."""
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
transformer_arch: str = "ZImageTransformer2DModel"
vae_arch: str = "AutoencoderKL"
text_encoder_archs: tuple[str, ...] = ("Qwen3Model", )
tokenizer_archs: tuple[str, ...] = ("Qwen2Tokenizer", )
dit_config: ZImageDiTConfig = field(default_factory=ZImageDiTConfig)
vae_config: AutoencoderKLVAEConfig = field(default_factory=AutoencoderKLVAEConfig)
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (Qwen3TextConfig(chat_template_enable_thinking=True), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (_zimage_text_postprocess, ))
dit_precision: str = "bf16"
vae_precision: str = "fp32"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
vae_tiling: bool = False
vae_sp: bool = False
embedded_cfg_scale: float = 0.0
flow_shift: float | None = 3.0
scheduler_step_in_fp32: bool = True
scheduler_sigma_min: float = 0.0
scheduler_use_reference_discrete_timesteps: bool = True
+45 -1
View File
@@ -112,6 +112,40 @@ def _infer_latent_batch_size(batch: ForwardBatch) -> int:
return latent_batch_size
def _resolve_output_size(
samples: torch.Tensor,
fallback: tuple[int, int, int],
*,
pixel_output: bool,
) -> tuple[int, int, int]:
"""Report the final decoded video's `(height, width, frames)`.
Refiner stages can produce a different resolution from the base request, so
pixel outputs use the final `[batch, channels, frames, height, width]` tensor.
Latent and audio outputs keep the requested fallback because their tensor
dimensions do not describe decoded pixels.
"""
if pixel_output and samples.ndim == 5:
return (int(samples.shape[-2]), int(samples.shape[-1]), int(samples.shape[-3]))
return fallback
def _validate_request_stage_overrides(model_path: str, request: GenerationRequest) -> None:
"""Validate typed stage overrides against the model's registered preset."""
if not request.stage_overrides:
return
from fastvideo.api.presets import validate_preset_selection
from fastvideo.registry import get_preset_selection
preset_name, model_family = get_preset_selection(model_path)
if preset_name is None or model_family is None:
raise ValueError(f"Model {model_path!r} has no preset for stage override validation")
validate_preset_selection(
preset_name,
model_family,
stage_overrides=request.stage_overrides,
)
class VideoGenerator:
"""
A unified class for generating videos using diffusion models.
@@ -443,6 +477,7 @@ class VideoGenerator:
self,
request: GenerationRequest,
) -> GenerationResult | list[GenerationResult]:
_validate_request_stage_overrides(self.fastvideo_args.model_path, request)
if isinstance(request.prompt, list):
if request.inputs.prompt_path is not None:
raise ValueError("request.prompt list cannot be combined with request.inputs.prompt_path")
@@ -808,6 +843,15 @@ class VideoGenerator:
# 2. Audio-only workload — `samples` is a 1×3×1×8×8 placeholder
# no caller will use; skip the grid loop and save a `.wav`.
# 3. Pixel video / image — the historical happy path.
# `GenerationResult.size` describes the produced media, not only the
# base-stage request. Refiner pipelines can change the final pixel
# dimensions, so derive this result metadata from the decoded output.
output_size = _resolve_output_size(
samples,
(target_height, target_width, batch.num_frames),
pixel_output=not is_latent_output and not audio_only,
)
postprocess_start = time.perf_counter()
frames: list[np.ndarray] | None
if is_latent_output or audio_only:
@@ -916,7 +960,7 @@ class VideoGenerator:
"audio": output_batch.extra.get("audio"),
"audio_sample_rate": output_batch.extra.get("audio_sample_rate"),
"ltx2_audio_latents": output_batch.extra.get("ltx2_audio_latents"),
"size": (target_height, target_width, batch.num_frames),
"size": output_size,
"generation_time": gen_time,
"e2e_latency": e2e_time,
"logging_info": logging_info,
+133
View File
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 sound tokenizer (AVAE) — decode path.
The Cosmos3 ``sound_tokenizer`` is an AVAE (audio VAE). Its shipped diffusers
checkpoint is **decoder-only** (``decoder.*``) in ``AutoencoderOobleck`` naming
with SnakeBeta activations and ``weight_g``/``weight_v`` weight-norm — exactly
FastVideo's native :class:`~fastvideo.models.vaes.oobleck.OobleckDecoder`
(verified bit-exact vs the framework in ``test_cosmos3_avae_parity``). Text-to-
video+sound (t2vs) only needs DECODE: the DiT generates the sound latent and
this module decodes it to a waveform, so only the decoder is ported (the
SpectrogramConvNeXt encoder is not exported in the checkpoint).
Mirrors the framework ``AVAEModel.decode``: run the Oobleck decoder, then clamp
to [-1, 1]. The VAE bottleneck's decode is the identity (the DiT already emits
the post-bottleneck latent), so there is no bottleneck step here.
"""
from __future__ import annotations
import json
import os
from dataclasses import dataclass, field
import numpy as np
import torch
import torch.nn as nn
from fastvideo.logger import init_logger
from fastvideo.models.vaes.oobleck import OobleckDecoder
logger = init_logger(__name__)
@dataclass
class Cosmos3SoundVAEArchConfig:
"""Cosmos3 AVAE decoder constants (from ``sound_tokenizer/config.json``)."""
dec_dim: int = 320 # decoder base channels
vocoder_input_dim: int = 64 # latent channels in
dec_c_mults: list[int] = field(default_factory=lambda: [1, 2, 4, 8, 16])
dec_strides: list[int] = field(default_factory=lambda: [2, 4, 5, 6, 8])
audio_channels: int = 2 # stereo
sampling_rate: int = 48000
@property
def hop_size(self) -> int:
return int(np.prod(self.dec_strides)) # 1920
class Cosmos3SoundVAE(nn.Module):
"""Decoder-only Cosmos3 AVAE: latent ``[B, z, T]`` -> waveform ``[B, C, N]``."""
def __init__(self, arch: Cosmos3SoundVAEArchConfig | None = None) -> None:
super().__init__()
self.arch = arch or Cosmos3SoundVAEArchConfig()
self.decoder = OobleckDecoder(
channels=self.arch.dec_dim,
input_channels=self.arch.vocoder_input_dim,
audio_channels=self.arch.audio_channels,
# The framework builds decoder blocks from ``reversed(dec_strides)``
# (deepest first), so block strides are e.g. [8,6,5,4,2].
upsampling_ratios=list(reversed(self.arch.dec_strides)),
channel_multiples=list(self.arch.dec_c_mults),
)
@property
def sample_rate(self) -> int:
return self.arch.sampling_rate
@property
def audio_channels(self) -> int:
return self.arch.audio_channels
@property
def hop_size(self) -> int:
return self.arch.hop_size
def get_latent_num_samples(self, num_audio_samples: int) -> int:
"""Latent length for a given audio length (``AVAEInterface``: ``N // hop``)."""
return int(num_audio_samples) // self.arch.hop_size
@torch.no_grad()
def decode(self, latent: torch.Tensor) -> torch.Tensor:
"""Decode normalized latent ``[B, z, T]`` to waveform ``[B, C, N]`` in [-1, 1].
Matches ``AVAEModel.decode``: Oobleck decoder then clamp to [-1, 1] (the
VAE bottleneck decode is identity).
"""
audio = self.decoder(latent) # [B, C, N]
return audio.clamp(-1.0, 1.0)
@classmethod
def from_pretrained(
cls,
model_path: str,
*,
torch_dtype: torch.dtype | None = None,
) -> "Cosmos3SoundVAE":
"""Build + load the decoder from a ``sound_tokenizer`` directory.
Reads ``config.json`` (``dec_dim`` / ``vocoder_input_dim`` /
``dec_c_mults`` / ``dec_strides`` / ``sampling_rate`` / ``stereo``) and
loads the ``decoder.*`` weights (the checkpoint is decoder-only).
"""
from safetensors.torch import load_file
cfg_path = os.path.join(model_path, "config.json")
with open(cfg_path) as f:
cfg = json.load(f)
arch = Cosmos3SoundVAEArchConfig(
dec_dim=int(cfg["dec_dim"]),
vocoder_input_dim=int(cfg["vocoder_input_dim"]),
dec_c_mults=list(cfg["dec_c_mults"]),
dec_strides=list(cfg["dec_strides"]),
audio_channels=2 if cfg.get("stereo", True) else 1,
sampling_rate=int(cfg.get("sampling_rate", 48000)),
)
model = cls(arch)
weights_path = os.path.join(model_path, "diffusion_pytorch_model.safetensors")
state = load_file(weights_path)
# Decoder-only checkpoint: strip the ``decoder.`` prefix.
dec_state = {k[len("decoder."):]: v for k, v in state.items() if k.startswith("decoder.")}
model.decoder.load_state_dict(dec_state, strict=True)
logger.info("Loaded Cosmos3 sound AVAE decoder (%d params) from %s",
sum(p.numel() for p in model.parameters()), model_path)
if torch_dtype is not None:
model = model.to(dtype=torch_dtype)
model.eval()
return model
EntryClass = Cosmos3SoundVAE
File diff suppressed because it is too large Load Diff
+811
View File
@@ -0,0 +1,811 @@
# SPDX-License-Identifier: Apache-2.0
"""FastVideo-native LingBot-Video Dense and MoE diffusion transformers."""
from __future__ import annotations
import math
from dataclasses import dataclass
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.dits.lingbot_video import LingBotVideoConfig, _is_lingbot_video_block
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_all_to_all_4D,
sequence_model_parallel_shard,
)
from fastvideo.distributed.parallel_state import get_sp_world_size, model_parallel_is_initialized
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
@dataclass
class LingBotVideoTransformerOutput:
"""Output container matching the released transformer contract."""
sample: torch.Tensor
_FP32_MODULE_NAMES = (
"time_embedder",
"time_modulation",
"scale_shift_table",
"norm",
"norm1",
"norm2",
"norm_q",
"norm_k",
"norm_post_attn",
"norm_post_ffn",
"norm_out",
"norm_out_modulation",
"router",
)
def _keep_in_fp32(name: str) -> bool:
"""Return whether a released checkpoint module keeps fp32 parameters."""
return any(module_name in name.split(".") for module_name in _FP32_MODULE_NAMES)
def _sequence_parallel_world_size() -> int:
"""Use standalone single-rank behavior before distributed initialization."""
return get_sp_world_size() if model_parallel_is_initialized() else 1
class LingBotVideoLinear(ReplicatedLinear):
"""Replicated FastVideo linear with the tensor-only official call contract."""
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Return the projected tensor while preserving the normal weight surface."""
output, _ = super().forward(hidden_states)
return output
class LingBotVideoRMSNorm(nn.Module):
"""RMSNorm with fp32 accumulation and input-dtype output."""
def __init__(self, dim: int, eps: float = 1e-6) -> None:
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.variance_epsilon = eps
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Normalize the last dimension using the official accumulation order."""
input_dtype = hidden_states.dtype
hidden_states = hidden_states.to(torch.float32)
variance = hidden_states.pow(2).mean(-1, keepdim=True)
hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
return (self.weight * hidden_states).to(input_dtype)
def _apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
"""Apply complex 3D rotary embeddings to `(B, S, H, D)` tensors."""
with torch.amp.autocast("cuda", enabled=False):
x_complex = torch.view_as_complex(x.float().reshape(*x.shape[:-1], -1, 2))
output = torch.view_as_real(x_complex * freqs_cis.unsqueeze(2)).flatten(3)
return output.type_as(x)
class LingBotVideoRotaryEmbedding(nn.Module):
"""Complex64 rotary table indexed by temporal and spatial positions."""
def __init__(self, axes_dims: tuple[int, ...], axes_lens: tuple[int, ...], theta: float) -> None:
super().__init__()
self.axes_dims = tuple(axes_dims)
self.axes_lens = list(axes_lens)
self.theta = theta
self.freqs_cis: list[torch.Tensor] | None = None
@staticmethod
def _precompute(dims: tuple[int, ...], lengths: tuple[int, ...], theta: float) -> list[torch.Tensor]:
"""Build the per-axis complex frequency tables on CPU."""
tables: list[torch.Tensor] = []
for dim, length in zip(dims, lengths, strict=True):
frequencies = 1.0 / (theta**(torch.arange(0, dim, 2, dtype=torch.float64, device="cpu") / dim))
positions = torch.arange(length, device=frequencies.device, dtype=torch.float64)
phases = torch.outer(positions, frequencies).float()
tables.append(torch.polar(torch.ones_like(phases), phases).to(torch.complex64))
return tables
def forward(
self,
position_ids: torch.Tensor,
maxima: tuple[int, ...] | None = None,
) -> torch.Tensor:
"""Gather and concatenate rotary frequencies for `(S, 3)` positions."""
device = position_ids.device
if maxima is None:
maxima = tuple(int(value) for value in position_ids.max(dim=0).values.tolist())
rebuild = self.freqs_cis is None or any(maximum >= length
for maximum, length in zip(maxima, self.axes_lens, strict=True))
if rebuild:
for index, maximum in enumerate(maxima):
if maximum >= self.axes_lens[index]:
self.axes_lens[index] = int(maximum * 1.5) + 1
self.freqs_cis = self._precompute(self.axes_dims, tuple(self.axes_lens), self.theta)
self.freqs_cis = [table.to(device) for table in self.freqs_cis]
elif self.freqs_cis[0].device != device:
self.freqs_cis = [table.to(device) for table in self.freqs_cis]
return torch.cat(
[self.freqs_cis[index][position_ids[:, index]] for index in range(len(self.axes_dims))],
dim=-1,
)
def _make_joint_position_ids(
text_len: int,
grid_t: int,
grid_h: int,
grid_w: int,
device: torch.device,
) -> torch.Tensor:
"""Create official `[video; text]` 3D positions for one sample."""
temporal = torch.arange(grid_t, device=device, dtype=torch.int32) + text_len + 1
height = torch.arange(grid_h, device=device, dtype=torch.int32)
width = torch.arange(grid_w, device=device, dtype=torch.int32)
video_positions = torch.stack(torch.meshgrid(temporal, height, width, indexing="ij"), dim=-1).flatten(0, 2)
text_temporal = torch.arange(text_len, device=device, dtype=torch.int32) + 1
text_positions = torch.stack(
[text_temporal, torch.zeros_like(text_temporal),
torch.zeros_like(text_temporal)], dim=-1)
return torch.cat([video_positions, text_positions], dim=0)
class LingBotVideoTextEmbedder(nn.Module):
"""Project Qwen3-VL hidden states into the DiT hidden dimension."""
def __init__(self, text_dim: int, hidden_size: int) -> None:
super().__init__()
self.norm = LingBotVideoRMSNorm(text_dim, eps=1e-6)
self.linear_1 = LingBotVideoLinear(text_dim, hidden_size, bias=True)
self.linear_2 = LingBotVideoLinear(hidden_size, hidden_size, bias=True)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Apply RMSNorm followed by the released two-layer SiLU projection."""
hidden_states = self.norm(hidden_states)
return self.linear_2(F.silu(self.linear_1(hidden_states)))
class LingBotVideoAttention(nn.Module):
"""Joint video-text attention shared by the Dense and MoE variants."""
def __init__(
self,
hidden_size: int,
num_heads: int,
norm_eps: float,
qkv_bias: bool,
out_bias: bool,
) -> None:
"""Create released QKV projections, per-head norms, and output projection."""
super().__init__()
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.to_q = LingBotVideoLinear(hidden_size, hidden_size, bias=qkv_bias)
self.to_k = LingBotVideoLinear(hidden_size, hidden_size, bias=qkv_bias)
self.to_v = LingBotVideoLinear(hidden_size, hidden_size, bias=qkv_bias)
self.norm_q = LingBotVideoRMSNorm(self.head_dim, norm_eps)
self.norm_k = LingBotVideoRMSNorm(self.head_dim, norm_eps)
self.to_out = LingBotVideoLinear(hidden_size, hidden_size, bias=out_bias)
def forward(
self,
hidden_states: torch.Tensor,
rotary_emb: torch.Tensor,
attention_mask: torch.Tensor | None = None,
original_seq_len: int | None = None,
) -> torch.Tensor:
"""Project QKV, apply rotary embeddings, and run non-causal SDPA."""
batch, sequence, _ = hidden_states.shape
query = self.to_q(hidden_states).unflatten(2, (self.num_heads, self.head_dim))
key = self.to_k(hidden_states).unflatten(2, (self.num_heads, self.head_dim))
value = self.to_v(hidden_states).unflatten(2, (self.num_heads, self.head_dim))
query = _apply_rotary_emb(self.norm_q(query), rotary_emb)
key = _apply_rotary_emb(self.norm_k(key), rotary_emb)
if original_seq_len is not None:
# Attention needs the full unpadded joint sequence while projections
# and residual blocks stay sharded over tokens.
qkv = torch.cat([query, key, value], dim=0)
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
padded_seq_len = qkv.shape[1]
query, key, value = qkv[:, :original_seq_len].chunk(3, dim=0)
output = F.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
attn_mask=attention_mask,
dropout_p=0.0,
is_causal=False,
).transpose(1, 2)
if original_seq_len is not None:
output = F.pad(output, (0, 0, 0, 0, 0, padded_seq_len - original_seq_len))
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
return self.to_out(output.reshape(batch, sequence, -1).type_as(hidden_states))
class LingBotVideoMLP(nn.Module):
"""Dense SwiGLU feed-forward network."""
def __init__(self, hidden_size: int, intermediate_size: int) -> None:
super().__init__()
self.gate_proj = LingBotVideoLinear(hidden_size, intermediate_size, bias=False)
self.up_proj = LingBotVideoLinear(hidden_size, intermediate_size, bias=False)
self.down_proj = LingBotVideoLinear(intermediate_size, hidden_size, bias=False)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Apply the released SiLU-gated MLP ordering."""
return self.down_proj(F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states))
class LingBotVideoRouter(nn.Module):
"""Released token-choice router with bias-only expert selection correction."""
def __init__(
self,
hidden_size: int,
num_experts: int,
top_k: int,
score_func: str,
norm_topk_prob: bool,
n_group: int | None,
topk_group: int | None,
route_scale: float,
) -> None:
"""Create fp32-routed expert scores with the released persistent bias."""
super().__init__()
self.num_experts = num_experts
self.top_k = top_k
self.score_func = score_func
self.norm_topk_prob = norm_topk_prob
self.n_group = n_group
self.topk_group = topk_group
self.route_scale = route_scale
self.weight = nn.Parameter(torch.empty(num_experts, hidden_size))
self.register_buffer("e_score_correction_bias", torch.zeros(num_experts), persistent=True)
def _group_limited_topk(self, scores_for_choice: torch.Tensor) -> torch.Tensor:
"""Restrict token choices to groups with the two strongest expert scores."""
sequence_length = scores_for_choice.shape[0]
experts_per_group = self.num_experts // self.n_group
grouped = scores_for_choice.view(sequence_length, self.n_group, experts_per_group)
group_scores = grouped.topk(2, dim=-1)[0].sum(dim=-1)
group_indices = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1]
group_mask = torch.zeros_like(group_scores)
group_mask.scatter_(1, group_indices, 1)
score_mask = (group_mask.unsqueeze(-1).expand(sequence_length, self.n_group,
experts_per_group).reshape(sequence_length, -1))
masked_scores = scores_for_choice.masked_fill(~score_mask.bool(), float("-inf"))
return torch.topk(masked_scores, k=self.top_k, dim=-1, sorted=False)[1]
def forward(self,
tokens: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Score in fp32, select with correction bias, and weight without it."""
with torch.amp.autocast(tokens.device.type, enabled=False):
logits = F.linear(tokens.float(), self.weight.float())
scores = F.softmax(logits, dim=-1) if self.score_func == "softmax" else logits.sigmoid()
scores_for_choice = scores + self.e_score_correction_bias.unsqueeze(0)
if self.n_group is not None and self.n_group > 1:
top_indices = self._group_limited_topk(scores_for_choice)
else:
top_indices = torch.topk(scores_for_choice, k=self.top_k, dim=-1, sorted=False)[1]
top_scores = scores.gather(1, top_indices)
if self.top_k > 1 and self.norm_topk_prob:
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-20)
top_scores = top_scores * self.route_scale
return top_indices, top_scores.to(tokens.dtype), logits, scores, scores_for_choice
class LingBotVideoGroupedExperts(nn.Module):
"""Released grouped-expert parameter layout: w1/w3 `[E,I,H]`, w2 `[E,H,I]`."""
def __init__(self, num_experts: int, hidden_size: int, intermediate_size: int) -> None:
super().__init__()
self.num_experts = num_experts
self.w1 = nn.Parameter(torch.empty(num_experts, intermediate_size, hidden_size))
self.w2 = nn.Parameter(torch.empty(num_experts, hidden_size, intermediate_size))
self.w3 = nn.Parameter(torch.empty(num_experts, intermediate_size, hidden_size))
def _round_up_to_multiple(value: int, multiple: int) -> int:
"""Round an integer up to the next multiple."""
return ((value + multiple - 1) // multiple) * multiple
class LingBotVideoSparseMoeBlock(nn.Module):
"""Token-choice sparse feed-forward block matching the released state surface."""
def __init__(
self,
hidden_size: int,
intermediate_size: int,
num_experts: int,
top_k: int,
moe_intermediate_size: int,
score_func: str,
norm_topk_prob: bool,
n_group: int | None,
topk_group: int | None,
routed_scaling_factor: float,
n_shared_experts: int | None,
) -> None:
"""Create routed and optional shared experts with released parameter names."""
super().__init__()
del intermediate_size
self.hidden_size = hidden_size
self.num_experts = num_experts
self.router = LingBotVideoRouter(
hidden_size,
num_experts,
top_k,
score_func,
norm_topk_prob,
n_group,
topk_group,
routed_scaling_factor,
)
self.experts = LingBotVideoGroupedExperts(num_experts, hidden_size, moe_intermediate_size)
self.shared_experts = None
if n_shared_experts is not None and n_shared_experts > 0:
self.shared_experts = LingBotVideoMLP(hidden_size, moe_intermediate_size * n_shared_experts)
@staticmethod
def _reorder_tokens(
tokens: torch.Tensor,
top_scores: torch.Tensor,
top_indices: torch.Tensor,
num_experts: int,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
"""Pack active token choices into stable expert-major order."""
num_tokens = tokens.shape[0]
top_k = top_indices.shape[1]
flat_scores = top_scores.reshape(-1)
flat_indices = top_indices.reshape(-1)
active_positions = torch.where(flat_scores != 0)[0]
active_experts = flat_indices[active_positions]
counts = torch.zeros(num_experts, device=tokens.device, dtype=torch.int64)
counts.scatter_add_(0, active_experts, torch.ones_like(active_experts, dtype=torch.int64))
sort_order = torch.argsort(active_experts, stable=True)
sorted_positions = active_positions[sort_order]
sorted_scores = flat_scores[sorted_positions]
original_token_indices = sorted_positions // top_k
permuted_tokens = tokens[original_token_indices]
return permuted_tokens, counts, sorted_positions, sorted_scores, num_tokens, top_k
@staticmethod
def _pad_grouped_tokens(tokens: torch.Tensor,
counts: torch.Tensor,
align: int = 8) -> tuple[torch.Size, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Align each expert segment for `torch._grouped_mm` and retain unpad indices."""
num_tokens = tokens.shape[0]
num_experts = int(counts.shape[0])
max_length = _round_up_to_multiple(num_tokens + num_experts * align, align)
# Build the small padding index on CPU after one synchronization instead
# of reading three CUDA scalars for every expert.
counts_cpu = counts.to(device="cpu", dtype=torch.int64)
total_per_expert = torch.clamp_min(counts_cpu, align)
aligned_counts_cpu = ((total_per_expert + align - 1) // align * align).to(torch.int32)
write_offsets = torch.cumsum(aligned_counts_cpu, dim=0) - aligned_counts_cpu
start_indices = torch.cumsum(counts_cpu, dim=0) - counts_cpu
permuted_indices_cpu = torch.full((max_length, ), num_tokens, dtype=torch.int64, device="cpu")
for expert_index in range(num_experts):
length = int(counts_cpu[expert_index])
if length == 0:
continue
write_start = int(write_offsets[expert_index])
start = int(start_indices[expert_index])
permuted_indices_cpu[write_start:write_start + length] = torch.arange(start,
start + length,
device="cpu",
dtype=torch.int64)
permuted_indices = permuted_indices_cpu.to(tokens.device)
aligned_counts = aligned_counts_cpu.to(tokens.device)
tokens_with_pad = torch.vstack((tokens, tokens.new_zeros((tokens.shape[-1], ))))
input_shape = tokens_with_pad.shape
return input_shape, tokens_with_pad[permuted_indices], permuted_indices, aligned_counts
@staticmethod
def _unpad_grouped_tokens(output: torch.Tensor, input_shape: torch.Size,
permuted_indices: torch.Tensor) -> torch.Tensor:
"""Undo per-expert alignment while dropping the shared padding row."""
unpermuted = output.new_empty(input_shape)
unpermuted[permuted_indices, :] = output
return unpermuted[:-1]
def _run_grouped_experts(self, tokens: torch.Tensor, counts: torch.Tensor) -> torch.Tensor:
"""Use the released bf16 grouped matmuls on CUDA and an eager CPU fallback."""
if tokens.device.type == "cpu" or not hasattr(torch, "_grouped_mm"):
return self._run_experts_for_loop(tokens, counts)
input_shape, padded_tokens, permuted_indices, aligned_counts = self._pad_grouped_tokens(tokens, counts)
offsets = torch.cumsum(aligned_counts, dim=0, dtype=torch.int32)
hidden = F.silu(
torch._grouped_mm(
padded_tokens.bfloat16(),
self.experts.w1.bfloat16().transpose(-2, -1),
offs=offsets,
))
hidden = hidden * torch._grouped_mm(
padded_tokens.bfloat16(),
self.experts.w3.bfloat16().transpose(-2, -1),
offs=offsets,
)
output = torch._grouped_mm(
hidden,
self.experts.w2.bfloat16().transpose(-2, -1),
offs=offsets,
).type_as(padded_tokens)
return self._unpad_grouped_tokens(output, input_shape, permuted_indices)
def _run_experts_for_loop(self, tokens: torch.Tensor, counts: torch.Tensor) -> torch.Tensor:
"""Evaluate contiguous expert segments eagerly for CPU correctness tests."""
splits = torch.split(tokens, counts.tolist(), dim=0)
outputs: list[torch.Tensor] = []
for expert_index, expert_tokens in enumerate(splits):
if expert_tokens.numel() == 0:
continue
hidden = F.silu(expert_tokens @ self.experts.w1[expert_index].transpose(-2, -1))
hidden = hidden * (expert_tokens @ self.experts.w3[expert_index].transpose(-2, -1))
outputs.append(hidden @ self.experts.w2[expert_index].transpose(-2, -1))
if not outputs:
return tokens.new_zeros(tokens.shape)
return torch.cat(outputs, dim=0)
@staticmethod
def _restore_tokens(
expert_output: torch.Tensor,
sorted_positions: torch.Tensor,
sorted_scores: torch.Tensor,
num_tokens: int,
top_k: int,
) -> torch.Tensor:
"""Restore token order and combine expert outputs with fp32 weighted sums."""
hidden_size = expert_output.shape[-1]
unsorted = torch.zeros(
(num_tokens * top_k, hidden_size),
dtype=expert_output.dtype,
device=expert_output.device,
)
unsorted[sorted_positions] = expert_output
unsorted = unsorted.reshape(num_tokens, top_k, hidden_size)
scores_unsorted = torch.zeros(
num_tokens * top_k,
dtype=sorted_scores.dtype,
device=sorted_scores.device,
)
scores_unsorted[sorted_positions] = sorted_scores
scores_unsorted = scores_unsorted.reshape(num_tokens, top_k, 1)
return (unsorted.float() * scores_unsorted).sum(dim=1).to(expert_output.dtype)
def _run_selected_experts(
self,
tokens: torch.Tensor,
top_scores: torch.Tensor,
top_indices: torch.Tensor,
) -> torch.Tensor:
"""Dispatch routed choices, execute experts, and restore token-major order."""
permuted_tokens, counts, sorted_positions, sorted_scores, num_tokens, top_k = self._reorder_tokens(
tokens, top_scores, top_indices, self.router.num_experts)
expert_output = self._run_grouped_experts(permuted_tokens, counts)
return self._restore_tokens(expert_output, sorted_positions, sorted_scores, num_tokens, top_k)
def forward(self, hidden_states: torch.Tensor, padding_mask: torch.Tensor | None = None) -> torch.Tensor:
"""Route token choices, zero padded routes, and add optional shared experts."""
batch = hidden_states.shape[0]
tokens = hidden_states.view(-1, self.hidden_size)
top_indices, top_scores, logits, scores, scores_for_choice = self.router(tokens)
del logits, scores, scores_for_choice
if padding_mask is not None:
mask = padding_mask.unsqueeze(-1).to(top_scores.dtype)
top_scores = top_scores * mask
top_scores = top_scores / (top_scores.sum(dim=-1, keepdim=True) + 1e-9)
top_scores = top_scores * self.router.route_scale
output = self._run_selected_experts(tokens, top_scores, top_indices)
output = output.view(batch, -1, self.hidden_size)
if self.shared_experts is not None:
output = output + self.shared_experts(hidden_states)
return output
class LingBotVideoBlock(nn.Module):
"""One Dense or sparse LingBot-Video transformer block."""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
intermediate_size: int,
norm_eps: float,
qkv_bias: bool,
out_bias: bool,
num_experts: int,
num_experts_per_tok: int,
moe_intermediate_size: int,
decoder_sparse_step: int,
mlp_only_layers: tuple[int, ...] | list[int],
n_shared_experts: int | None,
score_func: str,
norm_topk_prob: bool,
n_group: int | None,
topk_group: int | None,
routed_scaling_factor: float,
layer_idx: int,
) -> None:
"""Select the released Dense or sparse feed-forward structure for one layer."""
super().__init__()
self.layer_idx = layer_idx
self.scale_shift_table = nn.Parameter(torch.zeros(1, 6 * hidden_size))
self.norm1 = LingBotVideoRMSNorm(hidden_size, norm_eps)
self.attn = LingBotVideoAttention(hidden_size, num_attention_heads, norm_eps, qkv_bias, out_bias)
self.norm_post_attn = LingBotVideoRMSNorm(hidden_size, norm_eps)
self.norm2 = LingBotVideoRMSNorm(hidden_size, norm_eps)
if layer_idx not in mlp_only_layers and (num_experts > 0 and (layer_idx + 1) % decoder_sparse_step == 0):
self.ffn = LingBotVideoSparseMoeBlock(
hidden_size,
intermediate_size,
num_experts,
num_experts_per_tok,
moe_intermediate_size,
score_func,
norm_topk_prob,
n_group,
topk_group,
routed_scaling_factor,
n_shared_experts,
)
else:
self.ffn = LingBotVideoMLP(hidden_size, intermediate_size)
self.norm_post_ffn = LingBotVideoRMSNorm(hidden_size, norm_eps)
def forward(
self,
hidden_states: torch.Tensor,
temb6: torch.Tensor,
rotary_emb: torch.Tensor,
attention_mask: torch.Tensor | None = None,
moe_padding_mask: torch.Tensor | None = None,
original_seq_len: int | None = None,
) -> torch.Tensor:
"""Run attention and configured feed-forward residual branches with fp32 AdaLN."""
expected_tokens = hidden_states.shape[0] * hidden_states.shape[1]
if temb6.ndim != 2 or temb6.shape[0] != expected_tokens:
raise ValueError("LingBotVideoBlock expects token-level temb6 with shape "
f"(B*S, 6D); got {tuple(temb6.shape)} for {tuple(hidden_states.shape)}.")
modulation = temb6.view(*hidden_states.shape[:2], -1) + self.scale_shift_table.unsqueeze(0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = modulation.chunk(6, dim=-1)
gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh()
scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp
bulk_dtype = self.attn.to_q.weight.dtype
attention_input = (self.norm1(hidden_states) * scale_msa + shift_msa).to(bulk_dtype)
attention_output = self.attn(attention_input, rotary_emb, attention_mask, original_seq_len)
hidden_states = hidden_states + (gate_msa * self.norm_post_attn(attention_output)).to(hidden_states.dtype)
mlp_input = (self.norm2(hidden_states) * scale_mlp + shift_mlp).to(bulk_dtype)
if isinstance(self.ffn, LingBotVideoSparseMoeBlock):
mlp_output = self.ffn(mlp_input, padding_mask=moe_padding_mask)
else:
mlp_output = self.ffn(mlp_input)
mlp_output = self.norm_post_ffn(mlp_output)
return hidden_states + (gate_mlp * mlp_output).to(hidden_states.dtype)
class LingBotVideoTimestepEmbedding(nn.Module):
"""Two-layer timestep embedding with released parameter names."""
def __init__(self, input_dim: int, hidden_size: int, bias: bool) -> None:
super().__init__()
self.linear_1 = LingBotVideoLinear(input_dim, hidden_size, bias=bias)
self.linear_2 = LingBotVideoLinear(hidden_size, hidden_size, bias=bias)
def forward(self, sample: torch.Tensor) -> torch.Tensor:
"""Apply the official linear-SiLU-linear timestep projection."""
return self.linear_2(F.silu(self.linear_1(sample)))
class LingBotVideoTransformer3DModel(BaseDiT):
"""LingBot-Video DiT with a source-compatible Dense or MoE state surface."""
_fsdp_shard_conditions = [_is_lingbot_video_block]
_compile_conditions = [_is_lingbot_video_block]
_supported_attention_backends = (
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.FLASH_ATTN,
)
param_names_mapping = {r"^(.*)$": r"\1"}
reverse_param_names_mapping = {r"^(.*)$": r"\1"}
def _get_parameter_dtype(self, name: str, default_dtype: torch.dtype) -> torch.dtype:
"""Select the released mixed-precision dtype while loading each parameter."""
return torch.float32 if _keep_in_fp32(name) else default_dtype
def __init__(self, config: LingBotVideoConfig, hf_config: dict[str, Any]) -> None:
"""Construct the Dense or MoE variant from its released transformer config."""
config.update_model_arch(hf_config)
super().__init__(config, hf_config)
head_dim = config.hidden_size // config.num_attention_heads
if head_dim != sum(config.axes_dims):
raise ValueError(f"head_dim {head_dim} != sum(axes_dims) {sum(config.axes_dims)}")
sp_world_size = _sequence_parallel_world_size()
assert config.num_attention_heads % sp_world_size == 0, (
f"The number of attention heads ({config.num_attention_heads}) must be divisible by "
f"the sequence parallel size ({sp_world_size})")
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.in_channels
self.patch_embedder = LingBotVideoLinear(
config.in_channels * math.prod(config.patch_size),
config.hidden_size,
bias=config.patch_embed_bias,
)
self.time_proj = Timesteps(config.freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
self.time_embedder = LingBotVideoTimestepEmbedding(config.freq_dim, config.hidden_size,
config.timestep_mlp_bias)
self.time_modulation = nn.Sequential(nn.SiLU(), LingBotVideoLinear(config.hidden_size, 6 * config.hidden_size))
self.text_embedder = LingBotVideoTextEmbedder(config.text_dim, config.hidden_size)
self.rope = LingBotVideoRotaryEmbedding(tuple(config.axes_dims), tuple(config.axes_lens), config.rope_theta)
self.blocks = nn.ModuleList([
LingBotVideoBlock(
config.hidden_size,
config.num_attention_heads,
config.intermediate_size,
config.norm_eps,
config.qkv_bias,
config.out_bias,
config.num_experts,
config.num_experts_per_tok,
config.moe_intermediate_size,
config.decoder_sparse_step,
config.mlp_only_layers,
config.n_shared_experts,
config.score_func,
config.norm_topk_prob,
config.n_group,
config.topk_group,
config.routed_scaling_factor,
layer_index,
) for layer_index in range(config.depth)
])
self.norm_out = nn.LayerNorm(config.hidden_size, elementwise_affine=False, eps=config.norm_eps)
self.norm_out_modulation = nn.Sequential(nn.SiLU(),
LingBotVideoLinear(config.hidden_size, 2 * config.hidden_size))
self.proj_out = LingBotVideoLinear(config.hidden_size, math.prod(config.patch_size) * config.out_channels)
self.__post_init__()
def to(self, *args: Any, **kwargs: Any) -> LingBotVideoTransformer3DModel:
"""Cast bulk weights while retaining the released fp32-sensitive modules."""
device, dtype, non_blocking, _ = torch._C._nn._parse_to(*args, **kwargs)
if dtype is None or dtype == torch.float32:
return super().to(*args, **kwargs)
if not torch.is_floating_point(torch.empty((), dtype=dtype)):
return super().to(*args, **kwargs)
if device is not None:
super().to(device=device, non_blocking=non_blocking)
for name, parameter in self.named_parameters():
if torch.is_floating_point(parameter):
target_dtype = torch.float32 if _keep_in_fp32(name) else dtype
parameter.data = parameter.data.to(target_dtype, non_blocking=non_blocking)
if parameter.grad is not None:
parameter.grad.data = parameter.grad.data.to(target_dtype, non_blocking=non_blocking)
for name, buffer in self.named_buffers():
if torch.is_floating_point(buffer):
target_dtype = torch.float32 if _keep_in_fp32(name) else dtype
buffer.data = buffer.data.to(target_dtype, non_blocking=non_blocking)
return self
def forward(
self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance: torch.Tensor | None = None,
encoder_attention_mask: torch.Tensor | None = None,
return_dict: bool = True,
**kwargs: Any,
) -> LingBotVideoTransformerOutput | tuple[torch.Tensor]:
"""Denoise video latents with joint video-text attention."""
del encoder_hidden_states_image, guidance, kwargs
if isinstance(encoder_hidden_states, list):
if len(encoder_hidden_states) != 1:
raise ValueError("LingBot-Video expects one text-encoder output tensor.")
encoder_hidden_states = encoder_hidden_states[0]
batch, channels, frames, height, width = hidden_states.shape
patch_t, patch_h, patch_w = self.config.patch_size
grid_t, grid_h, grid_w = frames // patch_t, height // patch_h, width // patch_w
video_tokens = grid_t * grid_h * grid_w
text_tokens = encoder_hidden_states.shape[1]
device = hidden_states.device
if encoder_attention_mask is None:
encoder_attention_mask = torch.ones((batch, text_tokens), device=device, dtype=torch.bool)
text_lengths = encoder_attention_mask.sum(dim=-1).long()
patches = hidden_states.reshape(batch, channels, grid_t, patch_t, grid_h, patch_h, grid_w, patch_w)
patches = patches.permute(0, 2, 4, 6, 3, 5, 7, 1).reshape(batch, video_tokens,
patch_t * patch_h * patch_w * channels)
video_hidden = self.patch_embedder(patches)
text_hidden = self.text_embedder(encoder_hidden_states)
joint = torch.cat([video_hidden, text_hidden], dim=1)
rotary_parts: list[torch.Tensor] = []
for index in range(batch):
real_text_length = int(text_lengths[index].item())
positions = _make_joint_position_ids(real_text_length, grid_t, grid_h, grid_w, device)
maxima = (real_text_length + grid_t, grid_h - 1, grid_w - 1)
rotary = self.rope(positions, maxima=maxima)
if real_text_length < text_tokens:
padding = torch.zeros(
text_tokens - real_text_length,
rotary.shape[-1],
device=device,
dtype=rotary.dtype,
)
rotary = torch.cat([rotary, padding], dim=0)
rotary_parts.append(rotary)
rotary_emb = torch.stack(rotary_parts, dim=0)
attention_mask = None
moe_padding_mask = None
if bool((text_lengths < text_tokens).any()):
key_mask = torch.cat(
[
torch.ones(batch, video_tokens, dtype=torch.bool, device=device),
encoder_attention_mask.bool(),
],
dim=1,
)
attention_mask = key_mask[:, None, None, :]
moe_padding_mask = key_mask
timestep_projection = self.time_proj(timestep.float())
timestep_embedding = self.time_embedder(timestep_projection)
token_embedding = timestep_embedding.unsqueeze(1).expand(batch, joint.shape[1], -1)
original_joint_length: int | None = None
if _sequence_parallel_world_size() > 1:
# Match the official CP order: project token modulation before
# placing every token-aligned tensor on the same padded shard.
temb6 = self.time_modulation(token_embedding.reshape(-1,
self.hidden_size)).reshape(batch, joint.shape[1], -1)
joint, original_joint_length = sequence_model_parallel_shard(joint, dim=1)
rotary_emb, _ = sequence_model_parallel_shard(rotary_emb, dim=1)
token_embedding, _ = sequence_model_parallel_shard(token_embedding, dim=1)
temb6, _ = sequence_model_parallel_shard(temb6, dim=1)
if moe_padding_mask is None:
moe_padding_mask = torch.ones(batch, original_joint_length, dtype=torch.bool, device=device)
moe_padding_mask, _ = sequence_model_parallel_shard(moe_padding_mask, dim=1)
moe_padding_mask = moe_padding_mask.reshape(-1)
temb6 = temb6.reshape(-1, 6 * self.hidden_size)
else:
temb6 = self.time_modulation(token_embedding.reshape(-1, self.hidden_size))
if moe_padding_mask is not None:
moe_padding_mask = moe_padding_mask.reshape(-1)
for block in self.blocks:
joint = block(
joint,
temb6,
rotary_emb,
attention_mask,
moe_padding_mask,
original_joint_length,
)
final_modulation = self.norm_out_modulation(token_embedding.reshape(-1, self.hidden_size))
shift, scale = final_modulation.reshape(*joint.shape[:2], -1).chunk(2, dim=-1)
final_hidden = self.norm_out(joint) * (1.0 + scale) + shift
projected = self.proj_out(final_hidden.to(self.proj_out.weight.dtype))
if original_joint_length is not None:
projected = sequence_model_parallel_all_gather_with_unpad(projected, original_joint_length, dim=1)
projected = projected[:, :video_tokens]
output_channels = self.config.out_channels
output = projected.reshape(batch, grid_t, grid_h, grid_w, patch_t, patch_h, patch_w, output_channels)
output = output.permute(0, 7, 1, 4, 2, 5, 3, 6).reshape(batch, output_channels, frames, height, width)
if not return_dict:
return (output, )
return LingBotVideoTransformerOutput(sample=output)
EntryClass = LingBotVideoTransformer3DModel
@@ -0,0 +1,5 @@
from .causal_fast import LingBotWorld2CausalFastTransformer3DModel
__all__ = ["LingBotWorld2CausalFastTransformer3DModel"]
EntryClass = LingBotWorld2CausalFastTransformer3DModel
@@ -0,0 +1,204 @@
# SPDX-License-Identifier: Apache-2.0
import numpy as np
import os
import torch
from scipy.interpolate import interp1d
from scipy.spatial.transform import Rotation, Slerp
# --- Official Code (Leave Unchanged) ---
def interpolate_camera_poses(
src_indices: np.ndarray,
src_rot_mat: np.ndarray,
src_trans_vec: np.ndarray,
tgt_indices: np.ndarray,
) -> torch.Tensor:
# interpolate translation
interp_func_trans = interp1d(
src_indices,
src_trans_vec,
axis=0,
kind='linear',
bounds_error=False,
fill_value="extrapolate",
)
interpolated_trans_vec = interp_func_trans(tgt_indices)
# interpolate rotation
src_quat_vec = Rotation.from_matrix(src_rot_mat)
# ensure there is no sudden change in qw
quats = src_quat_vec.as_quat().copy() # [N, 4]
for i in range(1, len(quats)):
if np.dot(quats[i], quats[i-1]) < 0:
quats[i] = -quats[i]
src_quat_vec = Rotation.from_quat(quats)
slerp_func_rot = Slerp(src_indices, src_quat_vec)
interpolated_rot_quat = slerp_func_rot(tgt_indices)
interpolated_rot_mat = interpolated_rot_quat.as_matrix()
poses = np.zeros((len(tgt_indices), 4, 4))
poses[:, :3, :3] = interpolated_rot_mat
poses[:, :3, 3] = interpolated_trans_vec
poses[:, 3, 3] = 1.0
return torch.from_numpy(poses).float()
def SE3_inverse(T: torch.Tensor) -> torch.Tensor:
Rot = T[:, :3, :3] # [B,3,3]
trans = T[:, :3, 3:] # [B,3,1]
R_inv = Rot.transpose(-1, -2)
t_inv = -torch.bmm(R_inv, trans)
T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat(T.shape[0], 1, 1)
T_inv[:, :3, :3] = R_inv
T_inv[:, :3, 3:] = t_inv
return T_inv
def compute_relative_poses(
c2ws_mat: torch.Tensor,
framewise: bool = False,
normalize_trans: bool = True,
) -> torch.Tensor:
ref_w2cs = SE3_inverse(c2ws_mat[0:1])
relative_poses = torch.matmul(ref_w2cs, c2ws_mat)
# ensure identity matrix for 1st frame
relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype)
if framewise:
# compute pose between i and i+1
relative_poses_framewise = torch.bmm(SE3_inverse(relative_poses[:-1]), relative_poses[1:])
relative_poses[1:] = relative_poses_framewise
if normalize_trans: # note refer to camctrl2: "we scale the coordinate inputs to roughly 1 standard deviation to simplify model learning."
translations = relative_poses[:, :3, 3] # [f, 3]
max_norm = torch.norm(translations, dim=-1).max()
# only normlaize when moving
if max_norm > 0:
relative_poses[:, :3, 3] = translations / max_norm
return relative_poses
@torch.no_grad()
def create_meshgrid(n_frames: int, height: int, width: int, bias: float = 0.5, device='cuda', dtype=torch.float32) -> torch.Tensor:
x_range = torch.arange(width, device=device, dtype=dtype)
y_range = torch.arange(height, device=device, dtype=dtype)
grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing='ij')
grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias # [h*w, 2]
grid_xy = grid_xy[None, ...].repeat(n_frames, 1, 1) # [f, h*w, 2]
return grid_xy
def get_plucker_embeddings(
c2ws_mat: torch.Tensor,
Ks: torch.Tensor,
height: int,
width: int,
):
n_frames = c2ws_mat.shape[0]
device = c2ws_mat.device
dtype = c2ws_mat.dtype
grid_y, grid_x = torch.meshgrid(
torch.arange(height, device=device, dtype=dtype) + 0.5,
torch.arange(width, device=device, dtype=dtype) + 0.5,
indexing='ij',
)
x_flat = grid_x.reshape(-1)
y_flat = grid_y.reshape(-1)
fx, fy, cx, cy = Ks[0, 0], Ks[0, 1], Ks[0, 2], Ks[0, 3]
dirs = torch.stack([
(x_flat - cx) / fx,
(y_flat - cy) / fy,
torch.ones_like(x_flat),
], dim=-1)
dirs = dirs / dirs.norm(dim=-1, keepdim=True)
rays_d = (c2ws_mat[:, :3, :3] @ dirs.T).transpose(1, 2)
rays_o = c2ws_mat[:, :3, 3].unsqueeze(1).expand_as(rays_d)
plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1)
return plucker_embeddings.view(n_frames, height, width, 6)
def get_Ks_transformed(
Ks: torch.Tensor,
height_org: int,
width_org: int,
height_resize: int,
width_resize: int,
height_final: int,
width_final: int,
):
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
scale_x = width_resize / width_org
scale_y = height_resize / height_org
fx_resize = fx * scale_x
fy_resize = fy * scale_y
cx_resize = cx * scale_x
cy_resize = cy * scale_y
crop_offset_x = (width_resize - width_final) / 2
crop_offset_y = (height_resize - height_final) / 2
cx_final = cx_resize - crop_offset_x
cy_final = cy_resize - crop_offset_y
Ks_transformed = torch.zeros_like(Ks)
Ks_transformed[:, 0:1] = fx_resize
Ks_transformed[:, 1:2] = fy_resize
Ks_transformed[:, 2:3] = cx_final
Ks_transformed[:, 3:4] = cy_final
return Ks_transformed
# --- Custom ---
def prepare_camera_embedding(
action_path: str,
num_frames: int,
height: int,
width: int,
spatial_scale: int = 8,
) -> tuple[torch.Tensor, int]:
c2ws = np.load(os.path.join(action_path, "poses.npy"))
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
num_frames = min(num_frames, len_c2ws)
c2ws = c2ws[:num_frames]
Ks = torch.from_numpy(
np.load(os.path.join(action_path, "intrinsics.npy"))
).float()
Ks = get_Ks_transformed(
Ks,
height_org=480,
width_org=832,
height_resize=height,
width_resize=width,
height_final=height,
width_final=width,
)
Ks = Ks[0] # use first frame
len_c2ws = len(c2ws)
num_latent_frames = (len_c2ws - 1) // 4 + 1
c2ws_infer = interpolate_camera_poses(
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
src_rot_mat=c2ws[:, :3, :3],
src_trans_vec=c2ws[:, :3, 3],
tgt_indices=np.linspace(0, len_c2ws - 1, num_latent_frames),
)
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
Ks = Ks.repeat(num_latent_frames, 1)
plucker = get_plucker_embeddings(c2ws_infer, Ks, height, width) # [F, H, W, 6]
# reshpae
latent_height = height // spatial_scale
latent_width = width // spatial_scale
plucker = plucker.view(num_latent_frames, latent_height, spatial_scale, latent_width, spatial_scale, 6)
plucker = plucker.permute(0, 1, 3, 5, 2, 4).contiguous()
plucker = plucker.view(num_latent_frames, latent_height, latent_width, 6 * spatial_scale * spatial_scale)
c2ws_plucker_emb = plucker.permute(3, 0, 1, 2).contiguous().unsqueeze(0)
return c2ws_plucker_emb, num_frames
@@ -0,0 +1,776 @@
# SPDX-License-Identifier: Apache-2.0
"""LingBot World 2 causal-fast DiT implemented inside FastVideo."""
import math
import warnings
from typing import Any
from einops import rearrange
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.dits.lingbotworld2 import (
LingBotWorld2CausalFastVideoConfig,
)
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather,
sequence_model_parallel_all_to_all_4D,
)
from fastvideo.distributed.parallel_state import get_sp_parallel_rank, get_sp_world_size
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
try:
import flash_attn_interface
FLASH_ATTN_3_AVAILABLE = True
except ModuleNotFoundError:
FLASH_ATTN_3_AVAILABLE = False
try:
import flash_attn
FLASH_ATTN_2_AVAILABLE = True
except ModuleNotFoundError:
FLASH_ATTN_2_AVAILABLE = False
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
def sinusoidal_embedding_1d(dim: int, position: torch.Tensor) -> torch.Tensor:
"""Build Wan/LingBot World 2 sinusoidal timestep embeddings."""
assert dim % 2 == 0
half = dim // 2
position = position.type(torch.float64)
sinusoid = torch.outer(
position,
torch.pow(10000, -torch.arange(half, device=position.device).to(position).div(half)),
)
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
@torch.amp.autocast("cuda", enabled=False)
def rope_params(max_seq_len: int, dim: int, theta: int = 10000) -> torch.Tensor:
"""Return complex RoPE frequencies used by the released LingBot World 2 model."""
assert dim % 2 == 0
freqs = torch.outer(
torch.arange(max_seq_len),
1.0 / torch.pow(theta, torch.arange(0, dim, 2).to(torch.float64).div(dim)),
)
return torch.polar(torch.ones_like(freqs), freqs)
def flash_attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q_lens: torch.Tensor | None = None,
k_lens: torch.Tensor | None = None,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
q_scale: float | None = None,
causal: bool = False,
window_size: tuple[int, int] = (-1, -1),
deterministic: bool = False,
dtype: torch.dtype = torch.bfloat16,
version: int | None = None,
) -> torch.Tensor:
"""Run LingBot World 2-compatible FlashAttention on packed varlen inputs."""
half_dtypes = (torch.float16, torch.bfloat16)
assert dtype in half_dtypes
assert q.device.type == "cuda" and q.size(-1) <= 256
b, lq, lk, out_dtype = q.size(0), q.size(1), k.size(1), q.dtype
def half(x: torch.Tensor) -> torch.Tensor:
return x if x.dtype in half_dtypes else x.to(dtype)
if q_lens is None:
q = half(q.flatten(0, 1))
q_lens = torch.tensor([lq] * b, dtype=torch.int32, device=q.device)
else:
q = half(torch.cat([u[:v] for u, v in zip(q, q_lens, strict=True)]))
if k_lens is None:
k = half(k.flatten(0, 1))
v = half(v.flatten(0, 1))
k_lens = torch.tensor([lk] * b, dtype=torch.int32, device=k.device)
else:
k = half(torch.cat([u[:v] for u, v in zip(k, k_lens, strict=True)]))
v = half(torch.cat([u[:v] for u, v in zip(v, k_lens, strict=True)]))
q = q.to(v.dtype)
k = k.to(v.dtype)
if q_scale is not None:
q = q * q_scale
if version == 3 and not FLASH_ATTN_3_AVAILABLE:
warnings.warn("FlashAttention 3 is not available; using FlashAttention 2.")
cu_q = torch.cat([q_lens.new_zeros([1]), q_lens]).cumsum(
0, dtype=torch.int32
).to(q.device, non_blocking=True)
cu_k = torch.cat([k_lens.new_zeros([1]), k_lens]).cumsum(
0, dtype=torch.int32
).to(k.device, non_blocking=True)
if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE:
x = flash_attn_interface.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=cu_q,
cu_seqlens_k=cu_k,
seqused_q=None,
seqused_k=None,
max_seqlen_q=lq,
max_seqlen_k=lk,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
).unflatten(0, (b, lq))
else:
assert FLASH_ATTN_2_AVAILABLE
x = flash_attn.flash_attn_varlen_func(
q=q,
k=k,
v=v,
cu_seqlens_q=cu_q,
cu_seqlens_k=cu_k,
max_seqlen_q=lq,
max_seqlen_k=lk,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic,
).unflatten(0, (b, lq))
return x.type(out_dtype)
def attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q_lens: torch.Tensor | None = None,
k_lens: torch.Tensor | None = None,
dropout_p: float = 0.0,
softmax_scale: float | None = None,
q_scale: float | None = None,
causal: bool = False,
window_size: tuple[int, int] = (-1, -1),
deterministic: bool = False,
dtype: torch.dtype = torch.bfloat16,
fa_version: int | None = None,
) -> torch.Tensor:
"""Dispatch LingBot World 2 attention to FlashAttention when available."""
if FLASH_ATTN_2_AVAILABLE or FLASH_ATTN_3_AVAILABLE:
return flash_attention(
q=q,
k=k,
v=v,
q_lens=q_lens,
k_lens=k_lens,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
q_scale=q_scale,
causal=causal,
window_size=window_size,
deterministic=deterministic,
dtype=dtype,
version=fa_version,
)
if q_lens is not None or k_lens is not None:
warnings.warn("Padding masks are disabled without FlashAttention.")
q = q.transpose(1, 2).to(dtype)
k = k.transpose(1, 2).to(dtype)
v = v.transpose(1, 2).to(dtype)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=None, is_causal=causal, dropout_p=dropout_p
)
return out.transpose(1, 2).contiguous()
@torch.amp.autocast("cuda", enabled=False)
def causal_rope_apply(
x: torch.Tensor,
grid_sizes: torch.Tensor,
freqs: torch.Tensor,
start_frame: int = 0,
) -> torch.Tensor:
"""Apply LingBot World 2 causal RoPE with the current chunk frame offset."""
n, c = x.size(2), x.size(3) // 2
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(seq_len, n, -1, 2))
freqs_i = torch.cat(
[
freqs[0][start_frame : start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1),
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1),
],
dim=-1,
).reshape(seq_len, 1, -1)
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
output.append(x_i)
return torch.stack(output).type_as(x)
class WanRMSNorm(nn.Module):
"""RMSNorm used by Wan/LingBot World 2 attention projections."""
def __init__(self, dim: int, eps: float = 1e-5):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Normalize the last dimension in fp32 and restore input dtype."""
return self._norm(x.float()).type_as(x) * self.weight
def _norm(self, x: torch.Tensor) -> torch.Tensor:
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
class WanLayerNorm(nn.LayerNorm):
"""LayerNorm variant that computes in fp32 and returns the input dtype."""
def __init__(self, dim: int, eps: float = 1e-6, elementwise_affine: bool = False):
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Apply layer norm in fp32 for LingBot World 2 numerical parity."""
return super().forward(x.float()).type_as(x)
class CausalWanSelfAttention(nn.Module):
"""LingBot World 2 causal self-attention with rolling KV cache."""
def __init__(
self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm: bool = True,
eps: float = 1e-6,
):
super().__init__()
assert dim % num_heads == 0
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.q = nn.Linear(dim, dim)
self.k = nn.Linear(dim, dim)
self.v = nn.Linear(dim, dim)
self.o = nn.Linear(dim, dim)
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(
self,
x: torch.Tensor,
seq_lens: torch.Tensor,
grid_sizes: torch.Tensor,
freqs: torch.Tensor,
kv_cache: dict,
current_start: int = 0,
max_attention_size: int = 1_000_000,
frame_seqlen: int | None = None,
seq_lens_int: int | None = None,
) -> torch.Tensor:
"""Project QKV, update the rolling cache, and attend to its active window."""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
if frame_seqlen is None:
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
current_start_frame = current_start // frame_seqlen
if seq_lens_int is None:
seq_lens_int = int(seq_lens[0].item() if seq_lens.dim() > 0 else seq_lens.item())
sp_size = get_sp_world_size()
if sp_size > 1:
q = sequence_model_parallel_all_to_all_4D(q, scatter_dim=2, gather_dim=1)
k = sequence_model_parallel_all_to_all_4D(k, scatter_dim=2, gather_dim=1)
v = sequence_model_parallel_all_to_all_4D(v, scatter_dim=2, gather_dim=1)
padded_seq_len = s * sp_size
roped_query = causal_rope_apply(q, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
roped_key = causal_rope_apply(k, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
roped_query = roped_query[:, :seq_lens_int]
roped_key = roped_key[:, :seq_lens_int]
v = v[:, :seq_lens_int]
num_new_tokens = seq_lens_int
else:
padded_seq_len = s
roped_query = causal_rope_apply(q, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
roped_key = causal_rope_apply(k, grid_sizes, freqs, start_frame=current_start_frame).type_as(v)
num_new_tokens = roped_query.shape[1]
current_end = current_start + num_new_tokens
sink_tokens = self.sink_size * frame_seqlen
kv_cache_size = kv_cache["k"].shape[1]
if self.local_attn_size == -1:
local_end_index = current_start + num_new_tokens
local_start_index = current_start
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
elif (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size
):
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache["k"][:, sink_tokens : sink_tokens + num_rolled_tokens] = kv_cache["k"][
:, sink_tokens + num_evicted_tokens : sink_tokens + num_evicted_tokens + num_rolled_tokens
].clone()
kv_cache["v"][:, sink_tokens : sink_tokens + num_rolled_tokens] = kv_cache["v"][
:, sink_tokens + num_evicted_tokens : sink_tokens + num_evicted_tokens + num_rolled_tokens
].clone()
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
else:
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
k_cache = kv_cache["k"][:, max(0, local_end_index - max_attention_size) : local_end_index]
v_cache = kv_cache["v"][:, max(0, local_end_index - max_attention_size) : local_end_index]
x = attention(roped_query, k_cache, v_cache)
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
if sp_size > 1:
sp_pad = padded_seq_len - seq_lens_int
if sp_pad > 0:
x = torch.cat([x, x.new_zeros(b, sp_pad, x.size(2), d)], dim=1)
x = sequence_model_parallel_all_to_all_4D(x, scatter_dim=1, gather_dim=2)
return self.o(x.flatten(2))
class WanCrossAttention(CausalWanSelfAttention):
"""LingBot World 2 cross-attention with reusable text K/V cache."""
def forward(
self,
x: torch.Tensor,
context: torch.Tensor,
context_lens: torch.Tensor | None,
crossattn_cache: dict | None = None,
cross_attn_first_call: bool | None = None,
) -> torch.Tensor:
"""Attend hidden states to text context, populating cache on first use."""
b, n, d = x.size(0), self.num_heads, self.head_dim
q = self.norm_q(self.q(x)).view(b, -1, n, d)
if crossattn_cache is not None:
is_first = crossattn_cache["is_init"].item() == 0 if cross_attn_first_call is None else cross_attn_first_call
if is_first:
crossattn_cache["is_init"].fill_(1)
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
crossattn_cache["k"].copy_(k)
crossattn_cache["v"].copy_(v)
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x = flash_attention(q, k, v, k_lens=context_lens)
return self.o(x.flatten(2))
class CausalWanAttentionBlock(nn.Module):
"""One LingBot World 2 causal transformer block including camera injection."""
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm: bool = True,
cross_attn_norm: bool = False,
eps: float = 1e-6,
):
super().__init__()
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.local_attn_size = local_attn_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.norm1 = WanLayerNorm(dim, eps)
self.self_attn = CausalWanSelfAttention(dim, num_heads, local_attn_size, sink_size, qk_norm, eps)
self.norm3 = WanLayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.cross_attn = WanCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
self.norm2 = WanLayerNorm(dim, eps)
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(approximate="tanh"), nn.Linear(ffn_dim, dim))
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
self.cam_injector_layer1 = nn.Linear(dim, dim)
self.cam_injector_layer2 = nn.Linear(dim, dim)
self.cam_scale_layer = nn.Linear(dim, dim)
self.cam_shift_layer = nn.Linear(dim, dim)
def forward(
self,
x: torch.Tensor,
e: torch.Tensor,
seq_lens: torch.Tensor,
grid_sizes: torch.Tensor,
freqs: torch.Tensor,
context: torch.Tensor,
context_lens: torch.Tensor | None,
dit_cond_dict: dict[str, Any] | None = None,
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
max_attention_size: int = 1_000_000,
frame_seqlen: int | None = None,
cross_attn_first_call: bool | None = None,
seq_lens_int: int | None = None,
) -> torch.Tensor:
"""Apply self-attention, camera modulation, cross-attention, and FFN."""
assert kv_cache is not None
assert e.dtype == torch.float32
with torch.amp.autocast("cuda", dtype=torch.float32):
e = (self.modulation.unsqueeze(0) + e).chunk(6, dim=2)
y = self.self_attn(
self.norm1(x).float() * (1 + e[1].squeeze(2)) + e[0].squeeze(2),
seq_lens,
grid_sizes,
freqs,
kv_cache,
current_start,
max_attention_size,
frame_seqlen=frame_seqlen,
seq_lens_int=seq_lens_int,
)
with torch.amp.autocast("cuda", dtype=torch.float32):
x = x + y * e[2].squeeze(2)
if dit_cond_dict is not None and "c2ws_plucker_emb" in dit_cond_dict:
c2ws_plucker_emb = dit_cond_dict["c2ws_plucker_emb"]
c2ws_hidden_states = self.cam_injector_layer2(
F.silu(self.cam_injector_layer1(c2ws_plucker_emb))
)
c2ws_hidden_states = c2ws_hidden_states + c2ws_plucker_emb
x = (1.0 + self.cam_scale_layer(c2ws_hidden_states)) * x + self.cam_shift_layer(c2ws_hidden_states)
x = x + self.cross_attn(
self.norm3(x),
context,
context_lens,
crossattn_cache=crossattn_cache,
cross_attn_first_call=cross_attn_first_call,
)
y = self.ffn(self.norm2(x).float() * (1 + e[4].squeeze(2)) + e[3].squeeze(2))
with torch.amp.autocast("cuda", dtype=torch.float32):
x = x + y * e[5].squeeze(2)
return x
class CausalHead(nn.Module):
"""Output projection head for LingBot World 2 causal-fast DiT."""
def __init__(self, dim: int, out_dim: int, patch_size: tuple[int, int, int], eps: float = 1e-6):
super().__init__()
self.dim = dim
self.out_dim = out_dim
self.patch_size = patch_size
self.eps = eps
self.norm = WanLayerNorm(dim, eps)
self.head = nn.Linear(dim, math.prod(patch_size) * out_dim)
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
def forward(self, x: torch.Tensor, e: torch.Tensor) -> torch.Tensor:
"""Normalize, modulate, and project hidden states to latent patches."""
assert e.dtype == torch.float32
with torch.amp.autocast("cuda", dtype=torch.float32):
e = (self.modulation.unsqueeze(0) + e.unsqueeze(2)).chunk(2, dim=2)
x = self.head(self.norm(x) * (1 + e[1].squeeze(2)) + e[0].squeeze(2))
return x
class LingBotWorld2CausalFastTransformer3DModel(BaseDiT):
"""Released LingBot World 2 14B causal-fast model with native FastVideo loading."""
_fsdp_shard_conditions = [is_blocks]
_compile_conditions: list = []
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
param_names_mapping: dict = {}
reverse_param_names_mapping: dict = {}
lora_param_names_mapping: dict = {}
def __init__(self, config: LingBotWorld2CausalFastVideoConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
self.model_type = config.model_type
self.patch_size = tuple(config.patch_size)
self.text_len = config.text_len
self.in_dim = config.in_dim
self.dim = config.dim
self.hidden_size = config.dim
self.ffn_dim = config.ffn_dim
self.freq_dim = config.freq_dim
self.text_dim = config.text_dim
self.out_dim = config.out_dim
self.out_channels = config.out_dim
self.num_heads = config.num_heads
self.num_attention_heads = config.num_heads
self.attention_head_dim = config.dim // config.num_heads
self.num_layers = config.num_layers
self.local_attn_size = config.local_attn_size
self.sink_size = config.sink_size
self.qk_norm = config.qk_norm
self.cross_attn_norm = config.cross_attn_norm
self.eps = config.eps
self.num_channels_latents = config.out_dim
control_dim = 6
self.patch_embedding = nn.Conv3d(self.in_dim, self.dim, kernel_size=self.patch_size, stride=self.patch_size)
self.patch_embedding_wancamctrl = nn.Linear(
control_dim * 64 * self.patch_size[0] * self.patch_size[1] * self.patch_size[2],
self.dim,
)
self.c2ws_hidden_states_layer1 = nn.Linear(self.dim, self.dim)
self.c2ws_hidden_states_layer2 = nn.Linear(self.dim, self.dim)
self.text_embedding = nn.Sequential(
nn.Linear(self.text_dim, self.dim),
nn.GELU(approximate="tanh"),
nn.Linear(self.dim, self.dim),
)
self.time_embedding = nn.Sequential(
nn.Linear(self.freq_dim, self.dim),
nn.SiLU(),
nn.Linear(self.dim, self.dim),
)
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(self.dim, self.dim * 6))
self.blocks = nn.ModuleList(
[
CausalWanAttentionBlock(
self.dim,
self.ffn_dim,
self.num_heads,
self.local_attn_size,
self.sink_size,
self.qk_norm,
self.cross_attn_norm,
self.eps,
)
for _ in range(self.num_layers)
]
)
self.head = CausalHead(self.dim, self.out_dim, self.patch_size, self.eps)
self.freqs: torch.Tensor | None = None
self.init_weights()
self.__post_init__()
def _get_freqs(self, device: torch.device) -> torch.Tensor:
"""Materialize the non-persistent RoPE frequency table outside meta init."""
if self.freqs is None or self.freqs.is_meta or self.freqs.device != device:
d = self.dim // self.num_heads
self.freqs = torch.cat(
[
rope_params(1024, d - 4 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
],
dim=1,
).to(device)
return self.freqs
def forward(
self,
hidden_states: torch.Tensor | list[torch.Tensor] | None = None,
encoder_hidden_states: torch.Tensor | list[torch.Tensor] | None = None,
timestep: torch.Tensor | None = None,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance=None,
*,
x: list[torch.Tensor] | None = None,
t: torch.Tensor | None = None,
context: list[torch.Tensor] | torch.Tensor | None = None,
seq_len: int | None = None,
y: list[torch.Tensor] | None = None,
dit_cond_dict: dict[str, Any] | None = None,
kv_cache: list[dict] | None = None,
crossattn_cache: list[dict] | None = None,
current_start: int = 0,
max_attention_size: int = 1_000_000,
frame_seqlen: int | None = None,
cross_attn_first_call: bool | None = None,
**kwargs,
) -> list[torch.Tensor]:
"""Run one cached causal-fast DiT forward using the released LingBot World 2 ABI."""
del encoder_hidden_states_image, guidance, kwargs
if x is None:
assert isinstance(hidden_states, torch.Tensor)
x = [hidden_states[0]]
if t is None:
assert timestep is not None
t = timestep
if context is None:
context = encoder_hidden_states
if isinstance(context, torch.Tensor):
context = [u for u in context]
assert context is not None
assert seq_len is not None
assert kv_cache is not None
assert crossattn_cache is not None
if self.model_type == "i2v":
assert y is not None
device = self.patch_embedding.weight.device
freqs = self._get_freqs(device)
if y is not None:
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y, strict=True)]
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
grid_sizes = torch.stack([torch.tensor(u.shape[2:], dtype=torch.long, device=u.device) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long, device=device)
assert seq_lens.max() <= seq_len
x = torch.cat(x)
seq_lens_int = int(seq_lens[0].item())
sp_size = get_sp_world_size()
sp_rank = get_sp_parallel_rank()
padded_seq_len = ((seq_lens_int + sp_size - 1) // sp_size) * sp_size
sp_pad_len = padded_seq_len - seq_lens_int
if sp_pad_len > 0:
x = torch.cat([x, x.new_zeros(x.size(0), sp_pad_len, x.size(2))], dim=1)
if t.dim() == 1:
t = t.expand(t.size(0), padded_seq_len)
with torch.amp.autocast("cuda", dtype=torch.float32):
bt = t.size(0)
t = t.flatten()
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t).unflatten(0, (bt, padded_seq_len)).float()
)
e0 = self.time_projection(e).unflatten(2, (6, self.dim))
context_lens = None
context = self.text_embedding(
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context])
)
if dit_cond_dict is not None and "c2ws_plucker_emb" in dit_cond_dict:
c2ws_plucker_emb = dit_cond_dict["c2ws_plucker_emb"]
c2ws_plucker_emb = [
rearrange(
i,
"1 c (f c1) (h c2) (w c3) -> 1 (f h w) (c c1 c2 c3)",
c1=self.patch_size[0],
c2=self.patch_size[1],
c3=self.patch_size[2],
)
for i in c2ws_plucker_emb
]
c2ws_plucker_emb = torch.cat(c2ws_plucker_emb, dim=1)
c2ws_plucker_emb = self.patch_embedding_wancamctrl(c2ws_plucker_emb)
c2ws_hidden_states = self.c2ws_hidden_states_layer2(
F.silu(self.c2ws_hidden_states_layer1(c2ws_plucker_emb))
)
c2ws_plucker_emb = c2ws_plucker_emb + c2ws_hidden_states
cam_len = c2ws_plucker_emb.size(1)
if cam_len < padded_seq_len:
c2ws_plucker_emb = torch.cat(
[
c2ws_plucker_emb,
c2ws_plucker_emb.new_zeros(
c2ws_plucker_emb.size(0),
padded_seq_len - cam_len,
c2ws_plucker_emb.size(2),
),
],
dim=1,
)
elif cam_len > padded_seq_len:
c2ws_plucker_emb = c2ws_plucker_emb[:, :padded_seq_len, :]
if sp_size > 1:
c2ws_plucker_emb = torch.chunk(c2ws_plucker_emb, sp_size, dim=1)[sp_rank]
dit_cond_dict = dict(dit_cond_dict)
dit_cond_dict["c2ws_plucker_emb"] = c2ws_plucker_emb
if sp_size > 1:
x = torch.chunk(x, sp_size, dim=1)[sp_rank]
e = torch.chunk(e, sp_size, dim=1)[sp_rank]
e0 = torch.chunk(e0, sp_size, dim=1)[sp_rank]
for block_index, block in enumerate(self.blocks):
x = block(
x,
e=e0,
seq_lens=seq_lens,
grid_sizes=grid_sizes,
freqs=freqs,
context=context,
context_lens=context_lens,
dit_cond_dict=dit_cond_dict,
kv_cache=kv_cache[block_index],
crossattn_cache=crossattn_cache[block_index],
current_start=current_start,
max_attention_size=max_attention_size,
frame_seqlen=frame_seqlen,
cross_attn_first_call=cross_attn_first_call,
seq_lens_int=seq_lens_int,
)
x = self.head(x, e)
if sp_size > 1:
x = sequence_model_parallel_all_gather(x, dim=1)
return [u.float() for u in self.unpatchify(x, grid_sizes)]
def unpatchify(self, x: torch.Tensor, grid_sizes: torch.Tensor) -> list[torch.Tensor]:
"""Reconstruct latent videos from flattened patch tokens."""
c = self.out_dim
out = []
for u, v in zip(x, grid_sizes.tolist(), strict=True):
u = u[: math.prod(v)].view(*v, *self.patch_size, c)
u = torch.einsum("fhwpqrc->cfphqwr", u)
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size, strict=True)])
out.append(u)
return out
def init_weights(self) -> None:
"""Initialize modules for non-meta construction; checkpoint load overwrites them."""
if self.patch_embedding.weight.is_meta:
return
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
for m in self.text_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.02)
for m in self.time_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=0.02)
nn.init.zeros_(self.head.head.weight)
EntryClass = LingBotWorld2CausalFastTransformer3DModel
+570
View File
@@ -0,0 +1,570 @@
# SPDX-License-Identifier: Apache-2.0
"""FastVideo-native Z-Image transformer.
Z-Image attends over padded variable-length image/text streams and requires a
key-padding mask. FastVideo's distributed attention wrappers do not expose that
mask contract yet, so this implementation uses torch SDPA and is SP=1 only.
"""
from __future__ import annotations
import math
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.nn.utils.rnn import pad_sequence
from fastvideo.configs.models.dits.zimage import ZImageDiTConfig
from fastvideo.distributed.parallel_state import get_sp_world_size, model_parallel_is_initialized
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
def _linear(layer: ReplicatedLinear, x: torch.Tensor) -> torch.Tensor:
return layer(x)[0]
def _prepare_attention_mask(attention_mask: torch.Tensor | None, dtype: torch.dtype) -> torch.Tensor | None:
if attention_mask is None:
return None
if attention_mask.ndim == 2:
attention_mask = attention_mask[:, None, None, :]
if attention_mask.dtype == torch.bool:
additive_mask = torch.zeros_like(attention_mask, dtype=dtype)
additive_mask.masked_fill_(~attention_mask, float("-inf"))
return additive_mask
return attention_mask
class TimestepEmbedder(nn.Module):
def __init__(
self,
out_size: int,
mid_size: int,
frequency_embedding_size: int,
max_period: int,
) -> None:
super().__init__()
self.mlp = nn.ModuleList([
ReplicatedLinear(frequency_embedding_size, mid_size, bias=True),
nn.SiLU(),
ReplicatedLinear(mid_size, out_size, bias=True),
])
self.frequency_embedding_size = frequency_embedding_size
self.max_period = max_period
@staticmethod
def timestep_embedding(t: torch.Tensor, dim: int, max_period: int) -> torch.Tensor:
with torch.amp.autocast("cuda", enabled=False):
half = dim // 2
freqs = torch.exp(
-math.log(max_period) * torch.arange(half, dtype=torch.float32, device=t.device) / half)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
return embedding
def forward(self, t: torch.Tensor) -> torch.Tensor:
t_freq = self.timestep_embedding(t, self.frequency_embedding_size, self.max_period)
weight_dtype = self.mlp[0].weight.dtype
if weight_dtype.is_floating_point:
t_freq = t_freq.to(weight_dtype)
return _linear(self.mlp[2], self.mlp[1](_linear(self.mlp[0], t_freq)))
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-5) -> None:
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
output = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return output * self.weight
class FeedForward(nn.Module):
def __init__(self, dim: int, hidden_dim: int) -> None:
super().__init__()
self.w1 = ReplicatedLinear(dim, hidden_dim, bias=False)
self.w2 = ReplicatedLinear(hidden_dim, dim, bias=False)
self.w3 = ReplicatedLinear(dim, hidden_dim, bias=False)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return _linear(self.w2, F.silu(_linear(self.w1, x)) * _linear(self.w3, x))
def apply_rotary_emb(x_in: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
with torch.amp.autocast("cuda", enabled=False):
x = torch.view_as_complex(x_in.float().reshape(*x_in.shape[:-1], -1, 2))
x_out = torch.view_as_real(x * freqs_cis.unsqueeze(2)).flatten(3)
return x_out.type_as(x_in)
class ZImageAttention(nn.Module):
def __init__(self, dim: int, n_heads: int, n_kv_heads: int, qk_norm: bool = True, eps: float = 1e-5) -> None:
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.head_dim = dim // n_heads
self.to_q = ReplicatedLinear(dim, n_heads * self.head_dim, bias=False)
self.to_k = ReplicatedLinear(dim, n_kv_heads * self.head_dim, bias=False)
self.to_v = ReplicatedLinear(dim, n_kv_heads * self.head_dim, bias=False)
self.to_out = nn.ModuleList([ReplicatedLinear(n_heads * self.head_dim, dim, bias=False)])
self.norm_q = RMSNorm(self.head_dim, eps=eps) if qk_norm else None
self.norm_k = RMSNorm(self.head_dim, eps=eps) if qk_norm else None
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
freqs_cis: torch.Tensor | None = None,
) -> torch.Tensor:
query = _linear(self.to_q, hidden_states).unflatten(-1, (self.n_heads, -1))
key = _linear(self.to_k, hidden_states).unflatten(-1, (self.n_kv_heads, -1))
value = _linear(self.to_v, hidden_states).unflatten(-1, (self.n_kv_heads, -1))
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
if freqs_cis is not None:
query = apply_rotary_emb(query, freqs_cis)
key = apply_rotary_emb(key, freqs_cis)
mask = _prepare_attention_mask(attention_mask, query.dtype)
hidden_states = F.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
attn_mask=mask,
dropout_p=0.0,
is_causal=False,
).transpose(1, 2).contiguous()
return _linear(self.to_out[0], hidden_states.flatten(2, 3).to(query.dtype))
class ZImageTransformerBlock(nn.Module):
def __init__(
self,
layer_id: int,
dim: int,
n_heads: int,
n_kv_heads: int,
norm_eps: float,
qk_norm: bool,
adaln_embed_dim: int,
modulation: bool = True,
) -> None:
super().__init__()
self.dim = dim
self.head_dim = dim // n_heads
self.layer_id = layer_id
self.modulation = modulation
self.attention = ZImageAttention(dim, n_heads, n_kv_heads, qk_norm, norm_eps)
self.feed_forward = FeedForward(dim=dim, hidden_dim=int(dim / 3 * 8))
self.attention_norm1 = RMSNorm(dim, eps=norm_eps)
self.ffn_norm1 = RMSNorm(dim, eps=norm_eps)
self.attention_norm2 = RMSNorm(dim, eps=norm_eps)
self.ffn_norm2 = RMSNorm(dim, eps=norm_eps)
if modulation:
self.adaLN_modulation = nn.ModuleList(
[ReplicatedLinear(min(dim, adaln_embed_dim), 4 * dim, bias=True)])
def forward(
self,
x: torch.Tensor,
attn_mask: torch.Tensor,
freqs_cis: torch.Tensor,
adaln_input: torch.Tensor | None = None,
) -> torch.Tensor:
if self.modulation:
assert adaln_input is not None
scale_msa, gate_msa, scale_mlp, gate_mlp = _linear(
self.adaLN_modulation[0], adaln_input).unsqueeze(1).chunk(4, dim=2)
gate_msa, gate_mlp = gate_msa.tanh(), gate_mlp.tanh()
scale_msa, scale_mlp = 1.0 + scale_msa, 1.0 + scale_mlp
attn_out = self.attention(
self.attention_norm1(x) * scale_msa,
attention_mask=attn_mask,
freqs_cis=freqs_cis,
)
x = x + gate_msa * self.attention_norm2(attn_out)
x = x + gate_mlp * self.ffn_norm2(self.feed_forward(self.ffn_norm1(x) * scale_mlp))
else:
attn_out = self.attention(
self.attention_norm1(x),
attention_mask=attn_mask,
freqs_cis=freqs_cis,
)
x = x + self.attention_norm2(attn_out)
x = x + self.ffn_norm2(self.feed_forward(self.ffn_norm1(x)))
return x
class FinalLayer(nn.Module):
def __init__(self, hidden_size: int, out_channels: int, adaln_embed_dim: int) -> None:
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
self.linear = ReplicatedLinear(hidden_size, out_channels, bias=True)
self.adaLN_modulation = nn.ModuleList([
nn.SiLU(),
ReplicatedLinear(min(hidden_size, adaln_embed_dim), hidden_size, bias=True),
])
def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
scale = 1.0 + _linear(self.adaLN_modulation[1], self.adaLN_modulation[0](c))
return _linear(self.linear, self.norm_final(x) * scale.unsqueeze(1))
class RopeEmbedder:
def __init__(self, theta: float, axes_dims: tuple[int, ...], axes_lens: tuple[int, ...]) -> None:
if len(axes_dims) != len(axes_lens):
raise ValueError("RoPE axes require matching dimensions and lengths")
self.theta = theta
self.axes_dims = axes_dims
self.axes_lens = axes_lens
self.freqs_cis: list[torch.Tensor] | None = None
@staticmethod
def precompute_freqs_cis(dim: tuple[int, ...], end: tuple[int, ...], theta: float) -> list[torch.Tensor]:
with torch.device("cpu"):
freqs_cis = []
for axis_dim, axis_end in zip(dim, end):
freqs = 1.0 / (theta**(torch.arange(0, axis_dim, 2, dtype=torch.float64) / axis_dim))
timestep = torch.arange(axis_end, dtype=torch.float64)
angles = torch.outer(timestep, freqs).float()
freqs_cis.append(torch.polar(torch.ones_like(angles), angles).to(torch.complex64))
return freqs_cis
def __call__(self, ids: torch.Tensor) -> torch.Tensor:
if ids.ndim != 2 or ids.shape[-1] != len(self.axes_dims):
raise ValueError("RoPE ids must have shape [sequence, number_of_axes]")
if self.freqs_cis is None:
self.freqs_cis = [
freqs.to(ids.device)
for freqs in self.precompute_freqs_cis(self.axes_dims, self.axes_lens, self.theta)
]
elif self.freqs_cis[0].device != ids.device:
self.freqs_cis = [freqs.to(ids.device) for freqs in self.freqs_cis]
return torch.cat([self.freqs_cis[i][ids[:, i]] for i in range(len(self.axes_dims))], dim=-1)
class ZImageTransformer2DModel(BaseDiT):
_default_config = ZImageDiTConfig()
_fsdp_shard_conditions = _default_config.arch_config._fsdp_shard_conditions
_compile_conditions = _default_config.arch_config._compile_conditions
_supported_attention_backends = (AttentionBackendEnum.TORCH_SDPA, )
param_names_mapping = _default_config.arch_config.param_names_mapping
reverse_param_names_mapping = _default_config.arch_config.reverse_param_names_mapping
def __init__(self, config: ZImageDiTConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
arch = config.arch_config
self.in_channels = arch.in_channels
self.out_channels = arch.in_channels
self.all_patch_size = tuple(arch.all_patch_size)
self.all_f_patch_size = tuple(arch.all_f_patch_size)
self.dim = arch.dim
self.n_heads = arch.n_heads
self.rope_theta = arch.rope_theta
self.t_scale = arch.t_scale
self.seq_multi_of = arch.seq_multi_of
self.hidden_size = arch.dim
self.num_attention_heads = arch.n_heads
self.num_channels_latents = arch.in_channels
self.all_x_embedder = nn.ModuleDict({
f"{patch_size}-{f_patch_size}": ReplicatedLinear(
f_patch_size * patch_size * patch_size * arch.in_channels, arch.dim, bias=True)
for patch_size, f_patch_size in zip(self.all_patch_size, self.all_f_patch_size)
})
self.all_final_layer = nn.ModuleDict({
f"{patch_size}-{f_patch_size}": FinalLayer(
arch.dim,
patch_size * patch_size * f_patch_size * self.out_channels,
arch.adaln_embed_dim,
)
for patch_size, f_patch_size in zip(self.all_patch_size, self.all_f_patch_size)
})
block_kwargs = {
"dim": arch.dim,
"n_heads": arch.n_heads,
"n_kv_heads": arch.n_kv_heads,
"norm_eps": arch.norm_eps,
"qk_norm": arch.qk_norm,
"adaln_embed_dim": arch.adaln_embed_dim,
}
self.noise_refiner = nn.ModuleList([
ZImageTransformerBlock(1000 + layer_id, modulation=True, **block_kwargs)
for layer_id in range(arch.n_refiner_layers)
])
self.context_refiner = nn.ModuleList([
ZImageTransformerBlock(layer_id, modulation=False, **block_kwargs)
for layer_id in range(arch.n_refiner_layers)
])
self.t_embedder = TimestepEmbedder(
min(arch.dim, arch.adaln_embed_dim),
mid_size=arch.timestep_mid_size,
frequency_embedding_size=arch.frequency_embedding_size,
max_period=arch.max_period,
)
self.cap_embedder = nn.ModuleList([
RMSNorm(arch.cap_feat_dim, eps=arch.norm_eps),
ReplicatedLinear(arch.cap_feat_dim, arch.dim, bias=True),
])
self.x_pad_token = nn.Parameter(torch.empty((1, arch.dim)))
self.cap_pad_token = nn.Parameter(torch.empty((1, arch.dim)))
self.layers = nn.ModuleList([
ZImageTransformerBlock(layer_id, modulation=True, **block_kwargs) for layer_id in range(arch.n_layers)
])
self.axes_dims = tuple(arch.axes_dims)
self.axes_lens = tuple(arch.axes_lens)
self.rope_embedder = RopeEmbedder(arch.rope_theta, self.axes_dims, self.axes_lens)
self.__post_init__()
def unpatchify(
self,
x: list[torch.Tensor],
size: list[tuple[int, int, int]],
patch_size: int,
f_patch_size: int,
) -> list[torch.Tensor]:
patch_height = patch_width = patch_size
patch_frames = f_patch_size
if len(x) != len(size):
raise ValueError("output batch and original sizes must have equal length")
for i, (frames, height, width) in enumerate(size):
original_length = (frames // patch_frames) * (height // patch_height) * (width // patch_width)
x[i] = (x[i][:original_length].view(
frames // patch_frames,
height // patch_height,
width // patch_width,
patch_frames,
patch_height,
patch_width,
self.out_channels,
).permute(6, 0, 3, 1, 4, 2, 5).reshape(self.out_channels, frames, height, width))
return x
@staticmethod
def create_coordinate_grid(
size: tuple[int, int, int],
start: tuple[int, int, int] | None = None,
device: torch.device | None = None,
) -> torch.Tensor:
start = start or (0, ) * len(size)
axes = [
torch.arange(axis_start, axis_start + span, dtype=torch.int32, device=device)
for axis_start, span in zip(start, size)
]
return torch.stack(torch.meshgrid(axes, indexing="ij"), dim=-1)
def patchify_and_embed(
self,
all_image: list[torch.Tensor],
all_cap_feats: list[torch.Tensor],
patch_size: int,
f_patch_size: int,
) -> tuple[
list[torch.Tensor],
list[torch.Tensor],
list[tuple[int, int, int]],
list[torch.Tensor],
list[torch.Tensor],
list[torch.Tensor],
list[torch.Tensor],
]:
patch_height = patch_width = patch_size
patch_frames = f_patch_size
device = all_image[0].device
image_out = []
image_sizes = []
image_pos_ids = []
image_pad_masks = []
cap_pos_ids = []
cap_pad_masks = []
cap_feats_out = []
for image, cap_feat in zip(all_image, all_cap_feats):
cap_length = len(cap_feat)
cap_padding = (-cap_length) % self.seq_multi_of
cap_pos_ids.append(
self.create_coordinate_grid(
(cap_length + cap_padding, 1, 1),
start=(1, 0, 0),
device=device,
).flatten(0, 2))
cap_pad_masks.append(
torch.cat([
torch.zeros(cap_length, dtype=torch.bool, device=device),
torch.ones(cap_padding, dtype=torch.bool, device=device),
]) if cap_padding else torch.zeros(cap_length, dtype=torch.bool, device=device))
cap_feats_out.append(
torch.cat([cap_feat, cap_feat[-1:].repeat(cap_padding, 1)]) if cap_padding else cap_feat)
channels, frames, height, width = image.size()
image_sizes.append((frames, height, width))
frame_tokens, height_tokens, width_tokens = (
frames // patch_frames,
height // patch_height,
width // patch_width,
)
image = image.view(
channels,
frame_tokens,
patch_frames,
height_tokens,
patch_height,
width_tokens,
patch_width,
)
image = image.permute(1, 3, 5, 2, 4, 6, 0).reshape(
frame_tokens * height_tokens * width_tokens,
patch_frames * patch_height * patch_width * channels,
)
image_length = len(image)
image_padding = (-image_length) % self.seq_multi_of
original_pos_ids = self.create_coordinate_grid(
(frame_tokens, height_tokens, width_tokens),
start=(cap_length + cap_padding + 1, 0, 0),
device=device,
).flatten(0, 2)
if image_padding:
padding_pos_ids = self.create_coordinate_grid((1, 1, 1), device=device).flatten(0, 2).repeat(
image_padding, 1)
image_pos_ids.append(torch.cat([original_pos_ids, padding_pos_ids]))
else:
image_pos_ids.append(original_pos_ids)
image_pad_masks.append(
torch.cat([
torch.zeros(image_length, dtype=torch.bool, device=device),
torch.ones(image_padding, dtype=torch.bool, device=device),
]) if image_padding else torch.zeros(image_length, dtype=torch.bool, device=device))
image_out.append(
torch.cat([image, image[-1:].repeat(image_padding, 1)]) if image_padding else image)
return (
image_out,
cap_feats_out,
image_sizes,
image_pos_ids,
cap_pos_ids,
image_pad_masks,
cap_pad_masks,
)
@staticmethod
def _attention_mask(lengths: list[int], device: torch.device) -> torch.Tensor:
mask = torch.zeros((len(lengths), max(lengths)), dtype=torch.bool, device=device)
for i, length in enumerate(lengths):
mask[i, :length] = True
return mask
def forward(
self,
hidden_states: torch.Tensor | list[torch.Tensor],
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor] | None = None,
guidance=None,
patch_size: int = 2,
f_patch_size: int = 1,
**kwargs,
) -> tuple[list[torch.Tensor], dict]:
del encoder_hidden_states_image, guidance, kwargs
if model_parallel_is_initialized() and get_sp_world_size() != 1:
raise NotImplementedError(
"Z-Image masked SDPA does not support sequence parallelism; run with sp_size=1")
if patch_size not in self.all_patch_size or f_patch_size not in self.all_f_patch_size:
raise ValueError(f"unsupported patch sizes: spatial={patch_size}, temporal={f_patch_size}")
if isinstance(hidden_states, torch.Tensor):
hidden_states = list(hidden_states.unbind(0))
if isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = list(encoder_hidden_states.unbind(0))
device = hidden_states[0].device
timestep_embedding = self.t_embedder(timestep * self.t_scale)
(
hidden_states,
encoder_hidden_states,
image_sizes,
image_pos_ids,
cap_pos_ids,
image_inner_pad_masks,
cap_inner_pad_masks,
) = self.patchify_and_embed(hidden_states, encoder_hidden_states, patch_size, f_patch_size)
image_lengths = [len(item) for item in hidden_states]
if not all(length % self.seq_multi_of == 0 for length in image_lengths):
raise ValueError("padded image sequence lengths must be aligned")
hidden_states = torch.cat(hidden_states)
hidden_states = _linear(self.all_x_embedder[f"{patch_size}-{f_patch_size}"], hidden_states)
adaln_input = timestep_embedding.type_as(hidden_states)
hidden_states[torch.cat(image_inner_pad_masks)] = self.x_pad_token
hidden_states = list(hidden_states.split(image_lengths))
image_freqs_cis = list(
self.rope_embedder(torch.cat(image_pos_ids)).split([len(item) for item in image_pos_ids]))
hidden_states = pad_sequence(hidden_states, batch_first=True, padding_value=0.0)
image_freqs_cis = pad_sequence(image_freqs_cis, batch_first=True, padding_value=0.0)
image_freqs_cis = image_freqs_cis[:, :hidden_states.shape[1]]
image_attn_mask = self._attention_mask(image_lengths, device)
for layer in self.noise_refiner:
hidden_states = layer(hidden_states, image_attn_mask, image_freqs_cis, adaln_input)
cap_lengths = [len(item) for item in encoder_hidden_states]
if not all(length % self.seq_multi_of == 0 for length in cap_lengths):
raise ValueError("padded caption sequence lengths must be aligned")
encoder_hidden_states = torch.cat(encoder_hidden_states)
encoder_hidden_states = _linear(self.cap_embedder[1], self.cap_embedder[0](encoder_hidden_states))
encoder_hidden_states[torch.cat(cap_inner_pad_masks)] = self.cap_pad_token
encoder_hidden_states = list(encoder_hidden_states.split(cap_lengths))
cap_freqs_cis = list(self.rope_embedder(torch.cat(cap_pos_ids)).split([len(item) for item in cap_pos_ids]))
encoder_hidden_states = pad_sequence(encoder_hidden_states, batch_first=True, padding_value=0.0)
cap_freqs_cis = pad_sequence(cap_freqs_cis, batch_first=True, padding_value=0.0)
cap_freqs_cis = cap_freqs_cis[:, :encoder_hidden_states.shape[1]]
cap_attn_mask = self._attention_mask(cap_lengths, device)
for layer in self.context_refiner:
encoder_hidden_states = layer(encoder_hidden_states, cap_attn_mask, cap_freqs_cis)
unified = []
unified_freqs_cis = []
for i, (image_length, cap_length) in enumerate(zip(image_lengths, cap_lengths)):
unified.append(
torch.cat([hidden_states[i][:image_length], encoder_hidden_states[i][:cap_length]]))
unified_freqs_cis.append(
torch.cat([image_freqs_cis[i][:image_length], cap_freqs_cis[i][:cap_length]]))
unified_lengths = [image_length + cap_length for image_length, cap_length in zip(image_lengths, cap_lengths)]
unified = pad_sequence(unified, batch_first=True, padding_value=0.0)
unified_freqs_cis = pad_sequence(unified_freqs_cis, batch_first=True, padding_value=0.0)
unified_attn_mask = self._attention_mask(unified_lengths, device)
for layer in self.layers:
unified = layer(unified, unified_attn_mask, unified_freqs_cis, adaln_input)
unified = self.all_final_layer[f"{patch_size}-{f_patch_size}"](unified, adaln_input)
outputs = self.unpatchify(list(unified.unbind(0)), image_sizes, patch_size, f_patch_size)
return outputs, {}
EntryClass = ZImageTransformer2DModel
+221
View File
@@ -0,0 +1,221 @@
# SPDX-License-Identifier: Apache-2.0
"""Native LingBot-Video Qwen3-VL language model for text-only conditioning."""
from collections.abc import Iterable
from typing import Any
import torch
from torch import nn
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.encoders.qwen3 import (
Qwen3Attention,
Qwen3DecoderLayer,
Qwen3ForCausalLM,
Qwen3MLP,
)
class LingBotVideoQwen3VLAttention(Qwen3Attention):
"""Qwen3-VL attention with the official masked repeat-K/V SDPA path."""
def _apply_qwen3_vl_rope(
self,
positions: torch.Tensor,
query: torch.Tensor,
key: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Apply NeoX RoPE with Qwen3-VL's input-dtype multiply-add ordering."""
flat_positions = positions.flatten()
cos_sin = self.rotary_emb.cos_sin_cache.index_select(0, flat_positions)
cos_half, sin_half = cos_sin.chunk(2, dim=-1)
cos = torch.cat((cos_half, cos_half), dim=-1).to(query.dtype)
sin = torch.cat((sin_half, sin_half), dim=-1).to(query.dtype)
if flat_positions.numel() == query.shape[1]:
cos = cos.view(1, query.shape[1], 1, self.head_dim)
sin = sin.view(1, query.shape[1], 1, self.head_dim)
else:
cos = cos.view(*query.shape[:2], 1, self.head_dim)
sin = sin.view(*query.shape[:2], 1, self.head_dim)
def rotate(tensor: torch.Tensor) -> torch.Tensor:
first, second = tensor.chunk(2, dim=-1)
rotated = torch.cat((-second, first), dim=-1)
return tensor * cos + rotated * sin
return rotate(query), rotate(key)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""Apply fused projections, QK norm, RoPE, and causal grouped attention."""
qkv, _ = self.qkv_proj(hidden_states)
query, key, value = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
batch_size, sequence_length = query.shape[:2]
query = query.reshape(batch_size, sequence_length, self.num_heads, self.head_dim)
key = key.reshape(batch_size, sequence_length, self.num_kv_heads, self.head_dim)
value = value.reshape(batch_size, sequence_length, self.num_kv_heads, self.head_dim)
query = self.q_norm(query)
key = self.k_norm(key)
query, key = self._apply_qwen3_vl_rope(positions, query, key)
no_padding = attention_mask is None
if no_padding:
attention_output = torch.nn.functional.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
dropout_p=0.0,
is_causal=sequence_length > 1,
scale=self.scaling,
enable_gqa=self.num_heads != self.num_kv_heads,
).transpose(1, 2)
else:
groups = self.num_heads // self.num_kv_heads
key = (key[:, :, :, None, :].expand(-1, -1, -1, groups, -1).reshape(batch_size, sequence_length,
self.num_heads, self.head_dim))
value = (value[:, :, :, None, :].expand(-1, -1, -1, groups, -1).reshape(batch_size, sequence_length,
self.num_heads, self.head_dim))
causal_mask = torch.ones(sequence_length, sequence_length, device=query.device, dtype=torch.bool).tril()
key_mask = attention_mask.to(device=query.device, dtype=torch.bool)
sdpa_mask = causal_mask[None, None, :, :] & key_mask[:, None, None, :]
attention_output = torch.nn.functional.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
attn_mask=sdpa_mask,
dropout_p=0.0,
is_causal=False,
scale=self.scaling,
).transpose(1, 2)
output, _ = self.o_proj(attention_output.reshape(batch_size, sequence_length, -1))
return output
class LingBotVideoQwen3VLDecoderLayer(Qwen3DecoderLayer):
"""Qwen3-VL decoder layer with explicit official residual rounding order."""
def __init__(self, config: Any, prefix: str) -> None:
"""Build the final Qwen3-VL attention once to avoid orphan parameters."""
nn.Module.__init__(self)
self.hidden_size = config.hidden_size
quant_config = getattr(config, "quant_config", None)
self.self_attn = LingBotVideoQwen3VLAttention(
config=config,
hidden_size=self.hidden_size,
num_heads=config.num_attention_heads,
num_kv_heads=config.num_key_value_heads,
rope_theta=config.rope_theta,
rope_scaling=config.rope_scaling,
max_position_embeddings=config.max_position_embeddings,
quant_config=quant_config,
bias=config.attention_bias,
prefix=f"{prefix}.self_attn",
)
self.mlp = Qwen3MLP(
hidden_size=self.hidden_size,
intermediate_size=config.intermediate_size,
hidden_act=config.hidden_act,
quant_config=quant_config,
bias=getattr(config, "mlp_bias", False),
prefix=f"{prefix}.mlp",
)
self.input_layernorm = RMSNorm(self.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = RMSNorm(self.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
"""Run attention and MLP with each residual sum rounded before normalization."""
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.self_attn(positions, hidden_states, attention_mask)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
return residual + hidden_states
class LingBotVideoQwen3VLTextModel(Qwen3ForCausalLM):
"""Load the Qwen3-VL language-model subset without its vision tower or LM head."""
supports_hf_from_pretrained = False
def __init__(self, config) -> None:
"""Construct the exact Qwen3-VL module graph without replacing base layers."""
TextEncoder.__init__(self, config)
self.quant_config = getattr(config, "quant_config", None)
if getattr(config, "lora_config", None) is not None:
max_loras = getattr(config.lora_config, "max_loras", 1)
lora_vocab_size = getattr(config.lora_config, "lora_extra_vocab_size", 1)
lora_vocab = lora_vocab_size * max_loras
else:
lora_vocab = 0
self.vocab_size = config.vocab_size + lora_vocab
self.org_vocab_size = config.vocab_size
self.embed_tokens = VocabParallelEmbedding(
self.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size,
quant_config=self.quant_config,
)
self.layers = nn.ModuleList(
LingBotVideoQwen3VLDecoderLayer(config, prefix=f"{config.prefix}.layers.{index}")
for index in range(config.num_hidden_layers))
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
input_ids: torch.Tensor | None = None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs: Any,
) -> BaseEncoderOutput:
"""Run explicit Qwen3-VL layers and return the requested hidden-state tuple."""
del kwargs
output_hidden_states = (output_hidden_states
if output_hidden_states is not None else self.config.output_hidden_states)
if inputs_embeds is None:
if input_ids is None:
raise ValueError("input_ids or inputs_embeds is required")
hidden_states = self.get_input_embeddings(input_ids)
else:
hidden_states = inputs_embeds
if position_ids is None:
position_ids = torch.arange(hidden_states.shape[1], device=hidden_states.device).unsqueeze(0)
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
attention_mask = None
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
for layer in self.layers:
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
hidden_states = layer(position_ids, hidden_states, attention_mask)
hidden_states = self.norm(hidden_states)
if all_hidden_states is not None:
all_hidden_states += (hidden_states, )
return BaseEncoderOutput(
last_hidden_state=hidden_states,
hidden_states=all_hidden_states,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
"""Accept either official compound keys or converted native keys."""
prefix = "model.language_model."
language_weights = ((name[len(prefix):] if name.startswith(prefix) else name, tensor)
for name, tensor in weights
if name.startswith(prefix) or not name.startswith(("model.", "lm_head.")))
return super().load_weights(language_weights)
EntryClass = LingBotVideoQwen3VLTextModel
@@ -0,0 +1,269 @@
# SPDX-License-Identifier: Apache-2.0
"""LingBot World 2 UMT5 encoder with the released checkpoint's module names."""
from collections.abc import Iterable
import html
import math
import string
import ftfy
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.encoders import BaseEncoderOutput
from fastvideo.configs.models.encoders.lingbotworld2_t5 import LingBotWorld2UMT5Config
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
def basic_clean(text: str) -> str:
"""Apply LingBot World 2's ftfy/html cleanup before tokenization."""
text = ftfy.fix_text(text)
text = html.unescape(html.unescape(text))
return text.strip()
def whitespace_clean(text: str) -> str:
"""Collapse all whitespace runs to single spaces."""
return " ".join(text.split())
def canonicalize(text: str, keep_punctuation_exact_string: str | None = None) -> str:
"""Normalize prompts with LingBot World 2's optional punctuation handling."""
text = text.replace("_", " ")
if keep_punctuation_exact_string:
text = keep_punctuation_exact_string.join(
part.translate(str.maketrans("", "", string.punctuation))
for part in text.split(keep_punctuation_exact_string)
)
else:
text = text.translate(str.maketrans("", "", string.punctuation))
return " ".join(text.lower().split())
def lingbotworld2_whitespace_preprocess(prompt: str) -> str:
"""Match the LingBot World 2 source tokenizer's `clean='whitespace'` behavior."""
return whitespace_clean(basic_clean(prompt))
class GELU(nn.Module):
"""T5 gated GELU approximation used by the source checkpoint."""
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Apply the tanh GELU approximation."""
return 0.5 * x * (1.0 + torch.tanh(math.sqrt(2.0 / math.pi) * (x + 0.044715 * torch.pow(x, 3.0))))
class T5LayerNorm(nn.Module):
"""T5 RMS-style layer norm with source-compatible parameter name."""
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Normalize in fp32 and apply the learned scale."""
x = x * torch.rsqrt(x.float().pow(2).mean(dim=-1, keepdim=True) + self.eps)
if self.weight.dtype in (torch.float16, torch.bfloat16):
x = x.type_as(self.weight)
return self.weight * x
class T5Attention(nn.Module):
"""LingBot World 2 source T5 attention block."""
def __init__(self, dim: int, dim_attn: int, num_heads: int, dropout: float = 0.1):
super().__init__()
assert dim_attn % num_heads == 0
self.dim = dim
self.dim_attn = dim_attn
self.num_heads = num_heads
self.head_dim = dim_attn // num_heads
self.q = nn.Linear(dim, dim_attn, bias=False)
self.k = nn.Linear(dim, dim_attn, bias=False)
self.v = nn.Linear(dim, dim_attn, bias=False)
self.o = nn.Linear(dim_attn, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(
self,
x: torch.Tensor,
context: torch.Tensor | None = None,
mask: torch.Tensor | None = None,
pos_bias: torch.Tensor | None = None,
) -> torch.Tensor:
"""Project QKV, add relative bias/mask, and return attended states."""
context = x if context is None else context
b, n, c = x.size(0), self.num_heads, self.head_dim
q = self.q(x).view(b, -1, n, c)
k = self.k(context).view(b, -1, n, c)
v = self.v(context).view(b, -1, n, c)
attn_bias = x.new_zeros(b, n, q.size(1), k.size(1))
if pos_bias is not None:
attn_bias += pos_bias
if mask is not None:
assert mask.ndim in (2, 3)
mask = mask.view(b, 1, 1, -1) if mask.ndim == 2 else mask.unsqueeze(1)
attn_bias.masked_fill_(mask == 0, torch.finfo(x.dtype).min)
attn = torch.einsum("binc,bjnc->bnij", q, k) + attn_bias
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
x = torch.einsum("bnij,bjnc->binc", attn, v)
return self.dropout(self.o(x.reshape(b, -1, n * c)))
class T5FeedForward(nn.Module):
"""LingBot World 2 source T5 gated feed-forward block."""
def __init__(self, dim: int, dim_ffn: int, dropout: float = 0.1):
super().__init__()
self.dim = dim
self.dim_ffn = dim_ffn
self.gate = nn.Sequential(nn.Linear(dim, dim_ffn, bias=False), GELU())
self.fc1 = nn.Linear(dim, dim_ffn, bias=False)
self.fc2 = nn.Linear(dim_ffn, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Apply the gated feed-forward projection."""
x = self.fc1(x) * self.gate(x)
x = self.dropout(x)
x = self.fc2(x)
return self.dropout(x)
class T5RelativeEmbedding(nn.Module):
"""Per-block relative position embedding used by LingBot World 2 UMT5."""
def __init__(self, num_buckets: int, num_heads: int, bidirectional: bool, max_dist: int = 128):
super().__init__()
self.num_buckets = num_buckets
self.num_heads = num_heads
self.bidirectional = bidirectional
self.max_dist = max_dist
self.embedding = nn.Embedding(num_buckets, num_heads)
def forward(self, lq: int, lk: int) -> torch.Tensor:
"""Build a relative-position bias tensor for attention logits."""
device = self.embedding.weight.device
rel_pos = torch.arange(lk, device=device).unsqueeze(0) - torch.arange(lq, device=device).unsqueeze(1)
rel_pos = self._relative_position_bucket(rel_pos)
rel_pos_embeds = self.embedding(rel_pos)
return rel_pos_embeds.permute(2, 0, 1).unsqueeze(0).contiguous()
def _relative_position_bucket(self, rel_pos: torch.Tensor) -> torch.Tensor:
"""Map token offsets to T5 relative-position buckets."""
if self.bidirectional:
num_buckets = self.num_buckets // 2
rel_buckets = (rel_pos > 0).long() * num_buckets
rel_pos = torch.abs(rel_pos)
else:
num_buckets = self.num_buckets
rel_buckets = 0
rel_pos = -torch.min(rel_pos, torch.zeros_like(rel_pos))
max_exact = num_buckets // 2
rel_pos_large = max_exact + (
torch.log(rel_pos.float() / max_exact) / math.log(self.max_dist / max_exact) * (num_buckets - max_exact)
).long()
rel_pos_large = torch.min(rel_pos_large, torch.full_like(rel_pos_large, num_buckets - 1))
rel_buckets += torch.where(rel_pos < max_exact, rel_pos, rel_pos_large)
return rel_buckets
class T5SelfAttention(nn.Module):
"""One source-compatible UMT5 encoder block."""
def __init__(
self,
dim: int,
dim_attn: int,
dim_ffn: int,
num_heads: int,
num_buckets: int,
shared_pos: bool = False,
dropout: float = 0.1,
):
super().__init__()
self.norm1 = T5LayerNorm(dim)
self.attn = T5Attention(dim, dim_attn, num_heads, dropout)
self.norm2 = T5LayerNorm(dim)
self.ffn = T5FeedForward(dim, dim_ffn, dropout)
self.pos_embedding = None if shared_pos else T5RelativeEmbedding(num_buckets, num_heads, bidirectional=True)
def forward(self, x: torch.Tensor, mask: torch.Tensor | None = None, pos_bias: torch.Tensor | None = None) -> torch.Tensor:
"""Run self-attention and feed-forward residual updates."""
e = pos_bias if self.pos_embedding is None else self.pos_embedding(x.size(1), x.size(1))
x = self._fp16_clamp(x + self.attn(self.norm1(x), mask=mask, pos_bias=e))
return self._fp16_clamp(x + self.ffn(self.norm2(x)))
@staticmethod
def _fp16_clamp(x: torch.Tensor) -> torch.Tensor:
if x.dtype == torch.float16 and torch.isinf(x).any():
clamp = torch.finfo(x.dtype).max - 1000
x = torch.clamp(x, min=-clamp, max=clamp)
return x
class LingBotWorld2T5EncoderModel(TextEncoder):
"""FastVideo-native LingBot World 2 UMT5 encoder for the released `.pth` weights."""
fall_back_to_pt_during_load = True
allow_patterns_overrides = ["*.pt"]
def __init__(self, config: LingBotWorld2UMT5Config, prefix: str = ""):
super().__init__(config)
del prefix
arch = config.arch_config
self.token_embedding = nn.Embedding(arch.vocab_size, arch.dim)
self.dropout = nn.Dropout(arch.dropout)
self.blocks = nn.ModuleList(
[
T5SelfAttention(
arch.dim,
arch.dim_attn,
arch.dim_ffn,
arch.num_heads,
arch.num_buckets,
shared_pos=False,
dropout=arch.dropout,
)
for _ in range(arch.num_layers)
]
)
self.norm = T5LayerNorm(arch.dim)
def forward(
self,
input_ids: torch.Tensor | None,
position_ids: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
inputs_embeds: torch.Tensor | None = None,
output_hidden_states: bool | None = None,
**kwargs,
) -> BaseEncoderOutput:
"""Encode token ids and return source-compatible hidden states."""
del position_ids, inputs_embeds, output_hidden_states, kwargs
assert input_ids is not None
x = self.dropout(self.token_embedding(input_ids))
for block in self.blocks:
x = block(x, attention_mask)
x = self.dropout(self.norm(x))
return BaseEncoderOutput(last_hidden_state=x, attention_mask=attention_mask)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
"""Load source `.pth` weights whose names already match this module."""
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
EntryClass = LingBotWorld2T5EncoderModel
+79 -29
View File
@@ -327,17 +327,15 @@ class Qwen3ForCausalLM(TextEncoder):
dtype: torch.dtype,
device: torch.device,
) -> nn.Module:
from transformers import AutoModelForCausalLM
from transformers import AutoModel
if device.type == "cpu" and torch.cuda.is_available():
from fastvideo.distributed import get_local_torch_device
device = get_local_torch_device()
return AutoModelForCausalLM.from_pretrained(
# FastVideo uses Qwen3 only as a text encoder. Loading the body avoids
# materializing an unused LM head and full-vocabulary logits, including
# for checkpoints whose metadata names Qwen3ForCausalLM.
return AutoModel.from_pretrained(
model_path,
local_files_only=True,
torch_dtype=dtype,
dtype=dtype,
low_cpu_mem_usage=True,
).eval().to(device)
@@ -368,9 +366,18 @@ class Qwen3ForCausalLM(TextEncoder):
residual = None
if position_ids is None:
# Expand to [batch_size, seq_len]: the rotary layer flattens
# positions to ``num_tokens`` and reshapes q/k to
# ``(num_tokens, -1, head_dim)``. A bare [1, seq_len] only matches
# ``num_tokens`` when batch_size == 1; for batched inputs it folds
# the batch dim into the head dim and misaligns RoPE. Expanding to
# ``batch_size * seq_len`` tokens keeps the layout correct.
position_ids = torch.arange(
0, hidden_states.shape[1], device=hidden_states.device
).unsqueeze(0)
0,
hidden_states.shape[1],
device=hidden_states.device,
dtype=torch.long,
).unsqueeze(0).expand(hidden_states.shape[0], -1)
all_hidden_states: tuple[Any, ...] | None = (
() if output_hidden_states else None
@@ -405,6 +412,20 @@ class Qwen3ForCausalLM(TextEncoder):
) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
stacked_params_mapping = self.config.arch_config.stacked_params_mapping
# A fused destination is initialized either by one already-fused tensor
# or after every split source projection has been loaded. Include
# auxiliary quantization parameters (for example scale_weight) rather
# than limiting completeness checks to weight/bias tensors.
expected_stacked_shards = {
(name, shard_id)
for name in params_dict
for param_name, _, shard_id in stacked_params_mapping
if param_name in name
}
fused_param_names = {name for name, _ in expected_stacked_shards}
loaded_stacked_shards: set[tuple[str, str | int]] = set()
loaded_fused_params: set[str] = set()
for name, loaded_weight in weights:
if name.startswith("model."):
@@ -423,37 +444,66 @@ class Qwen3ForCausalLM(TextEncoder):
continue
name = kv_scale_name
for (
param_name,
weight_name,
shard_id,
) in self.config.arch_config.stacked_params_mapping:
matched_stacked_param = False
for param_name, weight_name, shard_id in stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
matched_stacked_param = True
target_name = name.replace(weight_name, param_name)
if name.endswith(".bias") and name not in params_dict:
continue
if target_name.endswith(".bias") and target_name not in params_dict:
break
if name not in params_dict:
continue
if target_name not in params_dict:
break
param = params_dict[name]
param = params_dict[target_name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
loaded_params.add(target_name)
loaded_stacked_shards.add((target_name, shard_id))
break
if matched_stacked_param:
continue
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
if name in fused_param_names and name.endswith(".scale_weight"):
# Merged scale loaders interpret a missing shard id as shard 0.
# An exact fused key is already a complete vector, so copy it
# atomically and retain the default loader's shape validation.
default_weight_loader(param, loaded_weight)
else:
if name.endswith(".bias") and name not in params_dict:
continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
if name in fused_param_names:
loaded_fused_params.add(name)
required_split_shards = {
(name, shard_id)
for name, shard_id in expected_stacked_shards
if name not in loaded_fused_params
}
missing_stacked_shards = required_split_shards - loaded_stacked_shards
if missing_stacked_shards:
formatted_missing = ", ".join(
f"{name}[{shard_id}]"
for name, shard_id in sorted(
missing_stacked_shards,
key=lambda item: (item[0], str(item[1])),
)
)
raise ValueError(
"Missing required stacked checkpoint shards: "
f"{formatted_missing}"
)
return loaded_params
+48 -12
View File
@@ -92,6 +92,8 @@ class ComponentLoader(ABC):
"tokenizer": (TokenizerLoader, "transformers"),
"tokenizer_2": (TokenizerLoader, "transformers"),
"tokenizer_3": (TokenizerLoader, "transformers"),
# Cosmos3's model_index names its Qwen2 tokenizer "text_tokenizer".
"text_tokenizer": (TokenizerLoader, "transformers"),
"image_processor": (ImageProcessorLoader, "transformers"),
"feature_extractor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
@@ -391,18 +393,22 @@ class TextEncoderLoader(ComponentLoader):
model_config.quant_config = quant_cls()
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
with target_device:
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
model_config, # type: ignore[arg-type]
dtype=PRECISION_TO_TYPE[dtype],
device=target_device,
)
return model.eval()
architectures = getattr(model_config, "architectures", [])
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
if getattr(model_cls, "supports_hf_from_pretrained", False):
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
model_path,
model_config, # type: ignore[arg-type]
dtype=PRECISION_TO_TYPE[dtype],
device=target_device,
)
# HF passthrough encoders return before FastVideo's FSDP
# wrapping path, so the text stage needs their placement to
# put token tensors on the same device.
model._fastvideo_input_device = target_device
return model.eval()
with target_device:
model = model_cls(model_config) # type: ignore
weights_to_load = {name for name, _ in model.named_parameters()}
@@ -820,6 +826,24 @@ class VAELoader(ComponentLoader):
vae.load_state_dict(sd, strict=False)
return vae.eval()
if class_name == "LingBotWorld2WanVAE":
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
config.pop("_class_name", None)
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
weight_path = os.path.join(model_path, "Wan2.1_VAE.pth")
if not os.path.exists(weight_path):
raise FileNotFoundError(
f"Missing LingBot World 2 VAE weights: {weight_path}"
)
vae = vae_cls(
vae_config,
checkpoint_path=weight_path,
dtype=dtype,
).to(target_device)
return vae.eval()
# LTX-2 uses CausalVideoAutoencoder with nested "vae" config
if class_name == "CausalVideoAutoencoder" and "vae" in config:
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
@@ -1137,7 +1161,19 @@ class SchedulerLoader(ComponentLoader):
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
scheduler = scheduler_cls(**config)
# Diffusers checkpoints can carry newer scheduler config keys than the
# vendored scheduler accepts (e.g. shift_terminal / sigma_min / sigma_max
# from a newer diffusers release). Filter to the class's __init__ params,
# mirroring diffusers' ``from_config``, so loading is robust to schema
# drift instead of crashing on an unexpected kwarg.
import inspect
valid_params = set(inspect.signature(scheduler_cls.__init__).parameters)
filtered_config = {k: v for k, v in config.items() if k in valid_params}
dropped = sorted(set(config) - set(filtered_config))
if dropped:
logger.warning("Scheduler %s: dropping unsupported config keys %s", class_name, dropped)
scheduler = scheduler_cls(**filtered_config)
if fastvideo_args.pipeline_config.flow_shift is not None:
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
return scheduler
+94 -68
View File
@@ -15,13 +15,11 @@ import torch
from torch import nn
from torch.distributed import DeviceMesh, init_device_mesh
from torch.distributed._tensor import distribute_tensor
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule,
MixedPrecisionPolicy, fully_shard)
from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule, MixedPrecisionPolicy, fully_shard)
from torch.nn.modules.module import _IncompatibleKeys
from fastvideo.logger import init_logger
from fastvideo.models.loader.utils import (get_param_names_mapping,
hf_to_custom_state_dict)
from fastvideo.models.loader.utils import (get_param_names_mapping, hf_to_custom_state_dict)
from fastvideo.models.loader.weight_utils import safetensors_weights_iterator
from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available
@@ -43,13 +41,16 @@ def _maybe_quantize_model(model: nn.Module) -> None:
"""
# Defer imports: these modules pull in heavy symbols at module-load time.
from fastvideo.layers.quantization.nvfp4_config import (
NVFP4QuantizeMethod, convert_model_to_nvfp4,
NVFP4QuantizeMethod,
convert_model_to_nvfp4,
)
from fastvideo.layers.quantization.nvfp4_qat_config import (
NVFP4QATQuantizeMethod, convert_model_to_fp4,
NVFP4QATQuantizeMethod,
convert_model_to_fp4,
)
from fastvideo.layers.quantization.fp8_config import (
FP8QuantizeMethod, convert_model_to_fp8,
FP8QuantizeMethod,
convert_model_to_fp8,
)
for mod in model.modules():
@@ -121,10 +122,7 @@ def maybe_load_fsdp_model(
"""
# NOTE(will): cast_forward_inputs=True shouldn't be needed as we are
# manually casting the inputs to the model
mp_policy = MixedPrecisionPolicy(param_dtype,
reduce_dtype,
output_dtype,
cast_forward_inputs=False)
mp_policy = MixedPrecisionPolicy(param_dtype, reduce_dtype, output_dtype, cast_forward_inputs=False)
set_mixed_precision_policy(
param_dtype=param_dtype,
@@ -137,6 +135,13 @@ def maybe_load_fsdp_model(
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
dtype_selector = getattr(model, "_get_parameter_dtype", None)
has_mixed_parameter_dtypes = callable(dtype_selector) and any(
dtype_selector(name, param_dtype) != param_dtype for name, _ in model.named_parameters())
if training_mode and has_mixed_parameter_dtypes:
raise NotImplementedError("FSDP training with model-selected mixed parameter dtypes requires "
"separate gradient synchronization for replicated parameters.")
# Check if we should use FSDP
use_fsdp = training_mode or fsdp_inference
@@ -152,7 +157,7 @@ def maybe_load_fsdp_model(
if not training_mode and not fsdp_inference:
hsdp_replicate_dim = world_size
hsdp_shard_dim = 1
if current_platform.is_npu():
with torch.device("cpu"):
device_mesh = init_device_mesh(
@@ -163,11 +168,11 @@ def maybe_load_fsdp_model(
)
else:
device_mesh = init_device_mesh(
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
shard_model(model,
cpu_offload=cpu_offload,
reshard_after_forward=True,
@@ -188,12 +193,10 @@ def maybe_load_fsdp_model(
param_names_mapping=param_names_mapping_fn,
)
if hasattr(model, "materialize_non_persistent_buffers"):
model.materialize_non_persistent_buffers(
device=device, dtype=default_dtype)
model.materialize_non_persistent_buffers(device=device, dtype=default_dtype)
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(
f"Unexpected param or buffer {n} on meta device.")
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
# Avoid unintended computation graph accumulation during inference
if isinstance(p, torch.nn.Parameter):
p.requires_grad = False
@@ -209,8 +212,7 @@ def maybe_load_fsdp_model(
compile_in_loader = enable_torch_compile and training_mode
if compile_in_loader:
compile_kwargs = torch_compile_kwargs or {}
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s",
compile_kwargs)
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
return model
@@ -254,58 +256,81 @@ def shard_model(
"""
# Check if we should use size-based filtering
use_size_filtering = os.environ.get("FASTVIDEO_FSDP2_AUTOWRAP", "0") == "1"
if not fsdp_shard_conditions:
logger.warning("No FSDP shard conditions provided; nothing will be sharded.")
return
default_param_dtype = getattr(mp_policy, "param_dtype", None)
dtype_selector = getattr(model, "_get_parameter_dtype", None)
ignored_params: set[nn.Parameter] = set()
if callable(dtype_selector) and default_param_dtype is not None:
ignored_params = {
parameter
for name, parameter in model.named_parameters()
if dtype_selector(name, default_param_dtype) != default_param_dtype
}
named_modules = list(model.named_modules())
ignored_params_by_module = {
id(module): ignored_params.intersection(set(module.parameters()))
for _, module in named_modules
}
fsdp_kwargs = {
"reshard_after_forward": reshard_after_forward,
"mesh": mesh,
"mp_policy": mp_policy,
}
if cpu_offload:
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
pin_memory=pin_cpu_memory)
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(pin_memory=pin_cpu_memory)
# iterating in reverse to start with
# lowest-level modules first
num_layers_sharded = 0
if use_size_filtering:
# Size-based filtering mode
min_params = int(os.environ.get("FASTVIDEO_FSDP2_MIN_PARAMS", "10000000"))
logger.info("Using size-based filtering with threshold: %.2fM", min_params / 1e6)
for n, m in reversed(list(model.named_modules())):
for n, m in reversed(named_modules):
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
# Count all parameters
param_count = sum(p.numel() for p in m.parameters(recurse=True))
# Skip small modules
if param_count < min_params:
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)",
n, param_count / 1e6, min_params / 1e6)
logger.info("Skipping module %s (%.2fM params < %.2fM threshold)", n, param_count / 1e6,
min_params / 1e6)
continue
# Shard this module
logger.info("Sharding module %s (%.2fM params)", n, param_count / 1e6)
fully_shard(m, **fsdp_kwargs)
module_kwargs = fsdp_kwargs
local_ignored_params = ignored_params_by_module[id(m)]
if local_ignored_params:
module_kwargs = {**fsdp_kwargs, "ignored_params": local_ignored_params}
fully_shard(m, **module_kwargs)
num_layers_sharded += 1
else:
# Shard all modules matching conditions
for n, m in reversed(list(model.named_modules())):
# Shard all modules matching conditions
for n, m in reversed(named_modules):
if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]):
fully_shard(m, **fsdp_kwargs)
module_kwargs = fsdp_kwargs
local_ignored_params = ignored_params_by_module[id(m)]
if local_ignored_params:
module_kwargs = {**fsdp_kwargs, "ignored_params": local_ignored_params}
fully_shard(m, **module_kwargs)
num_layers_sharded += 1
if num_layers_sharded == 0:
raise ValueError(
"No layer modules were sharded. Please check if shard conditions are working as expected."
)
raise ValueError("No layer modules were sharded. Please check if shard conditions are working as expected.")
# Finally shard the entire model to account for any stragglers
fully_shard(model, **fsdp_kwargs)
root_kwargs = fsdp_kwargs
if ignored_params:
root_kwargs = {**fsdp_kwargs, "ignored_params": ignored_params}
fully_shard(model, **root_kwargs)
# TODO(PY): device mesh for cfg parallel
@@ -341,17 +366,17 @@ def load_model_from_full_model_state_dict(
"""
meta_sd = model.state_dict()
named_parameters = dict(model.named_parameters())
named_buffers = dict(model.named_buffers())
sharded_sd = {}
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
full_sd_iterator, param_names_mapping) # type: ignore
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(full_sd_iterator,
param_names_mapping) # type: ignore
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
# Some checkpoints include extra entries that are not part of the
# instantiated model's state_dict (e.g. `_extra_state` keys from
# some FSDP checkpoint formats). These can be safely skipped.
if (target_param_name.endswith("._extra_state")
or target_param_name.endswith("_extra_state")):
if (target_param_name.endswith("._extra_state") or target_param_name.endswith("_extra_state")):
logger.warning(
"Skipping non-parameter checkpoint key: %s",
target_param_name,
@@ -370,8 +395,12 @@ def load_model_from_full_model_state_dict(
raise ValueError(
f"Parameter {target_param_name} not found in custom model state dict. The hf to custom mapping may be incorrect."
)
target_dtype = param_dtype
dtype_selector = getattr(model, "_get_parameter_dtype", None)
if callable(dtype_selector):
target_dtype = dtype_selector(target_param_name, param_dtype)
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
target_param = named_parameters.get(target_param_name)
weight_loader = getattr(target_param, "weight_loader", None)
# Gated on a shape mismatch: only fused/stacked params with a custom
@@ -380,9 +409,7 @@ def load_model_from_full_model_state_dict(
# fall through to the original `sharded_tensor = full_tensor` below.
if target_param is not None and callable(weight_loader) and tuple(target_param.shape) != tuple(
full_tensor.shape):
loaded_param = nn.Parameter(torch.empty(tuple(target_param.shape),
device=device,
dtype=param_dtype),
loaded_param = nn.Parameter(torch.empty(tuple(target_param.shape), device=device, dtype=target_dtype),
requires_grad=False)
for attr_name, attr_value in vars(target_param).items():
setattr(loaded_param, attr_name, attr_value)
@@ -392,7 +419,7 @@ def load_model_from_full_model_state_dict(
# In cases where parts of the model aren't sharded, some parameters will be plain tensors.
sharded_tensor = full_tensor
else:
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
full_tensor = full_tensor.to(device=device, dtype=target_dtype)
sharded_tensor = distribute_tensor(
full_tensor,
meta_sharded_param.device_mesh,
@@ -400,36 +427,35 @@ def load_model_from_full_model_state_dict(
)
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
if target_param_name in named_buffers:
sharded_sd[target_param_name] = sharded_tensor
else:
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
model.reverse_param_names_mapping = reverse_param_names_mapping
unused_keys = set(meta_sd.keys()) - set(sharded_sd.keys())
if unused_keys:
logger.warning("Found unloaded parameters in meta state dict: %s",
unused_keys)
logger.warning("Found unloaded parameters in meta state dict: %s", unused_keys)
# List of allowed parameter name patterns
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
for new_param_name in unused_keys:
if not any(pattern in new_param_name
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
logger.error("Unsupported new parameter: %s. Allowed patterns: %s",
new_param_name, ALLOWED_NEW_PARAM_PATTERNS)
raise ValueError(
f"New parameter '{new_param_name}' is not supported. "
f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed."
)
if not any(pattern in new_param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS):
logger.error("Unsupported new parameter: %s. Allowed patterns: %s", new_param_name,
ALLOWED_NEW_PARAM_PATTERNS)
raise ValueError(f"New parameter '{new_param_name}' is not supported. "
f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed.")
meta_sharded_param = meta_sd.get(new_param_name)
target_dtype = param_dtype
dtype_selector = getattr(model, "_get_parameter_dtype", None)
if callable(dtype_selector):
target_dtype = dtype_selector(new_param_name, param_dtype)
if not hasattr(meta_sharded_param, "device_mesh"):
# Initialize with zeros
sharded_tensor = torch.zeros_like(meta_sharded_param,
device=device,
dtype=param_dtype)
sharded_tensor = torch.zeros_like(meta_sharded_param, device=device, dtype=target_dtype)
else:
# Initialize with zeros and distribute
full_tensor = torch.zeros_like(meta_sharded_param,
device=device,
dtype=param_dtype)
full_tensor = torch.zeros_like(meta_sharded_param, device=device, dtype=target_dtype)
sharded_tensor = distribute_tensor(
full_tensor,
meta_sharded_param.device_mesh,
+4 -2
View File
@@ -18,8 +18,10 @@ def set_default_torch_dtype(dtype: torch.dtype):
"""Sets the default torch dtype to the given dtype."""
old_dtype = torch.get_default_dtype()
torch.set_default_dtype(dtype)
yield
torch.set_default_dtype(old_dtype)
try:
yield
finally:
torch.set_default_dtype(old_dtype)
def get_param_names_mapping(
+19
View File
@@ -37,11 +37,19 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
# Cosmos3-Nano's checkpoint model_index names the DiT "Cosmos3OmniTransformer";
# map that HF class name to FastVideo's native Cosmos3VFMTransformer.
"Cosmos3OmniTransformer": ("dits", "cosmos3", "Cosmos3VFMTransformer"),
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
"LingBotWorld2CausalFastTransformer3DModel": (
"dits",
"lingbotworld2",
"LingBotWorld2CausalFastTransformer3DModel",
),
"Gen3CTransformer3DModel": ("dits", "gen3c", "Gen3CTransformer3DModel"),
"Kandinsky5Transformer3DModel": ("dits", "kandinsky5", "Kandinsky5Transformer3DModel"),
"Flux2Transformer2DModel": ("dits", "flux_2", "Flux2Transformer2DModel"),
@@ -53,6 +61,11 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"LingBotWorld2CausalFastTransformer3DModel": (
"dits",
"lingbotworld2",
"LingBotWorld2CausalFastTransformer3DModel",
),
"MatrixGame2WanModel": ("dits", "matrixgame2", "MatrixGame2WanModel"),
"CausalMatrixGame2WanModel": ("dits", "matrixgame2", "CausalMatrixGame2WanModel"),
# Legacy aliases for older HF model_index.json files
@@ -64,6 +77,7 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
# Text-to-image DiT models (2D image generation)
_TEXT_TO_IMAGE_DIT_MODELS = {
"GlmImageTransformer2DModel": ("dits", "glm_image", "GlmImageTransformer2DModel"),
"ZImageTransformer2DModel": ("dits", "zimage", "ZImageTransformer2DModel"),
}
_TEXT_ENCODER_MODELS = {
@@ -72,12 +86,16 @@ _TEXT_ENCODER_MODELS = {
("encoders", "clip", "CLIPTextModelWithProjection"),
"LlamaModel": ("encoders", "llama", "LlamaModel"),
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
"LingBotWorld2T5EncoderModel": ("encoders", "lingbotworld2_t5", "LingBotWorld2T5EncoderModel"),
"T5EncoderModel": ("encoders", "t5_hf", "T5EncoderModel"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
"Reason1TextEncoder": ("encoders", "reason1", "Reason1TextEncoder"),
"Qwen2_5_VLForConditionalGeneration":
("encoders", "reason1", "Reason1TextEncoder"),
# Z-Image-Turbo's text_encoder/config.json declares architecture
# "Qwen3Model"; route it to the shared Qwen3 encoder (added for Flux2 Klein).
"Qwen3Model": ("encoders", "qwen3", "Qwen3ForCausalLM"),
"LTX2GemmaTextEncoderModel": ("encoders", "gemma", "LTX2GemmaTextEncoderModel"),
"Qwen3ForCausalLM": ("encoders", "qwen3", "Qwen3ForCausalLM"),
"Mistral3ForConditionalGeneration":
@@ -98,6 +116,7 @@ _VAE_MODELS = {
"AutoencoderKLHYWorld": ("vaes", "hyworldvae", "AutoencoderKLHYWorld"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"LingBotWorld2WanVAE": ("vaes", "lingbotworld2_wanvae", "LingBotWorld2WanVAE"),
"AutoencoderKL": ("vaes", "autoencoder_kl", "AutoencoderKL"),
"AutoencoderKLGen3CTokenizer":
("vaes", "gen3c_tokenizer_vae", "AutoencoderKLGen3CTokenizer"),
@@ -96,6 +96,14 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
The minimum sigma value for the noise schedule.
sigma_data (`float`, *optional*):
The sigma data value for scaling.
use_reference_discrete_timesteps (`bool`, defaults to False):
Some reference schedulers (e.g. Z-Image) construct the timestep
schedule by linspacing `num_inference_steps + 1` points from
`t_max` to `t_min` and dropping the terminal point. Default
(`False`) preserves the original `np.linspace(t_max, t_min,
num_inference_steps)` (float64) behaviour used by every existing
model. Enable this flag only when matching a reference scheduler
that expects the +1 + drop-terminal construction.
"""
_compatibles: list[Any] = []
@@ -122,6 +130,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
sigma_max: float | None = None,
sigma_min: float | None = None,
sigma_data: float | None = None,
use_reference_discrete_timesteps: bool = False,
):
if sum([
self.config.use_beta_sigmas, self.config.use_exponential_sigmas,
@@ -155,7 +164,7 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
self.sigmas = sigmas.to(
"cpu") # to avoid too much CPU/GPU communication
self.sigma_min = self.sigmas[-1].item()
self.sigma_min = sigma_min if sigma_min is not None else self.sigmas[-1].item()
self.sigma_max = self.sigmas[0].item()
BaseScheduler.__init__(self)
@@ -350,7 +359,19 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
if timesteps_array is None:
t_max = self._sigma_to_t(self.sigma_max)
t_min = self._sigma_to_t(self.sigma_min)
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
if self.config.use_reference_discrete_timesteps:
# Some reference schedulers (for example Z-Image) build a
# float64 num_steps+1 linspace and drop the terminal point.
timesteps_array = np.linspace(
t_max,
t_min,
num_inference_steps + 1,
)[:-1]
else:
# Preserve the original numpy default (float64) here —
# casting to float32 silently shifts rounded timestep
# values for every existing model that uses this branch.
timesteps_array = np.linspace(t_max, t_min, num_inference_steps)
sigmas_array = timesteps_array / self.config.num_train_timesteps
else:
sigmas_array = np.array(sigmas).astype(np.float32)
@@ -0,0 +1,722 @@
import logging
from types import SimpleNamespace
import torch
import torch.cuda.amp as amp
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
__all__ = [
'Wan2_1_VAE',
'LingBotWorld2WanVAE',
]
CACHE_T = 2
class CausalConv3d(nn.Conv3d):
"""
Causal 3d convolusion.
"""
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self._padding = (self.padding[2], self.padding[2], self.padding[1],
self.padding[1], 2 * self.padding[0], 0)
self.padding = (0, 0, 0)
def forward(self, x, cache_x=None):
padding = list(self._padding)
if cache_x is not None and self._padding[4] > 0:
cache_x = cache_x.to(x.device)
x = torch.cat([cache_x, x], dim=2)
padding[4] -= cache_x.shape[2]
x = F.pad(x, padding)
return super().forward(x)
class RMS_norm(nn.Module):
def __init__(self, dim, channel_first=True, images=True, bias=False):
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.
def forward(self, x):
return F.normalize(
x, dim=(1 if self.channel_first else
-1)) * self.scale * self.gamma + self.bias
class Upsample(nn.Upsample):
def forward(self, x):
"""
Fix bfloat16 support for nearest neighbor interpolation.
"""
return super().forward(x.float()).type_as(x)
class Resample(nn.Module):
def __init__(self, dim, mode):
assert mode in ('none', 'upsample2d', 'upsample3d', 'downsample2d',
'downsample3d')
super().__init__()
self.dim = dim
self.mode = mode
# layers
if mode == 'upsample2d':
self.resample = nn.Sequential(
Upsample(scale_factor=(2., 2.), mode='nearest-exact'),
nn.Conv2d(dim, dim // 2, 3, padding=1))
elif mode == 'upsample3d':
self.resample = nn.Sequential(
Upsample(scale_factor=(2., 2.), mode='nearest-exact'),
nn.Conv2d(dim, dim // 2, 3, padding=1))
self.time_conv = CausalConv3d(
dim, dim * 2, (3, 1, 1), padding=(1, 0, 0))
elif mode == 'downsample2d':
self.resample = nn.Sequential(
nn.ZeroPad2d((0, 1, 0, 1)),
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
elif mode == 'downsample3d':
self.resample = nn.Sequential(
nn.ZeroPad2d((0, 1, 0, 1)),
nn.Conv2d(dim, dim, 3, stride=(2, 2)))
self.time_conv = CausalConv3d(
dim, dim, (3, 1, 1), stride=(2, 1, 1), padding=(0, 0, 0))
else:
self.resample = nn.Identity()
def forward(self, x, feat_cache=None, feat_idx=[0]):
b, c, t, h, w = x.size()
if self.mode == 'upsample3d':
if feat_cache is not None:
idx = feat_idx[0]
if feat_cache[idx] is None:
feat_cache[idx] = 'Rep'
feat_idx[0] += 1
else:
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[
idx] is not None and feat_cache[idx] != 'Rep':
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
if cache_x.shape[2] < 2 and feat_cache[
idx] is not None and feat_cache[idx] == 'Rep':
cache_x = torch.cat([
torch.zeros_like(cache_x).to(cache_x.device),
cache_x
],
dim=2)
if feat_cache[idx] == 'Rep':
x = self.time_conv(x)
else:
x = self.time_conv(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
x = x.reshape(b, 2, c, t, h, w)
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
3)
x = x.reshape(b, c, t * 2, h, w)
t = x.shape[2]
x = rearrange(x, 'b c t h w -> (b t) c h w')
x = self.resample(x)
x = rearrange(x, '(b t) c h w -> b c t h w', t=t)
if self.mode == 'downsample3d':
if feat_cache is not None:
idx = feat_idx[0]
if feat_cache[idx] is None:
feat_cache[idx] = x.clone()
feat_idx[0] += 1
else:
cache_x = x[:, :, -1:, :, :].clone()
# if cache_x.shape[2] < 2 and feat_cache[idx] is not None and feat_cache[idx]!='Rep':
# # cache last frame of last two chunk
# cache_x = torch.cat([feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(cache_x.device), cache_x], dim=2)
x = self.time_conv(
torch.cat([feat_cache[idx][:, :, -1:, :, :], x], 2))
feat_cache[idx] = cache_x
feat_idx[0] += 1
return x
def init_weight(self, conv):
conv_weight = conv.weight
nn.init.zeros_(conv_weight)
c1, c2, t, h, w = conv_weight.size()
one_matrix = torch.eye(c1, c2)
init_matrix = one_matrix
nn.init.zeros_(conv_weight)
#conv_weight.data[:,:,-1,1,1] = init_matrix * 0.5
conv_weight.data[:, :, 1, 0, 0] = init_matrix #* 0.5
conv.weight.data.copy_(conv_weight)
nn.init.zeros_(conv.bias.data)
def init_weight2(self, conv):
conv_weight = conv.weight.data
nn.init.zeros_(conv_weight)
c1, c2, t, h, w = conv_weight.size()
init_matrix = torch.eye(c1 // 2, c2)
#init_matrix = repeat(init_matrix, 'o ... -> (o 2) ...').permute(1,0,2).contiguous().reshape(c1,c2)
conv_weight[:c1 // 2, :, -1, 0, 0] = init_matrix
conv_weight[c1 // 2:, :, -1, 0, 0] = init_matrix
conv.weight.data.copy_(conv_weight)
nn.init.zeros_(conv.bias.data)
class ResidualBlock(nn.Module):
def __init__(self, in_dim, out_dim, dropout=0.0):
super().__init__()
self.in_dim = in_dim
self.out_dim = out_dim
# layers
self.residual = nn.Sequential(
RMS_norm(in_dim, images=False), nn.SiLU(),
CausalConv3d(in_dim, out_dim, 3, padding=1),
RMS_norm(out_dim, images=False), nn.SiLU(), nn.Dropout(dropout),
CausalConv3d(out_dim, out_dim, 3, padding=1))
self.shortcut = CausalConv3d(in_dim, out_dim, 1) \
if in_dim != out_dim else nn.Identity()
def forward(self, x, feat_cache=None, feat_idx=[0]):
h = self.shortcut(x)
for layer in self.residual:
if isinstance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x + h
class AttentionBlock(nn.Module):
"""
Causal self-attention with a single head.
"""
def __init__(self, dim):
super().__init__()
self.dim = dim
# layers
self.norm = RMS_norm(dim)
self.to_qkv = nn.Conv2d(dim, dim * 3, 1)
self.proj = nn.Conv2d(dim, dim, 1)
# zero out the last layer params
nn.init.zeros_(self.proj.weight)
def forward(self, x):
identity = x
b, c, t, h, w = x.size()
x = rearrange(x, 'b c t h w -> (b t) c h w')
x = self.norm(x)
# compute query, key, value
q, k, v = self.to_qkv(x).reshape(b * t, 1, c * 3,
-1).permute(0, 1, 3,
2).contiguous().chunk(
3, dim=-1)
# apply attention
x = F.scaled_dot_product_attention(
q,
k,
v,
)
x = x.squeeze(1).permute(0, 2, 1).reshape(b * t, c, h, w)
# output
x = self.proj(x)
x = rearrange(x, '(b t) c h w-> b c t h w', t=t)
return x + identity
class Encoder3d(nn.Module):
def __init__(self,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[True, True, False],
dropout=0.0):
super().__init__()
self.dim = dim
self.z_dim = z_dim
self.dim_mult = dim_mult
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
# dimensions
dims = [dim * u for u in [1] + dim_mult]
scale = 1.0
# init block
self.conv1 = CausalConv3d(3, dims[0], 3, padding=1)
# downsample blocks
downsamples = []
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
for _ in range(num_res_blocks):
downsamples.append(ResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
downsamples.append(AttentionBlock(out_dim))
in_dim = out_dim
# downsample block
if i != len(dim_mult) - 1:
mode = 'downsample3d' if temperal_downsample[
i] else 'downsample2d'
downsamples.append(Resample(out_dim, mode=mode))
scale /= 2.0
self.downsamples = nn.Sequential(*downsamples)
# middle blocks
self.middle = nn.Sequential(
ResidualBlock(out_dim, out_dim, dropout), AttentionBlock(out_dim),
ResidualBlock(out_dim, out_dim, dropout))
# output blocks
self.head = nn.Sequential(
RMS_norm(out_dim, images=False), nn.SiLU(),
CausalConv3d(out_dim, z_dim, 3, padding=1))
def forward(self, x, feat_cache=None, feat_idx=[0]):
if feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv1(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = self.conv1(x)
## downsamples
for layer in self.downsamples:
if feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## middle
for layer in self.middle:
if isinstance(layer, ResidualBlock) and feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## head
for layer in self.head:
if isinstance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x
class Decoder3d(nn.Module):
def __init__(self,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_upsample=[False, True, True],
dropout=0.0):
super().__init__()
self.dim = dim
self.z_dim = z_dim
self.dim_mult = dim_mult
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_upsample = temperal_upsample
# dimensions
dims = [dim * u for u in [dim_mult[-1]] + dim_mult[::-1]]
scale = 1.0 / 2**(len(dim_mult) - 2)
# init block
self.conv1 = CausalConv3d(z_dim, dims[0], 3, padding=1)
# middle blocks
self.middle = nn.Sequential(
ResidualBlock(dims[0], dims[0], dropout), AttentionBlock(dims[0]),
ResidualBlock(dims[0], dims[0], dropout))
# upsample blocks
upsamples = []
for i, (in_dim, out_dim) in enumerate(zip(dims[:-1], dims[1:])):
# residual (+attention) blocks
if i == 1 or i == 2 or i == 3:
in_dim = in_dim // 2
for _ in range(num_res_blocks + 1):
upsamples.append(ResidualBlock(in_dim, out_dim, dropout))
if scale in attn_scales:
upsamples.append(AttentionBlock(out_dim))
in_dim = out_dim
# upsample block
if i != len(dim_mult) - 1:
mode = 'upsample3d' if temperal_upsample[i] else 'upsample2d'
upsamples.append(Resample(out_dim, mode=mode))
scale *= 2.0
self.upsamples = nn.Sequential(*upsamples)
# output blocks
self.head = nn.Sequential(
RMS_norm(out_dim, images=False), nn.SiLU(),
CausalConv3d(out_dim, 3, 3, padding=1))
def forward(self, x, feat_cache=None, feat_idx=[0]):
## conv1
if feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = self.conv1(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = self.conv1(x)
## middle
for layer in self.middle:
if isinstance(layer, ResidualBlock) and feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## upsamples
for layer in self.upsamples:
if feat_cache is not None:
x = layer(x, feat_cache, feat_idx)
else:
x = layer(x)
## head
for layer in self.head:
if isinstance(layer, CausalConv3d) and feat_cache is not None:
idx = feat_idx[0]
cache_x = x[:, :, -CACHE_T:, :, :].clone()
if cache_x.shape[2] < 2 and feat_cache[idx] is not None:
# cache last frame of last two chunk
cache_x = torch.cat([
feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
cache_x.device), cache_x
],
dim=2)
x = layer(x, feat_cache[idx])
feat_cache[idx] = cache_x
feat_idx[0] += 1
else:
x = layer(x)
return x
def count_conv3d(model):
count = 0
for m in model.modules():
if isinstance(m, CausalConv3d):
count += 1
return count
class WanVAE_(nn.Module):
def __init__(self,
dim=128,
z_dim=4,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[True, True, False],
dropout=0.0):
super().__init__()
self.dim = dim
self.z_dim = z_dim
self.dim_mult = dim_mult
self.num_res_blocks = num_res_blocks
self.attn_scales = attn_scales
self.temperal_downsample = temperal_downsample
self.temperal_upsample = temperal_downsample[::-1]
# modules
self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks,
attn_scales, self.temperal_downsample, dropout)
self.conv1 = CausalConv3d(z_dim * 2, z_dim * 2, 1)
self.conv2 = CausalConv3d(z_dim, z_dim, 1)
self.decoder = Decoder3d(dim, z_dim, dim_mult, num_res_blocks,
attn_scales, self.temperal_upsample, dropout)
def forward(self, x):
mu, log_var = self.encode(x)
z = self.reparameterize(mu, log_var)
x_recon = self.decode(z)
return x_recon, mu, log_var
def encode(self, x, scale):
self.clear_cache()
## cache
t = x.shape[2]
iter_ = 1 + (t - 1) // 4
for i in range(iter_):
self._enc_conv_idx = [0]
if i == 0:
out = self.encoder(
x[:, :, :1, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
else:
out_ = self.encoder(
x[:, :, 1 + 4 * (i - 1):1 + 4 * i, :, :],
feat_cache=self._enc_feat_map,
feat_idx=self._enc_conv_idx)
out = torch.cat([out, out_], 2)
mu, log_var = self.conv1(out).chunk(2, dim=1)
if isinstance(scale[0], torch.Tensor):
mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view(
1, self.z_dim, 1, 1, 1)
else:
mu = (mu - scale[0]) * scale[1]
self.clear_cache()
return mu
def decode(self, z, scale):
self.clear_cache()
# z: [b,c,t,h,w]
if isinstance(scale[0], torch.Tensor):
z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view(
1, self.z_dim, 1, 1, 1)
else:
z = z / scale[1] + scale[0]
iter_ = z.shape[2]
x = self.conv2(z)
for i in range(iter_):
self._conv_idx = [0]
if i == 0:
out = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
else:
out_ = self.decoder(
x[:, :, i:i + 1, :, :],
feat_cache=self._feat_map,
feat_idx=self._conv_idx)
out = torch.cat([out, out_], 2)
self.clear_cache()
return out
def reparameterize(self, mu, log_var):
std = torch.exp(0.5 * log_var)
eps = torch.randn_like(std)
return eps * std + mu
def sample(self, imgs, deterministic=False):
mu, log_var = self.encode(imgs)
if deterministic:
return mu
std = torch.exp(0.5 * log_var.clamp(-30.0, 20.0))
return mu + std * torch.randn_like(std)
def clear_cache(self):
self._conv_num = count_conv3d(self.decoder)
self._conv_idx = [0]
self._feat_map = [None] * self._conv_num
#cache encode
self._enc_conv_num = count_conv3d(self.encoder)
self._enc_conv_idx = [0]
self._enc_feat_map = [None] * self._enc_conv_num
def _video_vae(pretrained_path=None, z_dim=None, device='cpu', **kwargs):
"""
Autoencoder3d adapted from Stable Diffusion 1.x, 2.x and XL.
"""
# params
cfg = dict(
dim=96,
z_dim=z_dim,
dim_mult=[1, 2, 4, 4],
num_res_blocks=2,
attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0)
cfg.update(**kwargs)
# init model
with torch.device('meta'):
model = WanVAE_(**cfg)
# load checkpoint
logging.info(f'loading {pretrained_path}')
model.load_state_dict(
torch.load(pretrained_path, map_location=device), assign=True)
return model
class Wan2_1_VAE:
def __init__(self,
z_dim=16,
vae_pth='cache/vae_step_411000.pth',
dtype=torch.float,
device="cuda"):
self.dtype = dtype
self.device = device
mean = [
-0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508,
0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921
]
std = [
2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743,
3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160
]
self.mean = torch.tensor(mean, dtype=dtype, device=device)
self.std = torch.tensor(std, dtype=dtype, device=device)
self.scale = [self.mean, 1.0 / self.std]
# init model
self.model = _video_vae(
pretrained_path=vae_pth,
z_dim=z_dim,
).eval().requires_grad_(False).to(device)
def encode(self, videos):
"""
videos: A list of videos each with shape [C, T, H, W].
"""
with amp.autocast(dtype=self.dtype):
return [
self.model.encode(u.unsqueeze(0), self.scale).float().squeeze(0)
for u in videos
]
def decode(self, zs):
with amp.autocast(dtype=self.dtype):
return [
self.model.decode(u.unsqueeze(0),
self.scale).float().clamp_(-1, 1).squeeze(0)
for u in zs
]
class LingBotWorld2WanVAE(nn.Module):
"""FastVideo-facing wrapper around the exact LingBot World 2 Wan2.1 VAE computation."""
handles_latent_denorm = True
def __init__(self, config, checkpoint_path=None, dtype=torch.float):
"""Load the official LingBot World 2 VAE weights and expose FastVideo VAE APIs."""
super().__init__()
self.config = config
self.dtype = dtype
z_dim = int(getattr(config, "z_dim", 16))
mean = torch.tensor(getattr(config, "latents_mean"), dtype=dtype)
std = torch.tensor(getattr(config, "latents_std"), dtype=dtype)
self.register_buffer("shift_factor", mean.view(1, z_dim, 1, 1, 1), persistent=False)
self.register_buffer("scaling_factor", (1.0 / std).view(1, z_dim, 1, 1, 1), persistent=False)
self.scale = [mean, 1.0 / std]
if checkpoint_path is None:
self.model = WanVAE_(dim=96, z_dim=z_dim, dim_mult=[1, 2, 4, 4],
num_res_blocks=2, attn_scales=[],
temperal_downsample=[False, True, True],
dropout=0.0)
else:
self.model = _video_vae(pretrained_path=checkpoint_path, z_dim=z_dim)
self.model.eval().requires_grad_(False)
def _scale_for(self, device: torch.device) -> list[torch.Tensor]:
"""Return source VAE scale tensors on the active device."""
return [u.to(device) for u in self.scale]
def encode(self, videos: torch.Tensor):
"""Encode `[B,C,T,H,W]` videos and return a FastVideo-style mean tensor."""
if videos.ndim != 5:
raise ValueError(f"LingBotWorld2WanVAE.encode expects 5D input, got {videos.shape}")
scale = self._scale_for(videos.device)
latents = []
with amp.autocast(dtype=self.dtype):
for video in videos:
latent = self.model.encode(video.unsqueeze(0), scale).float().squeeze(0)
latents.append(latent)
normalized = torch.stack(latents, dim=0)
return SimpleNamespace(mean=normalized)
def decode(self, latents: torch.Tensor) -> torch.Tensor:
"""Decode normalized LingBot World 2 latents to clamped `[-1,1]` video tensors."""
if latents.ndim != 5:
raise ValueError(f"LingBotWorld2WanVAE.decode expects 5D input, got {latents.shape}")
scale = self._scale_for(latents.device)
videos = []
with amp.autocast(dtype=self.dtype):
for latent in latents:
video = self.model.decode(latent.unsqueeze(0), scale).float().clamp_(-1, 1).squeeze(0)
videos.append(video)
return torch.stack(videos, dim=0)
EntryClass = LingBotWorld2WanVAE
+3
View File
@@ -94,6 +94,9 @@ class OobleckDecoderBlock(nn.Module):
input_dim, output_dim,
kernel_size=2 * stride, stride=stride,
padding=math.ceil(stride / 2),
# Clean L*stride upsample for both parities; a no-op (0) for even
# strides (Stable Audio), needed for odd strides (Cosmos3: 5).
output_padding=stride % 2,
))
self.res_unit1 = OobleckResidualUnit(output_dim, dilation=1)
self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3)
+1 -1
View File
@@ -35,7 +35,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs,
"""
# Get pipeline type
model_path = fastvideo_args.model_path
model_path = maybe_download_model(model_path)
model_path = maybe_download_model(model_path, revision=fastvideo_args.revision)
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
@@ -0,0 +1,708 @@
# SPDX-License-Identifier: Apache-2.0
"""FastVideo-native Cosmos3 video pipeline (T2V / I2V / T2I).
This replaces the earlier vllm-omni-derived skeleton with a native, stage-based
:class:`ComposedPipelineBase` pipeline that wires the framework-parity-verified
Cosmos3 components:
* tokenizer: Qwen2 ``Qwen2TokenizerFast`` + chat template (the only allowed
third-party model-adjacent dependency; tokenizers are explicitly permitted),
* VAE: FastVideo-native ``AutoencoderKLWan`` (Wan2.2) via ``Cosmos3VAEConfig``;
encode normalizes ``(mu - mean) * inv_std`` and decode denormalizes + clamps,
* sequence-packing: :func:`pack_cosmos3_video_sequence` (native, parity-tested),
* DiT: ``Cosmos3VFMTransformer`` (native, bit-identical to the framework),
* scheduler: FastVideo-native ``UniPCMultistepScheduler`` configured for pure
flow matching (``flow_prediction`` + ``use_flow_sigmas``), numerically
equivalent to the framework's ``FlowUniPCMultistepScheduler`` (parity-tested
in ``test_cosmos3_scheduler_parity``).
The denoise/CFG glue is a faithful port of the framework's
``Cosmos3OmniDiffusersPipeline`` math (mirrored in the framework-equivalent
``diffusers_cosmos3.pipeline``): per UniPC timestep, run a SEQUENTIAL conditional
then unconditional pass (each repacks the sequence with the prompt / negative
prompt token ids, forwards the DiT, and zeros the prediction on conditioning
frames), then combine ``v = uncond + guidance * (cond - uncond)`` and take one
``scheduler.step(model_output=v, timestep, sample=latent)``. ``timestep_scale``
is applied to the per-token timesteps *inside* the DiT (its ``forward`` already
multiplies ``vision_timesteps * timestep_scale`` before the time embedder), so
the loop passes raw scheduler timesteps to the packer.
The pure denoise math lives in :class:`Cosmos3DenoiseEngine` and the free
function :func:`cosmos3_get_cfg_velocity` so it can be unit-/parity-tested
directly against the framework oracle without constructing the full pipeline.
No diffusers/transformers *model* classes are imported at runtime here; only the
Qwen2 tokenizer (loaded by the component loader) and the UniPC scheduler are
third-party, both explicitly allowed.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from typing import Any
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_unipc_multistep import (
UniPCMultistepScheduler, )
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
Cosmos3ActionItem,
Cosmos3SampleInputs,
Cosmos3SoundItem,
Cosmos3VisionItem,
pack_cosmos3_video_sequence,
)
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
logger = init_logger(__name__)
# System prompts, verbatim from the framework (diffusers_cosmos3.pipeline).
_SYSTEM_PROMPT_IMAGE = "You are a helpful assistant who will generate images from a give prompt."
_SYSTEM_PROMPT_VIDEO = "You are a helpful assistant who will generate videos from a give prompt."
# ===========================================================================
# Special-token resolution (Qwen2 chat tokenizer)
# ===========================================================================
def cosmos3_special_tokens(tokenizer: Any) -> dict[str, int]:
"""Resolve the Cosmos3 generation special tokens from a Qwen2 tokenizer.
Mirrors the framework's ``llm_special_tokens``:
``start_of_generation=<|vision_start|>``, ``end_of_generation=<|vision_end|>``,
``eos_token_id=tokenizer.eos_token_id``.
"""
return {
"start_of_generation": int(tokenizer.convert_tokens_to_ids("<|vision_start|>")),
"end_of_generation": int(tokenizer.convert_tokens_to_ids("<|vision_end|>")),
"eos_token_id": int(tokenizer.eos_token_id),
}
def cosmos3_tokenize_caption(
tokenizer: Any,
caption: str,
*,
is_video: bool = False,
use_system_prompt: bool = False,
) -> list[int]:
"""Tokenize a caption with the Qwen2 chat template (framework-faithful).
Optionally prepends an image/video system prompt; always adds the
generation prompt and disables ``add_vision_id`` (matching the framework's
``tokenize_caption``).
"""
conversations: list[dict[str, str]] = []
if use_system_prompt:
conversations.append({
"role": "system",
"content": _SYSTEM_PROMPT_VIDEO if is_video else _SYSTEM_PROMPT_IMAGE,
})
conversations.append({"role": "user", "content": caption})
token_ids = tokenizer.apply_chat_template(
conversations,
tokenize=True,
add_generation_prompt=True,
add_vision_id=False,
)
return list(token_ids)
# ===========================================================================
# Reasoning (VLM text generation) — und (causal) pathway + lm_head
# ===========================================================================
def cosmos3_generate_reasoner_text(
transformer: Any,
input_ids: list[int],
max_new_tokens: int,
*,
eos_token_id: int | list[int] | None = None,
) -> list[int]:
"""Greedy text reasoning via the und (causal) backbone + ``lm_head``.
Mirrors the framework ``generate_reasoner_text`` (text-only prefill, greedy):
only the und-pathway weights (no ``_moe_gen``) + ``embed_tokens`` / ``norm`` /
``lm_head`` participate; the generation pathway and the VFM multimodal
embedders are bypassed (no vision/sound/action tokens). Token-for-token
identical to the framework reasoner (``test_cosmos3_reasoning_parity``).
Re-prefills each step (no KV cache) — correctness-first; a KV-cache fast path
is a later optimization. Returns the newly generated token ids.
"""
device = next(transformer.parameters()).device
ids = [int(x) for x in input_ids]
eos: set[int] = set()
if eos_token_id is not None:
eos = {int(eos_token_id)} if isinstance(eos_token_id, int) else {int(x) for x in eos_token_id}
new_tokens: list[int] = []
for _ in range(int(max_new_tokens)):
n = len(ids)
pos = torch.arange(n).unsqueeze(0).expand(3, -1).contiguous().to(device)
out = transformer(
text_ids=torch.tensor(ids, device=device, dtype=torch.long),
text_indexes=torch.arange(n, device=device),
position_ids=pos,
sequence_length=n,
split_lens=[n],
attn_modes=["causal"],
vision_tokens=[],
vision_token_shapes=[],
vision_sequence_indexes=torch.empty(0, dtype=torch.long, device=device),
vision_timesteps=torch.empty(0, device=device),
vision_mse_loss_indexes=torch.empty(0, dtype=torch.long, device=device),
vision_noisy_frame_indexes=[],
)
logits = transformer.lm_head(out["last_hidden_state"][n - 1]) # [vocab]
nxt = int(logits.argmax().item())
ids.append(nxt)
new_tokens.append(nxt)
if nxt in eos:
break
return new_tokens
# ===========================================================================
# VAE encode/decode bridge (normalize / denormalize, matching the framework)
# ===========================================================================
@dataclass
class _VaeNorm:
"""Cached ``mean`` / ``inv_std`` for VAE (de)normalization."""
mean: torch.Tensor # [z_dim]
inv_std: torch.Tensor # [z_dim]
@classmethod
def from_vae(cls, vae: Any, dtype: torch.dtype) -> _VaeNorm:
mean = torch.tensor(list(vae.config.latents_mean), dtype=dtype)
std = torch.tensor(list(vae.config.latents_std), dtype=dtype)
return cls(mean=mean, inv_std=1.0 / std)
def cosmos3_vae_encode(vae: Any, video: torch.Tensor, norm: _VaeNorm) -> torch.Tensor:
"""Encode ``[B, 3, T, H, W]`` pixels in [-1, 1] to NORMALIZED latents.
Matches the framework ``DiffusersWan22VAE.encode``: take the posterior mode
and apply ``(mu - mean) * inv_std``. FastVideo's ``AutoencoderKLWan.encode``
returns a ``DiagonalGaussianDistribution``; we read ``.mode()``.
"""
in_dtype = video.dtype
device = video.device
mean = norm.mean.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
inv_std = norm.inv_std.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
raw_mu = vae.encode(video).mode()
return ((raw_mu - mean) * inv_std).to(in_dtype)
def cosmos3_vae_decode(vae: Any, latents: torch.Tensor, norm: _VaeNorm) -> torch.Tensor:
"""Decode NORMALIZED latents ``[B, z, T, H, W]`` to pixels ``[B, 3, T, H, W]``.
Inverts the normalization (``z / inv_std + mean``) then calls
``vae.decode`` (which already clamps to [-1, 1]).
"""
in_dtype = latents.dtype
device = latents.device
mean = norm.mean.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
inv_std = norm.inv_std.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
z_raw = latents / inv_std + mean
out = vae.decode(z_raw)
if isinstance(out, tuple):
out = out[0]
if hasattr(out, "sample"):
out = out.sample
return out.to(in_dtype)
# ===========================================================================
# Per-vision-item packing geometry
# ===========================================================================
@dataclass
class Cosmos3VisionSpec:
"""Geometry + conditioning for one vision item in a denoise run.
Args:
condition_frame_indexes: Latent-frame indices kept clean.
shape: ``(C, T, H, W)`` of the latent for this item.
"""
shape: tuple[int, int, int, int]
condition_frame_indexes: list[int]
@property
def numel(self) -> int:
return int(math.prod(self.shape))
# ===========================================================================
# Pure denoise/CFG math (parity oracle target)
# ===========================================================================
def _split_flat_latent(flat: torch.Tensor, specs: list[Any]) -> list[torch.Tensor]:
"""Split a flat vector into per-item tensors via each spec's ``numel``/``shape``.
Shared by vision (``[C, T, H, W]``), sound (``[C, T]``), and action
(``[T, D]``) specs — every spec exposes ``numel`` and ``shape``.
"""
out: list[torch.Tensor] = []
offset = 0
for spec in specs:
out.append(flat[offset:offset + spec.numel].reshape(spec.shape))
offset += spec.numel
return out
@dataclass
class Cosmos3SoundSpec:
"""Geometry + conditioning for one sound item in a denoise run.
Args:
shape: ``(C, T)`` of the sound latent (channels, temporal frames).
condition_frame_indexes: Latent-frame indices kept clean (``[]`` for t2vs).
fps: Sound latent FPS (``sound_latent_fps``); used iff fps modulation is on.
"""
shape: tuple[int, int]
condition_frame_indexes: list[int] = field(default_factory=list)
fps: float | None = None
@property
def numel(self) -> int:
return int(math.prod(self.shape))
@dataclass
class Cosmos3ActionSpec:
"""Geometry + conditioning for one action item in a denoise run.
Args:
shape: ``(T, action_dim)`` of the action latent.
condition_frame_indexes: Frame indices kept clean (conditioning actions).
domain_id: Embodiment domain id for the domain-aware action projection.
fps: Action FPS; used iff fps modulation is on.
"""
shape: tuple[int, int]
condition_frame_indexes: list[int] = field(default_factory=list)
domain_id: int = 0
fps: float | None = None
@property
def numel(self) -> int:
return int(math.prod(self.shape))
def cosmos3_get_cfg_velocity(
*,
transformer: Any,
flat_latent: torch.Tensor,
timestep: torch.Tensor,
guidance: float,
specs: list[Cosmos3VisionSpec],
cond_token_ids: list[int],
uncond_token_ids: list[int],
special_tokens: dict[str, int],
latent_patch_size: int,
temporal_modality_margin: int,
reset_spatial_ids: bool,
enable_fps_modulation: bool,
base_fps: float,
temporal_compression_factor: int,
include_end_of_generation_token: bool = False,
fps_per_item: list[float] | None = None,
normalize_cfg: bool = False,
sound_specs: list[Cosmos3SoundSpec] | None = None,
sound_fps_per_item: list[float] | None = None,
action_specs: list[Cosmos3ActionSpec] | None = None,
action_fps_per_item: list[float] | None = None,
) -> torch.Tensor:
"""Sequential-CFG velocity for one denoise step (framework math).
Replicates the framework ``get_cfg_velocity``:
1. split ``flat_latent`` into per-vision-item ``[C, T, H, W]`` latents,
2. run a conditional pass (prompt tokens) and an unconditional pass
(negative-prompt tokens); each repacks via
:func:`pack_cosmos3_video_sequence`, forwards the DiT to obtain
``preds_vision`` (a list of ``[1, C, T, H, W]`` unpatchified noisy-frame
predictions), and zeros the prediction on conditioning frames
(``pred * (1 - condition_mask)``),
3. combine ``v = uncond + guidance * (cond - uncond)`` (optionally
norm-rescaled), returned flattened to match ``flat_latent``.
``timestep`` is a scalar tensor (raw scheduler timestep); ``timestep_scale``
is applied inside the DiT, so it is passed through unscaled here.
"""
assert timestep.numel() == 1, "timestep must be a scalar"
timestep_value = float(timestep.reshape(()).item())
# Combined flat layout: [all vision | all action | all sound], matching the
# framework per-sample concat order ([vision_i | action_i | sound_i]); single
# sample here.
vision_total = sum(spec.numel for spec in specs)
action_total = sum(spec.numel for spec in action_specs) if action_specs else 0
noise_x_vision = _split_flat_latent(flat_latent[:vision_total], specs)
noise_x_action = (_split_flat_latent(flat_latent[vision_total:vision_total +
action_total], action_specs) if action_specs else None)
noise_x_sound = (_split_flat_latent(flat_latent[vision_total +
action_total:], sound_specs) if sound_specs else None)
device = next(transformer.parameters()).device
def _run(token_ids: list[int]) -> torch.Tensor:
sound_items: list[Cosmos3SoundItem] = []
if sound_specs is not None and noise_x_sound is not None:
sound_items = [
Cosmos3SoundItem(
latent=noise_x_sound[i],
condition_frame_indexes=list(ss.condition_frame_indexes),
fps=(sound_fps_per_item[i] if sound_fps_per_item is not None else None),
) for i, ss in enumerate(sound_specs)
]
action_items: list[Cosmos3ActionItem] = []
if action_specs is not None and noise_x_action is not None:
action_items = [
Cosmos3ActionItem(
latent=noise_x_action[i],
condition_frame_indexes=list(asp.condition_frame_indexes),
domain_id=asp.domain_id,
fps=(action_fps_per_item[i] if action_fps_per_item is not None else None),
) for i, asp in enumerate(action_specs)
]
samples = [
Cosmos3SampleInputs(
text_ids=list(token_ids),
vision=Cosmos3VisionItem(
latent=latent,
condition_frame_indexes=list(spec.condition_frame_indexes),
fps=(fps_per_item[i] if fps_per_item is not None else None),
),
sound=(sound_items[i] if i < len(sound_items) else None),
action=(action_items[i] if i < len(action_items) else None),
timestep=timestep_value,
) for i, (latent, spec) in enumerate(zip(noise_x_vision, specs, strict=False))
]
packed = pack_cosmos3_video_sequence(
samples,
special_tokens,
latent_patch_size=latent_patch_size,
include_end_of_generation_token=include_end_of_generation_token,
temporal_modality_margin=temporal_modality_margin,
reset_spatial_ids=reset_spatial_ids,
enable_fps_modulation=enable_fps_modulation,
base_fps=base_fps,
temporal_compression_factor=temporal_compression_factor,
)
out = transformer(**packed.to_dit_kwargs(device=device))
# Vision velocity: zero on conditioning frames, per item, flattened.
vision_vel = torch.zeros(vision_total, device=flat_latent.device, dtype=flat_latent.dtype)
preds = out.get("preds_vision")
if preds is not None:
items: list[torch.Tensor] = []
for pred, cond_mask in zip(preds, packed.vision_condition_mask, strict=False):
pred = pred.squeeze(0) if pred.dim() == 5 else pred # [C, T, H, W]
keep = (1.0 - cond_mask).to(dtype=pred.dtype, device=pred.device) # [T,1,1]
items.append(pred * keep if keep.sum() > 0 else torch.zeros_like(pred))
vision_vel = torch.cat([v.reshape(-1) for v in items]).to(flat_latent.dtype)
parts = [vision_vel]
if action_specs:
# Action velocity: preds_action are per-item [T, D], already zero on
# clean frames; zero on cond frames defensively.
action_vel = torch.zeros(action_total, device=flat_latent.device, dtype=flat_latent.dtype)
preds_a = out.get("preds_action")
if preds_a is not None:
a_items: list[torch.Tensor] = []
for pred, cond_mask in zip(preds_a, packed.action_condition_mask, strict=False):
pred = pred.squeeze(0) if pred.dim() == 3 else pred # [T, D]
keep = (1.0 - cond_mask).reshape(-1, 1).to(dtype=pred.dtype, device=pred.device) # [T, 1]
a_items.append(pred * keep)
action_vel = torch.cat([v.reshape(-1) for v in a_items]).to(flat_latent.dtype)
parts.append(action_vel)
if sound_specs:
# Sound velocity: preds_sound are per-item [C, T], already zero on clean
# frames (unpack fills only noisy frames); zero on cond frames defensively.
sound_total = sum(spec.numel for spec in sound_specs)
sound_vel = torch.zeros(sound_total, device=flat_latent.device, dtype=flat_latent.dtype)
preds_s = out.get("preds_sound")
if preds_s is not None:
s_items: list[torch.Tensor] = []
for pred, cond_mask in zip(preds_s, packed.sound_condition_mask, strict=False):
pred = pred.squeeze(0) if pred.dim() == 3 else pred # [C, T]
keep = (1.0 - cond_mask).reshape(1, -1).to(dtype=pred.dtype, device=pred.device) # [1, T]
s_items.append(pred * keep)
sound_vel = torch.cat([v.reshape(-1) for v in s_items]).to(flat_latent.dtype)
parts.append(sound_vel)
return vision_vel if len(parts) == 1 else torch.cat(parts)
cond_v = _run(cond_token_ids)
uncond_v = _run(uncond_token_ids)
v_pred = uncond_v + guidance * (cond_v - uncond_v)
if normalize_cfg:
scale = (torch.norm(cond_v) / (torch.norm(v_pred) + 1e-8)).clamp(min=0.0, max=1.0)
v_pred = v_pred * scale
return v_pred
class Cosmos3DenoiseEngine:
"""Stateless denoise driver tying CFG velocity to UniPC stepping.
Holds the transformer + scheduler + packing constants and runs the full
UniPC denoise loop. Kept separate from the pipeline so it can be exercised
in isolation (smoke + parity tests) with stub or real components.
"""
def __init__(
self,
*,
transformer: Any,
scheduler: Any,
special_tokens: dict[str, int],
latent_patch_size: int,
temporal_modality_margin: int,
reset_spatial_ids: bool,
enable_fps_modulation: bool,
base_fps: float,
temporal_compression_factor: int,
include_end_of_generation_token: bool = False,
) -> None:
self.transformer = transformer
self.scheduler = scheduler
self.special_tokens = special_tokens
self.latent_patch_size = latent_patch_size
self.temporal_modality_margin = temporal_modality_margin
self.reset_spatial_ids = reset_spatial_ids
self.enable_fps_modulation = enable_fps_modulation
self.base_fps = base_fps
self.temporal_compression_factor = temporal_compression_factor
self.include_end_of_generation_token = include_end_of_generation_token
def velocity(
self,
*,
flat_latent: torch.Tensor,
timestep: torch.Tensor,
guidance: float,
specs: list[Cosmos3VisionSpec],
cond_token_ids: list[int],
uncond_token_ids: list[int],
fps_per_item: list[float] | None = None,
sound_specs: list[Cosmos3SoundSpec] | None = None,
sound_fps_per_item: list[float] | None = None,
action_specs: list[Cosmos3ActionSpec] | None = None,
action_fps_per_item: list[float] | None = None,
) -> torch.Tensor:
return cosmos3_get_cfg_velocity(
transformer=self.transformer,
flat_latent=flat_latent,
timestep=timestep,
guidance=guidance,
specs=specs,
cond_token_ids=cond_token_ids,
uncond_token_ids=uncond_token_ids,
special_tokens=self.special_tokens,
latent_patch_size=self.latent_patch_size,
temporal_modality_margin=self.temporal_modality_margin,
reset_spatial_ids=self.reset_spatial_ids,
enable_fps_modulation=self.enable_fps_modulation,
base_fps=self.base_fps,
temporal_compression_factor=self.temporal_compression_factor,
include_end_of_generation_token=self.include_end_of_generation_token,
fps_per_item=fps_per_item,
sound_specs=sound_specs,
sound_fps_per_item=sound_fps_per_item,
action_specs=action_specs,
action_fps_per_item=action_fps_per_item,
)
def denoise(
self,
*,
flat_latent: torch.Tensor,
timesteps: torch.Tensor,
guidance: float,
specs: list[Cosmos3VisionSpec],
cond_token_ids: list[int],
uncond_token_ids: list[int],
fps_per_item: list[float] | None = None,
progress_bar: Any | None = None,
sound_specs: list[Cosmos3SoundSpec] | None = None,
sound_fps_per_item: list[float] | None = None,
action_specs: list[Cosmos3ActionSpec] | None = None,
action_fps_per_item: list[float] | None = None,
) -> torch.Tensor:
"""Run the full UniPC denoise loop, returning the final flat latent.
For each timestep: compute the sequential-CFG velocity, then
``scheduler.step(model_output=v, timestep, sample=latent.unsqueeze(0))``
(the framework steps with a leading batch axis), squeezing back to flat.
For t2vs the flat latent is ``[vision | sound]`` and the velocity covers
both; the scheduler steps the combined vector jointly.
"""
latent = flat_latent
iterator = progress_bar(timesteps) if progress_bar is not None else timesteps
for t in iterator:
v_pred = self.velocity(
flat_latent=latent,
timestep=t.reshape(1),
guidance=guidance,
specs=specs,
cond_token_ids=cond_token_ids,
uncond_token_ids=uncond_token_ids,
fps_per_item=fps_per_item,
sound_specs=sound_specs,
sound_fps_per_item=sound_fps_per_item,
action_specs=action_specs,
action_fps_per_item=action_fps_per_item,
)
stepped = self.scheduler.step(
model_output=v_pred,
timestep=t,
sample=latent.unsqueeze(0),
return_dict=False,
)[0]
latent = stepped.squeeze(0)
return latent
# ===========================================================================
# Pipeline (ComposedPipelineBase)
# ===========================================================================
class Cosmos3OmniDiffusersPipeline(ComposedPipelineBase):
"""Cosmos3 video generation pipeline (T2V / I2V / T2I).
Stage-based ``ComposedPipelineBase`` pipeline. The required modules
(``transformer`` / ``vae`` / ``scheduler`` / ``text_tokenizer``) are loaded
from the ``nvidia/Cosmos3-Nano`` checkpoint by the component loader. The
class name matches the checkpoint ``model_index.json`` ``_class_name`` so
the registry resolves it directly.
The denoise/CFG/VAE math is delegated to module-level helpers
(:func:`cosmos3_get_cfg_velocity`, :class:`Cosmos3DenoiseEngine`,
:func:`cosmos3_vae_encode` / :func:`cosmos3_vae_decode`) which are
framework-parity tested in ``tests/local_tests/cosmos3``.
"""
is_video_pipeline = True
# ``vision_encoder`` / ``sound_tokenizer`` ship in the checkpoint but the
# video path does not need them; they are intentionally omitted here.
_required_config_modules = ["text_tokenizer", "vae", "transformer", "scheduler"]
# Engine-init flow_shift (T2V/I2V); T2I overrides to 3.0 per request.
_engine_init_flow_shift: float = 1.0
# Class-attribute defaults so ``__new__``-based unit tests can read these
# before ``initialize_pipeline`` runs.
scheduler: Any = None
_base_scheduler_config: Any = None
_current_flow_shift: float | None = None
@staticmethod
def _flow_scheduler_config(config: Any) -> dict[str, Any]:
"""Coerce a loaded UniPC config to the framework's flow-matching setup.
The checkpoint ``scheduler_config.json`` carries diffusers-style fields
(``use_karras_sigmas=True``, ``sigma_min``/``sigma_max``, beta schedule)
that do not describe the framework sampler. The framework uses
``FlowUniPCMultistepScheduler`` (pure flow matching: ``shift`` +
``num_train_timesteps`` only). FastVideo's vendored UniPC checks
``use_karras_sigmas`` *before* ``use_flow_sigmas``, so leaving karras on
builds diffusion-style sigmas and the denoise diverges to NaN. Force the
flow config here (parity-verified in ``test_cosmos3_scheduler_parity``).
"""
cfg = dict(config)
cfg.update(
use_karras_sigmas=False,
use_exponential_sigmas=False,
use_beta_sigmas=False,
use_flow_sigmas=True,
prediction_type="flow_prediction",
predict_x0=True,
final_sigmas_type="zero",
)
return cfg
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
"""Bind the loaded scheduler + snapshot its config so per-request
flow_shift rebuilds are cheap and the engine-init shift is applied."""
pipeline_config = fastvideo_args.pipeline_config
engine_shift = getattr(pipeline_config, "flow_shift", None)
if engine_shift is not None:
self._engine_init_flow_shift = float(engine_shift)
scheduler = self.get_module("scheduler")
if scheduler is not None:
# Rebuild from a flow-coerced config so the runtime scheduler matches
# the framework sampler (the loaded checkpoint config is diffusers-style).
flow_config = self._flow_scheduler_config(scheduler.config)
self.scheduler = UniPCMultistepScheduler.from_config(flow_config)
if isinstance(self.modules, dict):
self.modules["scheduler"] = self.scheduler
self._base_scheduler_config = self.scheduler.config
self._current_flow_shift = float(getattr(self.scheduler.config, "flow_shift", 1.0))
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Wire the Cosmos3 stages.
The whole text->latent->denoise->decode flow is custom (sequential CFG
with per-pass repacking), so a single :class:`Cosmos3DenoisingStage`
owns it. ``InputValidationStage`` runs first for the standard checks.
"""
from fastvideo.pipelines.stages import InputValidationStage
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(
stage_name="denoising_stage",
stage=Cosmos3DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
tokenizer=self.get_module("text_tokenizer"),
pipeline=self,
),
)
# -- Scheduler control --------------------------------------------------
def _set_flow_shift(self, target_shift: float) -> None:
"""Set UniPC ``flow_shift`` to ``target_shift``.
Lazily builds a default UniPC scheduler when called before
``initialize_pipeline`` (e.g. the ``__new__``-based scheduler-parity
tests); otherwise rebuilds from the snapshotted base config only when
the target differs from the current shift.
"""
target = float(target_shift)
base_config = self._base_scheduler_config
if base_config is None:
self.scheduler = UniPCMultistepScheduler(
num_train_timesteps=1000,
solver_order=2,
prediction_type="flow_prediction",
use_flow_sigmas=True,
flow_shift=target,
)
self._base_scheduler_config = self.scheduler.config
self._current_flow_shift = target
return
current = self._current_flow_shift
if current is not None and target == float(current):
return
self.scheduler = UniPCMultistepScheduler.from_config(base_config, flow_shift=target)
if isinstance(self.modules, dict):
self.modules["scheduler"] = self.scheduler
self._current_flow_shift = target
# -- Tokenization -------------------------------------------------------
def tokenize_caption(self, caption: str, *, is_video: bool = False, use_system_prompt: bool = False) -> list[int]:
return cosmos3_tokenize_caption(self.get_module("text_tokenizer"),
caption,
is_video=is_video,
use_system_prompt=use_system_prompt)
# Entry point for the pipeline registry. The class name matches the checkpoint
# ``model_index.json`` ``_class_name`` so ``resolve_pipeline_cls`` finds it.
EntryClass = Cosmos3OmniDiffusersPipeline
@@ -0,0 +1,85 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 (Cosmos3-Nano) inference presets.
Defaults track the official ``cosmos-framework`` ``sample_args`` for the video
paths (``text2video`` / ``image2video``: guidance=6.0, num_steps=35, shift=10.0,
fps=24, num_frames=189) and ``text2image`` (guidance=4.0, num_steps=50,
shift=3.0). The default resolution is 16:9 at a VAE-aligned 704x1280 (spatial
compression 16 -> 44x80 latent grid).
"""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="Cosmos3 sequential-CFG UniPC denoising pass",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
# Framework video negative prompt (Cosmos quality prompt).
COSMOS3_VIDEO_NEGATIVE_PROMPT = (
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, "
"fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
"Overall, the video is of poor quality.")
COSMOS3_NANO = InferencePreset(
name="cosmos3_nano",
version=1,
model_family="cosmos3",
description="Cosmos3-Nano text-to-video",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 704,
"width": 1280,
"num_frames": 189,
"fps": 24,
"guidance_scale": 6.0,
"num_inference_steps": 35,
"negative_prompt": COSMOS3_VIDEO_NEGATIVE_PROMPT,
},
)
COSMOS3_NANO_I2V = InferencePreset(
name="cosmos3_nano_i2v",
version=1,
model_family="cosmos3",
description="Cosmos3-Nano image-to-video",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 704,
"width": 1280,
"num_frames": 189,
"fps": 24,
"guidance_scale": 6.0,
"num_inference_steps": 35,
"negative_prompt": COSMOS3_VIDEO_NEGATIVE_PROMPT,
},
)
COSMOS3_NANO_T2I = InferencePreset(
name="cosmos3_nano_t2i",
version=1,
model_family="cosmos3",
description="Cosmos3-Nano text-to-image",
workload_type="t2i",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 1024,
"width": 1024,
"num_frames": 1,
"fps": 24,
"guidance_scale": 4.0,
"num_inference_steps": 50,
"negative_prompt": "",
},
)
ALL_PRESETS = (COSMOS3_NANO, COSMOS3_NANO_I2V, COSMOS3_NANO_T2I)
@@ -0,0 +1,549 @@
# SPDX-License-Identifier: Apache-2.0
"""FastVideo-native Cosmos3 sequence packing (video subset).
Numerical-parity port of the official ``cosmos_framework`` data packer
(``cosmos_framework.data.vfm.sequence_packing.pack_input_sequence``) restricted
to the VIDEO generation path that the FastVideo Cosmos3 DiT consumes (T2V / I2V
/ T2I). It builds, per sample, two splits:
* a ``causal`` text split (prompt token ids, plus the trailing ``eos`` and
``start_of_generation`` markers the framework appends when a generation
modality follows), and
* a ``full`` vision split (VAE latent patch tokens).
The 3D-MRoPE position ids ``[3, seq]`` are produced exactly like the framework:
text tokens broadcast a single monotone id across the (t, h, w) axes, the
temporal offset is bumped by ``temporal_modality_margin`` at the text->vision
boundary, and vision tokens lay out a (T, H, W) grid with spatial ids reset per
segment. Condition frames (I2V cond frame 0, T2I single conditioned frame, ...)
are kept in the packed sequence and rope grid but excluded from the MSE-loss /
timestep bookkeeping, mirroring the framework.
The output ``Cosmos3PackedSequence`` maps 1:1 onto the
``Cosmos3VFMTransformer.forward`` kwargs via :meth:`to_dit_kwargs`. This module
is pure torch/python; it imports no diffusers/transformers model classes.
Reference of record: ``cosmos_framework`` (NVIDIA), the parity oracle used by
``tests/local_tests/cosmos3/test_cosmos3_packing_parity.py``.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from typing import Any
import torch
from fastvideo.models.dits.cosmos3 import (
compute_mrope_position_ids_text,
compute_mrope_position_ids_vision,
)
__all__ = [
"Cosmos3VisionItem",
"Cosmos3SampleInputs",
"Cosmos3PackedSequence",
"pack_cosmos3_video_sequence",
]
# ---------------------------------------------------------------------------
# Inputs
# ---------------------------------------------------------------------------
@dataclass
class Cosmos3VisionItem:
"""One vision latent for a sample.
Args:
latent: VAE latent ``[C, T, H, W]`` (a leading batch axis of size 1 is
accepted and squeezed).
condition_frame_indexes: Latent-frame indices that are *conditioned*
(clean) rather than noisy. ``[]`` for T2V, ``[0]`` for I2V, and the
single conditioned frame for T2I.
fps: Frames-per-second for this clip; only used when
``enable_fps_modulation`` is set.
"""
latent: torch.Tensor
condition_frame_indexes: list[int] = field(default_factory=list)
fps: float | None = None
@dataclass
class Cosmos3SoundItem:
"""One sound latent for a sample (t2vs).
Args:
latent: AVAE sound latent ``[C, T]`` (channels, temporal frames).
condition_frame_indexes: Latent-frame indices that are *conditioned*
(clean). ``[]`` for t2vs (all frames generated).
fps: Sound latent FPS (``sound_latent_fps``, e.g. 25); only used when
``enable_fps_modulation`` is set.
"""
latent: torch.Tensor
condition_frame_indexes: list[int] = field(default_factory=list)
fps: float | None = None
@dataclass
class Cosmos3ActionItem:
"""One action latent for a sample (action-conditioned world model).
Args:
latent: Action latent ``[T, action_dim]`` (per-frame action vectors).
condition_frame_indexes: Frame indices kept clean (conditioning actions).
domain_id: Embodiment domain id (scalar / ``[1]``) for the
domain-aware action projection.
fps: Action FPS; only used when ``enable_fps_modulation`` is set.
"""
latent: torch.Tensor
condition_frame_indexes: list[int] = field(default_factory=list)
domain_id: int = 0
fps: float | None = None
@dataclass
class Cosmos3SampleInputs:
"""Per-sample packing inputs (text prompt + vision item, +sound, +action)."""
text_ids: list[int]
vision: Cosmos3VisionItem
timestep: float
sound: Cosmos3SoundItem | None = None
action: Cosmos3ActionItem | None = None
# ---------------------------------------------------------------------------
# Output
# ---------------------------------------------------------------------------
@dataclass
class Cosmos3PackedSequence:
"""Packed-sequence inputs consumed by ``Cosmos3VFMTransformer.forward``.
Field names mirror the framework ``PackedSequence`` (+ its ``vision``
``ModalityData``) so the parity test can compare field-by-field.
"""
# Sequence structure.
sample_lens: list[int]
split_lens: list[int]
attn_modes: list[str]
sequence_length: int
is_image_batch: bool
# Text modality.
text_ids: torch.Tensor
text_indexes: torch.Tensor
position_ids: torch.Tensor # [3, sequence_length]
# Vision modality.
vision_tokens: list[torch.Tensor]
vision_token_shapes: list[tuple[int, int, int]]
vision_sequence_indexes: torch.Tensor
vision_timesteps: torch.Tensor
vision_mse_loss_indexes: torch.Tensor
vision_noisy_frame_indexes: list[torch.Tensor]
vision_condition_mask: list[torch.Tensor]
fps_vision: torch.Tensor | None = None
# Sound modality (t2vs); empty/None when no sound.
sound_tokens: list[torch.Tensor] = field(default_factory=list)
sound_token_shapes: list[tuple[int, int, int]] = field(default_factory=list)
sound_sequence_indexes: torch.Tensor | None = None
sound_timesteps: torch.Tensor | None = None
sound_mse_loss_indexes: torch.Tensor | None = None
sound_noisy_frame_indexes: list[torch.Tensor] = field(default_factory=list)
sound_condition_mask: list[torch.Tensor] = field(default_factory=list)
fps_sound: torch.Tensor | None = None
# Action modality (action-conditioned world model); empty/None when no action.
action_tokens: list[torch.Tensor] = field(default_factory=list)
action_token_shapes: list[tuple[int, ...]] = field(default_factory=list)
action_sequence_indexes: torch.Tensor | None = None
action_timesteps: torch.Tensor | None = None
action_mse_loss_indexes: torch.Tensor | None = None
action_noisy_frame_indexes: list[torch.Tensor] = field(default_factory=list)
action_condition_mask: list[torch.Tensor] = field(default_factory=list)
action_domain_id: list[torch.Tensor] = field(default_factory=list)
def to_dit_kwargs(self, device: torch.device | str | None = None) -> dict[str, Any]:
"""Return the kwargs dict for ``Cosmos3VFMTransformer.forward``.
Packing is device-agnostic (ids/indexes/position-ids are built on CPU).
When ``device`` is given, every tensor input is moved to it so the DiT
forward runs on a single device (e.g. the model's GPU at inference).
"""
def _mv(x: Any) -> Any:
return x.to(device) if (device is not None and torch.is_tensor(x)) else x
return dict(
text_ids=_mv(self.text_ids),
text_indexes=_mv(self.text_indexes),
position_ids=_mv(self.position_ids),
sequence_length=int(self.sequence_length),
split_lens=list(self.split_lens),
attn_modes=list(self.attn_modes),
vision_tokens=[_mv(t) for t in self.vision_tokens],
vision_token_shapes=list(self.vision_token_shapes),
vision_sequence_indexes=_mv(self.vision_sequence_indexes),
vision_timesteps=_mv(self.vision_timesteps),
vision_mse_loss_indexes=_mv(self.vision_mse_loss_indexes),
vision_noisy_frame_indexes=[_mv(t) for t in self.vision_noisy_frame_indexes],
fps_vision=self.fps_vision,
sound_tokens=[_mv(t) for t in self.sound_tokens],
sound_token_shapes=list(self.sound_token_shapes),
sound_sequence_indexes=_mv(self.sound_sequence_indexes),
sound_timesteps=_mv(self.sound_timesteps),
sound_mse_loss_indexes=_mv(self.sound_mse_loss_indexes),
sound_noisy_frame_indexes=[_mv(t) for t in self.sound_noisy_frame_indexes],
fps_sound=_mv(self.fps_sound),
action_tokens=[_mv(t) for t in self.action_tokens],
action_token_shapes=list(self.action_token_shapes),
action_sequence_indexes=_mv(self.action_sequence_indexes),
action_timesteps=_mv(self.action_timesteps),
action_mse_loss_indexes=_mv(self.action_mse_loss_indexes),
action_noisy_frame_indexes=[_mv(t) for t in self.action_noisy_frame_indexes],
action_domain_id=[_mv(t) for t in self.action_domain_id],
)
# ---------------------------------------------------------------------------
# Packing
# ---------------------------------------------------------------------------
def pack_cosmos3_video_sequence(
samples: list[Cosmos3SampleInputs],
special_tokens: dict[str, int],
*,
latent_patch_size: int = 2,
include_end_of_generation_token: bool = False,
temporal_modality_margin: int = 15_000,
reset_spatial_ids: bool = True,
enable_fps_modulation: bool = False,
base_fps: float = 24.0,
temporal_compression_factor: int = 4,
initial_mrope_temporal_offset: int | float = 0,
) -> Cosmos3PackedSequence:
"""Pack prompts + vision latents into the Cosmos3 DiT packed-sequence inputs.
Video subset of ``cosmos_framework`` ``pack_input_sequence`` under
``unified_3d_mrope``: each sample is ``[causal text, full vision]``.
Args:
samples: Per-sample text prompt token ids + vision item + timestep.
special_tokens: Must contain ``eos_token_id`` and
``start_of_generation`` (and ``end_of_generation`` if
``include_end_of_generation_token``). ``bos_token_id`` is honored if
present (prepended) to match the framework.
latent_patch_size: Latent patch size used by the DiT.
include_end_of_generation_token: Append the framework's end-of-generation
marker after the vision split.
temporal_modality_margin: Temporal-offset bump applied at the
text->vision boundary (``unified_3d_mrope_temporal_modality_margin``).
reset_spatial_ids: Reset vision spatial ids to 0 per segment.
enable_fps_modulation: Use float, fps-scaled temporal positions.
base_fps: Base FPS used when ``enable_fps_modulation``.
temporal_compression_factor: VAE temporal compression factor.
initial_mrope_temporal_offset: Per-sample starting temporal offset.
Returns:
A :class:`Cosmos3PackedSequence`.
"""
assert "eos_token_id" in special_tokens, "special_tokens must contain eos_token_id"
assert "start_of_generation" in special_tokens, "special_tokens must contain start_of_generation"
if latent_patch_size < 1:
raise ValueError(f"latent_patch_size must be >= 1, got {latent_patch_size}")
# Build-time accumulators (concatenated across samples).
sample_lens: list[int] = []
split_lens: list[int] = []
attn_modes: list[str] = []
text_ids: list[int] = []
text_indexes: list[int] = []
position_id_blocks: list[torch.Tensor] = [] # each [3, n]
vision_tokens: list[torch.Tensor] = []
vision_token_shapes: list[tuple[int, int, int]] = []
vision_sequence_indexes: list[int] = []
vision_timesteps: list[float] = []
vision_mse_loss_indexes: list[int] = []
vision_noisy_frame_indexes: list[torch.Tensor] = []
vision_condition_mask: list[torch.Tensor] = []
fps_values: list[float] = []
sound_tokens: list[torch.Tensor] = []
sound_token_shapes: list[tuple[int, int, int]] = []
sound_sequence_indexes: list[int] = []
sound_timesteps: list[float] = []
sound_mse_loss_indexes: list[int] = []
sound_noisy_frame_indexes: list[torch.Tensor] = []
sound_condition_mask: list[torch.Tensor] = []
sound_fps_values: list[float] = []
action_tokens: list[torch.Tensor] = []
action_token_shapes: list[tuple[int, ...]] = []
action_sequence_indexes: list[int] = []
action_timesteps: list[float] = []
action_mse_loss_indexes: list[int] = []
action_noisy_frame_indexes: list[torch.Tensor] = []
action_condition_mask: list[torch.Tensor] = []
action_domain_id: list[torch.Tensor] = []
curr = 0 # running position in the packed sequence
is_image_batch = True
for sample in samples:
temporal_offset: int | float = initial_mrope_temporal_offset
sample_len = 0
# ---- 1. Text split (causal) ----
if "bos_token_id" in special_tokens:
shifted_text_ids = [special_tokens["bos_token_id"], *sample.text_ids]
else:
shifted_text_ids = list(sample.text_ids)
# The video path always has a following generation modality, so the
# framework appends eos + start_of_generation.
shifted_text_ids = [*shifted_text_ids, special_tokens["eos_token_id"], special_tokens["start_of_generation"]]
text_split_len = len(shifted_text_ids)
text_ids.extend(shifted_text_ids)
text_indexes.extend(range(curr, curr + text_split_len))
text_mrope, temporal_offset = compute_mrope_position_ids_text(
num_tokens=text_split_len,
temporal_offset=int(temporal_offset),
)
position_id_blocks.append(text_mrope)
attn_modes.append("causal")
split_lens.append(text_split_len)
curr += text_split_len
sample_len += text_split_len
# End of text modality: bump temporal offset before vision.
temporal_offset += temporal_modality_margin
# Sound shares the vision temporal start (parallel temporal positions).
vision_start_temporal_offset = temporal_offset
# ---- 2. Vision split (full) ----
latent = sample.vision.latent
latent = latent.squeeze(0) if latent.dim() == 5 else latent # [C, T, H, W]
_c, latent_t, latent_h, latent_w = latent.shape
patch_h = math.ceil(latent_h / latent_patch_size)
patch_w = math.ceil(latent_w / latent_patch_size)
num_vision_tokens = latent_t * patch_h * patch_w
vision_tokens.append(sample.vision.latent)
vision_token_shapes.append((latent_t, patch_h, patch_w))
vision_sequence_indexes.extend(range(curr, curr + num_vision_tokens))
condition_set = {idx for idx in sample.vision.condition_frame_indexes if 0 <= idx < latent_t}
cond_mask = torch.zeros((latent_t, 1, 1), device=latent.device, dtype=latent.dtype)
for frame_idx in condition_set:
cond_mask[frame_idx, 0, 0] = 1.0
vision_condition_mask.append(cond_mask)
noisy_frames = torch.tensor(
[idx for idx in range(latent_t) if idx not in condition_set],
device=latent.device,
dtype=torch.long,
)
vision_noisy_frame_indexes.append(noisy_frames)
# MSE-loss indices + per-token timesteps cover only the noisy frames.
frame_token_stride = patch_h * patch_w
for frame_idx in range(latent_t):
if frame_idx in condition_set:
continue
frame_start = curr + frame_idx * frame_token_stride
vision_mse_loss_indexes.extend(range(frame_start, frame_start + frame_token_stride))
vision_timesteps.extend([float(sample.timestep)] * frame_token_stride)
vision_fps = sample.vision.fps if enable_fps_modulation else None
if vision_fps is not None:
fps_values.append(float(vision_fps))
vision_mrope, temporal_offset = compute_mrope_position_ids_vision(
grid_t=latent_t,
grid_h=patch_h,
grid_w=patch_w,
temporal_offset=temporal_offset,
fps=vision_fps,
base_fps=base_fps,
temporal_compression_factor=temporal_compression_factor,
enable_fps_modulation=enable_fps_modulation,
)
position_id_blocks.append(vision_mrope)
curr += num_vision_tokens
sample_len += num_vision_tokens
# ---- 2a2. Action split: shares the vision "full" split ----
# Mirrors framework ``_pack_action_tokens``: action latent [T, D] -> T
# tokens (token shape (T,)), domain-aware, 3D-MRoPE at the vision temporal
# offset with ``start_frame_offset=1`` (parallel to vision; tcf=1; does
# not advance the offset).
action_split_len = 0
if sample.action is not None:
action_latent = sample.action.latent # [T, D]
action_t = int(action_latent.shape[0])
action_split_len = action_t
action_tokens.append(action_latent)
action_token_shapes.append((action_t, ))
action_sequence_indexes.extend(range(curr, curr + action_t))
action_domain_id.append(torch.tensor([int(sample.action.domain_id)], dtype=torch.long))
a_cond_set = {idx for idx in sample.action.condition_frame_indexes if 0 <= idx < action_t}
a_cond_mask = torch.zeros((action_t, 1), device=action_latent.device, dtype=action_latent.dtype)
for fi in a_cond_set:
a_cond_mask[fi, 0] = 1.0
action_condition_mask.append(a_cond_mask)
a_noisy = torch.tensor([idx for idx in range(action_t) if idx not in a_cond_set],
device=action_latent.device,
dtype=torch.long)
action_noisy_frame_indexes.append(a_noisy)
for fi in range(action_t):
if fi in a_cond_set:
continue
action_mse_loss_indexes.append(curr + fi)
action_timesteps.append(float(sample.timestep))
action_fps = sample.action.fps if enable_fps_modulation else None
action_mrope, _ = compute_mrope_position_ids_vision(
grid_t=action_t,
grid_h=1,
grid_w=1,
temporal_offset=vision_start_temporal_offset,
fps=action_fps,
base_fps=base_fps,
temporal_compression_factor=1, # action is at frame rate
base_temporal_compression_factor=temporal_compression_factor,
enable_fps_modulation=enable_fps_modulation,
start_frame_offset=1,
)
position_id_blocks.append(action_mrope)
curr += action_t
sample_len += action_t
# ---- 2b. Sound split (t2vs): shares the vision "full" split ----
# Mirrors framework ``_pack_sound_tokens``: sound latent [C, T] -> T
# tokens (token shape (T,1,1)), packed right after vision, with 3D-MRoPE
# temporal positions starting at the vision temporal offset (parallel to
# vision, start_frame_offset=0, tcf=1) and NOT advancing it.
sound_split_len = 0
if sample.sound is not None:
sound_latent = sample.sound.latent
sound_latent = sound_latent.squeeze(0) if sound_latent.dim() == 3 else sound_latent # [C, T]
_sc, sound_t = sound_latent.shape
sound_split_len = sound_t
sound_tokens.append(sound_latent)
sound_token_shapes.append((sound_t, 1, 1))
sound_sequence_indexes.extend(range(curr, curr + sound_t))
s_cond_set = {idx for idx in sample.sound.condition_frame_indexes if 0 <= idx < sound_t}
s_cond_mask = torch.zeros((sound_t, 1), device=sound_latent.device, dtype=sound_latent.dtype)
for fi in s_cond_set:
s_cond_mask[fi, 0] = 1.0
sound_condition_mask.append(s_cond_mask)
s_noisy = torch.tensor([idx for idx in range(sound_t) if idx not in s_cond_set],
device=sound_latent.device,
dtype=torch.long)
sound_noisy_frame_indexes.append(s_noisy)
for fi in range(sound_t):
if fi in s_cond_set:
continue
sound_mse_loss_indexes.append(curr + fi) # 1 token per sound frame
sound_timesteps.append(float(sample.timestep))
sound_fps = sample.sound.fps if enable_fps_modulation else None
if sound_fps is not None:
sound_fps_values.append(float(sound_fps))
sound_mrope, _ = compute_mrope_position_ids_vision(
grid_t=sound_t,
grid_h=1,
grid_w=1,
temporal_offset=vision_start_temporal_offset,
fps=sound_fps,
base_fps=base_fps,
temporal_compression_factor=1, # sound latent already at sound_latent_fps
enable_fps_modulation=enable_fps_modulation,
start_frame_offset=0,
)
position_id_blocks.append(sound_mrope)
curr += sound_t
sample_len += sound_t
# ---- 3. Optional end-of-generation marker ----
eov_len = 0
if include_end_of_generation_token:
assert "end_of_generation" in special_tokens, ("special_tokens must contain end_of_generation when "
"include_end_of_generation_token=True")
text_ids.append(special_tokens["end_of_generation"])
text_indexes.append(curr)
eov_dtype = torch.float32 if enable_fps_modulation else torch.long
eov_ids = torch.full((3, 1), temporal_offset, dtype=eov_dtype)
position_id_blocks.append(eov_ids)
temporal_offset += 1
curr += 1
eov_len = 1
sample_len += 1
# Vision + action + sound + any trailing eov marker share one "full" split.
attn_modes.append("full")
split_lens.append(num_vision_tokens + action_split_len + sound_split_len + eov_len)
sample_lens.append(sample_len)
if latent_t != 1:
is_image_batch = False
sequence_length = sum(sample_lens)
# position_ids: float iff any block is float (fps modulation path).
any_float = any(b.dtype.is_floating_point for b in position_id_blocks)
if any_float:
position_id_blocks = [b.to(torch.float32) for b in position_id_blocks]
position_ids = torch.cat(position_id_blocks, dim=1) # [3, sequence_length]
timesteps_dtype = torch.float32
return Cosmos3PackedSequence(
sample_lens=sample_lens,
split_lens=split_lens,
attn_modes=attn_modes,
sequence_length=sequence_length,
is_image_batch=is_image_batch,
text_ids=torch.tensor(text_ids, dtype=torch.long),
text_indexes=torch.tensor(text_indexes, dtype=torch.long),
position_ids=position_ids,
vision_tokens=vision_tokens,
vision_token_shapes=vision_token_shapes,
vision_sequence_indexes=torch.tensor(vision_sequence_indexes, dtype=torch.long),
vision_timesteps=torch.tensor(vision_timesteps, dtype=timesteps_dtype),
vision_mse_loss_indexes=torch.tensor(vision_mse_loss_indexes, dtype=torch.long),
vision_noisy_frame_indexes=vision_noisy_frame_indexes,
vision_condition_mask=vision_condition_mask,
fps_vision=(torch.tensor(fps_values, dtype=torch.float32) if fps_values else None),
sound_tokens=sound_tokens,
sound_token_shapes=sound_token_shapes,
sound_sequence_indexes=(torch.tensor(sound_sequence_indexes, dtype=torch.long) if sound_tokens else None),
sound_timesteps=(torch.tensor(sound_timesteps, dtype=timesteps_dtype) if sound_tokens else None),
sound_mse_loss_indexes=(torch.tensor(sound_mse_loss_indexes, dtype=torch.long) if sound_tokens else None),
sound_noisy_frame_indexes=sound_noisy_frame_indexes,
sound_condition_mask=sound_condition_mask,
fps_sound=(torch.tensor(sound_fps_values, dtype=torch.float32) if sound_fps_values else None),
action_tokens=action_tokens,
action_token_shapes=action_token_shapes,
action_sequence_indexes=(torch.tensor(action_sequence_indexes, dtype=torch.long) if action_tokens else None),
action_timesteps=(torch.tensor(action_timesteps, dtype=timesteps_dtype) if action_tokens else None),
action_mse_loss_indexes=(torch.tensor(action_mse_loss_indexes, dtype=torch.long) if action_tokens else None),
action_noisy_frame_indexes=action_noisy_frame_indexes,
action_condition_mask=action_condition_mask,
action_domain_id=action_domain_id,
)
@@ -19,7 +19,7 @@ class GlmImageDecodingStage(DecodingStage):
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
vae_autocast = (vae_dtype != torch.float32 and not fastvideo_args.disable_autocast)
latents = self._denormalize_latents(latents)
latents = self._denormalize_latents(latents, fastvideo_args)
if latents.dim() == 5:
latents = latents.squeeze(2)
@@ -0,0 +1,5 @@
"""Dense LingBot-Video inference pipeline."""
from fastvideo.pipelines.basic.lingbot_video.lingbot_video_pipeline import LingBotVideoPipeline
__all__ = ["LingBotVideoPipeline"]
@@ -0,0 +1,106 @@
# SPDX-License-Identifier: Apache-2.0
"""Stage-composed LingBot-Video Dense and MoE/refiner T2V pipeline."""
from typing import Any
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.basic.lingbot_video.stages import (
LingBotVideoDenoisingStage,
LingBotVideoInputValidationStage,
LingBotVideoLatentPreparationStage,
LingBotVideoRefinerPreparationStage,
)
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
TextEncodingStage,
TimestepPreparationStage,
)
class LingBotVideoPipeline(LoRAPipeline, ComposedPipelineBase):
"""T2V pipeline with optional released MoE pixel-space refinement."""
is_video_pipeline = True
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
def load_modules(
self,
fastvideo_args: FastVideoArgs,
loaded_modules: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Load the optional refiner DiT and the VAE encoder only when declared."""
model_index = self._load_config(self.model_path)
required = list(type(self)._required_config_modules)
load_refiner = "transformer_2" in model_index and getattr(fastvideo_args, "refine_enabled", None) is not False
if load_refiner:
required.append("transformer_2")
fastvideo_args.pipeline_config.vae_config.load_encoder = True
self._required_config_modules = required
return super().load_modules(fastvideo_args, loaded_modules)
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
"""Apply the released runtime flow shift to the loaded scheduler."""
shift = fastvideo_args.pipeline_config.flow_shift
if shift is None:
raise ValueError("LingBot-Video requires a flow shift")
self.get_module("scheduler").set_shift(float(shift))
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Create base generation and the optional decoded-video refiner stages."""
refiner = self.get_module("transformer_2")
self.add_stage(
"input_validation_stage",
LingBotVideoInputValidationStage(refiner_enabled=refiner is not None),
)
self.add_stage(
"prompt_encoding_stage",
TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage("conditioning_stage", ConditioningStage())
self.add_stage(
"timestep_preparation_stage",
TimestepPreparationStage(scheduler=self.get_module("scheduler")),
)
self.add_stage(
"latent_preparation_stage",
LingBotVideoLatentPreparationStage(transformer=self.get_module("transformer")),
)
self.add_stage(
"denoising_stage",
LingBotVideoDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
"decoding_stage",
DecodingStage(vae=self.get_module("vae"), pipeline=self),
)
if refiner is not None:
self.add_stage(
"refiner_preparation_stage",
LingBotVideoRefinerPreparationStage(
vae=self.get_module("vae"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
"refiner_denoising_stage",
LingBotVideoDenoisingStage(
transformer=refiner,
scheduler=self.get_module("scheduler"),
refiner=True,
),
)
self.add_stage(
"refiner_decoding_stage",
DecodingStage(vae=self.get_module("vae"), pipeline=self),
)
EntryClass = LingBotVideoPipeline
@@ -0,0 +1,92 @@
# SPDX-License-Identifier: Apache-2.0
"""Official LingBot-Video T2V inference presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
DEFAULT_NEGATIVE_PROMPT = ('{"universal_negative": {"visual_quality": ["low quality", "worst quality", "blurry", '
'"pixelated", "jpeg artifacts", "low resolution", "unstable color", "color flicker", '
'"underexposed", "overexposed", "invisible subject", "subject hidden in darkness"], '
'"artistic_style": ["painting", "illustration", "drawing", "cartoon", "3d render", '
'"cgi", "sketch", "digital art"], "composition_and_content": ["text", "watermark", '
'"signature", "logo", "subtitles", "pillarboxed", "side bars", "portrait image in '
'landscape frame"], "temporal_and_motion_stability": ["flickering", "jittery", '
'"motion blur", "temporal inconsistency", "warping", "morphing", "incoherent motion", '
'"unnatural movement", "static object with sudden jump", "frame-to-frame inconsistency"], '
'"material_and_structure": ["plastic-like glass", "unrealistic texture", "deformed '
'bottle", "liquid freezing improperly", "distorted reflections"]}}')
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="LingBot-Video batched-CFG denoising",
allowed_overrides=frozenset({"num_inference_steps", "guidance_scale"}),
)
_REFINE_STAGE = PresetStageSpec(
name="refine",
kind="refinement",
description="LingBot-Video pixel-space resize, VAE re-encode, and refiner denoising",
allowed_overrides=frozenset({
"height_sr",
"width_sr",
"num_inference_steps_sr",
"guidance_scale_2",
"t_thresh",
}),
)
LINGBOT_VIDEO_DENSE_T2V = InferencePreset(
name="lingbot_video_dense_t2v",
version=1,
model_family="lingbot_video",
description="LingBot-Video Dense 1.3B text-to-video at 480p",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 480,
"width": 832,
"num_frames": 121,
"fps": 24,
"num_inference_steps": 40,
"guidance_scale": 3.0,
"batch_cfg": True,
"seed": 42,
"negative_prompt": DEFAULT_NEGATIVE_PROMPT,
},
)
LINGBOT_VIDEO_MOE_REFINER_T2V = InferencePreset(
name="lingbot_video_moe_refiner_t2v",
version=1,
model_family="lingbot_video",
description="LingBot-Video MoE 30B-A3B T2V with 1080p refiner",
workload_type="t2v",
stage_schemas=(_DENOISE_STAGE, _REFINE_STAGE),
defaults={
"height": 480,
"width": 832,
"height_sr": 1088,
"width_sr": 1920,
"num_frames": 121,
"fps": 24,
"num_inference_steps": 40,
"num_inference_steps_sr": 8,
"guidance_scale": 3.0,
"guidance_scale_2": 3.0,
"batch_cfg": True,
"t_thresh": 0.85,
"seed": 42,
"negative_prompt": DEFAULT_NEGATIVE_PROMPT,
},
stage_defaults={
"refine": {
"height_sr": 1088,
"width_sr": 1920,
"num_inference_steps_sr": 8,
"guidance_scale_2": 3.0,
"t_thresh": 0.85,
},
},
)
ALL_PRESETS = (LINGBOT_VIDEO_DENSE_T2V, LINGBOT_VIDEO_MOE_REFINER_T2V)
@@ -0,0 +1,345 @@
# SPDX-License-Identifier: Apache-2.0
"""LingBot-Video stages whose contracts differ from shared Wan behavior."""
from __future__ import annotations
import numpy as np
import torch
import torch.nn.functional as F
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.input_validation import InputValidationStage
LINGBOT_VIDEO_REFINER_TAIL_STEPS = 2
def _compute_refiner_sigmas(
sigma_max: float,
sigma_min: float,
num_inference_steps: int,
shift: float,
t_thresh: float,
) -> np.ndarray:
"""Build the released truncated schedule plus its two-step low-noise tail."""
if not 0.0 < t_thresh <= 1.0:
raise ValueError(f"LingBot-Video refiner t_thresh must be in (0, 1], got {t_thresh}")
if num_inference_steps < 1:
raise ValueError("LingBot-Video refiner requires at least one inference step")
base = np.linspace(sigma_max, sigma_min, num_inference_steps + 1).copy()[:-1]
shifted = shift * base / (1.0 + (shift - 1.0) * base)
sigmas = shifted[shifted <= t_thresh + 1e-6]
if sigmas.size == 0 or abs(float(sigmas[0]) - t_thresh) > 1e-6:
sigmas = np.concatenate(([t_thresh], sigmas))
tail = np.linspace(
float(sigmas[-1]),
min(sigma_min, float(sigmas[-1])),
LINGBOT_VIDEO_REFINER_TAIL_STEPS + 2,
)[1:-1]
sigmas = np.concatenate((sigmas, tail))
if sigmas.size > 1 and not np.all(np.diff(sigmas) < 0.0):
raise ValueError(f"LingBot-Video refiner sigmas must descend strictly, got {sigmas.tolist()}")
return sigmas.astype(np.float32)
class LingBotVideoInputValidationStage(InputValidationStage):
"""Validate released shape constraints and construct the official CUDA RNG."""
def __init__(self, refiner_enabled: bool = False) -> None:
self.refiner_enabled = refiner_enabled
def _generate_seeds(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> None:
"""Use one device-local generator, matching the official production runner."""
del fastvideo_args
if batch.seed is None:
raise ValueError("LingBot-Video requires a seed")
if batch.num_videos_per_prompt != 1:
raise ValueError("LingBot-Video currently supports one video per prompt")
batch.seeds = [batch.seed]
batch.generator = torch.Generator(device=get_local_torch_device()).manual_seed(batch.seed)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Run shared validation, then enforce LingBot temporal and spatial geometry."""
batch = super().forward(batch, fastvideo_args)
if not isinstance(batch.num_frames, int):
raise TypeError("LingBot-Video num_frames must be an integer")
if batch.num_frames != 1 and (batch.num_frames - 1) % 4 != 0:
raise ValueError(f"num_frames must be 1 or 4n+1, got {batch.num_frames}")
if not isinstance(batch.height, int) or not isinstance(batch.width, int):
raise TypeError("LingBot-Video height and width must be integers")
if batch.height % 16 != 0 or batch.width % 16 != 0:
raise ValueError(f"height and width must be divisible by 16, got {batch.height}x{batch.width}")
if isinstance(batch.prompt, list) and len(batch.prompt) != 1:
raise ValueError("LingBot-Video currently supports prompt batch size one")
if self.refiner_enabled and fastvideo_args.output_type == "latent":
raise ValueError("LingBot-Video refinement requires decoded pixel output")
return batch
class LingBotVideoLatentPreparationStage(PipelineStage):
"""Prepare fp32 latents in the released 4x temporal and 8x spatial geometry."""
def __init__(self, transformer) -> None:
self.transformer = transformer
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Generate or validate one normalized fp32 latent video."""
del fastvideo_args
if not all(isinstance(value, int) for value in (batch.num_frames, batch.height, batch.width)):
raise TypeError("latent geometry must contain integer frames, height, and width")
shape = (
1,
self.transformer.num_channels_latents,
(batch.num_frames - 1) // 4 + 1,
batch.height // 8,
batch.width // 8,
)
device = get_local_torch_device()
if batch.latents is None:
batch.latents = torch.randn(
shape,
generator=batch.generator,
device=device,
dtype=torch.float32,
)
else:
if tuple(batch.latents.shape) != shape:
raise ValueError(f"supplied latent shape {tuple(batch.latents.shape)} does not match {shape}")
batch.latents = batch.latents.to(device=device, dtype=torch.float32)
batch.raw_latent_shape = shape
return batch
class LingBotVideoRefinerPreparationStage(PipelineStage):
"""Resize and encode the base video, then initialize the released refiner state."""
performance_component_metric = "vae_encode_time_s"
def __init__(self, vae, scheduler) -> None:
self.vae = vae
self.scheduler = scheduler
@staticmethod
def _resize_video(video: torch.Tensor, height: int, width: int) -> torch.Tensor:
"""Bicubic-resize every decoded frame using the released tensor layout."""
batch, channels, frames, source_height, source_width = video.shape
flat = video.permute(0, 2, 1, 3, 4).reshape(batch * frames, channels, source_height, source_width)
resized = F.interpolate(flat, size=(height, width), mode="bicubic", align_corners=False).clamp(0.0, 1.0)
return resized.reshape(batch, frames, channels, height, width).permute(0, 2, 1, 3, 4).contiguous()
def _encode_video(
self,
video: torch.Tensor,
generator: torch.Generator,
device: torch.device,
) -> torch.Tensor:
"""Encode `[0,1]` pixels and convert Wan VAE latents to normalized DiT space."""
video = video.to(device=device, dtype=torch.float32).mul(2.0).sub(1.0)
with torch.autocast(device_type=device.type, dtype=torch.bfloat16, enabled=device.type == "cuda"):
encoded = self.vae.encode(video)
if hasattr(encoded, "latent_dist"):
latents = encoded.latent_dist.sample(generator)
elif hasattr(encoded, "sample") and callable(encoded.sample):
latents = encoded.sample(generator)
elif isinstance(encoded, tuple | list):
latents = encoded[0]
else:
latents = encoded
mean = torch.tensor(self.vae.config.latents_mean, device=device, dtype=torch.float32).view(1, -1, 1, 1, 1)
std = torch.tensor(self.vae.config.latents_std, device=device, dtype=torch.float32).view(1, -1, 1, 1, 1)
return ((latents.float() - mean) / std).to(latents)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Prepare high-resolution refiner latents and its exact truncated sigma schedule."""
if batch.output is None or batch.output.ndim != 5:
raise ValueError("LingBot-Video refinement requires a decoded base video")
if not isinstance(batch.height_sr, int) or not isinstance(batch.width_sr, int):
raise TypeError("LingBot-Video refinement requires integer height_sr and width_sr")
if batch.height_sr % 16 != 0 or batch.width_sr % 16 != 0:
raise ValueError("LingBot-Video refiner height_sr and width_sr must be divisible by 16")
if batch.seed is None:
raise ValueError("LingBot-Video refinement requires a seed")
device = get_local_torch_device()
if isinstance(self.vae, torch.nn.Module):
self.vae.to(device)
generator = torch.Generator(device=device).manual_seed(batch.seed)
resized = self._resize_video(batch.output, batch.height_sr, batch.width_sr)
encoded = self._encode_video(resized, generator, device)
noise = torch.randn(encoded.shape, generator=generator, device=device, dtype=encoded.dtype)
batch.latents = ((1.0 - batch.t_thresh) * encoded + batch.t_thresh * noise).float()
batch.generator = generator
batch.height = batch.height_sr
batch.width = batch.width_sr
batch.raw_latent_shape = tuple(batch.latents.shape)
batch.extra["lingbot_video_base_shape"] = tuple(batch.output.shape)
batch.output = None
shift = fastvideo_args.pipeline_config.flow_shift
if shift is None:
raise ValueError("LingBot-Video refinement requires a flow shift")
sigmas = _compute_refiner_sigmas(
float(self.scheduler.sigma_max),
float(self.scheduler.sigma_min),
batch.num_inference_steps_sr,
float(shift),
float(batch.t_thresh),
)
self.scheduler.set_timesteps(len(sigmas), device=device, sigmas=sigmas, shift=1.0)
batch.timesteps = self.scheduler.timesteps
if getattr(fastvideo_args, "vae_cpu_offload", False):
self.vae.to("cpu")
return batch
class LingBotVideoDenoisingStage(PipelineStage):
"""Run the released batched-CFG bf16 DiT loop with fp32 scheduler state."""
performance_component_metric = "dit_time_s"
def __init__(self, transformer, scheduler, refiner: bool = False) -> None:
self.transformer = transformer
self.scheduler = scheduler
self.refiner = refiner
@staticmethod
def _pad_condition(
embeds: torch.Tensor,
mask: torch.Tensor,
length: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Right-pad one condition stream to a shared batched-CFG length."""
pad_length = length - embeds.shape[1]
if pad_length < 0 or embeds.shape[:2] != mask.shape:
raise ValueError("invalid LingBot-Video prompt embedding/mask shapes")
if pad_length == 0:
return embeds, mask
embed_padding = embeds.new_zeros(embeds.shape[0], pad_length, embeds.shape[2])
mask_padding = mask.new_zeros(mask.shape[0], pad_length)
return (
torch.cat((embeds, embed_padding), dim=1),
torch.cat((mask, mask_padding), dim=1),
)
@staticmethod
def _transformer_timestep(timestep: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
"""Reproduce the official divide-cast-multiply timestep rounding."""
sigma = timestep.float() / 1000.0
if dtype in (torch.bfloat16, torch.float16):
sigma = sigma.to(dtype)
return (sigma * 1000.0).float()
def _prepare_conditions(
self,
batch: ForwardBatch,
dtype: torch.dtype,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Pack conditional then unconditional text streams for one batched CFG call."""
prompt = batch.prompt_embeds[0].to(device=device, dtype=dtype)
if batch.prompt_attention_mask is None or not batch.prompt_attention_mask:
raise ValueError("LingBot-Video requires a prompt attention mask")
prompt_mask = batch.prompt_attention_mask[0].to(device=device)
if not self._uses_cfg(batch) or not batch.batch_cfg:
return prompt, prompt_mask
negative, negative_mask = self._negative_condition(batch, prompt, prompt_mask, dtype, device)
target_length = max(prompt.shape[1], negative.shape[1])
prompt, prompt_mask = self._pad_condition(prompt, prompt_mask, target_length)
negative, negative_mask = self._pad_condition(negative, negative_mask, target_length)
return (
torch.cat((prompt, negative), dim=0),
torch.cat((prompt_mask, negative_mask), dim=0),
)
def _negative_condition(
self,
batch: ForwardBatch,
prompt: torch.Tensor,
prompt_mask: torch.Tensor,
dtype: torch.dtype,
device: torch.device,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Use the refiner's zero-cloned null condition or the encoded negative prompt."""
if self.refiner:
return torch.zeros_like(prompt), prompt_mask.clone()
if batch.negative_prompt_embeds is None or not batch.negative_prompt_embeds:
raise ValueError("LingBot-Video CFG requires negative prompt embeddings")
if batch.negative_attention_mask is None or not batch.negative_attention_mask:
raise ValueError("LingBot-Video CFG requires a negative prompt mask")
return (
batch.negative_prompt_embeds[0].to(device=device, dtype=dtype),
batch.negative_attention_mask[0].to(device=device),
)
def _uses_cfg(self, batch: ForwardBatch) -> bool:
"""Enable guidance independently for the base or refiner scale."""
scale = batch.guidance_scale_2 if self.refiner else batch.guidance_scale
return scale is not None and scale > 1.0
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Denoise latents while keeping scheduler samples and predictions in fp32."""
if batch.latents is None or batch.timesteps is None:
raise ValueError("LingBot-Video denoising requires latents and timesteps")
device = get_local_torch_device()
transformer_dtype = next(self.transformer.parameters()).dtype
condition, condition_mask = self._prepare_conditions(batch, transformer_dtype, device)
latents = batch.latents.to(device=device, dtype=torch.float32)
do_cfg = self._uses_cfg(batch)
negative = negative_mask = None
if do_cfg and not batch.batch_cfg:
negative, negative_mask = self._negative_condition(batch, condition, condition_mask, transformer_dtype,
device)
trajectory: list[torch.Tensor] = []
trajectory_timesteps: list[torch.Tensor] = []
for timestep in batch.timesteps:
timestep_batch = self._transformer_timestep(timestep, transformer_dtype).expand(1).to(device)
latent_input = latents
if do_cfg and batch.batch_cfg:
latent_input = torch.cat((latents, latents), dim=0)
timestep_batch = torch.cat((timestep_batch, timestep_batch), dim=0)
autocast_enabled = device.type == "cuda" and transformer_dtype != torch.float32
with torch.autocast(device_type=device.type, dtype=transformer_dtype, enabled=autocast_enabled):
prediction = self.transformer(
latent_input,
timestep_batch,
condition,
encoder_attention_mask=condition_mask,
return_dict=False,
)[0].float()
if do_cfg:
if batch.batch_cfg:
conditional, unconditional = prediction.chunk(2, dim=0)
else:
with torch.autocast(
device_type=device.type,
dtype=transformer_dtype,
enabled=autocast_enabled,
):
unconditional = self.transformer(
latents,
timestep_batch,
negative,
encoder_attention_mask=negative_mask,
return_dict=False,
)[0].float()
conditional = prediction
guidance_scale = batch.guidance_scale_2 if self.refiner else batch.guidance_scale
if guidance_scale is None:
raise ValueError("LingBot-Video CFG requires a guidance scale")
prediction = unconditional + guidance_scale * (conditional - unconditional)
latents = self.scheduler.step(
prediction,
timestep,
latents,
return_dict=False,
generator=batch.generator,
)[0].float()
if batch.return_trajectory_latents:
trajectory.append(latents.detach().cpu())
trajectory_timesteps.append(timestep.detach().cpu())
batch.latents = latents
if trajectory:
batch.trajectory_latents = torch.stack(trajectory, dim=1)
batch.trajectory_timesteps = trajectory_timesteps
return batch
@@ -0,0 +1,5 @@
from .causal_fast_pipeline import LingBotWorld2CausalFastPipeline
__all__ = ["LingBotWorld2CausalFastPipeline"]
EntryClass = LingBotWorld2CausalFastPipeline
@@ -0,0 +1,365 @@
# SPDX-License-Identifier: Apache-2.0
"""LingBot World 2 causal-fast image-to-video pipeline."""
import math
import os
from einops import rearrange
import numpy as np
import torch
import torchvision.transforms.functional as TF
from fastvideo.distributed import get_local_torch_device
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.dits.lingbotworld2.cam_utils import (
compute_relative_poses,
get_Ks_transformed,
get_plucker_embeddings,
interpolate_camera_poses,
)
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler, )
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
InputValidationStage,
TextEncodingStage,
)
from fastvideo.pipelines.stages.base import PipelineStage
logger = init_logger(__name__)
class LingBotWorld2TextEncodingStage(TextEncodingStage):
"""Keep LingBot World 2 T5 attention masks so DiT context matches the source pipeline."""
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Initialize mask storage before running the shared text-encoding stage."""
if batch.prompt_attention_mask is None:
batch.prompt_attention_mask = []
return super().forward(batch, fastvideo_args)
class LingBotWorld2CausalFastGenerationStage(PipelineStage):
"""Prepare LingBot World 2 conditions and run the released causal-fast sampling loop."""
def __init__(self, transformer, scheduler, vae) -> None:
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
self.vae = vae
self._cross_attn_initialized = False
def _convert_flow_pred_to_x0(
self,
flow_pred: torch.Tensor,
xt: torch.Tensor,
timestep: torch.Tensor,
) -> torch.Tensor:
"""Convert LingBot World 2 flow prediction to x0 using the scheduler sigma."""
original_dtype = flow_pred.dtype
flow_pred, xt, sigmas, timesteps = map(
lambda x: x.double().to(flow_pred.device),
[flow_pred, xt, self.scheduler.sigmas, self.scheduler.timesteps],
)
timestep_id = torch.argmin((timesteps - timestep).abs())
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
return (xt - sigma_t * flow_pred).to(original_dtype)
def _initialize_self_kv_cache(
self,
batch_size: int,
kv_size: int,
dtype: torch.dtype,
device: torch.device,
) -> list[dict]:
"""Allocate per-block self-attention KV cache tensors."""
head_dim = self.transformer.dim // self.transformer.num_heads
num_heads = self.transformer.num_heads // get_sp_world_size()
shape = [batch_size, kv_size, num_heads, head_dim]
return [{
"k": torch.zeros(shape, dtype=dtype, device=device),
"v": torch.zeros(shape, dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device),
} for _ in range(self.transformer.num_layers)]
def _initialize_crossattn_cache(
self,
batch_size: int,
max_sequence_length: int,
dtype: torch.dtype,
device: torch.device,
) -> list[dict]:
"""Allocate per-block text cross-attention KV cache tensors."""
head_dim = self.transformer.dim // self.transformer.num_heads
shape = [batch_size, max_sequence_length, self.transformer.num_heads, head_dim]
return [{
"k": torch.zeros(shape, dtype=dtype, device=device),
"v": torch.zeros(shape, dtype=dtype, device=device),
"is_init": torch.tensor([0], dtype=torch.bool, device=device),
} for _ in range(self.transformer.num_layers)]
@staticmethod
def _prompt_context(batch: ForwardBatch, device: torch.device) -> list[torch.Tensor]:
"""Slice padded text encoder states back to LingBot World 2's unpadded context list."""
assert batch.prompt_embeds
context_tensor = batch.prompt_embeds[0].to(device)
if batch.prompt_attention_mask:
mask = batch.prompt_attention_mask[0].to(device)
seq_lens = mask.gt(0).sum(dim=1).long()
return [u[:v] for u, v in zip(context_tensor, seq_lens, strict=True)]
return [u for u in context_tensor]
def _prepare_image_tensor(self, batch: ForwardBatch, device: torch.device) -> torch.Tensor:
"""Return the source-style normalized image tensor `[C,H,W]`."""
image = batch.pil_image
if image is None:
raise ValueError("LingBot World 2 causal-fast requires `image_path` or `pil_image`.")
if isinstance(image, torch.Tensor):
if image.ndim == 5:
return image[0, :, 0].to(device)
if image.ndim == 4:
return image[0].to(device)
return image.to(device)
return TF.to_tensor(image).sub_(0.5).div_(0.5).to(device)
def _prepare_camera(
self,
action_path: str,
c2ws: np.ndarray,
h: int,
w: int,
lat_f: int,
lat_h: int,
lat_w: int,
chunk_size: int,
dtype: torch.dtype,
device: torch.device,
) -> torch.Tensor:
"""Build the LingBot World 2 camera Plucker tensor for latent chunks."""
Ks = torch.from_numpy(np.load(os.path.join(action_path, "intrinsics.npy"))).float()
Ks = get_Ks_transformed(
Ks,
height_org=480,
width_org=832,
height_resize=h,
width_resize=w,
height_final=h,
width_final=w,
)
Ks = Ks[0]
len_c2ws = len(c2ws)
len_c2ws_ = int((len_c2ws - 1) // 4) + 1
len_c2ws_ = int(len_c2ws_ - (len_c2ws_ % chunk_size))
c2ws_infer = interpolate_camera_poses(
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
src_rot_mat=c2ws[:, :3, :3],
src_trans_vec=c2ws[:, :3, 3],
tgt_indices=np.linspace(0, len_c2ws - 1, len_c2ws_),
)
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
Ks = Ks.repeat(len(c2ws_infer), 1)
c2ws_plucker_emb = get_plucker_embeddings(c2ws_infer.to(device), Ks.to(device), h, w)
c2ws_plucker_emb = rearrange(
c2ws_plucker_emb,
"f (h c1) (w c2) c -> (f h w) (c c1 c2)",
c1=int(h // lat_h),
c2=int(w // lat_w),
)
c2ws_plucker_emb = c2ws_plucker_emb[None, ...]
return rearrange(
c2ws_plucker_emb,
"b (f h w) c -> b c f h w",
f=lat_f,
h=lat_h,
w=lat_w,
).to(device=device, dtype=dtype)
def _encode_condition_video(
self,
img: torch.Tensor,
h: int,
w: int,
frames: int,
mask: torch.Tensor,
fastvideo_args: FastVideoArgs,
) -> torch.Tensor:
"""Encode the first-frame conditioning video and prepend mask channels."""
device = get_local_torch_device()
self.vae = self.vae.to(device)
video_condition = torch.concat(
[
torch.nn.functional.interpolate(
img[None].cpu(),
size=(h, w),
mode="bicubic",
).transpose(0, 1),
torch.zeros(3, frames - 1, h, w),
],
dim=1,
).to(device)
vae_dtype = torch.float32
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=False):
encoder_output = self.vae.encode(video_condition.unsqueeze(0).to(torch.float32))
latent_condition = encoder_output.mean
if not bool(getattr(self.vae, "handles_latent_denorm", False)):
if hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None:
latent_condition -= self.vae.shift_factor.to(latent_condition.device, latent_condition.dtype)
latent_condition = latent_condition * self.vae.scaling_factor.to(latent_condition.device,
latent_condition.dtype)
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
return torch.concat([mask, latent_condition[0]], dim=0)
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Execute LingBot World 2 causal-fast generation and store final latents on the batch."""
device = get_local_torch_device()
cfg = fastvideo_args.pipeline_config.dit_config.arch_config
chunk_size = int(cfg.chunk_size)
max_sequence_length = int(batch.max_sequence_length or cfg.text_len)
action_path = batch.action_path
if action_path is None:
raise ValueError("LingBot World 2 causal-fast requires `action_path`.")
c2ws = np.load(os.path.join(action_path, "poses.npy"))
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
frame_num = ((int(batch.num_frames) - 1) // 4) * 4 + 1
frame_num = min(frame_num, len_c2ws)
c2ws = c2ws[:frame_num]
img = self._prepare_image_tensor(batch, device)
h0, w0 = img.shape[1:]
aspect_ratio = h0 / w0
lat_h = round(np.sqrt(cfg.max_area * aspect_ratio) // 8 // cfg.patch_size[1] * cfg.patch_size[1])
lat_w = round(np.sqrt(cfg.max_area / aspect_ratio) // 8 // cfg.patch_size[2] * cfg.patch_size[2])
h = lat_h * 8
w = lat_w * 8
lat_f = (frame_num - 1) // 4 + 1
lat_f = int(lat_f - (lat_f % chunk_size))
frames = (lat_f - 1) * 4 + 1
batch.height = h
batch.width = w
batch.num_frames = frames
seed = int(batch.seed if batch.seed is not None else 42)
seed_g = torch.Generator(device=device)
seed_g.manual_seed(seed)
noise = torch.randn(16, lat_f, lat_h, lat_w, dtype=torch.float32, generator=seed_g, device=device)
mask = torch.ones(1, frames, lat_h, lat_w, device=device)
mask[:, 1:] = 0
mask = torch.concat([torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1), mask[:, 1:]], dim=1)
mask = mask.view(1, mask.shape[1] // 4, 4, lat_h, lat_w).transpose(1, 2)[0]
self.scheduler.set_timesteps(cfg.num_train_timesteps, shift=cfg.sample_shift)
timesteps = self.scheduler.timesteps[list(cfg.timesteps_index)].to(device)
context = self._prompt_context(batch, device)
c2ws_plucker_emb = self._prepare_camera(
action_path,
c2ws,
h,
w,
lat_f,
lat_h,
lat_w,
chunk_size,
torch.bfloat16,
device,
)
y = self._encode_condition_video(img, h, w, frames, mask, fastvideo_args).to(device=device,
dtype=torch.bfloat16)
transformer_dtype = torch.bfloat16
frame_seqlen = int(noise.shape[-2] * noise.shape[-1] // 4)
kv_size = frame_seqlen * cfg.local_attn_size if cfg.local_attn_size > -1 else frame_seqlen * lat_f
self_kv_cache = self._initialize_self_kv_cache(1, kv_size, transformer_dtype, device)
cross_kv_cache = self._initialize_crossattn_cache(1, max_sequence_length, transformer_dtype, device)
self.transformer = self.transformer.to(device)
self._cross_attn_initialized = False
pred_latent_chunks = []
latents_chunk = noise.split(chunk_size, dim=1)
condition_chunk = y.split(chunk_size, dim=1)
c2ws_plucker_emb_chunk = c2ws_plucker_emb.split(chunk_size, dim=2)
max_seq_len = int(math.ceil(chunk_size * lat_h * lat_w // 4))
with torch.amp.autocast("cuda", dtype=transformer_dtype):
for chunk_id, current_latent in enumerate(latents_chunk):
current_condition = condition_chunk[chunk_id]
current_c2ws_plucker_emb = c2ws_plucker_emb_chunk[chunk_id]
dit_cond_dict = {"c2ws_plucker_emb": current_c2ws_plucker_emb.chunk(1, dim=0)}
kwargs = {
"context": [context[0]],
"seq_len": max_seq_len,
"y": [current_condition],
"dit_cond_dict": dit_cond_dict,
"kv_cache": self_kv_cache,
"crossattn_cache": cross_kv_cache,
"current_start": chunk_id * chunk_size * frame_seqlen,
"max_attention_size": kv_size,
"frame_seqlen": frame_seqlen,
}
x0 = current_latent
for timestep_idx, timestep_value in enumerate(timesteps):
timestep = torch.stack([timestep_value]).to(device)
noise_pred = self.transformer(
x=[current_latent.to(device)],
t=timestep,
cross_attn_first_call=not self._cross_attn_initialized,
**kwargs,
)[0]
self._cross_attn_initialized = True
x0 = self._convert_flow_pred_to_x0(noise_pred, current_latent, timestep_value)
if timestep_idx < len(timesteps) - 1:
next_timestep = timesteps[timestep_idx + 1].reshape(1)
current_latent = self.scheduler.add_noise(
x0,
torch.randn(x0.shape, generator=seed_g, device=x0.device, dtype=x0.dtype),
next_timestep,
)
pred_latent_chunks.append(x0)
context_timestep = torch.stack([timesteps[-1] * 0.0]).to(device)
self.transformer(x=[x0], t=context_timestep, cross_attn_first_call=False, **kwargs)
batch.latents = torch.cat(pred_latent_chunks, dim=1).unsqueeze(0)
return batch
class LingBotWorld2CausalFastPipeline(LoRAPipeline, ComposedPipelineBase):
"""FastVideo pipeline for LingBot World 2 14B causal-fast I2V generation."""
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
if "scheduler" not in self.modules:
self.modules["scheduler"] = FlowUniPCMultistepScheduler(shift=1.0, use_dynamic_shifting=False)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up the LingBot World 2 causal-fast pipeline stages."""
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(
stage_name="prompt_encoding_stage",
stage=LingBotWorld2TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
self.add_stage(
stage_name="lingbotworld2_causal_fast_generation_stage",
stage=LingBotWorld2CausalFastGenerationStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
),
)
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = LingBotWorld2CausalFastPipeline
@@ -0,0 +1,35 @@
# SPDX-License-Identifier: Apache-2.0
"""LingBotWorld2 causal-fast pipeline preset."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="Causal-fast denoising pass",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
LINGBOTWORLD2_CAUSAL_FAST_I2V = InferencePreset(
name="lingbotworld2_causal_fast_i2v",
version=1,
model_family="lingbotworld2",
description="LingBot World 2 14B causal-fast I2V",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"guidance_scale": 1.0,
"num_inference_steps": 4,
"fps": 16,
"seed": 42,
"num_frames": 65,
"height": 480,
"width": 832,
"negative_prompt": "",
},
)
ALL_PRESETS = (LINGBOTWORLD2_CAUSAL_FAST_I2V, )
@@ -0,0 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
"""Z-Image pipeline."""
from fastvideo.pipelines.basic.zimage.zimage_pipeline import ZImagePipeline
__all__ = ["ZImagePipeline"]
@@ -0,0 +1,40 @@
# SPDX-License-Identifier: Apache-2.0
"""Z-Image inference presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="Z-Image denoising pass",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
"cfg_normalization",
"cfg_truncation",
}),
)
ZIMAGE_TURBO = InferencePreset(
name="zimage_turbo",
version=1,
model_family="zimage",
description="Z-Image-Turbo text-to-image generation",
workload_type="t2i",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 1024,
"width": 1024,
"num_frames": 1,
"fps": 1,
"seed": 42,
"guidance_scale": 0.0,
"num_inference_steps": 8,
"negative_prompt": "",
"max_sequence_length": 512,
"cfg_normalization": False,
"cfg_truncation": 1.0,
},
)
ALL_PRESETS = (ZIMAGE_TURBO, )
+337
View File
@@ -0,0 +1,337 @@
# SPDX-License-Identifier: Apache-2.0
"""Pipeline stages for the native Z-Image text-to-image path."""
from __future__ import annotations
import inspect
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.hooks.activation_trace import trace_step
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.utils import PRECISION_TO_TYPE
class ZImageInputValidationStage(InputValidationStage):
"""Validate the image geometry and reproduce the official device RNG."""
def _generate_seeds(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> None:
del fastvideo_args
assert batch.seed is not None
batch.seeds = [batch.seed]
device = get_local_torch_device()
batch.generator = torch.Generator(device=device).manual_seed(batch.seed)
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.do_classifier_free_guidance and batch.negative_prompt is None and not batch.negative_prompt_embeds:
batch.negative_prompt = ""
batch = super().forward(batch, fastvideo_args)
if batch.num_frames != 1:
raise ValueError(f"Z-Image is text-to-image and requires num_frames=1, got {batch.num_frames}")
if batch.height is None or batch.width is None:
raise ValueError("Z-Image requires height and width")
if batch.height % 16 or batch.width % 16:
raise ValueError("Z-Image height and width must be divisible by 16; "
f"got {batch.height}x{batch.width}")
return batch
class ZImageConditioningStage(PipelineStage):
"""Trim padded Qwen states and materialize variable-length CFG streams."""
@staticmethod
def _trim_embeddings(
embeds: torch.Tensor,
attention_mask: torch.Tensor | None,
) -> list[torch.Tensor]:
if attention_mask is None:
return list(embeds.unbind(0))
return [
sample[mask.to(device=sample.device, dtype=torch.bool)]
for sample, mask in zip(embeds, attention_mask, strict=True)
]
@staticmethod
def _repeat(items: list[torch.Tensor], count: int) -> list[torch.Tensor]:
return [item for item in items for _ in range(count)]
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
del fastvideo_args
if len(batch.prompt_embeds) != 1:
raise ValueError(f"Z-Image expects one text encoder, got {len(batch.prompt_embeds)}")
prompt_mask = batch.prompt_attention_mask[0] if batch.prompt_attention_mask else None
prompt_embeds = self._trim_embeddings(batch.prompt_embeds[0], prompt_mask)
batch.extra["zimage_prompt_embeds"] = self._repeat(prompt_embeds, batch.num_videos_per_prompt)
if batch.do_classifier_free_guidance:
if not batch.negative_prompt_embeds:
raise ValueError("Z-Image CFG requires negative prompt embeddings")
negative_mask = batch.negative_attention_mask[0] if batch.negative_attention_mask else None
negative_embeds = self._trim_embeddings(batch.negative_prompt_embeds[0], negative_mask)
batch.extra["zimage_negative_prompt_embeds"] = self._repeat(
negative_embeds,
batch.num_videos_per_prompt,
)
else:
batch.extra["zimage_negative_prompt_embeds"] = []
return batch
class ZImageLatentPreparationStage(PipelineStage):
"""Create the official fp32 image latents on the transformer device."""
def __init__(self, transformer) -> None:
self.transformer = transformer
@staticmethod
def _randn(
shape: tuple[int, ...],
generators: torch.Generator | list[torch.Generator] | None,
device: torch.device,
) -> torch.Tensor:
if isinstance(generators, list):
if len(generators) != shape[0]:
raise ValueError(f"generator list length {len(generators)} does not match batch size {shape[0]}")
sample_shape = (1, *shape[1:])
return torch.cat([
torch.randn(sample_shape, generator=generator, device=device, dtype=torch.float32)
for generator in generators
])
return torch.randn(shape, generator=generators, device=device, dtype=torch.float32)
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.height is None or batch.width is None:
raise ValueError("Z-Image requires height and width before latent preparation")
prompt_embeds = batch.extra.get("zimage_prompt_embeds")
if not isinstance(prompt_embeds, list) or not prompt_embeds:
raise ValueError("Z-Image conditioning must run before latent preparation")
channels = int(getattr(self.transformer, "in_channels", 16))
spatial_ratio = int(fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio)
shape = (
len(prompt_embeds),
channels,
1,
batch.height // spatial_ratio,
batch.width // spatial_ratio,
)
device = get_local_torch_device()
if batch.latents is None:
latents = self._randn(shape, batch.generator, device)
else:
latents = batch.latents
if latents.ndim == 4:
latents = latents.unsqueeze(2)
if tuple(latents.shape) != shape:
raise ValueError(f"Expected Z-Image latents with shape {shape}, got {tuple(latents.shape)}")
latents = latents.to(device=device, dtype=torch.float32)
batch.latents = latents
batch.raw_latent_shape = shape
return batch
class ZImageTimestepPreparationStage(PipelineStage):
"""Apply the native scheduler's zero endpoint and discrete schedule."""
def __init__(self, scheduler) -> None:
self.scheduler = scheduler
@staticmethod
def _calculate_shift(
image_seq_len: int,
base_seq_len: int,
max_seq_len: int,
base_shift: float,
max_shift: float,
) -> float:
slope = (max_shift - base_shift) / (max_seq_len - base_seq_len)
return image_seq_len * slope + base_shift - slope * base_seq_len
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is None:
raise ValueError("Z-Image latents must be prepared before timesteps")
if batch.timesteps is not None and batch.sigmas is not None:
raise ValueError("Only one of timesteps or sigmas may be supplied")
scheduler = self.scheduler
sigma_min = float(fastvideo_args.pipeline_config.scheduler_sigma_min)
use_reference_timesteps = bool(fastvideo_args.pipeline_config.scheduler_use_reference_discrete_timesteps)
scheduler.sigma_min = sigma_min
scheduler.register_to_config(
sigma_min=sigma_min,
use_reference_discrete_timesteps=use_reference_timesteps,
)
config = scheduler.config
image_seq_len = (batch.latents.shape[-2] // 2) * (batch.latents.shape[-1] // 2)
mu = self._calculate_shift(
image_seq_len,
int(config.get("base_image_seq_len", 256)),
int(config.get("max_image_seq_len", 4096)),
float(config.get("base_shift", 0.5)),
float(config.get("max_shift", 1.15)),
)
device = get_local_torch_device()
if batch.timesteps is not None:
if "timesteps" not in inspect.signature(scheduler.set_timesteps).parameters:
raise ValueError(f"{type(scheduler).__name__} does not accept custom timesteps")
scheduler.set_timesteps(timesteps=batch.timesteps, device=device, mu=mu)
elif batch.sigmas is not None:
if "sigmas" not in inspect.signature(scheduler.set_timesteps).parameters:
raise ValueError(f"{type(scheduler).__name__} does not accept custom sigmas")
scheduler.set_timesteps(sigmas=batch.sigmas, device=device, mu=mu)
else:
scheduler.set_timesteps(batch.num_inference_steps, device=device, mu=mu)
batch.timesteps = scheduler.timesteps
batch.num_inference_steps = len(batch.timesteps)
return batch
class ZImageDenoisingStage(PipelineStage):
"""Run the native Z-Image flow-matching loop."""
performance_component_metric = "transformer_time_s"
def __init__(self, transformer, scheduler) -> None:
self.transformer = transformer
self.scheduler = scheduler
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is None or batch.timesteps is None:
raise ValueError("Z-Image denoising requires latents and timesteps")
latents = batch.latents.float()
positive = batch.extra.get("zimage_prompt_embeds")
negative = batch.extra.get("zimage_negative_prompt_embeds", [])
if not isinstance(positive, list) or not positive:
raise ValueError("Z-Image denoising requires prompt embeddings")
device = get_local_torch_device()
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
batch_size = latents.shape[0]
for index, timestep_value in enumerate(batch.timesteps):
if timestep_value.item() == 0 and index == len(batch.timesteps) - 1:
continue
timestep = timestep_value.expand(batch_size).to(device=device, dtype=torch.float32)
timestep = (1000.0 - timestep) / 1000.0
current_guidance_scale = float(batch.guidance_scale)
if (batch.do_classifier_free_guidance and batch.cfg_truncation is not None and batch.cfg_truncation <= 1.0
and timestep[0].item() > batch.cfg_truncation):
current_guidance_scale = 0.0
apply_cfg = batch.do_classifier_free_guidance and current_guidance_scale > 0.0
if apply_cfg:
if not isinstance(negative, list) or len(negative) != batch_size:
raise ValueError("Z-Image CFG requires one negative embedding per image")
model_latents = latents.to(target_dtype).repeat(2, 1, 1, 1, 1)
model_embeddings = positive + negative
model_timestep = timestep.repeat(2)
else:
model_latents = latents.to(target_dtype)
model_embeddings = positive
model_timestep = timestep
with (
torch.autocast(
device_type=device.type,
enabled=False,
),
trace_step(index),
set_forward_context(
current_timestep=index,
attn_metadata=None,
forward_batch=batch,
),
):
model_outputs = self.transformer(
hidden_states=model_latents,
encoder_hidden_states=model_embeddings,
timestep=model_timestep,
)[0]
if apply_cfg:
positive_outputs = model_outputs[:batch_size]
negative_outputs = model_outputs[batch_size:]
guided_outputs = []
for positive_output, negative_output in zip(
positive_outputs,
negative_outputs,
strict=True,
):
positive_fp32 = positive_output.float()
prediction = positive_fp32 + current_guidance_scale * (positive_fp32 - negative_output.float())
if batch.cfg_normalization:
positive_norm = torch.linalg.vector_norm(positive_fp32)
prediction_norm = torch.linalg.vector_norm(prediction)
if prediction_norm > positive_norm:
prediction = prediction * (positive_norm / prediction_norm)
guided_outputs.append(prediction)
noise_pred = torch.stack(guided_outputs)
else:
noise_pred = torch.stack([output.float() for output in model_outputs])
noise_pred = -noise_pred
latents = self.scheduler.step(
noise_pred,
timestep_value,
latents,
return_dict=False,
)[0].float()
batch.latents = latents
return batch
class ZImageDecodingStage(PipelineStage):
"""Apply the official latent transform and decode one image frame."""
performance_component_metric = "vae_decode_time_s"
def __init__(self, vae) -> None:
self.vae = vae
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.latents is None:
raise ValueError("Z-Image decoding requires latents")
if fastvideo_args.output_type == "latent":
# FastVideo standardizes image and video latents as [B,C,T,H,W].
# Tongyi's image-only API returns the equivalent tensor with T squeezed.
batch.output = batch.latents
return batch
latents = batch.latents
if latents.ndim != 5 or latents.shape[2] != 1:
raise ValueError(f"Expected Z-Image latents [B,C,1,H,W], got {tuple(latents.shape)}")
device = get_local_torch_device()
self.vae = self.vae.to(device)
vae_dtype = getattr(self.vae, "dtype", None)
if vae_dtype is None:
vae_dtype = next(self.vae.parameters()).dtype
config = self.vae.config
scaling_factor = float(config.scaling_factor)
shift_factor = float(config.shift_factor or 0.0)
latents_2d = latents.squeeze(2).to(device=device, dtype=vae_dtype)
latents_2d = latents_2d / scaling_factor + shift_factor
decoded = self.vae.decode(latents_2d, return_dict=False)[0]
decoded = (decoded / 2 + 0.5).clamp(0, 1)
batch.output = decoded.unsqueeze(2).float()
if fastvideo_args.vae_cpu_offload:
self.vae.to("cpu")
return batch
@@ -0,0 +1,63 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from fastvideo.configs.pipelines.zimage import ZImagePipelineConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
from .stages import (
ZImageConditioningStage,
ZImageDecodingStage,
ZImageDenoisingStage,
ZImageInputValidationStage,
ZImageLatentPreparationStage,
ZImageTimestepPreparationStage,
)
class ZImagePipeline(ComposedPipelineBase):
"""Native Z-Image text-to-image pipeline."""
pipeline_config_cls: type[ZImagePipelineConfig] = ZImagePipelineConfig
_required_config_modules = [
"scheduler",
"text_encoder",
"tokenizer",
"transformer",
"vae",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
scheduler = self.get_module("scheduler")
transformer = self.get_module("transformer")
self.add_stage("input_validation_stage", ZImageInputValidationStage())
self.add_stage(
"text_encoding_stage",
TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage("zimage_conditioning_stage", ZImageConditioningStage())
self.add_stage(
"latent_preparation_stage",
ZImageLatentPreparationStage(transformer=transformer),
)
self.add_stage(
"timestep_preparation_stage",
ZImageTimestepPreparationStage(scheduler=scheduler),
)
self.add_stage(
"denoising_stage",
ZImageDenoisingStage(transformer=transformer, scheduler=scheduler),
)
self.add_stage(
"decoding_stage",
ZImageDecodingStage(vae=self.get_module("vae")),
)
EntryClass = ZImagePipeline
@@ -303,7 +303,8 @@ class ComposedPipelineBase(ABC):
self.modules[module_name] = module
def _load_config(self, model_path: str) -> dict[str, Any]:
model_path = maybe_download_model(self.model_path)
revision = getattr(self.fastvideo_args, "revision", None)
model_path = maybe_download_model(self.model_path, revision=revision)
self.model_path = model_path
# fastvideo_args.downloaded_model_path = model_path
logger.info("Model path: %s", model_path)
+5 -1
View File
@@ -147,8 +147,9 @@ class ForwardBatch:
camera_trajectory: str | None = None # Camera trajectory file/identifier
action_list: list[str] | None = None # List of actions (e.g., ['forward', 'left'])
action_speed_list: list[float] | None = None # Speed for each action
# Camera control inputs (LingBotWorld)
# Camera control inputs (LingBotWorld and LingBotWorld2)
c2ws_plucker_emb: torch.Tensor | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
action_path: str | None = None # Directory containing poses.npy and intrinsics.npy
# Camera control inputs (GEN3C)
trajectory_type: str | None = None
@@ -177,7 +178,10 @@ class ForwardBatch:
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
batch_cfg: bool = False
guidance_scale_2: float | None = None
cfg_normalization: bool = False
cfg_truncation: float | None = 1.0
guidance_rescale: float = 0.0
eta: float = 0.0
sigmas: list[float] | None = None
@@ -0,0 +1,258 @@
# SPDX-License-Identifier: Apache-2.0
"""Preprocess LTX-2 overfit data into parquet format.
Encodes videos with the LTX-2 causal video VAE and captions with the
Gemma text encoder (feature extractor + embedding connector) into the
t2v parquet schema expected by the training framework.
The stored text embeddings are POST-connector: the connector replaces
pad positions with learnable registers and returns an all-valid mask,
so the parquet collate's ones/zeros mask stays semantically correct
and training needs no text encoder at all. Captions are encoded via
the encoder's forward() (the exact inference path), which handles both
LTX-2.0 (shared 3840-d features) and LTX-2.3 (separate 4096-d video /
2048-d audio feature extractors).
Videos are resampled to TRAIN_FPS and the preprocessed clip is also
saved as an mp4 next to the parquet so overfit tests can use it as
the SSIM reference.
Usage:
CUDA_VISIBLE_DEVICES=0 python fastvideo/pipelines/preprocess/preprocess_ltx2_overfit.py
"""
import json
import os
import shutil
from typing import Any
import av
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.utils import maybe_download_model, verify_model_config_and_directory
# --- Config ---
NUM_FRAMES = 81 # 8k+1 for temporal compression ratio 8
MAX_HEIGHT = 480 # divisible by 32 (spatial compression)
MAX_WIDTH = 832
TRAIN_FPS = 24.0 # matches the LTX-2 preset fps used at validation
DATA_DIR = os.environ.get("LTX2_OVERFIT_DATA_DIR", "data/cats")
CAPTION_JSON = os.environ.get("LTX2_OVERFIT_CAPTION_JSON", "videos2caption_1_sample.json")
VIDEO_SUBDIR = os.environ.get("LTX2_OVERFIT_VIDEO_SUBDIR", "video")
OUTPUT_DIR = os.environ.get("LTX2_OVERFIT_OUTPUT_DIR", "data/ltx2_overfit_preprocessed")
MODEL_REPO = os.environ.get("LTX2_OVERFIT_MODEL", "FastVideo/LTX2-Distilled-Diffusers")
# The train dataloader samples with drop_last=True across data-parallel
# groups, so the dataset must hold at least num_sp_groups * batch_size
# rows or every rank gets zero batches. Replicate the overfit sample so
# a 4-GPU FSDP run still sees one batch per rank.
NUM_COPIES = int(os.environ.get("LTX2_OVERFIT_NUM_COPIES", "4"))
def _init_single_process_distributed() -> None:
"""FastVideo component loaders expect an initialized distributed env."""
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29511")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("WORLD_SIZE", "1")
os.environ.setdefault("LOCAL_RANK", "0")
from fastvideo.distributed import (
maybe_init_distributed_environment_and_model_parallel, )
maybe_init_distributed_environment_and_model_parallel(1, 1)
def load_video(path: str, num_frames: int, target_fps: float, height: int,
width: int) -> tuple[torch.Tensor, np.ndarray]:
"""Load a video as [1, C, T, H, W] in [-1, 1], resampled to target_fps.
Also returns the uint8 RGB frames [T, H, W, C] for reference-video export.
"""
with av.open(path) as container:
if not container.streams.video:
raise RuntimeError(f"No video stream found in {path}")
stream = container.streams.video[0]
native_fps = float(stream.average_rate or target_fps)
decoded = [torch.from_numpy(frame.to_ndarray(format="rgb24")) for frame in container.decode(video=0)]
if not decoded:
raise RuntimeError(f"Could not read any frames from {path}")
raw = torch.stack(decoded).permute(0, 3, 1, 2) # [T, C, H, W] uint8
step = native_fps / target_fps
wanted = [min(int(round(i * step)), raw.shape[0] - 1) for i in range(num_frames)]
frames = raw[wanted].float() # [T, C, H, W] in [0, 255]
src_h, src_w = frames.shape[2], frames.shape[3]
scale = max(height / src_h, width / src_w)
new_h, new_w = int(round(src_h * scale)), int(round(src_w * scale))
frames = torch.nn.functional.interpolate(frames, size=(new_h, new_w), mode="bilinear", antialias=True)
top = (new_h - height) // 2
left = (new_w - width) // 2
frames = frames[:, :, top:top + height, left:left + width]
frames_np = (frames.permute(0, 2, 3, 1).clamp(0, 255).round().to(torch.uint8).numpy())
video = frames / 127.5 - 1.0 # [0,255] -> [-1,1]
video = video.permute(1, 0, 2, 3).unsqueeze(0) # [1,C,T,H,W]
return video, frames_np
def main() -> None:
_init_single_process_distributed()
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.models.loader.component_loader import (
PipelineComponentLoader, )
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
device = torch.device("cuda:0")
model_path = maybe_download_model(MODEL_REPO)
model_index = verify_model_config_and_directory(model_path)
os.makedirs(OUTPUT_DIR, exist_ok=True)
# The map-style dataset caches parquet file metadata; a stale cache
# next to a regenerated parquet can crash or serve old rows.
shutil.rmtree(os.path.join(OUTPUT_DIR, "map_style_cache"), ignore_errors=True)
with open(os.path.join(DATA_DIR, CAPTION_JSON)) as f:
caption_data = json.load(f)
pipeline_config = LTX2T2VConfig()
fastvideo_args = FastVideoArgs(
model_path=model_path,
pipeline_config=pipeline_config,
num_gpus=1,
tp_size=1,
sp_size=1,
hsdp_shard_dim=1,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
)
def load_component(name: str) -> Any:
transformers_or_diffusers, _ = model_index[name]
return PipelineComponentLoader.load_module(
module_name=name,
component_model_path=os.path.join(model_path, name),
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=fastvideo_args,
)
print("Loading LTX-2 VAE...")
vae = load_component("vae")
vae_dtype = next(vae.parameters()).dtype
print(f"VAE loaded ({sum(p.numel() for p in vae.parameters())/1e6:.0f}M, {vae_dtype})")
print("Loading Gemma text encoder + tokenizer...")
text_encoder = load_component("text_encoder")
tokenizer = load_component("tokenizer")
tokenizer.padding_side = "left"
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
tokenizer.pad_token = tokenizer.eos_token
encoder_config = pipeline_config.text_encoder_configs[0]
tokenizer_kwargs = dict(encoder_config.tokenizer_kwargs)
if "max_length" not in tokenizer_kwargs:
tokenizer_kwargs["max_length"] = encoder_config.arch_config.text_len
preprocess_text = pipeline_config.preprocess_text_funcs[0]
# --- Process each video ---
records = []
for idx, item in enumerate(caption_data):
video_name = item["path"]
record_id = f"{idx:04d}_{video_name}"
caption = item["cap"][0] if isinstance(item["cap"], list) else item["cap"]
video_path = os.path.join(DATA_DIR, VIDEO_SUBDIR, video_name)
print(f"\nProcessing: {video_name}")
print(f" Caption: {caption[:80]}...")
video, frames_np = load_video(video_path, NUM_FRAMES, TRAIN_FPS, MAX_HEIGHT, MAX_WIDTH)
video = video.to(device=device, dtype=vae_dtype)
print(f" Video shape: {video.shape}")
with torch.no_grad():
# LTX-2 encode() returns a deterministic distribution whose
# mean is already per-channel normalized; store as-is.
latent = vae.encode(video).mean.squeeze(0).float().cpu()
print(f" Latent shape: {latent.shape}")
with torch.no_grad():
# Encode through forward() — the inference text path. The
# two-step preprocess_text_embeddings + run_connectors route
# breaks on LTX-2.3, whose separate audio feature extractor
# is narrower than the video one.
text_inputs = tokenizer([preprocess_text(caption)], **tokenizer_kwargs)
encoder_out = text_encoder(
input_ids=text_inputs["input_ids"].to(device),
attention_mask=text_inputs["attention_mask"].to(device),
)
text_embedding = encoder_out.last_hidden_state.squeeze(0).float().cpu()
print(f" Text embedding shape: {text_embedding.shape}")
record = {
"id": record_id,
"vae_latent_bytes": latent.numpy().tobytes(),
"vae_latent_shape": list(latent.shape),
"vae_latent_dtype": str(latent.dtype).replace("torch.", ""),
"text_embedding_bytes": (text_embedding.numpy().tobytes()),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype).replace("torch.", ""),
"file_name": video_name,
"caption": caption,
"media_type": "video",
"width": MAX_WIDTH,
"height": MAX_HEIGHT,
"num_frames": NUM_FRAMES,
"duration_sec": NUM_FRAMES / TRAIN_FPS,
"fps": TRAIN_FPS,
}
records.append(record)
# Save the preprocessed clip so overfit tests can compare
# validation output against the memorization target.
import imageio
ref_path = os.path.join(OUTPUT_DIR, f"training_sample_{idx}.mp4")
with imageio.get_writer(ref_path, fps=TRAIN_FPS) as writer:
for frame in frames_np:
writer.append_data(frame)
print(f" Wrote reference clip to {ref_path}")
# Clean up
del text_encoder, tokenizer, vae
torch.cuda.empty_cache()
# Write parquet (replicated NUM_COPIES times; see comment at top)
replicated = []
for copy_idx in range(max(1, NUM_COPIES)):
for r in records:
row = dict(r)
row["id"] = f"{r['id']}_copy{copy_idx}"
replicated.append(row)
table = pa.table(
{k: [r[k] for r in replicated]
for k in replicated[0]},
schema=pyarrow_schema_t2v,
)
output_path = os.path.join(OUTPUT_DIR, "data_00000.parquet")
pq.write_table(table, output_path)
print(f"\nWrote {len(replicated)} records "
f"({len(records)} unique x {max(1, NUM_COPIES)} copies) to {output_path}")
# Write validation prompts for the validation callback
val_prompts = {
"data": [{
"caption": (item["cap"][0] if isinstance(item["cap"], list) else item["cap"]),
} for item in caption_data],
}
val_path = os.path.join(OUTPUT_DIR, "validation_prompts.json")
with open(val_path, "w") as f:
json.dump(val_prompts, f, indent=2)
print(f"Wrote validation prompts to {val_path}")
print("\nDone! Use data_path: " + OUTPUT_DIR + " in training config.")
if __name__ == "__main__":
main()

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