Compare commits

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

Adds examples/inference/basic/basic_cosmos3_new_api.py (new VideoGenerator
API). Verified: 1280x704, 29 frames, 35 steps on B200 -> coherent video
matching the prompt (no NaNs, real pixel variance). Full cosmos3 parity
suite green.
2026-07-13 03:23:27 -07:00
SolitaryThinker 4533447e58 [misc] cosmos3 PORT_STATUS: PR1 video core complete (framework parity) 2026-07-13 03:23:24 -07:00
SolitaryThinker c94d625844 [feat] cosmos3: native video pipeline + framework denoise parity 2026-07-13 03:23:23 -07:00
SolitaryThinker 7427b9d2d3 [feat] cosmos3: native MoT sequence-packing (video) + framework parity 2026-07-13 03:22:40 -07:00
SolitaryThinker 9dab81f180 [misc] cosmos3 PORT_STATUS: strict-load done; only pipeline remains 2026-07-13 03:22:39 -07:00
SolitaryThinker 840b43b4c3 [feat] cosmos3: strict-load verified (identity; needs_conversion=no) 2026-07-13 03:22:37 -07:00
SolitaryThinker ace34f866a [misc] cosmos3 PORT_STATUS: VAE component framework parity verified 2026-07-13 03:22:37 -07:00
SolitaryThinker b5f90a6b77 [feat] cosmos3: VAE config (Wan2.2 reuse) + framework parity test 2026-07-13 03:22:36 -07:00
SolitaryThinker ed8f118acd [misc] cosmos3 PORT_STATUS: DiT framework parity verified (both rope modes) 2026-07-13 03:22:36 -07:00
SolitaryThinker b379ea96b4 [test] cosmos3: DiT unified_3d_mrope framework parity (real checkpoint) 2026-07-13 03:22:34 -07:00
SolitaryThinker 6771f4ec2e [feat] cosmos3: native DiT + framework parity (Cosmos3VFMTransformer) 2026-07-13 03:22:33 -07:00
SolitaryThinker b48eafe515 [misc] cosmos3 PORT_STATUS: PR1 progress (arch config, framework parity ref) 2026-07-13 03:22:32 -07:00
SolitaryThinker 1d199050af [test] cosmos3: official framework DiT parity reference (CPU/SDPA) 2026-07-13 03:22:30 -07:00
SolitaryThinker 8b65f7ce6c [misc] cosmos3 PORT_STATUS: framework-reference pivot + Phase 1 findings 2026-07-13 03:22:30 -07:00
SolitaryThinker 2d2657821e [feat] cosmos3: arch config 1:1 with Cosmos3-Nano checkpoint 2026-07-13 03:22:29 -07:00
SolitaryThinker ded49ab1c8 [misc]: cosmos3 PORT_STATUS: post-rebase state + fv-cosmos3 env 2026-07-13 03:22:29 -07:00
SolitaryThinker 10f5372f42 [misc]: cosmos3 resume: PORT_STATUS + README for official diffusers ref 2026-07-13 03:22:27 -07:00
SolitaryThinker ed4125e7d8 [feat]: Cosmos3 checkpoint conversion script (cosmos3_convert.py)
scripts/checkpoint_conversion/cosmos3_convert.py wraps Cosmos3OmniDiffusersPipeline._remap_ckpt_key. convert_state_dict(src) returns (new_state, skipped, unmapped) for shard-by-shard remap; convert_checkpoint_dir reads *.safetensors shards and writes a consolidated FastVideo-format state dict (sharded index.json TODO when real weights land). smoke_test() exercises all 14 representative remap branches synthetically (no real weights needed) and returns nonzero exit on failure; runs in CI as a remap-drift guard. Usage: python scripts/checkpoint_conversion/cosmos3_convert.py --smoke-test.
2026-07-13 03:22:26 -07:00
SolitaryThinker 26e78f8fb1 [feat]: Cosmos3Config + registry wire-up for nvidia/Cosmos3-Nano
Cosmos3Config(PipelineConfig): reuses Cosmos25VAEConfig (Cosmos3 uses DistributedAutoencoderKLWan); text_encoder_configs=() since Cosmos3LanguageModel lives inside the DiT; flow_shift=1.0 default (T2V/I2V engine init; T2I uses 3.0 per-request via _set_flow_shift). Registry entry registers BEFORE the generic cosmos detector to win path precedence (same pattern as GEN3C). Resolves nvidia/Cosmos3-Nano -> Cosmos3Config without regressing Cosmos25/GEN3C/Cosmos.
2026-07-13 03:22:26 -07:00
SolitaryThinker 12f94fd6e7 [feat]: Cosmos3 pipeline (Cosmos3OmniDiffusersPipeline)
Pipeline class with: _remap_ckpt_key static method (verbatim port from pipeline_cosmos3.py:319-409, 14-rule checkpoint remap UND/GEN split + lm_head skip); _set_flow_shift method with lazy UniPCMultistepScheduler construction; diffuse() with sequential 3-mode CFG denoising loop + I2V velocity_mask + image_latent re-injection (ported from pipeline_cosmos3.py:883-1033); forward() with T2I/T2V/I2V mode dispatch + flow_shift selection (ported from pipeline_cosmos3.py:1037-1206); 4 helper-method stubs raise NotImplementedError. Dual inheritance (nn.Module, ComposedPipelineBase) matches upstream pattern.
2026-07-13 03:21:44 -07:00
SolitaryThinker facd035b97 [feat]: Cosmos3 DiT skeleton (Cosmos3VFMTransformer + Cosmos3LanguageModel)
Ports math (compute_mrope_position_ids_text/_vision, patchify, unpatchify) verbatim from vllm-omni transformer_cosmos3.py:113-177,1009-1036. Module tree (language_model.layers.*.self_attn.{q,k,v,o}_proj + gen_layers.*.cross_attention.* + vae2llm + llm2vae + time_embedder + norm_moe_gen) exposes the parameter names the checkpoint converter targets. All layer forward() raise NotImplementedError until weights publish; module instantiation + state-dict key tests pass.
2026-07-13 03:21:44 -07:00
SolitaryThinker c3c6c0cb2d [test]: Cosmos3 local-test parity scaffold (Tier A, 15 tests)
Adds 8 CPU-only parity test files + conftest stubs (StubScheduler, StubCosmos3VAE, StubCosmos3Transformer) under tests/local_tests/cosmos3/ mirroring the vllm-omni reference at tests/diffusion/models/cosmos3/conftest.py. All 15 tests skip until the FastVideo modules land (subsequent commits). Reference: vllm-omni PR #3454 @ 8536f5b1421f.
2026-07-13 03:21:44 -07:00
Mac Lee 1ea2517e22 [ci]: extend LoRA training CI timeout (#1589) 2026-07-13 02:47:28 -07:00
MookandSolitaryThinker 0c63528c59 [perf] Cache RoPE position-embedding tables across denoising steps (#1442)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-13 02:34:51 -07:00
b063f8ca41 [feat] Fix FLUX.1-dev port: native RoPE, parity tests, SSIM reference (#1321)
Co-authored-by: Ishan Vaish <ivaish@ucsd.edu>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-13 01:49:42 -07:00
William Lin e7fff0173a [bugfix]: benchmark_weight_loading_comparison.py — iterate safe_open via .keys() (#1378) 2026-07-12 22:50:45 -07:00
Shreejith SGandH1yori233 d82abc271e [feat] Add GLM-Image inference support (#1030)
Co-authored-by: H1yori233 <k1kong@ucsd.edu>
2026-07-12 22:42:42 -07:00
Guian FangandSolitaryThinker 970409962f [feat] Add AnyFlow any-step video distillation (pretrain + on-policy) (#1371)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2026-07-12 02:32:08 +00:00
Raghav K 055586703d [perf]: register a real backward for FA2 default + masked/varlen custom ops (training-under-compile) (#1388) 2026-07-12 00:45:27 +00:00
Mac Lee 5d89f86675 [ci] Stop forcing FA4 in model-load lanes (#1561) 2026-07-11 14:10:36 -07:00
127 changed files with 18867 additions and 191 deletions
+1 -1
View File
@@ -436,7 +436,7 @@ steps:
- "pyproject.toml"
- "docker/Dockerfile"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
command: "timeout 25m .buildkite/scripts/pr_test.sh"
label: ":test_tube: LoRA Training Tests"
env:
- TEST_TYPE=training_lora
+8 -2
View File
@@ -37,6 +37,11 @@ logs/
official_weights/
converted_weights/
# Cosmos3 local parity assets (symlinked from main worktree)
/official_weights/
/converted_weights/
/cosmos-framework
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
**/.cache/**
@@ -72,8 +77,7 @@ docs/distillation/examples/
# Python pickle files
*.pkl
# Reference videos
!fastvideo/tests/ssim/reference_videos/**/*.mp4
# Reference videos (negations must come after the catch-all on line below)
# Static images
!docs/assets/images/**/*.png
@@ -127,6 +131,8 @@ apps/dreamverse/web/.env.production.local
.sisyphus/
openspec/
fastvideo/tests/ssim/reference_videos/**
!fastvideo/tests/ssim/reference_videos/**/*.mp4
!fastvideo/tests/ssim/reference_videos/**/*.png
# Editor logs and local Python version pins (accidentally committed)
*.nvimlog
@@ -458,6 +458,8 @@ surfaces:
guidance_scale: request.sampling.guidance_scale
guidance_scale_2: request.sampling.guidance_scale_2
guidance_rescale: request.sampling.guidance_rescale
use_embedded_guidance: request.sampling.use_embedded_guidance
true_cfg_scale: request.sampling.true_cfg_scale
boundary_ratio: request.sampling.boundary_ratio
sigmas: request.sampling.sigmas
enable_teacache: request.runtime.enable_teacache
+116
View File
@@ -0,0 +1,116 @@
# 🌊 AnyFlow Any-Step Video Distillation
**AnyFlow** ([paper](https://arxiv.org/abs/2605.13724), [project page](https://nvlabs.github.io/AnyFlow/), [official code](https://github.com/NVlabs/AnyFlow), [model weights](https://huggingface.co/collections/nvidia/anyflow)) is an any-step video diffusion framework built on flow maps. A single distilled checkpoint can be evaluated at NFE ∈ {1, 2, 4, 8, 16, 32} without retraining, and quality scales **monotonically** with steps — unlike consistency-based distillation, which often degrades as NFE grows.
The student network ``u_θ(x_t, t, r)`` predicts the *average velocity* from time ``t`` back to time ``r``, so one Euler step is
```
x_r = x_t - ((t - r) / N) · u_θ(x_t, t, r)
```
for any ``t > r``.
## 📊 Model Overview
NVIDIA publishes four checkpoints under [`nvidia/anyflow`](https://huggingface.co/collections/nvidia/anyflow):
- `nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers` — bidirectional T2V, Wan2.1 1.3B base
- `nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers` — bidirectional T2V, Wan2.1 14B base
- `nvidia/AnyFlow-FAR-Wan2.1-1.3B-Diffusers` — frame-autoregressive variant, 1.3B
- `nvidia/AnyFlow-FAR-Wan2.1-14B-Diffusers` — frame-autoregressive variant, 14B
FastVideo currently supports the bidirectional T2V variants for training; the FAR variants can be loaded for inference through the diffusers integration.
## ⚙️ Inference
For inference, load the published checkpoint directly through diffusers; FastVideo's training-side ``WanModel`` config maps the HF AnyFlow ``delta_embedder`` weights onto its internal layout via ``param_names_mapping`` so the same checkpoint can be used as the ``init_from`` for the on-policy YAML below.
## 🧠 Algorithm
Training runs in two stages. Both use the dual-timestep Wan backbone — enabled by ``pipeline.dit_config.r_embedder: true`` in the YAML, which allocates a sibling ``condition_embedder.delta_embedder`` and fuses its embedding with the standard timestep embedding via either an additive or a gated mixer.
### Stage 1 — Pretrain (flow-map central-difference)
Method: ``AnyFlowPretrainMethod`` (``fastvideo/train/methods/distribution_matching/anyflow_pretrain.py``)
For each batch, sample ``(t, r) ∈ [0, 1]`` as ``(max, min)`` of two uniform draws, then:
- a ``diffusion_ratio`` fraction (default 0.5) gets ``r = t`` — recovers plain flow matching;
- a ``consistency_ratio`` fraction (default 0.25) gets ``r = 0`` — forces consistency to clean data;
- the remainder is free.
The student forward at ``(t, r)`` is trained against the central-difference target
```
target = (eps - x_0) - (t - r) · dF/dt
```
where ``dF/dt`` is estimated from the student's own forward at ``(t ± δ, r)`` with the sample also moved along the flow trajectory by ``v_pred · (δ / N)``. Per-timestep weighting uses ``beta08`` (``w(t) = t · sqrt(1 - t)``, renormalized). A stop-gradient scale-balance keeps the non-diffusion branches' loss magnitude aligned with the diffusion branch.
### Stage 2 — On-policy DMD
Method: ``AnyFlowMethod`` (``fastvideo/train/methods/distribution_matching/anyflow.py``)
Inherits ``DMD2Method``. The student is rolled out for ``student_sample_steps`` Euler-flow steps from pure noise; one randomly-chosen step is gradient-enabled (broadcast from rank 0 so every worker agrees), the rest run under ``torch.no_grad``. With ``use_mean_velocity: true`` (default) the rollout uses ``r = t_next`` at each step, matching AnyFlow's ``WanAnyFlowPipeline.training_rollout``.
The inherited ``_dmd_loss`` (VSD with fake-score critic) consumes the rollout output and the teacher's CFG prediction. The optional pinned ``t_list_override`` lets configs reproduce the paper's hand-tuned 4-step schedule ``[999, 937, 833, 624, 0]``.
## 🚀 Training Scripts
### Stage 1 — pretrain
```bash
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml
```
**Key configuration** (in ``examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml``):
- Global batch size: 32 (8 GPUs × 4 per-GPU)
- Learning rate: 5e-5
- Flow shift: 5.0
- ``diffusion_ratio`` / ``consistency_ratio``: 0.5 / 0.25
- ``epsilon`` (finite-difference step): 5 (absolute train-timestep units)
- ``weight_type``: ``beta08``
- ``fuse_guidance_scale``: 3.0
- Training steps: 6000
### Stage 2 — on-policy
```bash
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml \
--models.student.init_from outputs/wan2.1_anyflow_pretrain/checkpoint-final
```
(Or point ``models.student.init_from`` directly at ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers`` to bootstrap from the paper weights and skip Stage 1.)
**Key configuration**:
- Global batch size: 8 (8 GPUs × 1 per-GPU)
- Learning rate: 2e-6
- Flow shift: 5.0
- ``student_sample_steps``: 4
- ``t_list_override``: ``[999, 937, 833, 624, 0]``
- ``use_mean_velocity``: ``true`` (i.e. ``r = t_next`` during rollout)
- ``real_score_guidance_scale``: 3.0
- ``generator_update_interval``: 5 (DMD2 alternation)
- Training steps: 4000
## 🔌 Loading published AnyFlow checkpoints
The HF AnyFlow checkpoints expose ``condition_embedder.delta_embedder.*`` weights that FastVideo internally maps onto its ``condition_embedder.delta_embedder.mlp.*`` layout. This rename happens automatically through the regex in ``WanVideoArchConfig.param_names_mapping`` — no separate adapter is needed. The same regex is a no-op on plain Wan checkpoints (which don't contain any ``delta_embedder`` keys).
Set the YAML's ``pipeline.dit_config.r_embedder: true`` to allocate the ``delta_embedder`` module on the FastVideo side; when initializing from a plain Wan checkpoint the delta weights are deep-copied from ``time_embedder`` (matching AnyFlow's ``setup_flowmap_model()`` behavior).
## 🧭 Note on ``fuse_guidance_scale``
Stage 1 optionally fuses classifier-free guidance into the training target so the resulting checkpoint can be sampled at ``guidance_scale=1.0`` (no extra forward pass at inference time). The transformation is
```
noise_pred ← (noise_pred - (1 - g) · noise_pred_uncond) / g
```
with ``g = fuse_guidance_scale``. The negative prompt embedding comes from ``WanModel``'s ``ensure_negative_conditioning()`` — i.e. the dataset's configured ``sampling_param.negative_prompt``. Setting ``fuse_guidance_scale: 1.0`` skips the extra unconditional forward entirely.
The on-policy stage's ``real_score_guidance_scale`` (inherited from DMD2) follows the same parameterization conventions documented in [``dmd.md``](dmd.md#-note-on-real_score_guidance_scale).
@@ -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()
+140
View File
@@ -0,0 +1,140 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
import contextlib
import os
import re
DEFAULT_PROMPTS = [
"a photo of a cat",
(
"a cinematic photo of a red panda wearing a tiny backpack, standing on a "
"rainy neon-lit street at night, shallow depth of field, sharp focus, "
"35mm, bokeh"
),
]
def _safe_filename(text: str, max_len: int = 100) -> str:
"""Make a stable, filesystem-friendly filename base."""
s = text[:max_len].strip()
s = s.replace(os.sep, "_")
if os.altsep:
s = s.replace(os.altsep, "_")
s = re.sub(r"\s+", " ", s)
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
s = s.strip(" .")
return s or "prompt"
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
"""Delete prior outputs so reruns do not get _1, _2 suffixes."""
if not os.path.isdir(out_dir):
return
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.(mp4|png)$")
for fn in os.listdir(out_dir):
if pattern.match(fn):
with contextlib.suppress(FileNotFoundError):
os.remove(os.path.join(out_dir, fn))
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Run FLUX.1-dev text-to-image with FastVideo VideoGenerator.",
)
p.add_argument(
"--model-path",
default="official_weights/FLUX.1-dev",
help="Local Diffusers checkpoint dir or HF repo id.",
)
p.add_argument(
"--out-dir",
"--outdir",
default="outputs/flux_dev/samples",
help="Directory for saved PNG outputs.",
)
p.add_argument(
"--prompt",
action="append",
default=None,
help="Prompt. Repeat for multiple images.",
)
p.add_argument(
"--backend",
default=None,
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA).",
)
p.add_argument("--seed", type=int, default=42, help="Base seed; each prompt uses seed + index.")
p.add_argument("--height", type=int, default=1024, help="Output height.")
p.add_argument("--width", type=int, default=1024, help="Output width.")
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
p.add_argument("--guidance", type=float, default=3.5, help="Guidance scale.")
p.add_argument("--num-gpus", type=int, default=1, help="GPU count.")
return p.parse_args()
def main() -> None:
args = parse_args()
prompts: list[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
if args.backend:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
from fastvideo import VideoGenerator
os.makedirs(args.out_dir, exist_ok=True)
init_kwargs = {
"num_gpus": args.num_gpus,
"workload_type": "t2i",
"sp_size": 1,
"tp_size": 1,
"dit_cpu_offload": False,
"dit_layerwise_offload": False,
"text_encoder_cpu_offload": False,
"vae_cpu_offload": False,
"image_encoder_cpu_offload": False,
"pin_cpu_memory": False,
"use_fsdp_inference": False,
}
generator = VideoGenerator.from_pretrained(
model_path=args.model_path,
**init_kwargs,
)
try:
for i, prompt in enumerate(prompts):
seed = args.seed + i
filename_base = (
f"flux_dev_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
)
_remove_existing_outputs(args.out_dir, filename_base)
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
print(f"[flux] prompt_idx={i} seed={seed} output_path={output_path}")
generation_kwargs = {
"output_path": output_path,
"height": args.height,
"width": args.width,
"num_frames": 1,
"fps": 1,
"num_inference_steps": args.steps,
"guidance_scale": args.guidance,
"use_embedded_guidance": True,
"true_cfg_scale": 1.0,
"seed": seed,
"save_video": True,
}
generator.generate_video(prompt, **generation_kwargs)
print(f"[flux] done. outputs written to: {args.out_dir}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+107
View File
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""Run GLM-Image text-to-image generation through FastVideo.
User story:
"I have the HF `zai-org/GLM-Image` checkpoint and want a minimal
text-to-image generation command, saved as a PNG."
"""
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run GLM-Image text-to-image generation.")
parser.add_argument(
"--model-path",
default="zai-org/GLM-Image",
help="HF id or local diffusers-format GLM-Image weights directory.",
)
parser.add_argument(
"--output",
default="image_output/landscape.png",
help="Output PNG path.",
)
parser.add_argument(
"--prompt",
default=("A beautiful landscape photography with rolling hills, "
"a winding river, and a vibrant sunset in the background. "
"Warm golden light, photorealistic style."),
help="Text prompt.",
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--guidance-scale", type=float, default=1.5)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
# pipeline class come from the model's registered defaults — don't override.
generator_config = GeneratorConfig(
model_path=args.model_path,
trust_remote_code=True,
engine=EngineConfig(
num_gpus=args.num_gpus,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
save_video=False,
return_frames=True,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
frames = result.frames
if frames is not None and len(frames):
Image.fromarray(frames[0]).save(output)
print(f"Saved image to {output}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+120
View File
@@ -0,0 +1,120 @@
# SPDX-License-Identifier: Apache-2.0
"""Run GLM-Image image-to-image (edit) generation through FastVideo.
User story:
"I have the HF `zai-org/GLM-Image` checkpoint and a condition image, and
want a minimal edit command (text + image -> edited image), saved as a PNG."
GLM-Image is a single unified pipeline: passing a condition image switches it
from text-to-image to the edit path (the condition enters the DiT via a KV-cache
write pass), so the generator config is identical to `basic_glm_image.py` — the
`inputs.pil_image` on the request is what selects the edit mode.
"""
import argparse
from pathlib import Path
from PIL import Image
from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
InputConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run GLM-Image image-to-image (edit) generation.")
parser.add_argument(
"--model-path",
default="zai-org/GLM-Image",
help="HF id or local diffusers-format GLM-Image weights directory.",
)
parser.add_argument(
"--image",
default="assets/images/couple.jpg",
help="Condition image to edit.",
)
parser.add_argument(
"--output",
default="image_output/edited.png",
help="Output PNG path.",
)
parser.add_argument(
"--prompt",
default="Change the background to a snowy mountain landscape at golden hour.",
help="Edit instruction.",
)
parser.add_argument("--height", type=int, default=1024)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--steps", type=int, default=50)
parser.add_argument("--guidance-scale", type=float, default=1.5)
parser.add_argument("--seed", type=int, default=1024)
parser.add_argument("--num-gpus", type=int, default=1)
parser.add_argument("--tp-size", type=int, default=None)
parser.add_argument("--sp-size", type=int, default=None)
return parser.parse_args()
def main() -> None:
args = parse_args()
output = Path(args.output)
output.parent.mkdir(parents=True, exist_ok=True)
condition = Image.open(args.image).convert("RGB")
tp_size = args.tp_size if args.tp_size is not None else (args.num_gpus if args.num_gpus > 1 else 1)
sp_size = args.sp_size if args.sp_size is not None else (1 if args.num_gpus > 1 else args.num_gpus)
# GLM-Image needs trust_remote_code for its AR encoder; offload and the
# pipeline class come from the model's registered defaults — don't override.
# The pipeline is registered as t2i; passing inputs.pil_image below switches
# it to the edit path.
generator_config = GeneratorConfig(
model_path=args.model_path,
trust_remote_code=True,
engine=EngineConfig(
num_gpus=args.num_gpus,
parallelism=ParallelismConfig(tp_size=tp_size, sp_size=sp_size),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
request = GenerationRequest(
prompt=args.prompt,
inputs=InputConfig(pil_image=condition),
sampling=SamplingConfig(
height=args.height,
width=args.width,
num_frames=1,
fps=1,
num_inference_steps=args.steps,
guidance_scale=args.guidance_scale,
seed=args.seed,
),
output=OutputConfig(
output_path=str(output.parent),
save_video=False,
return_frames=True,
),
)
result = generator.generate(request)
if isinstance(result, list):
result = result[0]
frames = result.frames
if frames is not None and len(frames):
Image.fromarray(frames[0]).save(output)
print(f"Saved image to {output}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,105 @@
# AnyFlow on-policy DMD — Wan 2.1 T2V 1.3B.
#
# Stage 2 of the AnyFlow two-stage recipe. Continues from the pretrain
# checkpoint; refines the student via DMD2 with a multi-step Euler-flow
# rollout from pure noise. Teacher provides the real score, critic
# learns the fake score; both inherited from DMD2Method.
#
# Replace <PATH_TO_PRETRAIN_CKPT> with the output of the pretrain stage,
# or with the NVIDIA-released checkpoint
# nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers to bootstrap directly from
# the paper weights (the delta_embedder rename is handled by the
# param_names_mapping in WanVideoArchConfig).
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: <PATH_TO_PRETRAIN_CKPT>
trainable: true
teacher:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-14B-Diffusers
trainable: false
disable_custom_init_weights: true
critic:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.distribution_matching.anyflow.AnyFlowMethod
rollout_mode: simulate
generator_update_interval: 5
real_score_guidance_scale: 3.0
dmd_denoising_steps: [999, 937, 833, 624]
warp_denoising_step: false
# AnyFlow rollout knobs.
student_sample_steps: 4
use_mean_velocity: true
t_list_override: [999.0, 937.0, 833.0, 624.0, 0.0]
dmd_score_r_value: 0.0 # DMD scoring conditioning is at r=0 (consistency target).
# Critic optimizer (DMD2 inherited).
fake_score_learning_rate: 8.0e-6
fake_score_betas: [0.0, 0.999]
fake_score_lr_scheduler: constant
attn_kind: vsa
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/preprocessed
dataloader_num_workers: 4
train_batch_size: 1
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 2.0e-6
betas: [0.0, 0.999]
weight_decay: 0.01
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 4000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_anyflow_onpolicy
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: anyflow-wan
run_name: wan2.1_t2v_anyflow_onpolicy
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
pipeline:
flow_shift: 5.0
dit_config:
r_embedder: true
r_embedder_fusion: gated
r_embedder_gate_value: 0.25
r_embedder_deltatime_type: r
@@ -0,0 +1,83 @@
# AnyFlow pretrain (flow-map central-difference) — Wan 2.1 T2V 1.3B.
#
# Stage 1 of the AnyFlow two-stage recipe. Trains the dual-timestep
# u_θ(x_t, t, r) on the central-difference target so the same checkpoint
# can be sampled at arbitrary NFE in the on-policy stage.
#
# Initialize from base Wan 2.1 T2V 1.3B. No teacher or critic at this
# stage; AnyFlowPretrainMethod owns a single student + one optimizer.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
method:
_target_: fastvideo.train.methods.distribution_matching.anyflow_pretrain.AnyFlowPretrainMethod
diffusion_ratio: 0.5
consistency_ratio: 0.25
epsilon: 5 # finite-difference step in absolute train-timestep units
weight_type: beta08 # per-timestep loss weight = t * sqrt(1 - t), renormalized
fuse_guidance_scale: 3.0
# shift is taken from pipeline.flow_shift below.
training:
distributed:
num_gpus: 8
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/preprocessed
dataloader_num_workers: 4
train_batch_size: 4
training_cfg_rate: 0.0
seed: 1000
num_latent_t: 21
num_height: 480
num_width: 832
num_frames: 81
optimizer:
learning_rate: 5.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 6000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/wan2.1_anyflow_pretrain
training_state_checkpointing_steps: 500
checkpoints_total_limit: 3
resume_from_checkpoint: latest
tracker:
project_name: anyflow-wan
run_name: wan2.1_t2v_anyflow_pretrain
model:
enable_gradient_checkpointing_type: full
callbacks:
grad_clip:
_target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback
max_grad_norm: 1.0
pipeline:
flow_shift: 5.0
dit_config:
# Enable AnyFlow dual-timestep conditioning. The student loads from
# base Wan 2.1 — its checkpoint has no delta_embedder weights, so they
# get initialized identically to time_embedder via deep-copy in
# WanTimeTextImageEmbedding.__init__.
r_embedder: true
r_embedder_fusion: gated
r_embedder_gate_value: 0.25
r_embedder_deltatime_type: r
+27
View File
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from fastvideo.api.sampling_param import SamplingParam
@dataclass
class FluxSamplingParam(SamplingParam):
prompt: str | None = "a photo of a cat"
negative_prompt: str = ""
num_videos_per_prompt: int = 1
seed: int = 0
num_frames: int = 1
height: int = 1024
width: int = 1024
fps: int = 1
num_inference_steps: int = 28
guidance_scale: float = 3.5
use_embedded_guidance: bool = True
true_cfg_scale: float = 1.0
+16
View File
@@ -90,6 +90,10 @@ class SamplingParam:
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
# Embedded guidance (FLUX): do not treat ``guidance_scale > 1`` as classic CFG.
use_embedded_guidance: bool = False
# Diffusers-style true CFG for FLUX when > 1 (requires negative prompt encoding).
true_cfg_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
sigmas: list[float] | None = None
@@ -325,6 +329,18 @@ class SamplingParam:
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--use-embedded-guidance",
action="store_true",
default=SamplingParam.use_embedded_guidance,
help="Use embedded guidance scale (FLUX-style) instead of classic CFG",
)
parser.add_argument(
"--true-cfg-scale",
type=float,
default=SamplingParam.true_cfg_scale,
help="True CFG scale for FLUX when > 1 (requires negative prompt encoding)",
)
parser.add_argument(
"--boundary-ratio",
type=float,
+1
View File
@@ -150,6 +150,7 @@ class SamplingConfig:
guidance_scale_2: float | None = None
guidance_rescale: float = 0.0
true_cfg_scale: float | None = None
use_embedded_guidance: bool | None = None
boundary_ratio: float | None = None
sigmas: list[float] | None = None
+15 -118
View File
@@ -1,13 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
import importlib.util
import os
import torch
import torch.nn.functional as F
from dataclasses import dataclass
from fastvideo import envs
from fastvideo.attention.utils.flash_attn_default import (
fa_version,
flash_attn_func_compilable,
)
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
@@ -17,119 +19,6 @@ from fastvideo.attention.backends.abstract import (
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1, mirroring the
# kernel package's FASTVIDEO_VSA_CUTEDSL: its CuTeDSL kernels JIT-compile per
# shape family and can fail at runtime on some arch/shape combinations, so it
# is never auto-selected just because it is installed. Below sm90 a capability
# gate in flash_attn_cute routes to FA2 the calls FA4 cannot serve there:
# grad-enabled (its backward asserts sm90+) and GQA (pack_gqa fails CuTeDSL
# JIT, observed on sm_89).
if envs.FASTVIDEO_FA4:
try:
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
except ImportError as e:
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
fa_version = "4"
else:
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 no longer have a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "3"
except ImportError:
from flash_attn import flash_attn_func as flash_attn_2_func
flash_attn_func = flash_attn_2_func
fa_version = "2"
try:
if importlib.util.find_spec("flash_attn.cute") is not None:
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
"set FASTVIDEO_FA4=1 to use it for inference.")
except ImportError:
pass
# torch.compile traceability: the FA4/cute path (fa_version=="4") is
# already a registered torch.library custom op, so dynamo treats it as a
# graph node. The external FA2/FA3 `flash_attn_func` is NOT — dynamo
# breaks the graph at the call site (observed: wanvideo.py self-attn,
# once per layer every step), which fragments the compiled region and
# blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom
# op (mirrors the FP4 `flash_attn_cute` template) so it becomes an
# opaque-but-traceable node. The kernel still runs eager inside the op
# (correct — flash-attn must run eager); only dynamo's treatment of the
# boundary changes, so numerics are unchanged (SSIM-gate to confirm).
if fa_version in ("2", "3"):
_fa_default = flash_attn_func
# Scope: this op covers exactly the q/k/v + softmax_scale + causal
# call shape used by FlashAttentionImpl.forward's default branch
# (see `flash_attn_func_compilable(...)` call site below). The
# masked/no-pad and varlen / cross-attn paths use different
# entry points (`flash_attn_no_pad`, `flash_attn_varlen_*`) which
# are intentionally out of scope for this PR — wrapping them is a
# natural follow-up. The wrapper's signature is the contract: any
# extra kwarg (dropout_p, window_size, alibi_slopes, deterministic,
# return_attn_probs, ...) raises TypeError at the call site, so
# silent loss of kwargs is not a failure mode.
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
del softmax_scale, causal
# FA2/FA3 default path: [batch, seqlen_q, nheads, head_dim_v],
# same dtype/device as q (head dim taken from v).
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Autograd carve-out. The custom op above registers a forward + fake
# kernel but NO backward (register_autograd), so it is opaque to
# autograd. Inference runs under no_grad / inference_mode and routes
# through the traceable custom op — that is the torch.compile win, and
# the only path this PR claims. Training backprops through attention,
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
# (itself an autograd.Function, so backward is correct) at the cost of a
# dynamo graph break on the training path — i.e. pre-PR behavior, no
# regression. Full autograd parity for the custom op (mirroring the FP4
# cute template) is a tracked follow-up.
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
elif fa_version == "4":
# FA4 path: `flash_attn_func` (from `flash_attn_cute`) goes through a
# registered torch.library custom op (with an FA4 backward on sm90+;
# grad-enabled and GQA calls below sm90 route to FA2), so a passthrough
# is enough — no extra registration needed.
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
else:
# Defensive: the probe above only ever sets fa_version to "2", "3",
# or "4"; an unexpected value means an import/probe regression and
# we want a loud error at import, not a silent NameError later.
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
f"'2', '3', or '4' from the import probe above.")
logger.info("Using FlashAttention-%s backend", fa_version)
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
@@ -309,9 +198,17 @@ class FlashAttentionImpl(AttentionImpl):
attn_metadata: FlashAttnMetadata,
):
if (attn_metadata is not None and hasattr(attn_metadata, "attn_mask") and attn_metadata.attn_mask is not None):
# Route through the *_compilable wrappers so dynamo sees one
# traceable node for each masked entry point (the unpad/pad
# bookkeeping runs eager inside the custom op). On FA2 these
# wrappers go through ops with full register_autograd, so
# training also backprops through the op (no graph break on
# the training path); on FA3/FA4 they carve out to the
# autograd.Function for grad-enabled calls — see
# fastvideo/attention/utils/flash_attn_no_pad.py.
from fastvideo.attention.utils.flash_attn_no_pad import (
flash_attn_no_pad,
flash_attn_varlen_qk_no_pad,
flash_attn_no_pad_compilable as flash_attn_no_pad,
flash_attn_varlen_qk_no_pad_compilable as flash_attn_varlen_qk_no_pad,
)
attn_mask = attn_metadata.attn_mask
@@ -0,0 +1,251 @@
# SPDX-License-Identifier: Apache-2.0
"""torch.compile-traceable wrapper for the FA2/FA3/FA4 default attention path.
The FA4/cute path (`fa_version == "4"`) is already a registered
`torch.library.custom_op` in `fastvideo.attention.utils.flash_attn_cute`, so
dynamo treats it as a graph node. The external FA2/FA3 ``flash_attn_func`` is
NOT — dynamo breaks the graph at the call site (observed: wanvideo.py
self-attn, once per layer every step), which fragments the compiled region
and blocks CUDA-graph capture. Wrap the FA2/FA3 default call in a custom op
(mirrors the FP4 `flash_attn_cute` template) so it becomes an
opaque-but-traceable node. The kernel still runs eager inside the op
(correct — flash-attn must run eager); only dynamo's treatment of the
boundary changes, so numerics are unchanged (SSIM-gated).
Autograd: FA2 has full ``register_autograd`` parity — the custom op's
backward calls flash_attn's ``_flash_attn_backward`` directly, so training
backprops *through* the op (no graph break on the training path either).
FA3 currently keeps the no-backward + carve-out pattern from PR #1373
because FA3's private backward signature wants validation on a real Hopper
box (gated on Kuan-Hao's Modal FA3 setup PR). Once that lands the FA3 path
can mirror FA2.
Lives in `attention/utils/` (sibling of `flash_attn_cute.py` and
`flash_attn_no_pad.py`) so it can be imported by any backend that wants the
traceable FA default call without pulling in backend dispatch logic. The
backend (`attention/backends/flash_attn.py`) just imports
`flash_attn_func_compilable` and `fa_version` from here.
"""
import importlib.util
import torch
from fastvideo import envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# Pick the same backend the rest of FastVideo picked for `flash_attn_func`
# (FA4/cute → FA3 → FA2). Mirror the precedence used in
# `attention/utils/flash_attn_no_pad.py` so the two probes always agree.
#
# FA4 (flash_attn.cute) is explicit opt-in via FASTVIDEO_FA4=1: its CuTeDSL
# kernels JIT-compile per shape family and can fail at runtime on some
# arch/shape combinations, so it is never auto-selected just because it is
# installed.
if envs.FASTVIDEO_FA4:
try:
from fastvideo.attention.utils.flash_attn_cute import flash_attn_func
except ImportError as e:
raise RuntimeError(f"FASTVIDEO_FA4=1 but flash_attn.cute (FA4) is not usable ({e}); "
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
fa_version = "4"
else:
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
# flash_attn 3 no longer has a different API, see following commit:
# https://github.com/Dao-AILab/flash-attention/commit/ed209409acedbb2379f870bbd03abce31a7a51b7
flash_attn_func = flash_attn_3_func
fa_version = "3"
except ImportError:
from flash_attn import flash_attn_func as flash_attn_2_func
flash_attn_func = flash_attn_2_func
fa_version = "2"
try:
if importlib.util.find_spec("flash_attn.cute") is not None:
logger.info("flash_attn.cute (FA4) is installed but not enabled; "
"set FASTVIDEO_FA4=1 to use it for inference.")
except ImportError:
pass
if fa_version == "2":
# Scope: this op covers exactly the q/k/v + softmax_scale + causal call
# shape used by FlashAttentionImpl.forward's default branch (see
# `flash_attn_func_compilable(...)` call site in
# `attention/backends/flash_attn.py`). The masked/no-pad and varlen /
# cross-attn paths use different entry points
# (`flash_attn_no_pad`, `flash_attn_varlen_*`) which live in
# `attention/utils/flash_attn_no_pad.py`. The wrapper's signature is the
# contract: any extra kwarg (dropout_p, window_size, alibi_slopes,
# deterministic, return_attn_probs, ...) raises TypeError at the call
# site, so silent loss of kwargs is not a failure mode.
from flash_attn.flash_attn_interface import _flash_attn_backward as _fa2_backward
_fa_default = flash_attn_func
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
# `return_attn_probs=True` asks FA2 to also return softmax_lse +
# S_dmask. We need softmax_lse to feed the backward; S_dmask is the
# dropout mask (always None here since dropout_p is fixed at 0).
out, softmax_lse, _ = _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal, return_attn_probs=True)
return out, softmax_lse
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del softmax_scale, causal
# FA2 default path: out = [batch, seqlen_q, nheads, head_dim_v],
# softmax_lse = [batch, nheads, seqlen_q], fp32 regardless of q dtype.
b, sq, hq = q.shape[0], q.shape[1], q.shape[2]
out = q.new_empty(b, sq, hq, v.shape[-1])
lse = q.new_empty(b, hq, sq, dtype=torch.float32)
return out, lse
def _flash_attn_default_setup_context(ctx, inputs, output):
q, k, v, softmax_scale, causal = inputs
out, lse = output
ctx.save_for_backward(q, k, v, out, lse)
# `lse` is an auxiliary output we save to feed FA2's backward; nobody
# should differentiate through it. Mark it non-differentiable so
# autograd errors loudly if a caller wires it into a loss, rather
# than silently producing zero/None grads through the `del grad_lse`
# in our backward.
ctx.mark_non_differentiable(lse)
# FA2's *forward* substitutes `1 / sqrt(head_dim)` for `softmax_scale=None`
# internally; FA2's *backward* (`_flash_attn_backward`) demands a concrete
# float in its C++ schema and rejects None at the binding boundary. Resolve
# the default here so the value saved on ctx (and passed to backward) is
# always a real float — matches what FA2's own autograd.Function does.
if softmax_scale is None:
softmax_scale = q.shape[-1]**-0.5
ctx.softmax_scale = softmax_scale
ctx.causal = causal
def _flash_attn_default_backward(ctx, grad_out, grad_lse):
# We only differentiate `out`; softmax_lse is saved-for-backward, not
# a real differentiable output. (Mirrors the FP4 cute template.)
del grad_lse
q, k, v, out, lse = ctx.saved_tensors
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
# FA2's `_flash_attn_backward` writes into dq/dk/dv in place. The
# extra kwargs (window_size_*, softcap, alibi_slopes, deterministic,
# rng_state) are pinned to the same defaults the forward wrapper
# uses — flash-attn==2.8.1 (the version FastVideo pins) requires
# all of them explicitly. `rng_state=None` is correct for our
# `dropout_p=0` configuration.
_fa2_backward(
grad_out,
q,
k,
v,
out,
lse,
dq,
dk,
dv,
dropout_p=0.0,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
alibi_slopes=None,
deterministic=False,
rng_state=None,
)
return dq, dk, dv, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_default_forward",
_flash_attn_default_backward,
setup_context=_flash_attn_default_setup_context,
)
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Backward is registered: autograd flows through the op (training
# path is also traceable; no carve-out needed). Public API matches
# `flash_attn_func` — returns just `out`; we drop the saved-for-
# backward `lse` here so callers see the original single-tensor
# contract.
out, _ = torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
return out
elif fa_version == "3":
# FA3 path: same forward+fake custom op as the original PR #1373, with
# the autograd carve-out kept. The full backward (mirroring the FA2 leg
# above) wants a Hopper box for grad-check validation, which we don't
# have until Kuan-Hao's Modal FA3 setup PR lands. Until then this keeps
# inference traceable + training correct (via the original
# autograd.Function path + a pre-PR-style graph break on training).
_fa_default = flash_attn_func
@torch.library.custom_op(
"fastvideo::_flash_attn_default_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_default_forward(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
@torch.library.register_fake("fastvideo::_flash_attn_default_forward")
def _flash_attn_default_forward_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
softmax_scale: float | None,
causal: bool,
) -> torch.Tensor:
del softmax_scale, causal
return q.new_empty(q.shape[0], q.shape[1], q.shape[2], v.shape[-1])
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
# Autograd carve-out. The custom op above registers a forward + fake
# kernel but NO backward (register_autograd), so it is opaque to
# autograd. Inference runs under no_grad / inference_mode and routes
# through the traceable custom op — that is the torch.compile win, and
# the only path this PR claims. Training backprops through attention,
# so route grad-enabled calls to the original FA2/FA3 `flash_attn_func`
# (itself an autograd.Function, so backward is correct) at the cost of a
# dynamo graph break on the training path — i.e. pre-PR behavior, no
# regression. Full autograd parity for the custom op (mirroring the FP4
# cute template) is a tracked follow-up.
if torch.is_grad_enabled() and (q.requires_grad or k.requires_grad or v.requires_grad):
return _fa_default(q, k, v, softmax_scale=softmax_scale, causal=causal)
return torch.ops.fastvideo._flash_attn_default_forward(q, k, v, softmax_scale, causal)
elif fa_version == "4":
# FA4 path: `flash_attn_func` is already a torch.library custom op
# (registered in `fastvideo.attention.utils.flash_attn_cute`), so a
# passthrough is enough — no extra registration needed.
def flash_attn_func_compilable(q, k, v, softmax_scale=None, causal=False):
return flash_attn_func(q, k, v, softmax_scale=softmax_scale, causal=causal)
else:
# Defensive: the probe above only ever sets fa_version to "2", "3",
# or "4"; an unexpected value means an import/probe regression and
# we want a loud error at import, not a silent NameError later.
raise RuntimeError(f"Unsupported FlashAttention version: {fa_version!r} — expected "
f"'2', '3', or '4' from the import probe above.")
+483 -5
View File
@@ -24,7 +24,7 @@ from flash_attn.bert_padding import pad_input, unpad_input
from fastvideo import envs
def _resolve_flash_attn_varlen_func() -> Any:
def _resolve_flash_attn_varlen_func() -> tuple[Any, str]:
if envs.FASTVIDEO_FA4:
# FA4 cute is explicit opt-in (see fastvideo/attention/backends/
# flash_attn.py); with FASTVIDEO_FA4=1 an unimportable FA4 build must
@@ -39,20 +39,28 @@ def _resolve_flash_attn_varlen_func() -> Any:
"fix the FA4 install (see the flash-attn-4 pin in pyproject.toml) "
"or unset FASTVIDEO_FA4.") from e
return flash_attn_varlen_func_cute
return flash_attn_varlen_func_cute, "4"
try:
from flash_attn_interface import (
flash_attn_varlen_func as flash_attn_varlen_func_interface, )
return flash_attn_varlen_func_interface
return flash_attn_varlen_func_interface, "3"
except ImportError:
from flash_attn import (
flash_attn_varlen_func as flash_attn_varlen_func_flash, )
return flash_attn_varlen_func_flash
return flash_attn_varlen_func_flash, "2"
flash_attn_varlen_func_impl = _resolve_flash_attn_varlen_func()
flash_attn_varlen_func_impl, _FA_VARLEN_VERSION = _resolve_flash_attn_varlen_func()
# FA2-only: the private varlen backward we register against the custom ops
# below. FA3 / FA4 have different private signatures and validation paths
# (Hopper / Blackwell boxes) — those legs keep the autograd carve-out
# pattern from PR #1373 until their setup PRs land.
if _FA_VARLEN_VERSION == "2":
from flash_attn.flash_attn_interface import (
_flash_attn_varlen_backward as _fa2_varlen_backward, )
def flash_attn_no_pad(
@@ -191,3 +199,473 @@ def flash_attn_varlen_qk_no_pad(
h=nheads,
)
return output
# ---------------------------------------------------------------------------
# torch.compile traceability + register_autograd parity for the masked /
# varlen attention paths.
#
# Wraps the two entry points `FlashAttentionImpl.forward` calls
# (`flash_attn_no_pad`, `flash_attn_varlen_qk_no_pad`) as
# `torch.library.custom_op`s so dynamo sees one traceable node — the
# internal unpad / pad bookkeeping (data-dependent `nnz` shapes) runs
# eager inside the op, and the op's outputs are the statically-shaped
# padded tensors. This mirrors the FA2 default-path wrapper in
# `fastvideo/attention/backends/flash_attn.py`.
#
# Autograd: on FA2 we register a real backward (`register_autograd`)
# that calls FA2's `_flash_attn_varlen_backward` on the unpadded form
# — re-unpadding the saved padded tensors using the saved mask. The
# `softmax_lse` from the varlen forward is naturally unpadded
# (`[nheads, total_q]`); we pad it to `[batch, nheads, seqlen]` on
# the way out (statically shaped) and re-unpad in backward. So
# training backprops *through* the op (no graph break on the training
# path either).
#
# FA3 / FA4 keep the autograd carve-out pattern from PR #1373: the
# custom op has forward + fake only, and `*_compilable` falls back to
# the original autograd.Function for grad-enabled calls. Those legs
# are gated on Hopper-class / Blackwell-class boxes for backward
# validation and ship as separate follow-ups.
if _FA_VARLEN_VERSION == "2":
# ---------- masked self-attention: flash_attn_no_pad (FA2) ----------
@torch.library.custom_op(
"fastvideo::_flash_attn_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_no_pad_forward(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
b, s, _three, h, d = qkv.shape
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
out_unpad, lse_unpad, _ = flash_attn_varlen_qkvpacked_func(x_unpad,
cu_seqlens,
max_s,
dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
return_attn_probs=True)
# Pad out: [nnz, h, d] -> [b, s, h, d]
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), indices, b, s),
"b s (h d) -> b s h d",
h=h)
# Pad lse: FA2 varlen returns [nheads, total_q]. Transpose to [total_q,
# nheads], pad to [b, s, nheads], permute to [b, nheads, s] — statically
# shaped so register_fake matches.
lse_padded = pad_input(lse_unpad.t().contiguous(), indices, b, s).permute(0, 2, 1).contiguous()
return out_padded, lse_padded
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
def _flash_attn_no_pad_forward_fake(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
b, s, _three, h, d = qkv.shape
out = qkv.new_empty(b, s, h, d)
lse = qkv.new_empty(b, h, s, dtype=torch.float32)
return out, lse
def _flash_attn_no_pad_setup_context(ctx, inputs, output):
qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic = inputs
out, lse = output
ctx.save_for_backward(qkv, out, lse, key_padding_mask)
# Auxiliary output, not differentiable — see default-path note.
ctx.mark_non_differentiable(lse)
# FA2's varlen backward requires a concrete float for softmax_scale.
if softmax_scale is None:
softmax_scale = qkv.shape[-1]**-0.5 # head_dim from qkv's last dim
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.dropout_p = dropout_p
ctx.deterministic = deterministic
def _flash_attn_no_pad_backward(ctx, grad_out, grad_lse):
# lse is saved-for-backward, not differentiated.
del grad_lse
qkv, out_padded, lse_padded, key_padding_mask = ctx.saved_tensors
b, s, _three, h, d = qkv.shape
# One `unpad_input` call (on qkv) gives us indices + cu_seqlens + max_s;
# reuse those for out / dout / lse below via direct indexing instead
# of redundant `unpad_input` calls (each of which would re-run
# `nonzero` + `cumsum` + a `.max().item()` GPU→CPU sync).
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, _ = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=h)
q_unpad, k_unpad, v_unpad = (t.contiguous() for t in x_unpad.unbind(dim=1))
# Direct-index variants reuse `indices` (computed above).
out_unpad = out_padded.flatten(0, 1)[indices].view(-1, h, d).contiguous()
dout_unpad = grad_out.flatten(0, 1)[indices].view(-1, h, d).contiguous()
# lse_padded [b, h, s] -> [b, s, h] -> [nnz, h] -> [h, nnz].
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[indices].t().contiguous()
dq_unpad = torch.empty_like(q_unpad)
dk_unpad = torch.empty_like(k_unpad)
dv_unpad = torch.empty_like(v_unpad)
_fa2_varlen_backward(
dout_unpad,
q_unpad,
k_unpad,
v_unpad,
out_unpad,
lse_unpad,
dq_unpad,
dk_unpad,
dv_unpad,
cu_seqlens_q=cu_seqlens,
cu_seqlens_k=cu_seqlens,
max_seqlen_q=max_s,
max_seqlen_k=max_s,
dropout_p=ctx.dropout_p,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
alibi_slopes=None,
deterministic=ctx.deterministic,
rng_state=None,
)
# Re-pad each grad and stack into dqkv.
def _repad(dt_unpad: torch.Tensor) -> torch.Tensor:
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, b, s)
return rearrange(padded, "b s (h d) -> b s h d", h=h)
dqkv = torch.stack([_repad(dq_unpad), _repad(dk_unpad), _repad(dv_unpad)], dim=2)
# 6 inputs total: qkv, key_padding_mask, causal, dropout_p, softmax_scale, deterministic.
return dqkv, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_no_pad_forward",
_flash_attn_no_pad_backward,
setup_context=_flash_attn_no_pad_setup_context,
)
# ---------- cross-attention: flash_attn_varlen_qk_no_pad (FA2) ----------
@torch.library.custom_op(
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_varlen_qk_no_pad_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
b, sq, h, d = query.shape
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
query_padding_mask)
k_unpad, _, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
key_padding_mask)
v_unpad, _, _, _, _ = unpad_input(rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h)
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h)
v_unpad = rearrange(v_unpad, "nnz (h d) -> nnz h d", h=h)
out_unpad, lse_unpad, _ = flash_attn_varlen_func_impl(q_unpad,
k_unpad,
v_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_k,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
return_attn_probs=True)
# Pad out: [nnz_q, h, d] -> [b, sq, h, d]
out_padded = rearrange(pad_input(rearrange(out_unpad, "nnz h d -> nnz (h d)"), q_indices, b, sq),
"b s (h d) -> b s h d",
h=h)
# Pad lse: [h, nnz_q] -> [b, h, sq]
lse_padded = pad_input(lse_unpad.t().contiguous(), q_indices, b, sq).permute(0, 2, 1).contiguous()
return out_padded, lse_padded
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
def _flash_attn_varlen_qk_no_pad_forward_fake(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> tuple[torch.Tensor, torch.Tensor]:
del key, query_padding_mask, key_padding_mask
del causal, dropout_p, softmax_scale, deterministic
b, sq, h, _ = query.shape
# `out`'s head_dim comes from value (d_v), matching the real forward's
# out_padded ([b, sq, h, d_v]); it can differ from query's d_q.
out = query.new_empty(b, sq, h, value.shape[-1])
lse = query.new_empty(b, h, sq, dtype=torch.float32)
return out, lse
def _flash_attn_varlen_qk_no_pad_setup_context(ctx, inputs, output):
(query, key, value, query_padding_mask, key_padding_mask, causal, dropout_p, softmax_scale,
deterministic) = inputs
out, lse = output
ctx.save_for_backward(query, key, value, out, lse, query_padding_mask, key_padding_mask)
# Auxiliary output, not differentiable — see default-path note.
ctx.mark_non_differentiable(lse)
if softmax_scale is None:
softmax_scale = query.shape[-1]**-0.5
ctx.softmax_scale = softmax_scale
ctx.causal = causal
ctx.dropout_p = dropout_p
ctx.deterministic = deterministic
def _flash_attn_varlen_qk_no_pad_backward(ctx, grad_out, grad_lse):
del grad_lse
(query, key, value, out_padded, lse_padded, query_padding_mask, key_padding_mask) = ctx.saved_tensors
b, sq, h, d = query.shape
sk = key.shape[1]
# One `unpad_input` call per distinct mask; reuse the returned
# indices via direct indexing for everything else that shares
# the same mask (v with k_mask; out/dout/lse with q_mask; the
# final repad of dk/dv also reuses k_indices). Avoids ~4
# redundant `unpad_input` calls + their GPU→CPU `.max().item()`
# syncs.
q_unpad, q_indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(rearrange(query, "b s h d -> b s (h d)"),
query_padding_mask)
k_unpad, k_indices, cu_seqlens_k, max_seqlen_k, _ = unpad_input(rearrange(key, "b s h d -> b s (h d)"),
key_padding_mask)
q_unpad = rearrange(q_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
k_unpad = rearrange(k_unpad, "nnz (h d) -> nnz h d", h=h).contiguous()
v_unpad = value.flatten(0, 1)[k_indices].view(-1, h, d).contiguous()
# out / dout / lse follow q's shape, so index with q_indices.
out_unpad = out_padded.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
dout_unpad = grad_out.flatten(0, 1)[q_indices].view(-1, h, d).contiguous()
lse_unpad = lse_padded.permute(0, 2, 1).contiguous().flatten(0, 1)[q_indices].t().contiguous()
dq_unpad = torch.empty_like(q_unpad)
dk_unpad = torch.empty_like(k_unpad)
dv_unpad = torch.empty_like(v_unpad)
_fa2_varlen_backward(
dout_unpad,
q_unpad,
k_unpad,
v_unpad,
out_unpad,
lse_unpad,
dq_unpad,
dk_unpad,
dv_unpad,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=ctx.dropout_p,
softmax_scale=ctx.softmax_scale,
causal=ctx.causal,
window_size_left=-1,
window_size_right=-1,
softcap=0.0,
alibi_slopes=None,
deterministic=ctx.deterministic,
rng_state=None,
)
# k_indices is already available from the unpad_input above —
# no need to recompute it for the dk/dv repad.
def _repad(dt_unpad: torch.Tensor, indices: torch.Tensor, batch: int, seqlen: int) -> torch.Tensor:
padded = pad_input(rearrange(dt_unpad, "nnz h d -> nnz (h d)"), indices, batch, seqlen)
return rearrange(padded, "b s (h d) -> b s h d", h=h)
dq_padded = _repad(dq_unpad, q_indices, b, sq)
dk_padded = _repad(dk_unpad, k_indices, b, sk)
dv_padded = _repad(dv_unpad, k_indices, b, sk)
# 9 inputs total: query, key, value, q_mask, k_mask, causal, dropout_p,
# softmax_scale, deterministic.
return dq_padded, dk_padded, dv_padded, None, None, None, None, None, None
torch.library.register_autograd(
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
_flash_attn_varlen_qk_no_pad_backward,
setup_context=_flash_attn_varlen_qk_no_pad_setup_context,
)
# ---------- public dispatchers (FA2: autograd flows through the op) -----
def flash_attn_no_pad_compilable(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
"""dynamo-traceable wrapper around ``flash_attn_no_pad`` (registered op,
full register_autograd on FA2 — both inference and training go through
the op, no graph break on either)."""
out, _ = torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
deterministic)
return out
def flash_attn_varlen_qk_no_pad_compilable(query,
key,
value,
query_padding_mask,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
"""dynamo-traceable wrapper around ``flash_attn_varlen_qk_no_pad`` (registered
op, full register_autograd on FA2)."""
out, _ = torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
key_padding_mask, causal, dropout_p,
softmax_scale, deterministic)
return out
else:
# ---------- FA3 / FA4: carve-out (forward+fake only, no real backward) ---
# Same pattern as the parked varlen-extension and the FA3 default leg in
# `fastvideo/attention/backends/flash_attn.py`. Real backward for these
# versions is a follow-up gated on Hopper / Blackwell box validation.
@torch.library.custom_op(
"fastvideo::_flash_attn_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_no_pad_forward(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
return flash_attn_no_pad( # type: ignore[no-untyped-call]
qkv,
key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
@torch.library.register_fake("fastvideo::_flash_attn_no_pad_forward")
def _flash_attn_no_pad_forward_fake(
qkv: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
del key_padding_mask, causal, dropout_p, softmax_scale, deterministic
b, s, _three, h, d = qkv.shape
return qkv.new_empty(b, s, h, d)
@torch.library.custom_op(
"fastvideo::_flash_attn_varlen_qk_no_pad_forward",
mutates_args=(),
device_types="cuda",
)
def _flash_attn_varlen_qk_no_pad_forward(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
return flash_attn_varlen_qk_no_pad( # type: ignore[no-untyped-call]
query,
key,
value,
query_padding_mask=query_padding_mask,
key_padding_mask=key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
@torch.library.register_fake("fastvideo::_flash_attn_varlen_qk_no_pad_forward")
def _flash_attn_varlen_qk_no_pad_forward_fake(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
query_padding_mask: torch.Tensor,
key_padding_mask: torch.Tensor,
causal: bool,
dropout_p: float,
softmax_scale: float | None,
deterministic: bool,
) -> torch.Tensor:
del key, query_padding_mask, key_padding_mask
del causal, dropout_p, softmax_scale, deterministic
b, sq, h, _ = query.shape
# `out`'s head_dim comes from value (d_v), matching the real forward's
# output ([b, sq, h, d_v]); it can differ from query's d_q.
return query.new_empty(b, sq, h, value.shape[-1])
def flash_attn_no_pad_compilable(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
if torch.is_grad_enabled() and qkv.requires_grad:
return flash_attn_no_pad(qkv,
key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
return torch.ops.fastvideo._flash_attn_no_pad_forward(qkv, key_padding_mask, causal, dropout_p, softmax_scale,
deterministic)
def flash_attn_varlen_qk_no_pad_compilable(query,
key,
value,
query_padding_mask,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
if torch.is_grad_enabled() and (query.requires_grad or key.requires_grad or value.requires_grad):
return flash_attn_varlen_qk_no_pad(query,
key,
value,
query_padding_mask=query_padding_mask,
key_padding_mask=key_padding_mask,
causal=causal,
dropout_p=dropout_p,
softmax_scale=softmax_scale,
deterministic=deterministic)
return torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward(query, key, value, query_padding_mask,
key_padding_mask, causal, dropout_p,
softmax_scale, deterministic)
+5 -2
View File
@@ -1,7 +1,9 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
from fastvideo.configs.models.dits.flux import FluxDiTConfig
from fastvideo.configs.models.dits.flux_2 import Flux2Config
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
@@ -15,6 +17,7 @@ from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig"
]
+119
View File
@@ -0,0 +1,119 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 VFM Transformer FastVideo dataclass configs.
Architecture is 1:1 with the published ``nvidia/Cosmos3-Nano`` checkpoint
(``transformer/config.json``; class ``Cosmos3OmniTransformer`` / framework
``Cosmos3VFMNetwork``). Field values match that config so the FastVideo native
DiT builds a parameter tree matching the checkpoint's state-dict surface
(814 tensors / 44 patterns, validated 2026-06-06).
Reference of record: ``cosmos-framework`` (NVIDIA). The checkpoint is a single
``layers`` ModuleList of dual-pathway (understanding/text + generation/vision)
decoder blocks; per layer: ``self_attn`` with und (``to_{q,k,v}``/``to_out``)
and gen (``add_{q,k,v}_proj``/``to_add_out``) projections + QK-norms, plus
``mlp`` (und) and ``mlp_moe_gen`` (gen), and four RMSNorms. Top level adds
``embed_tokens``/``norm``/``norm_moe_gen``/``lm_head``/``proj_in``/``proj_out``/
``time_embedder`` and dormant ``action_*``/``audio_*`` heads. The checkpoint
remap lives in ``scripts/checkpoint_conversion/cosmos3_convert.py``.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_cosmos3_transformer_block(name: str, module) -> bool:
"""FSDP shard boundary: the dual-pathway decoder blocks ``layers.{i}``."""
del module
parts = name.split(".")
return "layers" in parts and parts[-1].isdigit()
@dataclass
class Cosmos3ArchConfig(DiTArchConfig):
"""Architecture config for the Cosmos3 omni DiT (Cosmos3-Nano).
1:1 with ``transformer/config.json``. The action/sound heads ship in the
checkpoint, so they are constructed for strict-load parity even though the
PR1 video path (T2V/I2V/T2I) leaves them dormant.
"""
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_cosmos3_transformer_block])
# Conversion is owned by scripts/checkpoint_conversion/cosmos3_convert.py;
# the native module tree is the source of truth for parameter names.
param_names_mapping: dict = field(default_factory=dict)
# ---- Backbone (Qwen3-VL-text) ----
hidden_size: int = 4096
num_hidden_layers: int = 36
num_attention_heads: int = 32
num_key_value_heads: int = 8 # GQA (4 query groups)
head_dim: int = 128
intermediate_size: int = 12288
hidden_act: str = "silu"
vocab_size: int = 151936
rms_norm_eps: float = 1e-6
attention_bias: bool = False
qk_norm_for_diffusion: bool = True
qk_norm_for_text: bool = True
use_moe: bool = True # dual-pathway weights; sparse routing unused
joint_attn_implementation: str = "two_way"
freeze_und: bool = False
# ---- Position embedding (unified 3D MRoPE) ----
position_embedding_type: str = "unified_3d_mrope"
rope_theta: float = 5_000_000.0
max_position_embeddings: int = 262144
mrope_section: list[int] = field(default_factory=lambda: [24, 20, 20])
mrope_interleaved: bool = True
unified_3d_mrope_reset_spatial_ids: bool = True
temporal_modality_margin: int = 15000 # unified_3d_mrope_temporal_modality_margin
# ---- VAE / patch geometry ----
latent_patch_size: int = 2
latent_channel: int = 48
patch_latent_dim: int = 192 # latent_patch_size**2 * latent_channel
# ---- Diffusion conditioning ----
timestep_scale: float = 0.001
# ---- Temporal / FPS modulation ----
base_fps: float = 24.0
temporal_compression_factor: int = 4
enable_fps_modulation: bool = True
video_temporal_causal: bool = False
# ---- Action generation head (dormant in PR1 video path) ----
action_gen: bool = True
action_dim: int = 64
max_action_dim: int = 64
num_embodiment_domains: int = 32
# ---- Sound generation head (dormant in PR1 video path) ----
sound_gen: bool = True
sound_dim: int = 64
sound_latent_fps: float = 25.0
temporal_compression_factor_sound: int = 1
# ---- BaseDiT bookkeeping ----
in_channels: int = 48
out_channels: int = 48
def __post_init__(self) -> None:
super().__post_init__()
# Video DiT contract: latent channels == VAE z_dim.
self.num_channels_latents = self.latent_channel
if not self.out_channels:
self.out_channels = self.in_channels
# Derived: patchify packs latent_patch_size**2 spatial patches * channels.
self.patch_latent_dim = self.latent_patch_size**2 * self.latent_channel
@dataclass
class Cosmos3VideoConfig(DiTConfig):
"""Pipeline-level Cosmos3 DiT config (T2V / I2V / T2I share this surface)."""
arch_config: DiTArchConfig = field(default_factory=Cosmos3ArchConfig)
prefix: str = "Cosmos3"
+27
View File
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@dataclass
class FluxTransformer2DArchConfig(DiTArchConfig):
patch_size: int = 1
in_channels: int = 64
out_channels: int | None = None
num_layers: int = 19
num_single_layers: int = 38
attention_head_dim: int = 128
num_attention_heads: int = 24
joint_attention_dim: int = 4096
pooled_projection_dim: int = 768
guidance_embeds: bool = True
axes_dims_rope: tuple[int, int, int] = (16, 56, 56)
@dataclass
class FluxDiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=FluxTransformer2DArchConfig)
prefix: str = "flux"
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class GlmImageDiTArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
hidden_size: int = 4096
num_attention_heads: int = 32
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
num_layers: int = 30
text_embed_dim: int = 1472
time_embed_dim: int = 512
condition_dim: int = 256
prior_vq_quantizer_codebook_size: int = 16384
patch_size: int = 2
max_height: int = 2048
max_width: int = 2048
qk_norm: str = "layer_norm"
eps: float = 1e-5
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["image_projector", "glyph_projector", "prior_token_embedding"])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^glyph_projector\.net\.0\.proj\.(.*)$": r"glyph_projector.fc_in.\1",
r"^glyph_projector\.net\.2\.(.*)$": r"glyph_projector.fc_out.\1",
r"^prior_projector\.net\.0\.proj\.(.*)$": r"prior_projector.fc_in.\1",
r"^prior_projector\.net\.2\.(.*)$": r"prior_projector.fc_out.\1",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$": r"transformer_blocks.\1.ff.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$": r"transformer_blocks.\1.ff.fc_out.\2",
})
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
def __post_init__(self):
super().__post_init__()
self.num_channels_latents = self.out_channels
@dataclass
class GlmImageDiTConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=GlmImageDiTArchConfig)
prefix: str = "GlmImage"
+14
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Literal
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
@@ -19,6 +20,11 @@ class WanVideoArchConfig(DiTArchConfig):
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
# AnyFlow dual-timestep checkpoints expose delta_embedder weights with the
# same internal layout as time_embedder. The regex is harmless on plain
# Wan checkpoints (no delta_embedder keys to match).
r"^condition_embedder\.delta_embedder\.linear_1\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_in.\1",
r"^condition_embedder\.delta_embedder\.linear_2\.(.*)$": r"condition_embedder.delta_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
@@ -86,6 +92,14 @@ class WanVideoArchConfig(DiTArchConfig):
# "relativistic" keeps long rollouts in-distribution; a no-op unless sink_size > 0 and local_attn_size > 0.
rope_cache_policy: str = "absolute"
# AnyFlow dual-timestep conditioning. Defaults preserve bit-identity with
# the legacy single-timestep forward (no delta_embedder allocated, no
# extra computation on the embedder forward path).
r_embedder: bool = False
r_embedder_fusion: Literal["additive", "gated"] = "additive"
r_embedder_gate_value: float = 0.25
r_embedder_deltatime_type: Literal["r", "t-r"] = "r"
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
@@ -1,7 +1,9 @@
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
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
@@ -15,10 +17,12 @@ __all__ = [
"WanVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Cosmos3VAEConfig",
"Gen3CVAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
"OobleckVAEArchConfig",
"OobleckVAEConfig",
"Flux2VAEConfig",
"GlmImageVAEConfig",
]
+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``.
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.autoencoder_kl import (AutoencoderKLArchConfig, AutoencoderKLVAEConfig)
_GLM_IMAGE_LATENTS_MEAN: tuple[float, ...] = (
-0.2080078125,
1.875,
-0.470703125,
-1.265625,
-1.421875,
0.77734375,
-0.3671875,
-0.9453125,
0.318359375,
0.7734375,
-0.1884765625,
-0.022216796875,
-0.220703125,
-1.59375,
-0.81640625,
-0.255859375,
)
_GLM_IMAGE_LATENTS_STD: tuple[float, ...] = (
3.0625,
2.203125,
2.265625,
4.84375,
2.5,
3.9375,
2.203125,
3.03125,
2.1875,
2.046875,
2.71875,
2.390625,
2.390625,
2.453125,
2.25,
2.15625,
)
@dataclass
class GlmImageVAEArchConfig(AutoencoderKLArchConfig):
act_fn: str = "silu"
block_out_channels: tuple[int, ...] = (128, 512, 1024, 1024)
down_block_types: tuple[str, ...] = (
"DownEncoderBlock2D",
"DownEncoderBlock2D",
"DownEncoderBlock2D",
"DownEncoderBlock2D",
)
up_block_types: tuple[str, ...] = (
"UpDecoderBlock2D",
"UpDecoderBlock2D",
"UpDecoderBlock2D",
"UpDecoderBlock2D",
)
force_upcast: bool = True
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 16
latents_mean: tuple[float, ...] = _GLM_IMAGE_LATENTS_MEAN
latents_std: tuple[float, ...] = _GLM_IMAGE_LATENTS_STD
layers_per_block: int = 3
mid_block_add_attention: bool = False
norm_num_groups: int = 32
sample_size: int = 1024
scaling_factor: float = 0.18215
shift_factor: float | None = None
use_quant_conv: bool = False
use_post_quant_conv: bool = False
temporal_compression_ratio: int = 1
spatial_compression_ratio: int = 8
@dataclass
class GlmImageVAEConfig(AutoencoderKLVAEConfig):
arch_config: GlmImageVAEArchConfig = field(default_factory=GlmImageVAEArchConfig)
use_tiling: bool = True
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
tile_sample_min_height: int = 512
tile_sample_min_width: int = 512
tile_sample_stride_height: int = 384
tile_sample_stride_width: int = 384
load_encoder: bool = True
load_decoder: bool = True
+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
+74
View File
@@ -0,0 +1,74 @@
# 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.flux import FluxDiTConfig
from fastvideo.configs.models.encoders import (
BaseEncoderOutput,
CLIPTextConfig,
T5LargeConfig,
)
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def _flux_clip_pooled_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
"""CLIP branch for FLUX: Diffusers uses pooled prompt embeddings only."""
if outputs.pooler_output is None:
raise RuntimeError(
"FLUX CLIP conditioning requires pooler_output. Ensure the CLIP text encoder returns pooled features.")
return outputs.pooler_output
def _flux_t5_sequence_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
if outputs.last_hidden_state is None:
raise RuntimeError("FLUX T5 conditioning requires last_hidden_state.")
return outputs.last_hidden_state
@dataclass
class FluxPipelineConfig(PipelineConfig):
"""Pipeline layout for Diffusers FLUX.1-dev (CLIP + T5 + packed DiT + FlowMatch)."""
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
transformer_arch: str = "FluxTransformer2DModel"
vae_arch: str = "AutoencoderKL"
text_encoder_archs: tuple[str, ...] = ("CLIPTextModel", "T5EncoderModel")
tokenizer_archs: tuple[str, ...] = ("CLIPTokenizer", "T5TokenizerFast")
dit_config: FluxDiTConfig = field(default_factory=FluxDiTConfig)
vae_config: AutoencoderKLVAEConfig = field(default_factory=AutoencoderKLVAEConfig)
embedded_cfg_scale: float = 3.5
flow_shift: float | None = None
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (CLIPTextConfig(), T5LargeConfig()))
preprocess_text_funcs: tuple[Callable[[str], str],
...] = field(default_factory=lambda: (preprocess_text, preprocess_text))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(
default_factory=lambda: (_flux_clip_pooled_postprocess, _flux_t5_sequence_postprocess))
dit_precision: str = "bf16"
vae_precision: str = "fp32"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", "bf16"))
def __post_init__(self) -> None:
te_cfgs = list(self.text_encoder_configs)
if len(te_cfgs) >= 1:
te_cfgs[0].tokenizer_kwargs.setdefault("padding", "max_length")
te_cfgs[0].tokenizer_kwargs.setdefault("max_length", 77)
te_cfgs[0].tokenizer_kwargs.setdefault("truncation", True)
te_cfgs[0].tokenizer_kwargs.setdefault("return_tensors", "pt")
if len(te_cfgs) >= 2:
cap = 512
te_cfgs[1].tokenizer_kwargs["max_length"] = min(int(te_cfgs[1].tokenizer_kwargs.get("max_length", cap)),
cap)
te_cfgs[1].tokenizer_kwargs.setdefault("padding", "max_length")
te_cfgs[1].tokenizer_kwargs.setdefault("truncation", True)
te_cfgs[1].tokenizer_kwargs.setdefault("return_tensors", "pt")
+46
View File
@@ -0,0 +1,46 @@
# SPDX-License-Identifier: Apache-2.0
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.glm_image import GlmImageDiTConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.configs.models.vaes.glm_image import GlmImageVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
def glm_image_t5_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.Tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0, "T5 hidden states contain NaN"
max_len = 512
prompt_embeds = [u[:min(v, max_len)] for u, v in zip(hidden_state, seq_lens, strict=True)]
prompt_embeds_tensor: torch.Tensor = torch.stack(
[torch.cat([u, u.new_zeros(max_len - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0)
return prompt_embeds_tensor
@dataclass
class GlmImageConfig(PipelineConfig):
dit_config: DiTConfig = field(default_factory=GlmImageDiTConfig)
dit_precision: str = "bf16"
vae_config: VAEConfig = field(default_factory=GlmImageVAEConfig)
vae_precision: str = "fp32"
vae_tiling: bool = True
vae_sp: bool = False
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=lambda: (T5Config(), ))
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("fp32", ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (glm_image_t5_postprocess, ))
flow_shift: float | None = 1.0
embedded_cfg_scale: float = 7.5
+54 -1
View File
@@ -295,6 +295,7 @@ def get_1d_rotary_pos_embed(
interpolation_factor: float = 1.0,
dtype: torch.dtype = torch.float32,
use_real: bool = True,
freqs_dtype: torch.dtype | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
@@ -319,12 +320,16 @@ def get_1d_rotary_pos_embed(
if isinstance(pos, int):
pos = torch.arange(pos).float()
# freqs_dtype is an alias for dtype (Diffusers-compatible calling convention).
if freqs_dtype is not None:
dtype = freqs_dtype
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
# has some connection to NTK literature
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor**(dim / (dim - 2))
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(dtype) / dim)) # [D/2]
freqs = 1.0 / (theta**(torch.arange(0, dim, 2, device=pos.device)[:(dim // 2)].to(dtype) / dim)) # [D/2]
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
freqs_cos = freqs.cos() # [S, D/2]
freqs_sin = freqs.sin() # [S, D/2]
@@ -445,6 +450,21 @@ def get_nd_rotary_pos_embed(
return cos, sin
_ROTARY_POS_EMBED_CACHE: dict[tuple, tuple[torch.Tensor, torch.Tensor]] = {}
# Bound the table cache so long-running servers / causal models (which vary
# start_frame per frame) cannot grow it without limit; entries are large float64
# tensors. Least-recently-used eviction keeps the active resolution(s) hot while
# capping memory.
_ROTARY_POS_EMBED_CACHE_MAXSIZE = 16
def _hashable(value: Any) -> Any:
"""Return a hashable view of a scalar or sequence for use in a cache key."""
if isinstance(value, list | tuple):
return tuple(value)
return value
def get_rotary_pos_embed(
rope_sizes,
hidden_size,
@@ -495,6 +515,31 @@ def get_rotary_pos_embed(
sp_rank = 0
sp_world_size = 1
# Memoize on every output-affecting argument; the table is constant across
# denoising steps, so this avoids recomputing the float64 cos/sin tables.
cache_key = (
_hashable(rope_sizes),
tuple(rope_dim_list),
rope_theta,
_hashable(theta_rescale_factor),
_hashable(interpolation_factor),
shard_dim,
sp_rank,
sp_world_size,
dtype,
start_frame,
use_real,
)
cached = _ROTARY_POS_EMBED_CACHE.get(cache_key)
if cached is not None:
# Move to most-recently-used position so the active table is not evicted
# when several resolutions / buckets share the process (LRU recency).
# Pop with a default: a concurrent eviction between the get() above and
# here would otherwise raise KeyError on the hit path.
if _ROTARY_POS_EMBED_CACHE.pop(cache_key, None) is not None:
_ROTARY_POS_EMBED_CACHE[cache_key] = cached
return cached
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
@@ -508,6 +553,14 @@ def get_rotary_pos_embed(
start_frame=start_frame,
use_real=use_real,
)
# The returned tensors are shared cache entries: callers must never mutate
# them in place. Note .to(device) is an identity alias when the tensor is
# already on the target device (e.g. CPU runs), so it does NOT guarantee a
# copy — treat the tables as read-only and copy before any in-place op.
# Reached only on a miss, so evict the least-recently-used entry at capacity.
if len(_ROTARY_POS_EMBED_CACHE) >= _ROTARY_POS_EMBED_CACHE_MAXSIZE:
_ROTARY_POS_EMBED_CACHE.pop(next(iter(_ROTARY_POS_EMBED_CACHE)))
_ROTARY_POS_EMBED_CACHE[cache_key] = (freqs_cos, freqs_sin)
return freqs_cos, freqs_sin
+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
+578
View File
@@ -0,0 +1,578 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from contextlib import nullcontext
from dataclasses import dataclass
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.layers.rotary_embedding import apply_rotary_emb, get_1d_rotary_pos_embed
from fastvideo.attention import DistributedAttention
from fastvideo.configs.models import DiTConfig
from fastvideo.forward_context import get_forward_context, set_forward_context
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.sd3 import (
CombinedTimestepTextProjEmbeddings,
SD3AdaLayerNormContinuous,
SD3AdaLayerNormZero,
SD3FeedForward,
SD3TextProjection,
SD3TimestepEmbedding,
)
from fastvideo.platforms import AttentionBackendEnum
@dataclass
class FluxTransformer2DModelOutput:
sample: torch.Tensor
class FluxPosEmbed(nn.Module):
"""1D RoPE axes concatenated per Diffusers `FluxPosEmbed`."""
def __init__(self, theta: int, axes_dim: list[int]) -> None:
super().__init__()
self.theta = theta
self.axes_dim = axes_dim
def forward(self, ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
n_axes = ids.shape[-1]
cos_out: list[torch.Tensor] = []
sin_out: list[torch.Tensor] = []
pos = ids.float()
is_mps = ids.device.type == "mps"
is_npu = ids.device.type == "npu"
freqs_dtype = torch.float32 if (is_mps or is_npu) else torch.float64
for i in range(n_axes):
cos, sin = get_1d_rotary_pos_embed(
self.axes_dim[i],
pos[:, i],
theta=self.theta,
use_real=True,
freqs_dtype=freqs_dtype,
)
cos_out.append(cos)
sin_out.append(sin)
freqs_cos = torch.cat(cos_out, dim=-1).to(ids.device)
freqs_sin = torch.cat(sin_out, dim=-1).to(ids.device)
return freqs_cos, freqs_sin
class FluxCombinedTimestepGuidanceTextProjEmbeddings(nn.Module):
def __init__(self, embedding_dim: int, pooled_projection_dim: int) -> None:
super().__init__()
self.time_proj = Timesteps(
num_channels=256,
flip_sin_to_cos=True,
downscale_freq_shift=0,
)
self.timestep_embedder = SD3TimestepEmbedding(
in_channels=256,
time_embed_dim=embedding_dim,
act_fn="silu",
)
self.guidance_embedder = SD3TimestepEmbedding(
in_channels=256,
time_embed_dim=embedding_dim,
act_fn="silu",
)
self.text_embedder = SD3TextProjection(
pooled_projection_dim,
embedding_dim,
act_fn="silu",
)
def forward(
self,
timestep: torch.Tensor,
guidance: torch.Tensor,
pooled_projection: torch.Tensor,
) -> torch.Tensor:
timesteps_proj = self.time_proj(timestep)
timesteps_emb = self.timestep_embedder(timesteps_proj.to(dtype=pooled_projection.dtype))
guidance_proj = self.time_proj(guidance)
guidance_emb = self.guidance_embedder(guidance_proj.to(dtype=pooled_projection.dtype))
time_guidance_emb = timesteps_emb + guidance_emb
pooled_projections = self.text_embedder(pooled_projection)
return time_guidance_emb + pooled_projections
class FluxAdaLayerNormZeroSingle(nn.Module):
def __init__(self, embedding_dim: int, bias: bool = True) -> None:
super().__init__()
self.silu = nn.SiLU()
self.linear = ReplicatedLinear(embedding_dim, 3 * embedding_dim, bias=bias)
self.norm = nn.LayerNorm(embedding_dim, elementwise_affine=False, eps=1e-6)
def forward(
self,
x: torch.Tensor,
emb: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
emb, _ = self.linear(self.silu(emb))
shift_msa, scale_msa, gate_msa = emb.chunk(3, dim=1)
x = self.norm(x) * (1 + scale_msa[:, None]) + shift_msa[:, None]
return x, gate_msa
class FluxJointAttention(nn.Module):
"""Joint attention: text tokens precede image tokens (Diffusers order)."""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
super().__init__()
self.heads = num_attention_heads
self.head_dim = attention_head_dim
self.inner_dim = num_attention_heads * attention_head_dim
self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
self.norm_added_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
self.norm_added_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.add_q_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.add_k_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.add_v_proj = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_out = nn.ModuleList(
[
ReplicatedLinear(self.inner_dim, dim, bias=True),
nn.Dropout(0.0),
]
)
self.to_add_out = ReplicatedLinear(self.inner_dim, dim, bias=True)
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=attention_head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
batch_size = hidden_states.shape[0]
text_seq_len = encoder_hidden_states.shape[1]
img_seq_len = hidden_states.shape[1]
q, _ = self.to_q(hidden_states)
k, _ = self.to_k(hidden_states)
v, _ = self.to_v(hidden_states)
q = q.view(batch_size, img_seq_len, self.heads, self.head_dim)
k = k.view(batch_size, img_seq_len, self.heads, self.head_dim)
v = v.view(batch_size, img_seq_len, self.heads, self.head_dim)
q = self.norm_q(q)
k = self.norm_k(k)
enc_q, _ = self.add_q_proj(encoder_hidden_states)
enc_k, _ = self.add_k_proj(encoder_hidden_states)
enc_v, _ = self.add_v_proj(encoder_hidden_states)
enc_q = enc_q.view(batch_size, text_seq_len, self.heads, self.head_dim)
enc_k = enc_k.view(batch_size, text_seq_len, self.heads, self.head_dim)
enc_v = enc_v.view(batch_size, text_seq_len, self.heads, self.head_dim)
enc_q = self.norm_added_q(enc_q)
enc_k = self.norm_added_k(enc_k)
q = torch.cat([enc_q, q], dim=1)
k = torch.cat([enc_k, k], dim=1)
v = torch.cat([enc_v, v], dim=1)
q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1)
k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1)
joint_out, _ = self.attn(q, k, v)
joint_out = joint_out.reshape(batch_size, text_seq_len + img_seq_len, self.inner_dim)
enc_out = joint_out[:, :text_seq_len]
img_out = joint_out[:, text_seq_len:]
img_out, _ = self.to_out[0](img_out)
img_out = self.to_out[1](img_out)
enc_out, _ = self.to_add_out(enc_out)
return img_out, enc_out
class FluxSingleStreamAttention(nn.Module):
"""Self-attention on concatenated text+image sequence (single blocks)."""
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
super().__init__()
self.heads = num_attention_heads
self.head_dim = attention_head_dim
self.inner_dim = num_attention_heads * attention_head_dim
self.norm_q = nn.RMSNorm(attention_head_dim, eps=1e-6)
self.norm_k = nn.RMSNorm(attention_head_dim, eps=1e-6)
self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=attention_head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
)
def forward(
self,
hidden_states: torch.Tensor,
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
batch_size, seq_len, _ = hidden_states.shape
q, _ = self.to_q(hidden_states)
k, _ = self.to_k(hidden_states)
v, _ = self.to_v(hidden_states)
q = q.view(batch_size, seq_len, self.heads, self.head_dim)
k = k.view(batch_size, seq_len, self.heads, self.head_dim)
v = v.view(batch_size, seq_len, self.heads, self.head_dim)
q = self.norm_q(q)
k = self.norm_k(k)
q = apply_rotary_emb(q, image_rotary_emb, sequence_dim=1)
k = apply_rotary_emb(k, image_rotary_emb, sequence_dim=1)
out, _ = self.attn(q, k, v)
return out.reshape(batch_size, seq_len, self.inner_dim)
class FluxTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
super().__init__()
self.norm1 = SD3AdaLayerNormZero(dim)
self.norm1_context = SD3AdaLayerNormZero(dim)
self.attn = FluxJointAttention(
dim=dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
supported_attention_backends=supported_attention_backends,
)
self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
self.ff = SD3FeedForward(
dim=dim,
dim_out=dim,
activation_fn="gelu-approximate",
)
self.norm2_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
self.ff_context = SD3FeedForward(
dim=dim,
dim_out=dim,
activation_fn="gelu-approximate",
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
joint_attention_kwargs: dict[str, Any] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
del joint_attention_kwargs
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb)
(norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp) = self.norm1_context(
encoder_hidden_states, emb=temb
)
attn_output, context_attn_output = self.attn(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
)
attn_output = gate_msa.unsqueeze(1) * attn_output
hidden_states = hidden_states + attn_output
norm_hidden_states = self.norm2(hidden_states)
norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output
context_attn_output = c_gate_msa.unsqueeze(1) * context_attn_output
encoder_hidden_states = encoder_hidden_states + context_attn_output
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None]
context_ff_output = self.ff_context(norm_encoder_hidden_states)
encoder_hidden_states = encoder_hidden_states + (c_gate_mlp.unsqueeze(1) * context_ff_output)
if encoder_hidden_states.dtype == torch.float16:
encoder_hidden_states = encoder_hidden_states.clip(-65504, 65504)
return encoder_hidden_states, hidden_states
class FluxSingleTransformerBlock(nn.Module):
def __init__(
self,
dim: int,
num_attention_heads: int,
attention_head_dim: int,
mlp_ratio: float = 4.0,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
) -> None:
super().__init__()
mlp_hidden_dim = int(dim * mlp_ratio)
self.norm = FluxAdaLayerNormZeroSingle(dim)
self.proj_mlp = ReplicatedLinear(dim, mlp_hidden_dim, bias=True)
self.act_mlp = nn.GELU(approximate="tanh")
self.proj_out = ReplicatedLinear(dim + mlp_hidden_dim, dim, bias=True)
self.attn = FluxSingleStreamAttention(
dim=dim,
num_attention_heads=num_attention_heads,
attention_head_dim=attention_head_dim,
supported_attention_backends=supported_attention_backends,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
image_rotary_emb: tuple[torch.Tensor, torch.Tensor],
joint_attention_kwargs: dict[str, Any] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
del joint_attention_kwargs
text_seq_len = encoder_hidden_states.shape[1]
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
residual = hidden_states
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states)[0])
attn_output = self.attn(
hidden_states=norm_hidden_states,
image_rotary_emb=image_rotary_emb,
)
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
gate = gate.unsqueeze(1)
hidden_states = gate * self.proj_out(hidden_states)[0]
hidden_states = residual + hidden_states
if hidden_states.dtype == torch.float16:
hidden_states = hidden_states.clip(-65504, 65504)
encoder_hidden_states = hidden_states[:, :text_seq_len]
hidden_states = hidden_states[:, text_seq_len:]
return encoder_hidden_states, hidden_states
class FluxTransformer2DModel(BaseDiT):
"""FastVideo FLUX transformer; load Diffusers FLUX safetensors 1:1."""
_fsdp_shard_conditions = [
lambda n, m: (n.startswith("transformer_blocks.") or n.startswith("single_transformer_blocks."))
and n.split(".")[-1].isdigit(),
]
_compile_conditions = _fsdp_shard_conditions
# HF weight names already match this module layout (cf. SGLang regex maps).
param_names_mapping: dict[str, Any] = {}
reverse_param_names_mapping: dict[str, Any] = {}
lora_param_names_mapping: dict[str, Any] = {}
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs) -> None:
del kwargs
super().__init__(config=config, hf_config=hf_config)
self.fastvideo_config = config
self.hf_config = hf_config
arch = config.arch_config
out_ch = arch.out_channels
self.out_channels = out_ch if out_ch is not None else arch.in_channels
self.inner_dim = arch.num_attention_heads * arch.attention_head_dim
self.hidden_size = self.inner_dim
self.num_attention_heads = arch.num_attention_heads
self.num_channels_latents = arch.in_channels
axes_list = list(arch.axes_dims_rope)
self.pos_embed = FluxPosEmbed(theta=10000, axes_dim=axes_list)
if arch.guidance_embeds:
self.time_text_embed = FluxCombinedTimestepGuidanceTextProjEmbeddings(
embedding_dim=self.inner_dim,
pooled_projection_dim=arch.pooled_projection_dim,
)
else:
self.time_text_embed = CombinedTimestepTextProjEmbeddings(
embedding_dim=self.inner_dim,
pooled_projection_dim=arch.pooled_projection_dim,
)
self.context_embedder = ReplicatedLinear(arch.joint_attention_dim, self.inner_dim)
self.x_embedder = ReplicatedLinear(arch.in_channels, self.inner_dim)
self.transformer_blocks = nn.ModuleList(
[
FluxTransformerBlock(
dim=self.inner_dim,
num_attention_heads=arch.num_attention_heads,
attention_head_dim=arch.attention_head_dim,
supported_attention_backends=self._supported_attention_backends,
)
for _ in range(arch.num_layers)
]
)
self.single_transformer_blocks = nn.ModuleList(
[
FluxSingleTransformerBlock(
dim=self.inner_dim,
num_attention_heads=arch.num_attention_heads,
attention_head_dim=arch.attention_head_dim,
supported_attention_backends=self._supported_attention_backends,
)
for _ in range(arch.num_single_layers)
]
)
self.norm_out = SD3AdaLayerNormContinuous(
self.inner_dim,
self.inner_dim,
elementwise_affine=False,
eps=1e-6,
bias=True,
norm_type="layer_norm",
)
self.proj_out = ReplicatedLinear(
self.inner_dim,
arch.patch_size * arch.patch_size * self.out_channels,
bias=True,
)
self.gradient_checkpointing = False
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | None = None,
pooled_projections: torch.Tensor | None = None,
timestep: torch.LongTensor | torch.Tensor | None = None,
img_ids: torch.Tensor | None = None,
txt_ids: torch.Tensor | None = None,
guidance: torch.Tensor | None = None,
joint_attention_kwargs: dict[str, Any] | None = None,
return_dict: bool = True,
controlnet_block_samples: Any | None = None,
controlnet_single_block_samples: Any | None = None,
controlnet_blocks_repeat: bool = False,
**kwargs: Any,
) -> FluxTransformer2DModelOutput | tuple[torch.Tensor, ...]:
del kwargs
if encoder_hidden_states is None:
raise ValueError("encoder_hidden_states must be provided")
if pooled_projections is None:
raise ValueError("pooled_projections must be provided")
if timestep is None:
raise ValueError("timestep must be provided")
if img_ids is None or txt_ids is None:
raise ValueError("img_ids and txt_ids must be provided")
arch = self.fastvideo_config.arch_config
if arch.guidance_embeds and guidance is None:
raise ValueError("guidance must be provided when guidance_embeds=True")
if timestep.dim() == 0:
timestep = timestep[None]
if timestep.dim() > 1:
timestep = timestep.reshape(-1)
if timestep.shape[0] == 1 and hidden_states.shape[0] > 1:
timestep = timestep.expand(hidden_states.shape[0])
try:
get_forward_context()
forward_context = nullcontext()
except AssertionError:
if timestep.numel() == 0:
ts0 = 0
elif torch.is_floating_point(timestep):
ts0 = int(round(timestep[0].item() * 1000))
else:
ts0 = int(timestep[0].item())
forward_context = set_forward_context(current_timestep=ts0, attn_metadata=None)
with forward_context:
hidden_states, _ = self.x_embedder(hidden_states)
ts = timestep.to(hidden_states.dtype) * 1000
g = None if guidance is None else guidance.to(hidden_states.dtype) * 1000
if arch.guidance_embeds:
assert g is not None
temb = self.time_text_embed(ts, g, pooled_projections)
else:
temb = self.time_text_embed(timestep=ts, pooled_projection=pooled_projections)
encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states)
if txt_ids.ndim == 3:
txt_ids = txt_ids[0]
if img_ids.ndim == 3:
img_ids = img_ids[0]
ids = torch.cat((txt_ids, img_ids), dim=0)
image_rotary_emb = self.pos_embed(ids)
jkwargs = joint_attention_kwargs or {}
for idx, block in enumerate(self.transformer_blocks):
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=temb,
image_rotary_emb=image_rotary_emb,
joint_attention_kwargs=jkwargs,
)
if controlnet_block_samples:
interval = len(self.transformer_blocks) / len(controlnet_block_samples)
interval = int(math.ceil(interval))
if controlnet_blocks_repeat:
hidden_states = hidden_states + controlnet_block_samples[idx % len(controlnet_block_samples)]
else:
hidden_states = hidden_states + controlnet_block_samples[idx // interval]
for idx, block in enumerate(self.single_transformer_blocks):
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
temb=temb,
image_rotary_emb=image_rotary_emb,
joint_attention_kwargs=jkwargs,
)
if controlnet_single_block_samples:
interval = len(self.single_transformer_blocks) / len(controlnet_single_block_samples)
interval = int(math.ceil(interval))
hidden_states = hidden_states + controlnet_single_block_samples[idx // interval]
hidden_states = self.norm_out(hidden_states, temb)
output, _ = self.proj_out(hidden_states)
if not return_dict:
return (output,)
return FluxTransformer2DModelOutput(sample=output)
EntryClass = FluxTransformer2DModel
+776
View File
@@ -0,0 +1,776 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.layers.mlp import MLP
from fastvideo.attention import LocalAttention
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
from fastvideo.layers.layernorm import ScaleResidualLayerNormScaleShift
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
from fastvideo.layers.visual_embedding import Timesteps
from fastvideo.models.dits.base import BaseDiT
from fastvideo.platforms import AttentionBackendEnum
class GlmImageLayerKVCache:
def __init__(self):
self.k_cache = None
self.v_cache = None
self.mode: Optional[str] = None
def store(self, k: torch.Tensor, v: torch.Tensor):
# Append along seq (dim=1).
if self.k_cache is None:
self.k_cache = k
self.v_cache = v
else:
self.k_cache = torch.cat([self.k_cache, k], dim=1)
self.v_cache = torch.cat([self.v_cache, v], dim=1)
def get(self):
return self.k_cache, self.v_cache
def clear(self):
self.k_cache = None
self.v_cache = None
self.mode = None
class GlmImageKVCache:
def __init__(self, num_layers: int):
self.num_layers = num_layers
self.caches = [GlmImageLayerKVCache() for _ in range(num_layers)]
def __getitem__(self, layer_idx: int) -> GlmImageLayerKVCache:
return self.caches[layer_idx]
def set_mode(self, mode: Optional[str]):
if mode is not None and mode not in ["write", "read", "skip"]:
raise ValueError(
f"Invalid mode: {mode}, must be one of 'write', 'read', 'skip'"
)
for cache in self.caches:
cache.mode = mode
def clear(self):
for cache in self.caches:
cache.clear()
# =============================================================================
# Timestep and Text Projection
# =============================================================================
class GlmImageTimestepEmbedding(nn.Module):
def __init__(
self,
in_channels: int,
time_embed_dim: int,
act_fn: str = "silu",
out_dim: int = None,
):
super().__init__()
if out_dim is None:
out_dim = time_embed_dim
self.linear_1 = ReplicatedLinear(in_channels, time_embed_dim, bias=True)
if act_fn == "silu":
self.act = nn.SiLU()
elif act_fn == "gelu":
self.act = nn.GELU(approximate="tanh")
else:
self.act = nn.SiLU()
self.linear_2 = ReplicatedLinear(time_embed_dim, out_dim, bias=True)
def forward(self, sample: torch.Tensor) -> torch.Tensor:
sample, _ = self.linear_1(sample)
sample = self.act(sample)
sample, _ = self.linear_2(sample)
return sample
class GlmImageTextProjection(nn.Module):
def __init__(
self,
in_features: int,
hidden_size: int,
out_features: int = None,
act_fn: str = "silu",
):
super().__init__()
if out_features is None:
out_features = hidden_size
self.linear_1 = ReplicatedLinear(in_features, hidden_size, bias=True)
if act_fn == "silu":
self.act_1 = nn.SiLU()
elif act_fn == "gelu_tanh":
self.act_1 = nn.GELU(approximate="tanh")
else:
self.act_1 = nn.SiLU()
self.linear_2 = ReplicatedLinear(hidden_size, out_features, bias=True)
def forward(self, caption: torch.Tensor) -> torch.Tensor:
hidden_states, _ = self.linear_1(caption)
hidden_states = self.act_1(hidden_states)
hidden_states, _ = self.linear_2(hidden_states)
return hidden_states
class GlmImageCombinedTimestepSizeEmbeddings(nn.Module):
def __init__(
self,
embedding_dim: int,
condition_dim: int,
pooled_projection_dim: int,
timesteps_dim: int = 256,
):
super().__init__()
self.time_proj = Timesteps(
num_channels=timesteps_dim, flip_sin_to_cos=True, downscale_freq_shift=0
)
self.condition_proj = Timesteps(
num_channels=condition_dim, flip_sin_to_cos=True, downscale_freq_shift=0
)
self.timestep_embedder = GlmImageTimestepEmbedding(
in_channels=timesteps_dim, time_embed_dim=embedding_dim
)
self.condition_embedder = GlmImageTextProjection(
pooled_projection_dim, embedding_dim, act_fn="silu"
)
def forward(
self,
timestep: torch.Tensor,
target_size: torch.Tensor,
crop_coords: torch.Tensor,
hidden_dtype: torch.dtype,
) -> torch.Tensor:
timesteps_proj = self.time_proj(timestep)
crop_coords_proj = self.condition_proj(crop_coords.flatten()).view(
crop_coords.size(0), -1
)
target_size_proj = self.condition_proj(target_size.flatten()).view(
target_size.size(0), -1
)
condition_proj = torch.cat([crop_coords_proj, target_size_proj], dim=1)
timesteps_emb = self.timestep_embedder(
timesteps_proj.to(dtype=hidden_dtype)
)
condition_emb = self.condition_embedder(
condition_proj.to(dtype=hidden_dtype)
)
conditioning = timesteps_emb + condition_emb
return conditioning
# =============================================================================
# Image Projector
# =============================================================================
class GlmImageImageProjector(nn.Module):
def __init__(
self,
in_channels: int = 16,
hidden_size: int = 2560,
patch_size: int = 2,
):
super().__init__()
self.patch_size = patch_size
self.proj = nn.Linear(in_channels * patch_size**2, hidden_size)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, channel, height, width = hidden_states.shape
post_patch_height = height // self.patch_size
post_patch_width = width // self.patch_size
hidden_states = hidden_states.reshape(
batch_size,
channel,
post_patch_height,
self.patch_size,
post_patch_width,
self.patch_size,
)
hidden_states = (
hidden_states.permute(0, 2, 4, 1, 3, 5).flatten(3, 5).flatten(1, 2)
)
hidden_states = self.proj(hidden_states)
return hidden_states
# =============================================================================
# AdaLayerNorm
# =============================================================================
class GlmImageAdaLayerNormZero(nn.Module):
def __init__(self, embedding_dim: int, dim: int) -> None:
super().__init__()
self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5)
self.norm_context = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-5)
self.linear = ReplicatedLinear(embedding_dim, 12 * dim, bias=True)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
) -> Tuple[torch.Tensor, ...]:
dtype = hidden_states.dtype
norm_hidden_states = self.norm(hidden_states).to(dtype=dtype)
norm_encoder_hidden_states = self.norm_context(encoder_hidden_states).to(
dtype=dtype
)
emb, _ = self.linear(temb)
(
shift_msa,
c_shift_msa,
scale_msa,
c_scale_msa,
gate_msa,
c_gate_msa,
shift_mlp,
c_shift_mlp,
scale_mlp,
c_scale_mlp,
gate_mlp,
c_gate_mlp,
) = emb.chunk(12, dim=1)
hidden_states = norm_hidden_states * (
1 + scale_msa.unsqueeze(1)
) + shift_msa.unsqueeze(1)
encoder_hidden_states = norm_encoder_hidden_states * (
1 + c_scale_msa.unsqueeze(1)
) + c_shift_msa.unsqueeze(1)
return (
hidden_states,
gate_msa,
shift_mlp,
scale_mlp,
gate_mlp,
encoder_hidden_states,
c_gate_msa,
c_shift_mlp,
c_scale_mlp,
c_gate_mlp,
)
# =============================================================================
# Attention
# =============================================================================
class GlmImageAttention(nn.Module):
def __init__(
self,
query_dim: int,
heads: int,
dim_head: int,
out_dim: int,
bias: bool = True,
qk_norm: str = "layer_norm",
elementwise_affine: bool = False,
eps: float = 1e-5,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
prefix: str = "",
):
super().__init__()
self.heads = out_dim // dim_head if out_dim is not None else heads
self.num_kv_heads = self.heads
self.dim_head = dim_head
self.inner_dim = out_dim if out_dim is not None else dim_head * heads
self.inner_kv_dim = self.inner_dim
self.out_dim = out_dim if out_dim is not None else query_dim
self.to_q = ReplicatedLinear(query_dim, self.inner_dim, bias=bias)
self.to_k = ReplicatedLinear(query_dim, self.inner_kv_dim, bias=bias)
self.to_v = ReplicatedLinear(query_dim, self.inner_kv_dim, bias=bias)
self.to_out = nn.ModuleList(
[ReplicatedLinear(self.inner_dim, self.out_dim, bias=True)]
)
if qk_norm is None:
self.norm_q = None
self.norm_k = None
elif qk_norm == "layer_norm":
self.norm_q = nn.LayerNorm(
dim_head, eps=eps, elementwise_affine=elementwise_affine
)
self.norm_k = nn.LayerNorm(
dim_head, eps=eps, elementwise_affine=elementwise_affine
)
else:
raise ValueError(f"unknown qk_norm: {qk_norm}")
self.attn = LocalAttention(
num_heads=self.heads,
head_size=dim_head,
num_kv_heads=self.heads,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
kv_cache: Optional[GlmImageLayerKVCache] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
dtype = encoder_hidden_states.dtype
batch_size, text_seq_length, embed_dim = encoder_hidden_states.shape
batch_size, image_seq_length, embed_dim = hidden_states.shape
hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1)
# 1. QKV projections
query, _ = self.to_q(hidden_states)
key, _ = self.to_k(hidden_states)
value, _ = self.to_v(hidden_states)
query = query.unflatten(2, (self.heads, -1))
key = key.unflatten(2, (self.heads, -1))
value = value.unflatten(2, (self.heads, -1))
# 2. QK normalization
if self.norm_q is not None:
query = self.norm_q(query).to(dtype=dtype)
if self.norm_k is not None:
key = self.norm_k(key).to(dtype=dtype)
# 3. Rotational positional embeddings applied to latent stream
if image_rotary_emb is not None:
cos, sin = image_rotary_emb
query[:, text_seq_length:, :, :] = _apply_rotary_emb(
query[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
)
key[:, text_seq_length:, :, :] = _apply_rotary_emb(
key[:, text_seq_length:, :, :], cos, sin, is_neox_style=True
)
# 4. KV Cache handling
if kv_cache is not None:
if kv_cache.mode == "write":
kv_cache.store(key, value)
elif kv_cache.mode == "read":
# Prepend cached condition k/v along seq (dim=1).
k_cache, v_cache = kv_cache.get()
key = torch.cat([k_cache, key], dim=1) if k_cache is not None else key
value = (
torch.cat([v_cache, value], dim=1) if v_cache is not None else value
)
elif kv_cache.mode == "skip":
pass
hidden_states = self.attn(query, key, value)
hidden_states = hidden_states.flatten(2, 3)
hidden_states = hidden_states.to(query.dtype)
# 6. Output projection
hidden_states, _ = self.to_out[0](hidden_states)
encoder_hidden_states, hidden_states = hidden_states.split(
[text_seq_length, hidden_states.size(1) - text_seq_length], dim=1
)
return hidden_states, encoder_hidden_states
# =============================================================================
# Transformer Block
# =============================================================================
class GlmImageTransformerBlock(nn.Module):
def __init__(
self,
dim: int = 2560,
num_attention_heads: int = 64,
attention_head_dim: int = 40,
time_embed_dim: int = 512,
supported_attention_backends: set[AttentionBackendEnum] | None = None,
prefix: str = "",
) -> None:
super().__init__()
# 1. Attention
self.norm1 = GlmImageAdaLayerNormZero(time_embed_dim, dim)
self.attn1 = GlmImageAttention(
query_dim=dim,
heads=num_attention_heads,
dim_head=attention_head_dim,
out_dim=dim,
bias=True,
qk_norm="layer_norm",
elementwise_affine=False,
eps=1e-5,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1",
)
# 2. Feedforward with fused ScaleResidualLayerNorm
self.norm2 = ScaleResidualLayerNormScaleShift(
dim, norm_type="layer", eps=1e-5, elementwise_affine=False
)
self.norm2_context = ScaleResidualLayerNormScaleShift(
dim, norm_type="layer", eps=1e-5, elementwise_affine=False
)
self.ff = MLP(input_dim=dim, mlp_hidden_dim=dim * 4, output_dim=dim, act_type="gelu_pytorch_tanh")
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: Optional[torch.Tensor] = None,
image_rotary_emb: Optional[
Union[
Tuple[torch.Tensor, torch.Tensor],
List[Tuple[torch.Tensor, torch.Tensor]],
]
] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
kv_cache: Optional[GlmImageLayerKVCache] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
# 1. Timestep conditioning
(
norm_hidden_states,
gate_msa,
shift_mlp,
scale_mlp,
gate_mlp,
norm_encoder_hidden_states,
c_gate_msa,
c_shift_mlp,
c_scale_mlp,
c_gate_mlp,
) = self.norm1(hidden_states, encoder_hidden_states, temb)
# 2. Attention
if attention_kwargs is None:
attention_kwargs = {}
attn_hidden_states, attn_encoder_hidden_states = self.attn1(
hidden_states=norm_hidden_states,
encoder_hidden_states=norm_encoder_hidden_states,
image_rotary_emb=image_rotary_emb,
kv_cache=kv_cache,
**attention_kwargs,
)
# 3. Feedforward (fused residual + norm + scale/shift)
norm_hidden_states, hidden_states = self.norm2(
hidden_states,
attn_hidden_states,
gate_msa.unsqueeze(1),
shift_mlp.unsqueeze(1),
scale_mlp.unsqueeze(1),
)
norm_encoder_hidden_states, encoder_hidden_states = self.norm2_context(
encoder_hidden_states,
attn_encoder_hidden_states,
c_gate_msa.unsqueeze(1),
c_shift_mlp.unsqueeze(1),
c_scale_mlp.unsqueeze(1),
)
ff_output = self.ff(norm_hidden_states)
ff_output_context = self.ff(norm_encoder_hidden_states)
hidden_states = hidden_states + ff_output * gate_mlp.unsqueeze(1)
encoder_hidden_states = (
encoder_hidden_states + ff_output_context * c_gate_mlp.unsqueeze(1)
)
return hidden_states, encoder_hidden_states
# =============================================================================
# Rotary Positional Embedding
# =============================================================================
class GlmImageRotaryPosEmbed(nn.Module):
def __init__(self, dim: int, patch_size: int, theta: float = 10000.0) -> None:
super().__init__()
self.dim = dim
self.patch_size = patch_size
self.theta = theta
self._cache_key: tuple | None = None
self._cache_value: tuple[torch.Tensor, torch.Tensor] | None = None
def forward(self, hidden_states: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
_, _, raw_h, raw_w = hidden_states.shape
height = raw_h // self.patch_size
width = raw_w // self.patch_size
device = hidden_states.device
cache_key = (height, width, device.type,
device.index if device.index is not None else -1)
if self._cache_key == cache_key and self._cache_value is not None:
return self._cache_value
dim_h, dim_w = self.dim // 2, self.dim // 2
h_inv_freq = 1.0 / (
self.theta
** (
torch.arange(0, dim_h, 2, dtype=torch.float32, device=device)[
: (dim_h // 2)
]
/ dim_h
)
)
w_inv_freq = 1.0 / (
self.theta
** (
torch.arange(0, dim_w, 2, dtype=torch.float32, device=device)[
: (dim_w // 2)
]
/ dim_w
)
)
h_seq = torch.arange(height, device=device)
w_seq = torch.arange(width, device=device)
freqs_h = torch.outer(h_seq, h_inv_freq).unsqueeze(1).expand(height, width, -1)
freqs_w = torch.outer(w_seq, w_inv_freq).unsqueeze(0).expand(height, width, -1)
freqs = torch.cat([freqs_h, freqs_w], dim=-1).reshape(height * width, -1)
result = (freqs.cos(), freqs.sin())
self._cache_key = cache_key
self._cache_value = result
return result
# =============================================================================
# Final AdaLayerNorm
# =============================================================================
class GlmImageAdaLayerNormContinuous(nn.Module):
def __init__(
self,
embedding_dim: int,
conditioning_embedding_dim: int,
elementwise_affine: bool = True,
eps: float = 1e-5,
bias: bool = True,
norm_type: str = "layer_norm",
):
super().__init__()
self.linear = nn.Linear(
conditioning_embedding_dim, embedding_dim * 2, bias=bias
)
if norm_type == "layer_norm":
self.norm = nn.LayerNorm(embedding_dim, eps, elementwise_affine, bias)
elif norm_type == "rms_norm":
self.norm = nn.RMSNorm(embedding_dim, eps, elementwise_affine)
else:
raise ValueError(f"unknown norm_type {norm_type}")
def forward(
self, x: torch.Tensor, conditioning_embedding: torch.Tensor
) -> torch.Tensor:
emb = self.linear(conditioning_embedding.to(x.dtype))
scale, shift = torch.chunk(emb, 2, dim=1)
x = self.norm(x) * (1 + scale)[:, None, :] + shift[:, None, :]
return x
# =============================================================================
# Main Model
# =============================================================================
class GlmImageTransformer2DModel(BaseDiT):
_fsdp_shard_conditions = GlmImageDiTConfig().arch_config._fsdp_shard_conditions
_compile_conditions = GlmImageDiTConfig().arch_config._compile_conditions
_supported_attention_backends = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
param_names_mapping = GlmImageDiTConfig().arch_config.param_names_mapping
reverse_param_names_mapping = {}
lora_param_names_mapping = {}
def __init__(
self,
config: GlmImageDiTConfig,
hf_config: dict[str, Any],
):
super().__init__(config=config, hf_config=hf_config)
arch_config = config.arch_config
self.in_channels = arch_config.in_channels
self.out_channels = arch_config.out_channels
self.patch_size = arch_config.patch_size
self.num_layers = arch_config.num_layers
self.attention_head_dim = arch_config.attention_head_dim
self.num_attention_heads = arch_config.num_attention_heads
self.text_embed_dim = arch_config.text_embed_dim
self.time_embed_dim = arch_config.time_embed_dim
# GlmImage uses 2 additional SDXL-like conditions - target_size, crop_coords
# Each of these are sincos embeddings of shape 2 * condition_dim
pooled_projection_dim = 2 * 2 * arch_config.condition_dim
inner_dim = arch_config.num_attention_heads * arch_config.attention_head_dim
self.hidden_size = inner_dim
self.num_channels_latents = arch_config.out_channels
# 1. RoPE
self.rotary_emb = GlmImageRotaryPosEmbed(
arch_config.attention_head_dim, arch_config.patch_size, theta=10000.0
)
# 2. Patch & Text-timestep embedding
self.image_projector = GlmImageImageProjector(
arch_config.in_channels, inner_dim, arch_config.patch_size
)
self.glyph_projector = MLP(
input_dim=arch_config.text_embed_dim,
mlp_hidden_dim=inner_dim,
output_dim=inner_dim,
act_type="gelu",
)
self.prior_token_embedding = nn.Embedding(
arch_config.prior_vq_quantizer_codebook_size, inner_dim
)
self.prior_projector = MLP(
input_dim=inner_dim,
mlp_hidden_dim=inner_dim,
output_dim=inner_dim,
act_type="silu",
)
self.time_condition_embed = GlmImageCombinedTimestepSizeEmbeddings(
embedding_dim=arch_config.time_embed_dim,
condition_dim=arch_config.condition_dim,
pooled_projection_dim=pooled_projection_dim,
timesteps_dim=arch_config.time_embed_dim,
)
# 3. Transformer blocks
self.transformer_blocks = nn.ModuleList(
[
GlmImageTransformerBlock(
inner_dim,
arch_config.num_attention_heads,
arch_config.attention_head_dim,
arch_config.time_embed_dim,
supported_attention_backends=self._supported_attention_backends,
prefix=f"transformer_blocks.{i}",
)
for i in range(arch_config.num_layers)
]
)
# 4. Output projection
self.norm_out = GlmImageAdaLayerNormContinuous(
inner_dim, arch_config.time_embed_dim, elementwise_affine=False
)
self.proj_out = nn.Linear(
inner_dim,
arch_config.patch_size * arch_config.patch_size * arch_config.out_channels,
bias=True,
)
self.gradient_checkpointing = False
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
prior_token_id: torch.Tensor,
prior_token_drop: torch.Tensor,
timestep: torch.LongTensor,
target_size: torch.Tensor,
crop_coords: torch.Tensor,
attention_kwargs: Optional[Dict[str, Any]] = None,
kv_caches: Optional[GlmImageKVCache] = None,
kv_caches_mode: Optional[str] = None,
freqs_cis: Optional[
Union[
Tuple[torch.Tensor, torch.Tensor],
List[Tuple[torch.Tensor, torch.Tensor]],
]
] = None,
guidance: torch.Tensor = None,
**kwargs,
) -> torch.Tensor:
if kv_caches is not None:
kv_caches.set_mode(kv_caches_mode)
batch_size, num_channels, height, width = hidden_states.shape
if isinstance(encoder_hidden_states, list):
encoder_hidden_states = encoder_hidden_states[0]
# 1. RoPE
image_rotary_emb = freqs_cis
if image_rotary_emb is None:
image_rotary_emb = self.rotary_emb(hidden_states)
# 2. Patch & Timestep embeddings
p = self.patch_size
post_patch_height = height // p
post_patch_width = width // p
hidden_states = self.image_projector(hidden_states)
encoder_hidden_states = self.glyph_projector(encoder_hidden_states)
prior_embedding = self.prior_token_embedding(prior_token_id)
# Zero dropped priors by multiply: boolean indexing + .any() syncs each step.
keep = (~prior_token_drop).to(device=prior_embedding.device, dtype=prior_embedding.dtype)
while keep.dim() < prior_embedding.dim():
keep = keep.unsqueeze(-1)
prior_embedding = prior_embedding * keep
prior_hidden_states = self.prior_projector(prior_embedding)
hidden_states = hidden_states + prior_hidden_states
temb = self.time_condition_embed(
timestep, target_size, crop_coords, hidden_states.dtype
)
temb = F.silu(temb)
# 3. Transformer blocks
for idx, block in enumerate(self.transformer_blocks):
hidden_states, encoder_hidden_states = block(
hidden_states,
encoder_hidden_states,
temb,
image_rotary_emb,
attention_kwargs,
kv_cache=kv_caches[idx] if kv_caches is not None else None,
)
# 4. Output norm & projection
hidden_states = self.norm_out(hidden_states, temb)
hidden_states = self.proj_out(hidden_states)
# 5. Unpatchify
hidden_states = hidden_states.reshape(
batch_size, post_patch_height, post_patch_width, -1, p, p
)
output = hidden_states.permute(0, 3, 1, 4, 2, 5).flatten(4, 5).flatten(2, 3)
return output.float()
EntryClass = GlmImageTransformer2DModel
+64 -1
View File
@@ -1,5 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import math
from typing import Any
@@ -59,6 +60,11 @@ class WanTimeTextImageEmbedding(nn.Module):
time_freq_dim: int,
text_embed_dim: int,
image_embed_dim: int | None = None,
*,
r_embedder: bool = False,
r_embedder_fusion: str = "additive",
r_embedder_gate_value: float = 0.25,
r_embedder_deltatime_type: str = "r",
):
super().__init__()
@@ -77,14 +83,57 @@ class WanTimeTextImageEmbedding(nn.Module):
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
# AnyFlow dual-timestep support. When r_embedder is False the forward
# path bypasses delta_embedder entirely and the output is byte-identical
# to the legacy single-timestep implementation.
self._r_embedder_enabled = bool(r_embedder)
self._r_embedder_fusion = r_embedder_fusion
self._r_embedder_deltatime_type = r_embedder_deltatime_type
if self._r_embedder_enabled:
if r_embedder_fusion not in ("additive", "gated"):
raise ValueError(
"r_embedder_fusion must be one of {additive, gated}, "
f"got {r_embedder_fusion!r}")
if r_embedder_deltatime_type not in ("r", "t-r"):
raise ValueError(
"r_embedder_deltatime_type must be one of {r, t-r}, "
f"got {r_embedder_deltatime_type!r}")
# Deep-copy preserves identical initialization with time_embedder,
# matching AnyFlow reference setup_flowmap_model() behavior.
self.delta_embedder = copy.deepcopy(self.time_embedder)
# Non-persistent buffer — gate is a hyperparameter, not learned.
self.register_buffer(
"_r_embedder_gate",
torch.tensor(float(r_embedder_gate_value)),
persistent=False,
)
else:
self.delta_embedder = None
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
r_timestep: torch.Tensor | None = None,
):
temb = self.time_embedder(timestep, timestep_seq_len)
if self._r_embedder_enabled and r_timestep is not None:
assert self.delta_embedder is not None
if self._r_embedder_deltatime_type == "r":
delta_input = r_timestep
else:
delta_input = timestep - r_timestep
delta_emb = self.delta_embedder(delta_input, timestep_seq_len)
gate = self._r_embedder_gate
if self._r_embedder_fusion == "gated":
temb = (1.0 - gate) * temb + gate * delta_emb
else:
# Additive (additional channel, no convex blend).
temb = temb + gate * delta_emb
timestep_proj = self.time_modulation(temb)
if self.text_embedder is not None:
@@ -595,6 +644,10 @@ class WanTransformer3DModel(BaseDiT):
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
r_embedder=config.r_embedder,
r_embedder_fusion=config.r_embedder_fusion,
r_embedder_gate_value=config.r_embedder_gate_value,
r_embedder_deltatime_type=config.r_embedder_deltatime_type,
)
# 3. Transformer blocks
@@ -636,6 +689,7 @@ class WanTransformer3DModel(BaseDiT):
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
r_timestep: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
@@ -684,8 +738,17 @@ class WanTransformer3DModel(BaseDiT):
else:
ts_seq_len = None
# AnyFlow dual-timestep — match timestep's flattening so embedder
# sees aligned shapes.
if r_timestep is not None and r_timestep.dim() == 2:
r_timestep = r_timestep.flatten()
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
timestep,
encoder_hidden_states,
encoder_hidden_states_image,
timestep_seq_len=ts_seq_len,
r_timestep=r_timestep)
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
@@ -0,0 +1,64 @@
# SPDX-License-Identifier: Apache-2.0
"""Sole HF-import boundary for GLM-Image's AR encoder (lazy-wrapper exception E001)."""
from __future__ import annotations
from typing import Any
import torch
import torch.nn as nn
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class GlmImageARLoader(nn.Module):
def __init__(self, model_path: str, processor_path: str | None = None,
*, torch_dtype: torch.dtype = torch.bfloat16,
trust_remote_code: bool = True) -> None:
super().__init__()
from transformers import (AutoProcessor,
GlmImageForConditionalGeneration)
logger.info("Loading GLM-Image AR encoder from %s", model_path)
self._model = GlmImageForConditionalGeneration.from_pretrained(
model_path,
torch_dtype=torch_dtype,
trust_remote_code=trust_remote_code,
)
if processor_path is not None:
logger.info("Loading GLM-Image processor from %s", processor_path)
self.processor = AutoProcessor.from_pretrained(
processor_path, trust_remote_code=trust_remote_code)
else:
self.processor = None
@torch.no_grad()
def generate(self, *args: Any, **kwargs: Any) -> torch.Tensor:
return self._model.generate(*args, **kwargs)
@torch.no_grad()
def get_image_features(self, pixel_values: torch.Tensor,
image_grid_thw: torch.Tensor) -> Any:
return self._model.get_image_features(pixel_values, image_grid_thw)
@torch.no_grad()
def get_image_tokens(self, image_embeds: torch.Tensor,
image_grid_thw: torch.Tensor) -> torch.Tensor:
return self._model.get_image_tokens(image_embeds, image_grid_thw)
@property
def config(self): # type: ignore[no-untyped-def]
return self._model.config
@property
def generation_config(self): # type: ignore[no-untyped-def]
return self._model.generation_config
def to(self, *args, **kwargs): # type: ignore[override]
self._model = self._model.to(*args, **kwargs)
return super().to(*args, **kwargs)
def eval(self): # type: ignore[override]
self._model = self._model.eval()
return super().eval()
@@ -42,7 +42,7 @@ def load_independent(files: list[str], device: str):
"""Before-PR behavior: every rank reads every tensor from disk to GPU."""
for st_file in files:
with safe_open(st_file, framework="pt", device=device) as f:
for name in f:
for name in f.keys(): # noqa: SIM118
param = f.get_tensor(name)
yield name, param
@@ -54,7 +54,7 @@ def load_broadcast(files: list[str], device: str, node_group,
handles = []
for st_file in files:
with safe_open(st_file, framework="pt", device=device) as f:
for name in f:
for name in f.keys(): # noqa: SIM118
if local_rank == 0:
param = f.get_tensor(name)
else:
+50 -1
View File
@@ -92,9 +92,13 @@ 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"),
"vision_language_encoder": (VisionLanguageEncoderLoader, "transformers"),
"processor": (ProcessorLoader, "transformers"),
"upsampler": (UpsamplerLoader, "diffusers"),
"upsampler_2": (UpsamplerLoader, "diffusers"),
# Stable Audio's `StableAudioMultiConditioner` bundles T5 +
@@ -523,6 +527,39 @@ class ImageEncoderLoader(TextEncoderLoader):
)
class VisionLanguageEncoderLoader(ComponentLoader):
"""Loader for vision-language autoregressive encoders."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
from fastvideo.distributed.parallel_state import get_local_torch_device
from fastvideo.models.encoders.glm_image_ar_loader import (
GlmImageARLoader)
logger.info("Loading vision-language encoder from %s", model_path)
target_device = get_local_torch_device()
loader = GlmImageARLoader(
model_path,
torch_dtype=torch.bfloat16,
trust_remote_code=fastvideo_args.trust_remote_code,
).to(target_device).eval()
return loader
class ProcessorLoader(ComponentLoader):
"""Loader for HF processors that pair with vision-language encoders."""
def load(self, model_path: str, fastvideo_args: FastVideoArgs):
from transformers import AutoProcessor
logger.info("Loading processor from %s", model_path)
processor = AutoProcessor.from_pretrained(
model_path,
trust_remote_code=fastvideo_args.trust_remote_code,
)
logger.info("Loaded processor: %s", processor.__class__.__name__)
return processor
class ImageProcessorLoader(ComponentLoader):
"""Loader for image processor."""
@@ -1102,7 +1139,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
+9
View File
@@ -37,6 +37,9 @@ _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"),
@@ -61,6 +64,11 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
"MatrixGame3WanModel": ("dits", "matrixgame3", "MatrixGame3WanModel"),
}
# Text-to-image DiT models (2D image generation)
_TEXT_TO_IMAGE_DIT_MODELS = {
"GlmImageTransformer2DModel": ("dits", "glm_image", "GlmImageTransformer2DModel"),
}
_TEXT_ENCODER_MODELS = {
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
"CLIPTextModelWithProjection":
@@ -134,6 +142,7 @@ _UPSAMPLERS = {
_LEGACY_FAST_VIDEO_MODELS = {
**_TEXT_TO_VIDEO_DIT_MODELS,
**_IMAGE_TO_VIDEO_DIT_MODELS,
**_TEXT_TO_IMAGE_DIT_MODELS,
**_TEXT_ENCODER_MODELS,
**_IMAGE_ENCODER_MODELS,
**_VAE_MODELS,
@@ -0,0 +1,202 @@
# SPDX-License-Identifier: Apache-2.0
"""Flow-map any-step Euler scheduler for AnyFlow.
The model predicts the *average* velocity ``u_θ(x_t, t, r)`` from time
``t`` back to time ``r``, so one Euler step is
x_r = x_t - ((t - r) / num_train_timesteps) * u_θ(x_t, t, r)
regardless of how far apart ``t`` and ``r`` are. The scheduler also
provides the AnyFlow training-time helpers ``apply_shift`` (flow-matching
shift transform) and ``get_train_weight`` (per-timestep loss weight,
including ``beta08``).
Standalone — does not depend on diffusers' ConfigMixin/SchedulerMixin.
"""
from __future__ import annotations
from typing import Literal
import torch
from fastvideo.models.schedulers.base import BaseScheduler
WeightType = Literal["uniform", "gaussian", "beta08"]
class FlowMapEulerDiscreteScheduler(BaseScheduler):
"""Minimal flow-map scheduler.
Parameters
----------
num_train_timesteps:
Discretization granularity for training. ``t`` is expressed in
absolute units in ``[0, num_train_timesteps]``.
shift:
Flow-matching shift parameter (Wan video default: ``5.0``). Set
to ``1.0`` for an identity shift.
"""
order: int = 1
def __init__(
self,
*,
num_train_timesteps: int = 1000,
shift: float = 1.0,
) -> None:
self.num_train_timesteps = int(num_train_timesteps)
self.shift = float(shift)
self.timesteps: torch.Tensor = torch.empty(0)
self.sigmas: torch.Tensor = torch.empty(0)
super().__init__()
# ------------------------------------------------------------------
# BaseScheduler abstract surface
def set_shift(self, shift: float) -> None:
self.shift = float(shift)
def scale_model_input(
self,
sample: torch.Tensor,
timestep: int | None = None,
) -> torch.Tensor:
# Flow-matching has no per-step input scaling; pass through.
del timestep
return sample
# ------------------------------------------------------------------
# Public helpers used by the AnyFlow pretrain method.
def apply_shift(
self,
t: torch.Tensor,
*,
shift: float | None = None,
) -> torch.Tensor:
"""Apply the flow-matching shift: ``t' = s * t / (1 + (s - 1) * t)``.
Operates in the normalized ``[0, 1]`` domain — callers should pass
``t / num_train_timesteps`` (or sample ``t`` directly from
``[0, 1]``).
"""
s = self.shift if shift is None else float(shift)
if s == 1.0:
return t
return s * t / (1.0 + (s - 1.0) * t)
def get_train_weight(
self,
t: torch.Tensor,
*,
weight_type: WeightType = "beta08",
) -> torch.Tensor:
"""Per-timestep training weight, renormalized so the total weight
mass equals ``num_train_timesteps`` (matching AnyFlow reference's
``scheduling_flowmap_euler_discrete.py``).
``beta08``: ``w(t) = t * sqrt(1 - t)`` (in normalized t-space).
"""
# Auto-detect domain: if t was given in absolute units, normalize.
t_f = t.float()
max_val = t_f.max() if t_f.numel() > 0 else torch.tensor(0.0)
if max_val > 1.0 + 1e-6:
t_norm = t_f / self.num_train_timesteps
else:
t_norm = t_f
t_norm = t_norm.clamp(min=0.0, max=1.0)
if weight_type == "uniform":
w = torch.ones_like(t_norm)
elif weight_type == "gaussian":
w = torch.exp(-0.5 * ((t_norm - 0.5) / 0.2) ** 2)
elif weight_type == "beta08":
w = t_norm.pow(1.0) * (1.0 - t_norm).clamp_min(0.0).pow(0.5)
else:
raise ValueError(f"Unknown weight_type: {weight_type!r}")
denom = w.sum().clamp_min(1e-8)
return w * (float(self.num_train_timesteps) / denom)
# ------------------------------------------------------------------
def set_timesteps(
self,
*,
num_inference_steps: int,
device: torch.device | str = "cpu",
custom_timesteps: list[float] | torch.Tensor | None = None,
) -> None:
"""Build a descending timestep schedule ending at 0.
With ``num_inference_steps=N`` the schedule has ``N + 1`` entries
``[T_max, ..., 0]`` so a rollout consumes ``N`` Euler steps.
``custom_timesteps`` overrides the linspace+shift schedule with
a pinned list (in absolute train-timestep units), useful for the
AnyFlow paper's hand-tuned ``[999, 937, 833, 624, 0]`` schedule.
"""
if num_inference_steps <= 0:
raise ValueError(
"num_inference_steps must be positive, "
f"got {num_inference_steps}")
device = torch.device(device)
if custom_timesteps is not None:
ts = torch.as_tensor(
custom_timesteps, dtype=torch.float32, device=device)
if ts.ndim != 1:
raise ValueError(
"custom_timesteps must be 1-D, got shape "
f"{tuple(ts.shape)}")
if not torch.all(ts[:-1] >= ts[1:]):
raise ValueError(
"custom_timesteps must be descending (largest first)")
else:
ts_norm = torch.linspace(
1.0, 0.0, num_inference_steps + 1, device=device)
ts_norm = self.apply_shift(ts_norm)
ts = ts_norm * self.num_train_timesteps
self.timesteps = ts
self.sigmas = ts / self.num_train_timesteps
def step(
self,
model_output: torch.Tensor,
*,
sample: torch.Tensor,
timestep: torch.Tensor,
r_timestep: torch.Tensor,
) -> torch.Tensor:
"""One Euler step from ``t`` to ``r``.
``model_output`` is the average-velocity prediction
``u_θ(x_t, t, r)``. Both ``timestep`` and ``r_timestep`` are in
absolute train-timestep units (``[0, num_train_timesteps]``).
"""
t = timestep.to(sample.device, dtype=sample.dtype)
r = r_timestep.to(sample.device, dtype=sample.dtype)
dt_norm = (t - r) / float(self.num_train_timesteps)
# Broadcast dt over channel/spatial dims.
view: list[int] = [-1] + [1] * (sample.ndim - 1)
return sample - dt_norm.view(*view) * model_output
def add_noise(
self,
original_samples: torch.Tensor,
noise: torch.Tensor,
timestep: torch.Tensor,
) -> torch.Tensor:
"""Linear flow-matching interpolation: ``x_t = (1 - σ) * x_0 + σ * ε``,
where ``σ = t / num_train_timesteps``.
"""
sigma = (timestep.to(original_samples.device,
dtype=original_samples.dtype)
/ float(self.num_train_timesteps))
view: list[int] = [-1] + [1] * (original_samples.ndim - 1)
sigma = sigma.view(*view)
return (1.0 - sigma) * original_samples + sigma * noise
+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)
@@ -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,
)
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages.flux_stages import (
FluxConditioningStage,
FluxDecodingStage,
FluxDenoisingStage,
FluxInputValidationStage,
FluxLatentPreparationStage,
FluxTimestepPreparationStage,
)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
class FluxPipeline(ComposedPipelineBase):
"""FLUX.1-dev T2I (Diffusers module layout, packed latents, embedded guidance)."""
_required_config_modules = [
"scheduler",
"transformer",
"vae",
"text_encoder",
"text_encoder_2",
"tokenizer",
"tokenizer_2",
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=FluxInputValidationStage())
self.add_stage(
stage_name="text_encoding_stage",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2"),
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2"),
],
),
)
self.add_stage(stage_name="flux_conditioning_stage", stage=FluxConditioningStage())
self.add_stage(
stage_name="timestep_preparation_stage",
stage=FluxTimestepPreparationStage(scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=FluxLatentPreparationStage(scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="denoising_stage",
stage=FluxDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
),
)
self.add_stage(
stage_name="decoding_stage",
stage=FluxDecodingStage(vae=self.get_module("vae")),
)
EntryClass = FluxPipeline
@@ -0,0 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
"""GLM-Image pipeline package."""
from fastvideo.pipelines.basic.glm_image.glm_image_pipeline import (
GlmImagePipeline, )
__all__ = ["GlmImagePipeline"]
@@ -0,0 +1,82 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.basic.glm_image.stages import (
GlmImageBeforeDenoisingStage,
GlmImageConditionEncodingStage,
GlmImageDecodingStage,
GlmImageDenoisingStage,
)
from fastvideo.pipelines.stages import InputValidationStage
logger = init_logger(__name__)
class GlmImagePipeline(LoRAPipeline, ComposedPipelineBase):
pipeline_name = "GlmImagePipeline"
_required_config_modules = [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
"vision_language_encoder",
"processor",
]
_optional_config_modules: list[str] = []
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=1.0)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(
stage_name="input_validation_stage",
stage=InputValidationStage(),
)
self.add_stage(
stage_name="glm_image_before_denoising_stage",
stage=GlmImageBeforeDenoisingStage(
vae=self.get_module("vae"),
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer"),
processor=self.get_module("processor"),
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vision_language_encoder=self.get_module("vision_language_encoder"),
),
)
self.add_stage(
stage_name="glm_image_condition_encoding_stage",
stage=GlmImageConditionEncodingStage(
vae=self.get_module("vae"),
transformer=self.get_module("transformer"),
),
)
self.add_stage(
stage_name="denoising_stage",
stage=GlmImageDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self,
),
)
self.add_stage(
stage_name="decoding_stage",
stage=GlmImageDecodingStage(
vae=self.get_module("vae"),
pipeline=self,
),
)
EntryClass = GlmImagePipeline
@@ -0,0 +1,14 @@
# SPDX-License-Identifier: Apache-2.0
"""GLM-Image pipeline stages."""
from fastvideo.pipelines.basic.glm_image.stages.before_denoising import (GlmImageBeforeDenoisingStage)
from fastvideo.pipelines.basic.glm_image.stages.condition_encoding import (GlmImageConditionEncodingStage)
from fastvideo.pipelines.basic.glm_image.stages.decoding import (GlmImageDecodingStage)
from fastvideo.pipelines.basic.glm_image.stages.denoising import (GlmImageDenoisingStage)
__all__ = [
"GlmImageBeforeDenoisingStage",
"GlmImageConditionEncodingStage",
"GlmImageDecodingStage",
"GlmImageDenoisingStage",
]
@@ -0,0 +1,269 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import re
from math import sqrt
import numpy as np
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
logger = init_logger(__name__)
def calculate_shift(
image_seq_len: int,
base_seq_len: int = 256,
base_shift: float = 0.25,
max_shift: float = 0.75,
) -> float:
return (image_seq_len / base_seq_len)**0.5 * max_shift + base_shift
def get_glyph_texts(prompt: str | list[str]) -> list[str] | list[list[str]]:
if isinstance(prompt, str):
prompts: list[str] = [prompt]
is_batch = False
else:
prompts = prompt
is_batch = True
out: list[list[str]] = []
for p in prompts:
out.append(
re.findall(r"'([^']*)'", p) + re.findall(r"“([^“”]*)”", p) + re.findall(r'"([^"]*)"', p) +
re.findall(r"「([^「」]*)」", p))
return out if is_batch else out[0]
def compute_glyph_embeds(
prompts: list[str],
tokenizer,
text_encoder,
device: torch.device,
dtype: torch.dtype,
max_sequence_length: int = 2048,
) -> torch.Tensor:
all_glyph_texts = get_glyph_texts(prompts)
all_glyph_embeds = []
for glyph_texts in all_glyph_texts:
if len(glyph_texts) == 0:
glyph_texts = [""]
input_ids = tokenizer(
glyph_texts,
max_length=max_sequence_length,
truncation=True,
).input_ids
input_ids = [[tokenizer.pad_token_id] * ((len(input_ids) + 1) % 2) + ids for ids in input_ids]
max_length = max(len(ids) for ids in input_ids)
attention_mask = torch.tensor(
[[1] * len(ids) + [0] * (max_length - len(ids)) for ids in input_ids],
device=device,
)
input_ids_t = torch.tensor(
[ids + [tokenizer.pad_token_id] * (max_length - len(ids)) for ids in input_ids],
device=device,
)
outputs = text_encoder(input_ids_t, attention_mask=attention_mask)
glyph_embeds = outputs.last_hidden_state[attention_mask.bool()].unsqueeze(0)
all_glyph_embeds.append(glyph_embeds)
max_seq_len = max(emb.size(1) for emb in all_glyph_embeds)
padded = []
for emb in all_glyph_embeds:
if emb.size(1) < max_seq_len:
pad = torch.zeros(emb.size(0), max_seq_len - emb.size(1), emb.size(2), device=device, dtype=emb.dtype)
emb = torch.cat([pad, emb], dim=1)
padded.append(emb)
return torch.cat(padded, dim=0).to(device=device, dtype=dtype)
def _grid_dims(height: int, width: int) -> tuple[int, int, int, int]:
th, tw = height // 32, width // 32
ratio = th / tw
pth = int(sqrt(ratio) * 16)
ptw = int(sqrt(1 / ratio) * 16)
return th, tw, pth, ptw
def _upsample_d32_to_d16(tokens: torch.Tensor, th: int, tw: int) -> torch.Tensor:
tokens = tokens.view(1, 1, th, tw).float()
tokens = torch.nn.functional.interpolate(tokens, scale_factor=2, mode="nearest").long()
return tokens.view(1, -1)
class GlmImageBeforeDenoisingStage(PipelineStage):
def __init__(self,
vae,
text_encoder,
tokenizer,
processor,
transformer,
scheduler,
vision_language_encoder=None) -> None:
super().__init__()
self.vae = vae
self.text_encoder = text_encoder
self.tokenizer = tokenizer
self.processor = processor
self.transformer = transformer
self.scheduler = scheduler
if isinstance(vision_language_encoder, tuple):
self.vision_language_encoder, self.vl_processor = (vision_language_encoder[0], vision_language_encoder[1]
or processor)
else:
self.vision_language_encoder = vision_language_encoder
self.vl_processor = processor
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
device = get_local_torch_device()
dtype = torch.bfloat16
th, tw, pth, ptw = _grid_dims(batch.height, batch.width)
if batch.seed is not None:
torch.manual_seed(batch.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(batch.seed)
# 1-3. AR token generation. I2I prepends the condition image and uses a
# single-scale target grid; T2I is multi-scale.
is_t2i = batch.pil_image is None
if self.vision_language_encoder is not None:
content = [{"type": "text", "text": batch.prompt}]
if not is_t2i:
content.insert(0, {"type": "image", "image": batch.pil_image})
messages = [{"role": "user", "content": content}]
inputs = self.vl_processor.apply_chat_template(messages,
tokenize=True,
target_h=batch.height,
target_w=batch.width,
return_dict=True,
return_tensors="pt").to(device)
if is_t2i:
up_h, up_w = th, tw
large_start, large_count = pth * ptw, th * tw
max_new = large_count + (pth * ptw) + 1
else:
# Condition grid(s) first, target grid last.
_, t_h, t_w = inputs["image_grid_thw"][-1].tolist()
up_h, up_w = int(t_h), int(t_w)
large_start, large_count = 0, up_h * up_w
max_new = large_count + 1
outputs = self.vision_language_encoder.generate(**inputs, max_new_tokens=max_new, do_sample=True)
gen_tokens = outputs[0][inputs.input_ids.shape[-1]:]
if gen_tokens.shape[0] >= large_start + large_count:
large_tokens = gen_tokens[large_start:large_start + large_count]
else:
available = gen_tokens[large_start:]
large_tokens = torch.zeros(large_count, dtype=gen_tokens.dtype, device=gen_tokens.device)
if available.shape[0] > 0:
large_tokens[:min(available.shape[0], large_count)] = available[:large_count]
logger.warning("AR generated %d tokens, expected %d. Padding with zeros.", gen_tokens.shape[0],
large_start + large_count)
batch.prior_token_id = _upsample_d32_to_d16(large_tokens, up_h, up_w)
batch.prior_token_drop = torch.zeros(batch.prior_token_id.shape, dtype=torch.bool, device=device)
if not is_t2i:
self._compute_source_prior_tokens(batch, inputs)
else:
num_prior_tokens = 4 * th * tw
logger.warning("No vision_language_encoder provided; using random dropped priors.")
batch.prior_token_id = torch.randint(0, 16384, (1, num_prior_tokens), device=device)
batch.prior_token_drop = torch.ones(batch.prior_token_id.shape, dtype=torch.bool, device=device)
# 4. Glyph T5 encoding.
prompts = [batch.prompt] if isinstance(batch.prompt, str) else list(batch.prompt)
prompt_embeds = compute_glyph_embeds(prompts, self.tokenizer, self.text_encoder, device, dtype)
# 5. CFG-side negative encoding.
if batch.do_classifier_free_guidance:
neg_prompts = [batch.negative_prompt or ""] * len(prompts)
neg_embeds = compute_glyph_embeds(neg_prompts, self.tokenizer, self.text_encoder, device, dtype)
L_pos, L_neg = prompt_embeds.shape[1], neg_embeds.shape[1]
max_L = max(L_pos, L_neg)
if L_pos < max_L:
pad = torch.zeros(prompt_embeds.shape[0],
max_L - L_pos,
prompt_embeds.shape[2],
device=device,
dtype=dtype)
prompt_embeds = torch.cat([pad, prompt_embeds], dim=1)
if L_neg < max_L:
pad = torch.zeros(neg_embeds.shape[0], max_L - L_neg, neg_embeds.shape[2], device=device, dtype=dtype)
neg_embeds = torch.cat([pad, neg_embeds], dim=1)
# Row 0 conditional (positive), row 1 unconditional (negative).
prompt_embeds = torch.cat([prompt_embeds, neg_embeds], dim=0)
att_pos = torch.ones((1, max_L), device=device)
att_neg = torch.ones((1, max_L), device=device)
if L_pos < max_L:
att_pos[:, :max_L - L_pos] = 0
if L_neg < max_L:
att_neg[:, :max_L - L_neg] = 0
attention_mask = torch.cat([att_pos, att_neg], dim=0)
else:
attention_mask = torch.ones((1, prompt_embeds.shape[1]), device=device)
batch.prompt_embeds = [prompt_embeds]
batch.attention_mask = attention_mask
# 6. Latents + dynamic flow shift.
if batch.seed is not None:
torch.manual_seed(batch.seed)
if torch.cuda.is_available():
torch.cuda.manual_seed(batch.seed)
batch.latents = torch.randn((1, 16, 1, batch.height // 8, batch.width // 8), device=device, dtype=dtype)
# Integer-cast linspace timesteps with resolution-dependent shift applied to
# sigmas only; the DiT is conditioned on the unshifted integer timesteps.
ntt = self.scheduler.config.num_train_timesteps
patch_size = self.transformer.patch_size
image_seq_len = ((batch.height // 8) * (batch.width // 8)) // (patch_size**2)
sched_timesteps = np.linspace(ntt, 1.0, batch.num_inference_steps + 1)[:-1].astype(np.int64).astype(np.float32)
sched_sigmas = sched_timesteps / ntt
self.scheduler.set_shift(calculate_shift(image_seq_len))
self.scheduler.set_timesteps(batch.num_inference_steps,
device=device,
sigmas=sched_sigmas.tolist(),
timesteps=sched_timesteps.tolist())
batch.timesteps = self.scheduler.timesteps
return batch
@torch.no_grad()
def _compute_source_prior_tokens(self, batch: ForwardBatch, inputs) -> None:
image_grid_thw = inputs["image_grid_thw"]
num_condition_images = image_grid_thw.shape[0] - 1
source_grids = image_grid_thw[:num_condition_images]
image_features = self.vision_language_encoder.get_image_features(inputs["pixel_values"], source_grids)
image_feature_parts = getattr(image_features, "pooler_output", image_features)
embed = torch.cat(image_feature_parts, dim=0)
src_ids_d32 = self.vision_language_encoder.get_image_tokens(embed, source_grids)
split_sizes = source_grids.prod(dim=-1).tolist()
upsampled = [
_upsample_d32_to_d16(ids, int(grid[1]), int(grid[2])).squeeze(0)
for ids, grid in zip(torch.split(src_ids_d32, split_sizes), source_grids, strict=False)
]
src_grids_up = source_grids.clone()
src_grids_up[:, 1] *= 2
src_grids_up[:, 2] *= 2
batch.extra["glm_prior_token_image_ids"] = torch.cat(upsampled, dim=0)
batch.extra["glm_source_image_grid_thw"] = src_grids_up
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_not_empty)
return result
def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("prior_token_id", batch.prior_token_id, V.is_tensor)
return result
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
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.image_processor import ImageProcessor
from fastvideo.models.dits.glm_image import GlmImageKVCache
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
_CONDITION_MULTIPLE_OF = 16 # vae_scale_factor (8) * DiT patch_size (2)
class GlmImageConditionEncodingStage(PipelineStage):
def __init__(self, vae, transformer) -> None:
super().__init__()
self.vae = vae
self.transformer = transformer
self.image_processor = ImageProcessor(vae_scale_factor=_CONDITION_MULTIPLE_OF)
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.pil_image is None:
return batch
device = get_local_torch_device()
dtype = torch.bfloat16
self.vae.to(device)
prior_ids = batch.extra["glm_prior_token_image_ids"].to(device)
if prior_ids.dim() == 1:
prior_ids = prior_ids.unsqueeze(0)
# Latent patch count must match the source prior tokens; mismatch is fatal.
src_grid = batch.extra["glm_source_image_grid_thw"][0]
cond_h = int(src_grid[1]) * _CONDITION_MULTIPLE_OF
cond_w = int(src_grid[2]) * _CONDITION_MULTIPLE_OF
cond_img = self.image_processor.preprocess(batch.pil_image, cond_h, cond_w).to(device=device,
dtype=torch.float32)
latent = self.vae.encode(cond_img).latent_dist.mode()
# NOTE: at runtime self.vae.config is a diffusers FrozenDict with flat
# latents_mean/latents_std fields. Access the flat fields directly.
cfg = self.vae.config
mean = torch.tensor(cfg.latents_mean, device=device, dtype=torch.float32).view(1, -1, 1, 1)
std = torch.tensor(cfg.latents_std, device=device, dtype=torch.float32).view(1, -1, 1, 1)
latent = ((latent - mean) / std).to(dtype)
kv_caches = GlmImageKVCache(num_layers=self.transformer.num_layers)
empty_text = batch.prompt_embeds[0][:1, :0, :].to(device=device, dtype=dtype)
with set_forward_context(current_timestep=0, attn_metadata=None, forward_batch=batch):
self.transformer(
hidden_states=latent,
encoder_hidden_states=empty_text,
prior_token_id=prior_ids,
prior_token_drop=torch.zeros((prior_ids.shape[0], ), dtype=torch.bool, device=device),
timestep=torch.zeros((1, ), device=device),
target_size=torch.tensor([tuple(cond_img.shape[-2:])], device=device, dtype=torch.long),
crop_coords=torch.zeros((1, 2), device=device, dtype=torch.long),
kv_caches=kv_caches,
kv_caches_mode="write",
)
batch.extra["glm_kv_caches"] = kv_caches
return batch
@@ -0,0 +1,35 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import torch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.utils import PRECISION_TO_TYPE
class GlmImageDecodingStage(DecodingStage):
@torch.no_grad()
def decode(self, latents: torch.Tensor, fastvideo_args: FastVideoArgs) -> torch.Tensor:
self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
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)
if latents.dim() == 5:
latents = latents.squeeze(2)
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
if not vae_autocast:
latents = latents.to(vae_dtype)
decoded = self.vae.decode(latents)
image = decoded.sample if hasattr(decoded, "sample") else decoded
image = (image / 2 + 0.5).clamp(0, 1)
return image.unsqueeze(2)
@@ -0,0 +1,192 @@
# SPDX-License-Identifier: Apache-2.0
"""CFG convention (both denoise paths): row 0 conditional (positive), row 1 unconditional."""
from __future__ import annotations
import torch
from fastvideo.attention.backends.sdpa import SDPAMetadata
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.platforms import AttentionBackendEnum
class GlmImageDenoisingStage(DenoisingStage):
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
result.add_check("timesteps", batch.timesteps, [V.is_tensor, V.min_dims(1)])
latents = getattr(batch, "latent", getattr(batch, "latents", None))
result.add_check("latents", latents, [V.is_tensor, V.with_dims(5)])
result.add_check("num_inference_steps", batch.num_inference_steps, V.positive_int)
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
return result
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
device = get_local_torch_device()
dtype = torch.bfloat16
guidance_scale = batch.guidance_scale
do_cfg = guidance_scale > 1.0
latents = getattr(batch, "latent", getattr(batch, "latents", None))
if latents is None:
raise ValueError("No latents found in batch.")
if latents.dim() == 5:
latents = latents.squeeze(2)
prompt_embeds = batch.prompt_embeds[0]
text_attention_mask = getattr(batch, "attention_mask", None)
timesteps = batch.timesteps
patch_size = self.transformer.patch_size
_, _, h, w = latents.shape
image_seq_length = (h // patch_size) * (w // patch_size)
text_seq_length = prompt_embeds.shape[1] if prompt_embeds.dim() >= 2 else 0
first_block = self.transformer.transformer_blocks[0]
backend = getattr(first_block.attn1.attn, "backend", None)
sdpa = backend == AttentionBackendEnum.TORCH_SDPA and text_attention_mask is not None
kv_caches = batch.extra.get("glm_kv_caches")
if kv_caches is None:
self._denoise_t2i(batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
text_seq_length, image_seq_length, sdpa, device, dtype)
else:
self._denoise_i2i(batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
text_seq_length, image_seq_length, sdpa, kv_caches, device, dtype)
return batch
def _denoise_t2i(self, batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
text_seq_length, image_seq_length, sdpa, device, dtype) -> None:
num_inference_steps = batch.num_inference_steps
bs = 2 if do_cfg else 1
target_size = torch.tensor([[batch.height, batch.width]], device=device, dtype=torch.long).repeat(bs, 1)
crop_coords = torch.zeros((bs, 2), device=device, dtype=torch.long)
prior_token_id = batch.prior_token_id
if do_cfg and prior_token_id.shape[0] == 1:
prior_token_id = prior_token_id.repeat(2, 1)
if do_cfg:
prior_token_drop = torch.tensor([False, True], device=device)
else:
prior_token_drop = getattr(batch, "prior_token_drop", torch.tensor([False], device=device))
attention_mask_kv = None
if sdpa:
if (text_attention_mask.shape[0] == 1 and bs > 1):
text_attention_mask = text_attention_mask.repeat(bs, 1)
mix_attn_mask = torch.ones((bs, text_seq_length + image_seq_length), device=device, dtype=torch.float32)
mix_attn_mask[:, :text_seq_length] = (text_attention_mask.float().to(device))
attention_mask_kv = (mix_attn_mask > 0).unsqueeze(1).unsqueeze(2)
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
latent_model_input = torch.cat([latents] * 2) if do_cfg else latents
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t).to(dtype)
t_expand = t.expand(latent_model_input.shape[0]) - 1
attn_metadata = (SDPAMetadata(current_timestep=i, attn_mask=attention_mask_kv)
if attention_mask_kv is not None else None)
with torch.no_grad(), set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
noise_pred = self.transformer(
latent_model_input,
prompt_embeds,
prior_token_id,
prior_token_drop,
t_expand,
target_size,
crop_coords,
)
if do_cfg:
noise_pred_cond, noise_pred_uncond = noise_pred.float().chunk(2)
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_cond - noise_pred_uncond)
guidance_rescale = getattr(batch, "guidance_rescale", 0.0)
if guidance_rescale > 0.0:
dims = list(range(1, noise_pred_cond.ndim))
std_text = noise_pred_cond.std(dim=dims, keepdim=True)
std_cfg = noise_pred.std(dim=dims, keepdim=True)
rescaled = noise_pred * (std_text / std_cfg)
noise_pred = (guidance_rescale * rescaled + (1 - guidance_rescale) * noise_pred)
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
progress_bar.update()
batch.latents = latents.unsqueeze(2)
def _denoise_i2i(self, batch, latents, prompt_embeds, text_attention_mask, timesteps, do_cfg, guidance_scale,
text_seq_length, image_seq_length, sdpa, kv_caches, device, dtype) -> None:
"""Two separate transformer calls (cond reads the cache, uncond skips it):
the cache mode is one global flag with batch-1 k/v, so a 2-row CFG call
cannot express both."""
num_inference_steps = batch.num_inference_steps
target_size = torch.tensor([[batch.height, batch.width]], device=device, dtype=torch.long)
crop_coords = torch.zeros((1, 2), device=device, dtype=torch.long)
prior_token_id = batch.prior_token_id[:1]
drop_keep = torch.zeros((1, ), dtype=torch.bool, device=device)
drop_all = torch.ones((1, ), dtype=torch.bool, device=device)
cache_len = kv_caches[0].k_cache.shape[1] if kv_caches[0].k_cache is not None else 0
def _mask(row: int, with_cache: bool):
if not sdpa:
return None
prefix = cache_len if with_cache else 0
m = torch.ones((1, prefix + text_seq_length + image_seq_length), device=device, dtype=torch.float32)
m[:, prefix:prefix + text_seq_length] = text_attention_mask[row:row + 1].float().to(device)
return (m > 0).unsqueeze(1).unsqueeze(2)
cond_mask = _mask(0, with_cache=True)
uncond_mask = _mask(1, with_cache=False) if do_cfg else None
pos_embeds = prompt_embeds[:1]
neg_embeds = prompt_embeds[1:2] if do_cfg else None
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
latent_model_input = self.scheduler.scale_model_input(latents, t).to(dtype)
t_expand = t.expand(1) - 1
cond_meta = (SDPAMetadata(current_timestep=i, attn_mask=cond_mask) if cond_mask is not None else None)
with torch.no_grad(), set_forward_context(current_timestep=i,
attn_metadata=cond_meta,
forward_batch=batch):
noise_pred = self.transformer(latent_model_input,
pos_embeds,
prior_token_id,
drop_keep,
t_expand,
target_size,
crop_coords,
kv_caches=kv_caches,
kv_caches_mode="read")
if do_cfg:
uncond_meta = (SDPAMetadata(current_timestep=i, attn_mask=uncond_mask)
if uncond_mask is not None else None)
with torch.no_grad(), set_forward_context(current_timestep=i,
attn_metadata=uncond_meta,
forward_batch=batch):
noise_pred_uncond = self.transformer(latent_model_input,
neg_embeds,
prior_token_id,
drop_all,
t_expand,
target_size,
crop_coords,
kv_caches=kv_caches,
kv_caches_mode="skip")
noise_pred = noise_pred_uncond.float() + guidance_scale * (noise_pred.float() -
noise_pred_uncond.float())
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
progress_bar.update()
kv_caches.clear()
batch.latents = latents.unsqueeze(2)
+9 -2
View File
@@ -109,6 +109,10 @@ class ForwardBatch:
max_sequence_length: int | None = None
prompt_template: dict[str, Any] | None = None
do_classifier_free_guidance: bool = False
# When True, ``guidance_scale`` is passed into models that use embedded guidance (e.g. FLUX)
# and must not imply classic dual-forward CFG. Use ``true_cfg_scale > 1`` for true CFG.
use_embedded_guidance: bool = False
true_cfg_scale: float = 1.0
# Batch info
batch_size: int | None = None
@@ -252,9 +256,12 @@ class ForwardBatch:
def __post_init__(self):
"""Initialize dependent fields after dataclass initialization."""
# Enable CFG for standard guidance_scale and LTX-2 text CFG scales.
# LTX-2 text CFG scales; FLUX uses ``use_embedded_guidance`` so ``guidance_scale > 1`` alone
# does not enable classifier-free guidance.
ltx2_text_cfg_enabled = (self.ltx2_cfg_scale_video != 1.0 or self.ltx2_cfg_scale_audio != 1.0)
if self.guidance_scale > 1.0 or ltx2_text_cfg_enabled:
if self.use_embedded_guidance:
self.do_classifier_free_guidance = (self.true_cfg_scale > 1.0) or ltx2_text_cfg_enabled
elif self.guidance_scale > 1.0 or ltx2_text_cfg_enabled:
self.do_classifier_free_guidance = True
if self.negative_prompt_embeds is None:
self.negative_prompt_embeds = []
@@ -0,0 +1,331 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 video denoising stage.
The Cosmos3 video path is monolithic by design: each CFG pass repacks the whole
text+vision sequence (the conditional pass carries prompt tokens, the
unconditional pass carries negative-prompt tokens), so the standard
encode/condition/denoise/decode stage split does not apply. This single stage
owns the full flow, delegating the framework-parity-tested math to
``fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline``:
1. resolve mode (T2I / I2V / T2V) + per-mode defaults, set ``flow_shift``;
2. tokenize the prompt + negative prompt with the Qwen2 chat template;
3. VAE-encode the conditioning frame(s) for I2V / T2I (kept clean), build the
initial noise (clean condition frames + pure noise elsewhere);
4. run the UniPC denoise loop with sequential CFG
(``Cosmos3DenoiseEngine.denoise``);
5. VAE-decode + ``(1 + x) / 2`` clamp to [0, 1].
This mirrors the framework ``Cosmos3OmniDiffusersPipeline.__call__``.
"""
from __future__ import annotations
import os
import weakref
from typing import Any
import torch
from diffusers.utils.torch_utils import randn_tensor
from tqdm import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3DenoiseEngine,
Cosmos3SoundSpec,
Cosmos3VisionSpec,
_VaeNorm,
cosmos3_special_tokens,
cosmos3_tokenize_caption,
cosmos3_vae_decode,
cosmos3_vae_encode,
)
from fastvideo.pipelines.basic.cosmos3.presets import (
COSMOS3_VIDEO_NEGATIVE_PROMPT, )
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
logger = init_logger(__name__)
class Cosmos3DenoisingStage(PipelineStage):
"""Full Cosmos3 video denoise: tokenize + encode + denoise + decode."""
def __init__(self, *, transformer, scheduler, vae, tokenizer, pipeline=None) -> None:
self.transformer = transformer
self.scheduler = scheduler
self.vae = vae
self.tokenizer = tokenizer
self.pipeline = weakref.ref(pipeline) if pipeline is not None else None
# ------------------------------------------------------------------
# Geometry helpers
# ------------------------------------------------------------------
@staticmethod
def _latent_frames(num_frames: int, temporal_factor: int) -> int:
return (int(num_frames) - 1) // int(temporal_factor) + 1
@staticmethod
def _flow_shift_for_resolution(height: int, width: int) -> float:
"""UniPC ``flow_shift`` for a given pixel resolution.
Mirrors the framework's ``_RESOLUTION_SHIFT_DEFAULTS`` (8B VLM backbone,
which Cosmos3-Nano uses): the shift is keyed by the named resolution
bucket the (H, W) belongs to, regardless of task (T2V/I2V/T2I):
"256" -> 3.0, "480" -> 5.0, "704"/"720"/"768" -> 10.0
We invert the framework's ``{IMAGE,VIDEO}_RES_SIZE_INFO`` tables by the
longest side: <=320 is the 256 bucket, 640-832 the 480 bucket, and
960-1360 the 704/720/768 buckets.
"""
long_side = max(int(height), int(width))
if long_side <= 480:
return 3.0
if long_side <= 896:
return 5.0
return 10.0
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
pipeline_config = fastvideo_args.pipeline_config
arch = pipeline_config.dit_config.arch_config
device = self.transformer.embed_tokens.weight.device
dtype = self.transformer.embed_tokens.weight.dtype
num_frames = int(batch.num_frames) if batch.num_frames is not None else 1
height = int(batch.height)
width = int(batch.width)
fps = float(batch.fps) if batch.fps is not None else float(arch.base_fps)
guidance = float(batch.guidance_scale)
is_t2i = num_frames == 1 and batch.preprocessed_image is None and batch.pil_image is None
is_i2v = (batch.preprocessed_image is not None or batch.pil_image is not None) and not is_t2i
# Resolution-based flow_shift, set on the owning pipeline (rebuilds the
# scheduler). The framework picks the UniPC shift purely from the named
# resolution bucket (``_RESOLUTION_SHIFT_DEFAULTS``), NOT from the task,
# so T2V/I2V/T2I at the same resolution share a shift.
pipe = self.pipeline() if self.pipeline is not None else None
flow_shift = self._flow_shift_for_resolution(height, width)
if pipe is not None and hasattr(pipe, "_set_flow_shift"):
pipe._set_flow_shift(flow_shift)
scheduler = pipe.scheduler
else:
scheduler = self.scheduler
# ---- Tokenize prompt + negative prompt ----
prompt = batch.prompt if isinstance(batch.prompt, str) else (batch.prompt[0] if batch.prompt else "")
negative_prompt = batch.negative_prompt
if negative_prompt is None:
negative_prompt = "" if is_t2i else COSMOS3_VIDEO_NEGATIVE_PROMPT
if isinstance(negative_prompt, list):
negative_prompt = negative_prompt[0] if negative_prompt else ""
special_tokens = cosmos3_special_tokens(self.tokenizer)
is_video = not is_t2i
cond_ids = cosmos3_tokenize_caption(self.tokenizer, prompt, is_video=is_video, use_system_prompt=False)
uncond_ids = cosmos3_tokenize_caption(self.tokenizer,
negative_prompt,
is_video=is_video,
use_system_prompt=False)
# ---- VAE normalization constants + geometry ----
norm = _VaeNorm.from_vae(self.vae, dtype)
temporal_factor = int(arch.temporal_compression_factor)
spatial_factor = int(self.vae.config.scale_factor_spatial)
latent_t = self._latent_frames(num_frames, temporal_factor)
latent_h = height // spatial_factor
latent_w = width // spatial_factor
latent_channel = int(arch.latent_channel)
latent_shape = (latent_channel, latent_t, latent_h, latent_w)
generator = batch.generator
if isinstance(generator, list):
generator = generator[0] if generator else None
# ---- Conditioning latent (I2V / T2I) + condition mask ----
condition_frame_indexes: list[int] = []
clean_latent: torch.Tensor | None = None
if is_i2v or (is_t2i and (batch.preprocessed_image is not None or batch.pil_image is not None)):
image = batch.preprocessed_image if batch.preprocessed_image is not None else batch.pil_image
cond_pixels = self._image_to_video_tensor(image, num_frames, height, width, device, dtype)
clean_latent = cosmos3_vae_encode(self.vae, cond_pixels, norm).squeeze(0).float() # [C, T, H, W]
condition_frame_indexes = [0]
# ---- Initial noise (clean condition frames + pure noise elsewhere) ----
pure_noise = randn_tensor(latent_shape, generator=generator, device=device, dtype=dtype).float()
if clean_latent is not None:
cond_mask = torch.zeros((latent_t, 1, 1), device=device, dtype=pure_noise.dtype)
for idx in condition_frame_indexes:
if 0 <= idx < latent_t:
cond_mask[idx, 0, 0] = 1.0
clean = clean_latent.to(device=device, dtype=pure_noise.dtype)
init_latent = cond_mask * clean + (1.0 - cond_mask) * pure_noise
else:
init_latent = pure_noise
spec = Cosmos3VisionSpec(
shape=latent_shape,
condition_frame_indexes=condition_frame_indexes,
)
# ---- Scheduler timesteps ----
scheduler.set_timesteps(int(batch.num_inference_steps), device=device)
timesteps = scheduler.timesteps
engine = Cosmos3DenoiseEngine(
transformer=self.transformer,
scheduler=scheduler,
special_tokens=special_tokens,
latent_patch_size=int(arch.latent_patch_size),
temporal_modality_margin=int(arch.temporal_modality_margin),
reset_spatial_ids=bool(arch.unified_3d_mrope_reset_spatial_ids),
enable_fps_modulation=bool(arch.enable_fps_modulation),
base_fps=float(arch.base_fps),
temporal_compression_factor=temporal_factor,
include_end_of_generation_token=False,
)
flat_latent = init_latent.reshape(-1)
fps_per_item = [fps] if bool(arch.enable_fps_modulation) else None
# ---- t2vs: jointly generate sound (combined [vision | sound] latent) ----
# Mirrors the framework: a placeholder audio sized to the video duration
# sets the sound latent length; sound shares the denoise/CFG with vision.
with_audio = is_video and os.environ.get("COSMOS3_T2VS", "") not in ("", "0")
sound_specs = None
sound_fps_per_item = None
sound_vae = None
sound_shape: tuple[int, int] | None = None
if with_audio:
sound_vae = self._get_sound_vae(pipe, device, dtype)
sound_dim = int(arch.sound_dim)
sound_latent_fps = float(arch.sound_latent_fps)
# Framework ``create_placeholder_audio`` + ``get_latent_num_samples``.
num_audio_samples = int(num_frames / fps * sound_vae.sample_rate)
sound_latent_t = max(1, sound_vae.get_latent_num_samples(num_audio_samples))
sound_shape = (sound_dim, sound_latent_t)
sound_noise = randn_tensor((sound_dim, sound_latent_t), generator=generator, device=device,
dtype=dtype).float()
flat_latent = torch.cat([flat_latent, sound_noise.reshape(-1)])
sound_specs = [Cosmos3SoundSpec(shape=sound_shape, condition_frame_indexes=[], fps=sound_latent_fps)]
sound_fps_per_item = [sound_latent_fps] if bool(arch.enable_fps_modulation) else None
final_flat = engine.denoise(
flat_latent=flat_latent,
timesteps=timesteps,
guidance=guidance,
specs=[spec],
cond_token_ids=cond_ids,
uncond_token_ids=uncond_ids,
fps_per_item=fps_per_item,
progress_bar=lambda it: tqdm(it, desc="Cosmos3 denoising"),
sound_specs=sound_specs,
sound_fps_per_item=sound_fps_per_item,
)
# ---- Decode vision: [C, T, H, W] -> pixels [B, 3, T, H, W] in [0, 1] ----
vision_flat = final_flat[:spec.numel]
result_latent = vision_flat.reshape(latent_shape).unsqueeze(0).to(device=device, dtype=dtype)
decoded = cosmos3_vae_decode(self.vae, result_latent, norm) # [B, 3, T, H, W] in [-1, 1]
video = ((1.0 + decoded) / 2.0).clamp(0.0, 1.0)
batch.latents = result_latent
batch.output = video
# ---- Decode sound: AVAE latent [C, T] -> waveform [C, N] in [-1, 1] ----
if with_audio and sound_vae is not None and sound_shape is not None:
sound_latent = final_flat[spec.numel:].reshape(sound_shape).unsqueeze(0).to(device=device, dtype=dtype)
waveform = sound_vae.decode(sound_latent) # [1, C_audio, N]
batch.extra["audio"] = waveform[0].detach().float().cpu() # [C_audio, N]
batch.extra["audio_sample_rate"] = int(sound_vae.sample_rate)
return batch
# ------------------------------------------------------------------
# Image preprocessing
# ------------------------------------------------------------------
@staticmethod
def _resize_and_center_crop(img: torch.Tensor, target_h: int, target_w: int) -> torch.Tensor:
"""Aspect-ratio-preserving resize + center crop, matching the framework
(``cosmos_framework.inference.vision._resize_and_center_crop``)."""
import math
import torchvision.transforms.functional as TF
orig_h, orig_w = img.shape[-2], img.shape[-1]
scaling_ratio = max(target_w / orig_w, target_h / orig_h)
resize_h = int(math.ceil(scaling_ratio * orig_h))
resize_w = int(math.ceil(scaling_ratio * orig_w))
img = TF.resize(img, [resize_h, resize_w])
return TF.center_crop(img, [target_h, target_w])
@staticmethod
def _get_sound_vae(pipe: Any, device: torch.device, dtype: torch.dtype) -> Any:
"""Lazily load + cache the Cosmos3 sound AVAE decoder from the checkpoint.
The video path does not load ``sound_tokenizer``; t2vs needs only its
decoder, so we load it on first use from ``<model_path>/sound_tokenizer``.
"""
cached = getattr(pipe, "_sound_vae", None) if pipe is not None else None
if cached is not None:
return cached
from fastvideo.models.audio.cosmos3_avae import Cosmos3SoundVAE
model_path = pipe.model_path
sound_dir = os.path.join(model_path, "sound_tokenizer")
sound_vae = Cosmos3SoundVAE.from_pretrained(sound_dir, torch_dtype=dtype).to(device)
if pipe is not None:
pipe._sound_vae = sound_vae
return sound_vae
@classmethod
def _image_to_video_tensor(
cls,
image: Any,
num_frames: int,
height: int,
width: int,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor:
"""Build the I2V conditioning pixel video ``[1, 3, T, H, W]`` in [-1, 1].
Faithful to the framework (``cosmos_framework.inference.vision``):
``load_conditioning_image`` (aspect-preserving resize + center crop +
uint8 quantization, then ``/127.5 - 1``) followed by
``build_conditioned_video_batch``, which fills frame 0 with the image and
**repeats the last conditioning frame** for the rest of the clip (a static
video), NOT zeros. The whole clip is VAE-encoded by the caller; only the
latent condition frame(s) are kept clean by the condition mask, but the
VAE is temporal, so the repeated (not zeroed) frames change the condition
latent — zero-filling here produces a wrong conditioning latent.
"""
import numpy as np
if hasattr(image, "convert"): # PIL.Image: framework-exact preprocessing.
arr = np.array(image.convert("RGB"))
img = torch.from_numpy(arr).permute(2, 0, 1).float() # [3, H, W] in [0, 255]
# Resize + center crop + uint8 quantization, then -> [-1, 1]
# (load_conditioning_image / load_conditioning_image_pixels).
img = cls._resize_and_center_crop(img.unsqueeze(0), height, width).squeeze(0)
img = img.round().clamp(0, 255) / 127.5 - 1.0 # [3, H, W] in [-1, 1]
elif isinstance(image, torch.Tensor): # already-preprocessed conditioning frame.
img = image.float()
if img.dim() == 5: # [B,3,T,H,W]
img = img[0]
if img.dim() == 4: # [3,T,H,W] or [B,3,H,W] -> first frame
img = img[:, 0]
if img.max() > 1.5: # [0, 255] -> [-1, 1]; otherwise assume already [-1, 1].
img = img / 127.5 - 1.0
if img.shape[-2:] != (height, width):
img = cls._resize_and_center_crop(img.unsqueeze(0), height, width).squeeze(0)
else:
raise TypeError(f"Unsupported conditioning image type: {type(image)}")
# Static-repeat video (build_conditioned_video_batch: frame 0 = image,
# remaining frames repeat the last conditioning frame). The whole clip is
# VAE-encoded by the caller; only the latent condition frame(s) are kept
# clean by the condition mask, but the VAE is temporal, so the repeated
# (not zeroed) frames change the condition latent — zero-filling here
# produces a wrong conditioning latent.
img = img.to(device=device, dtype=dtype)
video = img.unsqueeze(0).unsqueeze(2).expand(1, 3, num_frames, height, width)
return video.contiguous()
+424
View File
@@ -0,0 +1,424 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import inspect
from typing import Any
import torch
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
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.pipelines.stages.timestep_preparation import TimestepPreparationStage
from fastvideo.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
def _pack_latents(
latents: torch.Tensor,
batch_size: int,
num_channels_latents: int,
height: int,
width: int,
) -> torch.Tensor:
"""Diffusers ``_pack_latents`` for FLUX (2×2 spatial pack in latent space)."""
latents = latents.view(batch_size, num_channels_latents, height // 2, 2, width // 2, 2)
latents = latents.permute(0, 2, 4, 1, 3, 5)
return latents.reshape(batch_size, (height // 2) * (width // 2), num_channels_latents * 4)
def _unpack_latents(
latents: torch.Tensor,
batch_size: int,
num_channels_latents: int,
height: int,
width: int,
) -> torch.Tensor:
"""Inverse of ``_pack_latents``."""
latents = latents.reshape(batch_size, height // 2, width // 2, num_channels_latents, 2, 2)
latents = latents.permute(0, 3, 1, 4, 2, 5)
return latents.reshape(batch_size, num_channels_latents, height, width)
def _prepare_latent_image_ids(
patch_height: int,
patch_width: int,
device: torch.device,
dtype: torch.dtype = torch.long,
) -> torch.Tensor:
"""Match Diffusers ``FluxPipeline._prepare_latent_image_ids`` (no batch dim)."""
latent_image_ids = torch.zeros(patch_height, patch_width, 3, device=device, dtype=torch.float32)
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(patch_height, device=device)[:, None]
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(patch_width, device=device)[None, :]
h, w, c = latent_image_ids.shape
return latent_image_ids.reshape(h * w, c).to(dtype=dtype)
class FluxInputValidationStage(InputValidationStage):
"""Require height/width divisible by 16 (VAE scale × 2 for FLUX packing)."""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
if (batch.height is not None and batch.width is not None and (batch.height % 16 != 0 or batch.width % 16 != 0)):
raise ValueError("FLUX expects height and width divisible by 16 "
f"(VAE latent grid × 2× packing); got {batch.height}×{batch.width}.")
return super().forward(batch, fastvideo_args)
class FluxConditioningStage(PipelineStage):
"""Build CLIP pooled + T5 sequence + ``text_ids`` (and optional negative for true CFG)."""
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if len(batch.prompt_embeds) < 2:
raise ValueError("FluxConditioningStage expects 2 prompt_embeds (CLIP pooled, T5 sequence), "
f"got {len(batch.prompt_embeds)}")
device = get_local_torch_device()
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
pooled = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
enc = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
seq_len = enc.shape[1]
text_ids = torch.zeros(seq_len, 3, device=device, dtype=torch.long)
batch.extra["flux_pooled_projections"] = pooled
batch.extra["flux_encoder_hidden_states"] = enc
batch.extra["flux_text_ids"] = text_ids
if batch.do_classifier_free_guidance:
if not batch.negative_prompt_embeds or len(batch.negative_prompt_embeds) < 2:
raise ValueError("True CFG requires two negative_prompt_embeds (CLIP, T5).")
neg_pooled = batch.negative_prompt_embeds[0].to(device=device, dtype=target_dtype)
neg_enc = batch.negative_prompt_embeds[1].to(device=device, dtype=target_dtype)
batch.extra["flux_negative_pooled_projections"] = neg_pooled
batch.extra["flux_negative_encoder_hidden_states"] = neg_enc
return batch
class FluxTimestepPreparationStage(TimestepPreparationStage):
"""Flow Match with resolution-dependent ``mu`` from packed image sequence length."""
@staticmethod
def _calculate_mu(
image_seq_len: int,
base_seq_len: int = 256,
max_seq_len: int = 4096,
base_shift: float = 0.5,
max_shift: float = 1.15,
) -> float:
m = (max_shift - base_shift) / (max_seq_len - base_seq_len)
b = base_shift - m * base_seq_len
return float(image_seq_len) * m + b
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
sig = inspect.signature(self.scheduler.set_timesteps)
if "mu" not in sig.parameters:
logger.warning(
"FLUX timestep prep: scheduler %s.set_timesteps does not accept 'mu'; falling back to the base "
"timestep schedule. FLUX expects a FlowMatchEulerDiscreteScheduler with resolution-dependent "
"dynamic shifting — output quality may degrade.",
type(self.scheduler).__name__)
return super().forward(batch, fastvideo_args)
cfg = getattr(self.scheduler, "config", None)
use_dynamic = bool(getattr(cfg, "use_dynamic_shifting", False)) if cfg is not None else False
if not use_dynamic:
logger.warning(
"FLUX timestep prep: scheduler has use_dynamic_shifting=False; falling back to the base timestep "
"schedule and skipping the resolution-dependent 'mu' shift. FLUX requires dynamic shifting for "
"correct timesteps — output quality may degrade.")
return super().forward(batch, fastvideo_args)
if batch.height is None or batch.width is None:
raise ValueError("height/width must be set before FluxTimestepPreparationStage")
vae_arch = fastvideo_args.pipeline_config.vae_config.arch_config
spatial_ratio = int(getattr(vae_arch, "spatial_compression_ratio", 8))
h_lat = batch.height // spatial_ratio
w_lat = batch.width // spatial_ratio
if h_lat % 2 != 0 or w_lat % 2 != 0:
raise ValueError(
f"Latent spatial dims must be even for FLUX packing; got {h_lat}×{w_lat} from {batch.height}×{batch.width}."
)
image_seq_len = (h_lat // 2) * (w_lat // 2)
base_seq_len = int(getattr(cfg, "base_image_seq_len", 256))
max_seq_len = int(getattr(cfg, "max_image_seq_len", 4096))
base_shift = float(getattr(cfg, "base_shift", 0.5))
max_shift = float(getattr(cfg, "max_shift", 1.15))
device = get_local_torch_device()
mu = self._calculate_mu(
image_seq_len=image_seq_len,
base_seq_len=base_seq_len,
max_seq_len=max_seq_len,
base_shift=base_shift,
max_shift=max_shift,
)
self.scheduler.set_timesteps(batch.num_inference_steps, device=device, mu=mu)
batch.timesteps = self.scheduler.timesteps
return batch
class FluxLatentPreparationStage(PipelineStage):
def __init__(self, scheduler) -> None:
super().__init__()
self.scheduler = scheduler
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.height is None or batch.width is None:
raise ValueError("height/width required for FluxLatentPreparationStage")
if isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
else:
if not batch.prompt_embeds:
raise ValueError("prompt or prompt_embeds must be provided")
batch_size = batch.prompt_embeds[0].shape[0]
batch_size *= batch.num_videos_per_prompt
if isinstance(batch.generator, list) and len(batch.generator) != batch_size:
raise ValueError(f"generator list length {len(batch.generator)} does not match batch_size {batch_size}")
arch = fastvideo_args.pipeline_config.dit_config.arch_config
in_channels = int(getattr(arch, "in_channels", 64))
num_channels_latents = in_channels // 4
vae_arch = fastvideo_args.pipeline_config.vae_config.arch_config
spatial_ratio = int(getattr(vae_arch, "spatial_compression_ratio", 8))
h_lat = batch.height // spatial_ratio
w_lat = batch.width // spatial_ratio
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
device = get_local_torch_device()
shape = (batch_size, num_channels_latents, h_lat, w_lat)
latents = batch.latents
if latents is None:
latents = randn_tensor(shape, generator=batch.generator, device=device, dtype=dtype)
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
else:
latents = latents.to(device=device, dtype=dtype)
if latents.shape != shape:
raise ValueError(f"Expected latents shape {shape}, got {tuple(latents.shape)}")
if hasattr(self.scheduler, "init_noise_sigma"):
latents = latents * self.scheduler.init_noise_sigma
packed = _pack_latents(latents, batch_size, num_channels_latents, h_lat, w_lat)
patch_h, patch_w = h_lat // 2, w_lat // 2
img_ids = _prepare_latent_image_ids(patch_h, patch_w, device, dtype=torch.long)
batch.latents = packed
batch.raw_latent_shape = shape
batch.extra["flux_h_lat"] = h_lat
batch.extra["flux_w_lat"] = w_lat
batch.extra["flux_num_channels_latents"] = num_channels_latents
batch.extra["flux_latent_image_ids"] = img_ids
return batch
class FluxDenoisingStage(PipelineStage):
def __init__(self, transformer, scheduler) -> None:
super().__init__()
self.transformer = transformer
self.scheduler = scheduler
@staticmethod
def _step_kwargs(scheduler_step, batch: ForwardBatch) -> dict[str, Any]:
kwargs: dict[str, Any] = {}
sig = inspect.signature(scheduler_step)
if "generator" in sig.parameters:
gen = batch.generator[0] if isinstance(batch.generator, list) else batch.generator
kwargs["generator"] = gen
return kwargs
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
if batch.timesteps is None:
raise ValueError("timesteps must be set before FluxDenoisingStage")
if batch.latents is None:
raise ValueError("latents must be set before FluxDenoisingStage")
packed = batch.latents
timesteps = batch.timesteps
pooled = batch.extra["flux_pooled_projections"]
enc = batch.extra["flux_encoder_hidden_states"]
txt_ids = batch.extra["flux_text_ids"]
img_ids = batch.extra["flux_latent_image_ids"]
neg_pooled = batch.extra.get("flux_negative_pooled_projections")
neg_enc = batch.extra.get("flux_negative_encoder_hidden_states")
true_cfg_scale = float(batch.true_cfg_scale)
use_true_cfg = batch.do_classifier_free_guidance and true_cfg_scale > 1.0
# Prefer the loaded transformer's arch (HF ``guidance_embeds``), not static pipeline defaults.
tr_arch = self.transformer.fastvideo_config.arch_config
guidance_embeds = bool(getattr(tr_arch, "guidance_embeds", False))
device = get_local_torch_device()
target_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
autocast_enabled = (target_dtype != torch.float32) and not fastvideo_args.disable_autocast
bs = packed.shape[0]
if guidance_embeds:
guidance = torch.full((bs, ), float(batch.guidance_scale), device=device, dtype=torch.float32)
else:
guidance = None
step_extras = self._step_kwargs(self.scheduler.step, batch)
for t in timesteps:
t_scalar = t
if not isinstance(t_scalar, torch.Tensor):
t_scalar = torch.tensor([t_scalar], device=device, dtype=torch.float32)
t_scalar = t_scalar.to(device=device, dtype=torch.float32)
timestep_model = t_scalar.expand(bs).float() / 1000.0
timestep_model = timestep_model.to(dtype=target_dtype)
ts_ctx = int(t_scalar.reshape(-1)[0].item())
with (
torch.autocast(
device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled and device.type == "cuda",
),
set_forward_context(
current_timestep=ts_ctx,
attn_metadata=None,
forward_batch=batch,
),
):
if use_true_cfg:
assert neg_enc is not None and neg_pooled is not None
n_neg = self.transformer(
hidden_states=packed,
encoder_hidden_states=neg_enc,
pooled_projections=neg_pooled,
timestep=timestep_model,
guidance=guidance,
txt_ids=txt_ids,
img_ids=img_ids,
return_dict=False,
)[0]
n_pos = self.transformer(
hidden_states=packed,
encoder_hidden_states=enc,
pooled_projections=pooled,
timestep=timestep_model,
guidance=guidance,
txt_ids=txt_ids,
img_ids=img_ids,
return_dict=False,
)[0]
noise_pred = n_neg + true_cfg_scale * (n_pos - n_neg)
else:
noise_pred = self.transformer(
hidden_states=packed,
encoder_hidden_states=enc,
pooled_projections=pooled,
timestep=timestep_model,
guidance=guidance,
txt_ids=txt_ids,
img_ids=img_ids,
return_dict=False,
)[0]
packed = self.scheduler.step(
noise_pred,
t_scalar,
packed,
return_dict=False,
**step_extras,
)[0]
batch.latents = packed
return batch
class FluxDecodingStage(PipelineStage):
"""Unpack latents, apply VAE scaling/shift, decode to pixels (5D output ``B×3×1×H×W``)."""
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
@staticmethod
def _denormalize_latents(latents: torch.Tensor, vae: Any) -> torch.Tensor:
cfg = getattr(vae, "config", None)
sf = getattr(cfg, "scaling_factor", None) if cfg is not None else None
sh = getattr(cfg, "shift_factor", None) if cfg is not None else None
if sf is None and hasattr(vae, "scaling_factor"):
sf = vae.scaling_factor
if sh is None and hasattr(vae, "shift_factor"):
sh = vae.shift_factor
if sf is not None:
latents = latents / (sf.to(latents.device, latents.dtype) if isinstance(sf, torch.Tensor) else sf)
if sh is not None:
latents = latents + (sh.to(latents.device, latents.dtype) if isinstance(sh, torch.Tensor) else sh)
return latents
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
packed = batch.latents
if packed is None:
raise ValueError("latents must be set before FluxDecodingStage")
h_lat = int(batch.extra["flux_h_lat"])
w_lat = int(batch.extra["flux_w_lat"])
num_ch = int(batch.extra["flux_num_channels_latents"])
raw_shape = batch.raw_latent_shape
if raw_shape is None:
raise ValueError("raw_latent_shape missing; FluxLatentPreparationStage must run first.")
batch_size = int(raw_shape[0])
infer_device = get_local_torch_device()
packed = packed.to(infer_device)
latents_4d = _unpack_latents(packed, batch_size, num_ch, h_lat, w_lat)
latents_4d = self._denormalize_latents(latents_4d, self.vae)
vae_device = next(self.vae.parameters()).device
latents_4d = latents_4d.to(device=vae_device)
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
autocast_enabled = (vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
use_cuda_autocast = autocast_enabled and vae_device.type == "cuda"
with torch.autocast(
device_type="cuda",
dtype=vae_dtype,
enabled=use_cuda_autocast,
):
if not autocast_enabled:
latents_4d = latents_4d.to(dtype=vae_dtype)
dec = self.vae.decode(latents_4d)
image = dec.sample if hasattr(dec, "sample") else dec[0]
image = (image / 2 + 0.5).clamp(0, 1)
batch.output = image.unsqueeze(2).detach().float().cpu()
return batch
+50
View File
@@ -20,6 +20,7 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
Cosmos25Config,
Cosmos25_14BConfig,
)
from fastvideo.configs.pipelines.cosmos3 import Cosmos3Config
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
@@ -58,11 +59,14 @@ from fastvideo.configs.pipelines.wan import (
WanT2V480PConfig,
WanT2V720PConfig,
)
from fastvideo.configs.pipelines.glm_image import GlmImageConfig
from fastvideo.configs.pipelines.flux import FluxPipelineConfig
from fastvideo.configs.pipelines.sd35 import SD35Config
from fastvideo.configs.pipelines.stable_audio import (StableAudioOpenSmallConfig, StableAudioT2AConfig)
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.api.matrixgame2 import MatrixGame2SamplingParam
from fastvideo.api.matrixgame3 import MatrixGame3SamplingParam
from fastvideo.api.flux import FluxSamplingParam
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.logger import init_logger
@@ -769,6 +773,22 @@ def _register_configs() -> None:
default_preset="gen3c_cosmos_7b",
)
# Cosmos 3 (must register before Cosmos 2.5 and generic Cosmos detectors
# so the cosmos3 path-detection takes precedence)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=Cosmos3Config,
workload_types=(WorkloadType.T2V, ),
hf_model_paths=[
"nvidia/Cosmos3-Nano",
],
model_detectors=[
lambda path: "cosmos3" in path.lower() or "cosmos-3" in path.lower(),
],
model_family="cosmos3",
default_preset="cosmos3_nano",
)
# Cosmos 2.5 (2B)
register_configs(
sampling_param_cls=None,
@@ -1074,6 +1094,33 @@ def _register_configs() -> None:
default_preset="sd35_medium",
)
# GLM-Image
register_configs(
sampling_param_cls=None,
pipeline_config_cls=GlmImageConfig,
hf_model_paths=[
"zai-org/GLM-Image",
],
model_detectors=[lambda path: "glmimage" in path.lower() or "glm-image" in path.lower()],
workload_types=(WorkloadType.T2I, ),
model_family="glm_image",
)
# FLUX.1-dev (Diffusers)
register_configs(
sampling_param_cls=FluxSamplingParam,
pipeline_config_cls=FluxPipelineConfig,
workload_types=(WorkloadType.T2I, ),
hf_model_paths=[
"black-forest-labs/FLUX.1-dev",
],
model_detectors=[
lambda path: "fluxpipeline" in path,
lambda path: "flux.1-dev" in path or "flux_1_dev" in path,
lambda path: "/flux/" in path or path.endswith("/flux"),
],
)
# --- Part 3: Main Resolver ---
@@ -1164,6 +1211,8 @@ def _register_presets() -> None:
from fastvideo.api.presets import register_preset
from fastvideo.pipelines.basic.cosmos.presets import (
ALL_PRESETS as COSMOS_PRESETS, )
from fastvideo.pipelines.basic.cosmos3.presets import (
ALL_PRESETS as COSMOS3_PRESETS, )
from fastvideo.pipelines.basic.dreamx_world.presets import (
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
from fastvideo.pipelines.basic.gamecraft.presets import (
@@ -1201,6 +1250,7 @@ def _register_presets() -> None:
all_preset_groups = (
COSMOS_PRESETS,
COSMOS3_PRESETS,
DREAMX_WORLD_PRESETS,
FLUX2_PRESETS,
GAMECRAFT_PRESETS,
+1
View File
@@ -183,6 +183,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None:
"guidance_scale_2": None,
"guidance_rescale": 0.0,
"true_cfg_scale": None,
"use_embedded_guidance": None,
"boundary_ratio": None,
"sigmas": None,
},
@@ -31,23 +31,29 @@ def fa_default_impls():
if not torch.cuda.is_available():
pytest.skip("CUDA is required for the FA2/FA3 default custom-op tests")
# The registration + dispatcher live in `attention/utils/`, alongside the
# FP4 cute template and the masked/varlen wrappers. The backend
# (`attention/backends/flash_attn.py`) just imports `flash_attn_func_compilable`
# from there, so test references the utils module directly.
try:
from fastvideo.attention.backends import flash_attn as fa_backend
from fastvideo.attention.utils import flash_attn_default as fa_module
except ImportError as exc:
pytest.skip(f"flash_attn backend not importable: {exc}")
pytest.skip(f"flash_attn_default not importable: {exc}")
if fa_backend.fa_version not in ("2", "3"):
if fa_module.fa_version not in ("2", "3"):
pytest.skip(
f"FA2/FA3 default custom op only exists for fa_version in (2, 3); "
f"got {fa_backend.fa_version!r}"
f"got {fa_module.fa_version!r}"
)
# compilable dispatcher, the original FA wrapper it falls back to, and
# the raw custom op for opcheck.
# compilable dispatcher, the original FA wrapper it falls back to, the
# raw custom op for opcheck, and the fa_version (FA2 has full register_
# autograd; FA3 keeps the carve-out so some tests gate on this).
return (
fa_backend.flash_attn_func_compilable,
fa_backend._fa_default,
fa_module.flash_attn_func_compilable,
fa_module._fa_default,
torch.ops.fastvideo._flash_attn_default_forward,
fa_module.fa_version,
)
@@ -71,7 +77,7 @@ def test_default_compilable_inference_matches_original(fa_default_impls, dtype,
"""No-grad path routes through the custom op and is numerically identical."""
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
pytest.skip("bfloat16 is not supported on this GPU")
compilable, original, _ = fa_default_impls
compilable, original, _, _ = fa_default_impls
torch.manual_seed(0)
q, k, v = _qkv(dtype, requires_grad=False)
@@ -91,7 +97,7 @@ def test_default_compilable_training_backward_flows(fa_default_impls, dtype, cau
"""
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
pytest.skip("bfloat16 is not supported on this GPU")
compilable, original, _ = fa_default_impls
compilable, original, _, _ = fa_default_impls
torch.manual_seed(0)
q_ref, k_ref, v_ref = _qkv(dtype, requires_grad=True)
@@ -115,7 +121,65 @@ def test_default_compilable_training_backward_flows(fa_default_impls, dtype, cau
@pytest.mark.parametrize("causal", [False, True])
def test_default_forward_opcheck(fa_default_impls, causal):
"""Schema / fake-kernel consistency for the custom op (forward only)."""
_, _, op = fa_default_impls
_, _, op, _ = fa_default_impls
torch.manual_seed(0)
q, k, v = _qkv(torch.float16, requires_grad=False)
torch.library.opcheck(op, (q, k, v, None, causal))
# --------------------------------------------------------------------------- #
# FA2-only: backward is registered on the custom op itself. Exercises the #
# `register_autograd` wiring directly (not via the dispatcher's carve-out). #
# Skipped on FA3 until Kuan-Hao's Modal FA3 setup PR lands and we mirror the #
# pattern there. #
# --------------------------------------------------------------------------- #
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
@pytest.mark.parametrize("causal", [False, True])
def test_default_op_backward_through_registered_autograd(fa_default_impls, dtype, causal):
"""FA2: gradients flow through ``torch.ops.fastvideo._flash_attn_default_forward``
itself (no dispatcher carve-out involved) and match the original
``flash_attn_func``'s gradients."""
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
pytest.skip("bfloat16 is not supported on this GPU")
_, original, op, fa_version = fa_default_impls
if fa_version != "2":
pytest.skip(f"register_autograd is only wired for FA2 right now; got {fa_version!r}")
torch.manual_seed(0)
q_ref, k_ref, v_ref = _qkv(dtype, requires_grad=True)
q_test, k_test, v_test = _clone(q_ref, k_ref, v_ref)
# Reference grads via the original autograd.Function.
out_ref = original(q_ref, k_ref, v_ref, softmax_scale=None, causal=causal)
# Custom-op grads via the registered backward — unpack (out, lse), discard lse.
out_test, _ = op(q_test, k_test, v_test, None, causal)
torch.testing.assert_close(out_test, out_ref,
atol=0 if dtype == torch.float16 else 1e-3,
rtol=0 if dtype == torch.float16 else 1e-3)
dout = torch.randn_like(out_ref)
dq_ref, dk_ref, dv_ref = torch.autograd.grad(
(out_ref * dout).sum(), (q_ref, k_ref, v_ref))
dq_test, dk_test, dv_test = torch.autograd.grad(
(out_test * dout).sum(), (q_test, k_test, v_test))
atol = rtol = 6e-3 if dtype == torch.float16 else 2e-2
torch.testing.assert_close(dq_test, dq_ref, atol=atol, rtol=rtol)
torch.testing.assert_close(dk_test, dk_ref, atol=atol, rtol=rtol)
torch.testing.assert_close(dv_test, dv_ref, atol=atol, rtol=rtol)
@pytest.mark.parametrize("causal", [False, True])
def test_default_op_opcheck_with_grad_inputs(fa_default_impls, causal):
"""FA2: full ``opcheck`` including ``test_autograd_registration`` —
catches a missing/inconsistent backward at unit-test time, which was
exactly the gap that #1373's first revision shipped."""
_, _, op, fa_version = fa_default_impls
if fa_version != "2":
pytest.skip(f"autograd registration only wired for FA2; got {fa_version!r}")
torch.manual_seed(0)
q, k, v = _qkv(torch.float16, requires_grad=True)
torch.library.opcheck(op, (q, k, v, None, causal))
@@ -0,0 +1,239 @@
# SPDX-License-Identifier: Apache-2.0
"""Regression guard for the masked/varlen custom ops.
`flash_attn_no_pad_compilable` / `flash_attn_varlen_qk_no_pad_compilable` wrap
the whole masked-attention functions in `torch.library.custom_op`s so dynamo
sees one traceable node (the internal unpad/pad bookkeeping runs eager inside).
On FA2 these ops register a real backward — `softmax_lse` is padded back to a
statically-shaped `[batch, nheads, seqlen]` form on the way out and re-unpadded
in backward, which calls FA2's `_flash_attn_varlen_backward`. So training also
backprops through the op (no graph break on the training path either).
On FA3/FA4 these ops are forward+fake only and the `*_compilable` dispatchers
carve out to the autograd.Function for grad-enabled calls (PR #1373 pattern).
Tests gating on FA2 are skipped on FA3/FA4.
These tests pin:
- inference (no grad): output through the custom op is bit-identical to the
original function;
- training (requires_grad): gradients through the registered op match the
original autograd.Function;
- schema/fake-kernel consistency via torch.library.opcheck, both with and
without grad-requiring inputs (the latter exercises
`test_autograd_registration` — would have caught a missing/broken backward
at unit-test time).
GPU assumptions: requires CUDA and the FA2 `flash_attn` varlen package.
Skips on CPU and when flash_attn is unavailable.
"""
from __future__ import annotations
import pytest
import torch
@pytest.fixture(scope="module")
def no_pad_impls():
if not torch.cuda.is_available():
pytest.skip("CUDA is required for the masked/varlen custom-op tests")
try:
from fastvideo.attention.utils import flash_attn_no_pad as mod
except ImportError as exc:
pytest.skip(f"flash_attn_no_pad not importable (flash_attn missing?): {exc}")
return mod
def _dtype_skip(dtype):
if dtype == torch.bfloat16 and not torch.cuda.is_bf16_supported():
pytest.skip("bfloat16 is not supported on this GPU")
def _fa2_only(mod):
if mod._FA_VARLEN_VERSION != "2":
pytest.skip(
f"register_autograd is only wired for FA2; got "
f"_FA_VARLEN_VERSION={mod._FA_VARLEN_VERSION!r}"
)
def _padding_mask(batch, seqlen, valid_lens, device):
mask = torch.zeros(batch, seqlen, dtype=torch.bool, device=device)
for i, n in enumerate(valid_lens):
mask[i, :n] = True
return mask
# --------------------------------------------------------------------------- #
# flash_attn_no_pad (masked self-attention) #
# --------------------------------------------------------------------------- #
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_no_pad_inference_matches_original(no_pad_impls, dtype):
"""No-grad path through the custom op is bit-identical to the original."""
_dtype_skip(dtype)
mod = no_pad_impls
torch.manual_seed(0)
device = torch.device("cuda")
b, s, h, d = 2, 64, 4, 64
qkv = torch.randn(b, s, 3, h, d, device=device, dtype=dtype)
mask = _padding_mask(b, s, [64, 48], device)
with torch.inference_mode():
out_ref = mod.flash_attn_no_pad(qkv, mask, causal=False, dropout_p=0.0)
out_test = mod.flash_attn_no_pad_compilable(qkv, mask, causal=False, dropout_p=0.0)
torch.testing.assert_close(out_test, out_ref, atol=0, rtol=0)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_no_pad_training_backward_through_registered_autograd(no_pad_impls, dtype):
"""FA2: grads flow through the registered op and match the original."""
_dtype_skip(dtype)
mod = no_pad_impls
_fa2_only(mod)
torch.manual_seed(0)
device = torch.device("cuda")
b, s, h, d = 2, 64, 4, 64
mask = _padding_mask(b, s, [64, 48], device)
qkv_ref = torch.randn(b, s, 3, h, d, device=device, dtype=dtype, requires_grad=True)
qkv_test = qkv_ref.detach().clone().requires_grad_(True)
out_ref = mod.flash_attn_no_pad(qkv_ref, mask, causal=False, dropout_p=0.0)
# Go through the compilable wrapper (which on FA2 unconditionally routes
# to the registered op — no carve-out).
out_test = mod.flash_attn_no_pad_compilable(qkv_test, mask, causal=False, dropout_p=0.0)
dout = torch.randn_like(out_ref)
(dqkv_ref,) = torch.autograd.grad((out_ref * dout).sum(), (qkv_ref,))
(dqkv_test,) = torch.autograd.grad((out_test * dout).sum(), (qkv_test,))
atol = rtol = 6e-3 if dtype == torch.float16 else 2e-2
torch.testing.assert_close(dqkv_test, dqkv_ref, atol=atol, rtol=rtol)
def test_no_pad_forward_opcheck(no_pad_impls):
"""Schema/fake-kernel consistency for the custom op (forward only)."""
torch.manual_seed(0)
device = torch.device("cuda")
b, s, h, d = 2, 64, 4, 64
qkv = torch.randn(b, s, 3, h, d, device=device, dtype=torch.float16)
mask = _padding_mask(b, s, [64, 48], device)
torch.library.opcheck(
torch.ops.fastvideo._flash_attn_no_pad_forward,
(qkv, mask, False, 0.0, None, False),
)
def test_no_pad_opcheck_with_grad_inputs(no_pad_impls):
"""FA2: full opcheck including ``test_autograd_registration`` — catches a
missing/inconsistent backward at unit-test time."""
mod = no_pad_impls
_fa2_only(mod)
torch.manual_seed(0)
device = torch.device("cuda")
b, s, h, d = 2, 64, 4, 64
qkv = torch.randn(b, s, 3, h, d, device=device, dtype=torch.float16, requires_grad=True)
mask = _padding_mask(b, s, [64, 48], device)
torch.library.opcheck(
torch.ops.fastvideo._flash_attn_no_pad_forward,
(qkv, mask, False, 0.0, None, False),
)
# --------------------------------------------------------------------------- #
# flash_attn_varlen_qk_no_pad (cross-attn / unequal q-k seqlen) #
# --------------------------------------------------------------------------- #
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_varlen_qk_inference_matches_original(no_pad_impls, dtype):
_dtype_skip(dtype)
mod = no_pad_impls
# The varlen-qk custom op's real forward goes through FA's varlen func; off
# FA2 that path is the pre-existing FA3 `dropout_p` carve-out (out of scope
# here), so skip cleanly like the autograd tests until that's fixed.
_fa2_only(mod)
torch.manual_seed(1)
device = torch.device("cuda")
b, sq, sk, h, d = 2, 48, 64, 4, 64
q = torch.randn(b, sq, h, d, device=device, dtype=dtype)
k = torch.randn(b, sk, h, d, device=device, dtype=dtype)
v = torch.randn(b, sk, h, d, device=device, dtype=dtype)
qmask = _padding_mask(b, sq, [48, 40], device)
kmask = _padding_mask(b, sk, [64, 56], device)
with torch.inference_mode():
out_ref = mod.flash_attn_varlen_qk_no_pad(q, k, v, qmask, kmask, causal=False, dropout_p=0.0)
out_test = mod.flash_attn_varlen_qk_no_pad_compilable(q, k, v, qmask, kmask, causal=False, dropout_p=0.0)
torch.testing.assert_close(out_test, out_ref, atol=0, rtol=0)
@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16])
def test_varlen_qk_training_backward_through_registered_autograd(no_pad_impls, dtype):
"""FA2: grads through q, k, v all flow via the registered op and match."""
_dtype_skip(dtype)
mod = no_pad_impls
_fa2_only(mod)
torch.manual_seed(1)
device = torch.device("cuda")
b, sq, sk, h, d = 2, 48, 64, 4, 64
qmask = _padding_mask(b, sq, [48, 40], device)
kmask = _padding_mask(b, sk, [64, 56], device)
q_ref = torch.randn(b, sq, h, d, device=device, dtype=dtype, requires_grad=True)
k_ref = torch.randn(b, sk, h, d, device=device, dtype=dtype, requires_grad=True)
v_ref = torch.randn(b, sk, h, d, device=device, dtype=dtype, requires_grad=True)
q_test = q_ref.detach().clone().requires_grad_(True)
k_test = k_ref.detach().clone().requires_grad_(True)
v_test = v_ref.detach().clone().requires_grad_(True)
out_ref = mod.flash_attn_varlen_qk_no_pad(q_ref, k_ref, v_ref, qmask, kmask, causal=False, dropout_p=0.0)
out_test = mod.flash_attn_varlen_qk_no_pad_compilable(q_test, k_test, v_test, qmask, kmask,
causal=False, dropout_p=0.0)
dout = torch.randn_like(out_ref)
dq_ref, dk_ref, dv_ref = torch.autograd.grad((out_ref * dout).sum(), (q_ref, k_ref, v_ref))
dq_test, dk_test, dv_test = torch.autograd.grad((out_test * dout).sum(), (q_test, k_test, v_test))
atol = rtol = 6e-3 if dtype == torch.float16 else 2e-2
torch.testing.assert_close(dq_test, dq_ref, atol=atol, rtol=rtol)
torch.testing.assert_close(dk_test, dk_ref, atol=atol, rtol=rtol)
torch.testing.assert_close(dv_test, dv_ref, atol=atol, rtol=rtol)
def test_varlen_qk_forward_opcheck(no_pad_impls):
mod = no_pad_impls
_fa2_only(mod)
torch.manual_seed(1)
device = torch.device("cuda")
b, sq, sk, h, d = 2, 48, 64, 4, 64
q = torch.randn(b, sq, h, d, device=device, dtype=torch.float16)
k = torch.randn(b, sk, h, d, device=device, dtype=torch.float16)
v = torch.randn(b, sk, h, d, device=device, dtype=torch.float16)
qmask = _padding_mask(b, sq, [48, 40], device)
kmask = _padding_mask(b, sk, [64, 56], device)
torch.library.opcheck(
torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward,
(q, k, v, qmask, kmask, False, 0.0, None, False),
)
def test_varlen_qk_opcheck_with_grad_inputs(no_pad_impls):
"""FA2: full opcheck with requires_grad inputs (autograd-registration check)."""
mod = no_pad_impls
_fa2_only(mod)
torch.manual_seed(1)
device = torch.device("cuda")
b, sq, sk, h, d = 2, 48, 64, 4, 64
q = torch.randn(b, sq, h, d, device=device, dtype=torch.float16, requires_grad=True)
k = torch.randn(b, sk, h, d, device=device, dtype=torch.float16, requires_grad=True)
v = torch.randn(b, sk, h, d, device=device, dtype=torch.float16, requires_grad=True)
qmask = _padding_mask(b, sq, [48, 40], device)
kmask = _padding_mask(b, sk, [64, 56], device)
torch.library.opcheck(
torch.ops.fastvideo._flash_attn_varlen_qk_no_pad_forward,
(q, k, v, qmask, kmask, False, 0.0, None, False),
)
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
"""Guard the Modal FA4 defaults that keep CI lanes on their intended backend.
Pure text/AST analysis: no fastvideo imports, no torch, no Modal client.
"""
from __future__ import annotations
import ast
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[3]
MODAL_ROOT = REPO_ROOT / "fastvideo" / "tests" / "modal"
PR_TEST = MODAL_ROOT / "pr_test.py"
LAUNCH_L40S_JOB = MODAL_ROOT / "launch_l40s_job.py"
SSIM_TEST = MODAL_ROOT / "ssim_test.py"
def _function_strings(path: Path, function_name: str) -> str:
source = path.read_text(encoding="utf-8")
tree = ast.parse(source)
for node in tree.body:
if isinstance(node, ast.FunctionDef) and node.name == function_name:
return "\n".join(
child.value
for child in ast.walk(node)
if isinstance(child, ast.Constant)
and isinstance(child.value, str)
)
raise AssertionError(f"{function_name} not found in {path}")
def test_generic_l40s_launcher_defaults_fa4_off():
source = LAUNCH_L40S_JOB.read_text(encoding="utf-8")
assert '"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "0")' in source
def test_ssim_launcher_keeps_fa4_enabled_by_default():
source = SSIM_TEST.read_text(encoding="utf-8")
assert '"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1")' in source
def test_pr_model_load_and_training_lanes_disable_fa4():
lanes = {
"run_transformer_tests": "pytest ./fastvideo/tests/transformers -vs",
"run_training_tests": "pytest ./fastvideo/tests/training/Vanilla -srP",
"run_training_lora_tests": "pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP",
"run_training_tests_VSA": "pytest ./fastvideo/tests/training/VSA -srP",
"run_distill_dmd_tests": "pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs",
"run_self_forcing_tests": "pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs",
"run_train_framework_tests": "pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs",
"seed_grad_norm_references": "pytest ./fastvideo/tests/train/methods -vs -rs",
}
for function_name, pytest_command in lanes.items():
function_strings = _function_strings(PR_TEST, function_name)
assert "FASTVIDEO_FA4=0" in function_strings
assert pytest_command in function_strings
+4 -3
View File
@@ -97,9 +97,10 @@ image = (
"TOKENIZERS_PARALLELISM": "false",
**({"UV_TORCH_BACKEND": uv_torch_backend_override} if uv_torch_backend_override else {}),
"FASTVIDEO_ATTENTION_BACKEND": os.environ.get("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN"),
# FA4 is opt-in (FASTVIDEO_FA4); keep CI parity with the seeded
# references. Caller override wins.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
# FA4 is opt-in (FASTVIDEO_FA4). Generic ad hoc jobs should follow the
# product default unless a caller opts in through the local env or
# --env-vars.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "0"),
})
)
+18 -10
View File
@@ -73,8 +73,10 @@ ci_env_secret = modal.Secret.from_dict({
**({
"UV_TORCH_BACKEND": uv_torch_backend_override
} if uv_torch_backend_override else {}),
# FA4 is opt-in (FASTVIDEO_FA4); CI lanes keep it enabled to match the
# SSIM/perf baselines. Caller override wins.
# FA4 is opt-in (FASTVIDEO_FA4). Keep the default enabled for
# inference/perf parity; model-load and training lanes that do not exercise
# FA4 explicitly set FASTVIDEO_FA4=0 in their command strings below.
# Caller override wins.
"FASTVIDEO_FA4": os.environ.get("FASTVIDEO_FA4", "1"),
})
@@ -184,7 +186,8 @@ def run_vae_tests():
volumes={"/root/data": model_vol})
def run_transformer_tests():
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/transformers -vs"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/transformers -vs"
)
@@ -197,7 +200,8 @@ def run_transformer_tests():
volumes={"/root/data": model_vol})
def run_training_tests():
run_test(
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/Vanilla -srP"
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/Vanilla -srP"
)
@@ -210,7 +214,8 @@ def run_training_tests():
volumes={"/root/data": model_vol})
def run_training_lora_tests():
run_test(
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP"
"export HF_HOME='/root/data/.cache' && wandb login $WANDB_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/lora/test_lora_training.py -srP"
)
@@ -220,7 +225,7 @@ def run_training_lora_tests():
secrets=[wandb_secret, ci_env_secret])
def run_training_tests_VSA():
run_test(
"wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/VSA -srP"
"wandb login $WANDB_API_KEY && FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/VSA -srP"
)
@@ -254,7 +259,7 @@ def run_inference_lora_tests():
@app.function(gpu="L40S:2", image=image, timeout=900, secrets=[ci_env_secret])
def run_distill_dmd_tests():
run_test(
"pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
@app.function(gpu="L40S:2",
@@ -263,7 +268,8 @@ def run_distill_dmd_tests():
secrets=[wandb_secret, ci_env_secret])
def run_self_forcing_tests():
run_test(
"wandb login $WANDB_API_KEY && pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs"
"wandb login $WANDB_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/training/self-forcing/test_self_forcing.py -vs"
)
@@ -323,7 +329,8 @@ def run_dreamverse_app_tests():
volumes={"/root/data": model_vol})
def run_train_framework_tests():
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
"FASTVIDEO_FA4=0 pytest ./fastvideo/tests/train/models ./fastvideo/tests/train/methods -vs"
)
@@ -349,7 +356,8 @@ def seed_grad_norm_references():
the local command and the ``_DEVICE_MAPPINGS`` table.
"""
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && "
"FASTVIDEO_FA4=0 FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
)
@@ -0,0 +1,190 @@
# SPDX-License-Identifier: Apache-2.0
"""Unit tests for the get_rotary_pos_embed memoization cache."""
import pytest
import torch
from fastvideo.layers.rotary_embedding import (
_ROTARY_POS_EMBED_CACHE,
_ROTARY_POS_EMBED_CACHE_MAXSIZE,
get_rotary_pos_embed,
)
def _rope_dim_list(hidden_size: int, heads_num: int) -> list[int]:
"""Return the default 3-axis rope_dim_list used by the video DiTs."""
d = hidden_size // heads_num
return [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
def _call(
rope_sizes=(21, 30, 52),
hidden_size=1536,
heads_num=12,
rope_dim_list="default",
rope_theta=10000.0,
dtype=torch.float64,
start_frame=0,
use_real=True,
**kwargs,
):
"""Thin wrapper around get_rotary_pos_embed with DiT-like defaults."""
if rope_dim_list == "default":
rope_dim_list = _rope_dim_list(hidden_size, heads_num)
return get_rotary_pos_embed(
rope_sizes,
hidden_size,
heads_num,
rope_dim_list,
rope_theta,
dtype=dtype,
start_frame=start_frame,
use_real=use_real,
**kwargs,
)
@pytest.fixture(autouse=True)
def _clear_cache():
"""Isolate every test by clearing the module-level cache around it."""
_ROTARY_POS_EMBED_CACHE.clear()
yield
_ROTARY_POS_EMBED_CACHE.clear()
def test_repeated_call_hits_cache():
"""A second identical call returns the exact same tensor objects."""
cos1, sin1 = _call()
cos2, sin2 = _call()
assert cos1 is cos2 and sin1 is sin2
assert len(_ROTARY_POS_EMBED_CACHE) == 1
def test_many_identical_calls_keep_single_entry():
"""Many identical calls never grow the cache beyond one entry."""
for _ in range(10):
_call()
assert len(_ROTARY_POS_EMBED_CACHE) == 1
@pytest.mark.parametrize(
"rope_sizes,hidden_size,heads_num,dtype,use_real",
[
((21, 30, 52), 1536, 12, torch.float64, True),
((21, 45, 80), 5120, 40, torch.float64, True),
((1, 16, 16), 1536, 12, torch.float32, True),
((4, 8, 8), 1536, 12, torch.float64, False),
],
)
def test_cached_matches_fresh_recompute(rope_sizes, hidden_size, heads_num,
dtype, use_real):
"""Cached tensors are bitwise-equal to a fresh uncached recompute."""
cos_cached, sin_cached = _call(rope_sizes=rope_sizes,
hidden_size=hidden_size,
heads_num=heads_num,
dtype=dtype,
use_real=use_real)
_ROTARY_POS_EMBED_CACHE.clear()
cos_fresh, sin_fresh = _call(rope_sizes=rope_sizes,
hidden_size=hidden_size,
heads_num=heads_num,
dtype=dtype,
use_real=use_real)
assert torch.equal(cos_cached, cos_fresh)
assert torch.equal(sin_cached, sin_fresh)
@pytest.mark.parametrize(
"kwargs_a,kwargs_b",
[
({"rope_sizes": (21, 30, 52)}, {"rope_sizes": (21, 45, 80)}),
({"dtype": torch.float64}, {"dtype": torch.float32}),
({"use_real": True}, {"use_real": False}),
({"start_frame": 0}, {"start_frame": 3}),
({"rope_theta": 10000.0}, {"rope_theta": 5000.0}),
({"shard_dim": 0}, {"shard_dim": 1}),
],
)
def test_distinct_args_create_distinct_entries(kwargs_a, kwargs_b):
"""Any output-affecting argument difference yields a separate cache entry."""
_call(**kwargs_a)
_call(**kwargs_b)
assert len(_ROTARY_POS_EMBED_CACHE) == 2
def test_none_rope_dim_list_shares_key_with_equivalent_list():
"""None rope_dim_list and its derived explicit list map to one entry."""
# head_dim must be divisible by 3 for the None branch to stay valid.
hidden_size, heads_num = 1536, 16 # head_dim == 96 -> [32, 32, 32]
_call(rope_dim_list=None, hidden_size=hidden_size, heads_num=heads_num)
before = len(_ROTARY_POS_EMBED_CACHE)
_call(rope_dim_list=[32, 32, 32], hidden_size=hidden_size, heads_num=heads_num)
assert len(_ROTARY_POS_EMBED_CACHE) == before == 1
def test_use_real_controls_last_dim():
"""use_real=True spans full head_dim; use_real=False spans half."""
cos_full, _ = _call(use_real=True)
cos_half, _ = _call(use_real=False)
assert cos_full.shape[-1] == 128
assert cos_half.shape[-1] == 64
@pytest.mark.parametrize("rope_sizes", [(1, 1, 1), (1, 30, 52), (21, 1, 1)])
def test_degenerate_grid_shapes(rope_sizes):
"""Degenerate single-element axes still produce a correctly sized table."""
cos, sin = _call(rope_sizes=rope_sizes)
expected = rope_sizes[0] * rope_sizes[1] * rope_sizes[2]
assert cos.shape[0] == expected
assert sin.shape[0] == expected
def test_scalar_and_list_factors_are_hashable_and_distinct():
"""List-valued rescale factors are hashable and keyed apart from scalars."""
_call(theta_rescale_factor=1.0)
_call(theta_rescale_factor=[1.0, 1.0, 1.0])
assert len(_ROTARY_POS_EMBED_CACHE) == 2
def test_caller_device_copy_does_not_corrupt_cache():
"""The .to()/.float() copy callers perform must not mutate cached tensors."""
cos, _ = _call()
snapshot = cos.clone()
_ = cos.to("cpu").float()
cos_again, _ = _call()
assert torch.equal(cos_again, snapshot)
def test_start_frame_offsets_values():
"""A non-zero start_frame shifts the temporal positions, changing output."""
cos0, _ = _call(start_frame=0)
cos3, _ = _call(start_frame=3)
assert not torch.equal(cos0, cos3)
assert len(_ROTARY_POS_EMBED_CACHE) == 2
def test_cache_is_bounded_and_evicts_oldest():
"""The cache caps at the max size and evicts the oldest entry first."""
# Tiny grids keep this lightweight; each start_frame is a distinct key.
overshoot = _ROTARY_POS_EMBED_CACHE_MAXSIZE + 4
for frame in range(overshoot):
_call(rope_sizes=(2, 2, 2), start_frame=frame)
assert len(_ROTARY_POS_EMBED_CACHE) <= _ROTARY_POS_EMBED_CACHE_MAXSIZE
assert len(_ROTARY_POS_EMBED_CACHE) == _ROTARY_POS_EMBED_CACHE_MAXSIZE
# The earliest-inserted frames must have been evicted; the latest survive.
surviving = {key[-2] for key in _ROTARY_POS_EMBED_CACHE} # start_frame slot
assert overshoot - 1 in surviving
assert 0 not in surviving
def test_cache_hit_refreshes_recency():
"""Re-accessing an entry protects it from eviction over an untouched one."""
_call(rope_sizes=(2, 2, 2), start_frame=0) # entry we will keep hot
for frame in range(1, _ROTARY_POS_EMBED_CACHE_MAXSIZE):
_call(rope_sizes=(2, 2, 2), start_frame=frame)
assert len(_ROTARY_POS_EMBED_CACHE) == _ROTARY_POS_EMBED_CACHE_MAXSIZE
_call(rope_sizes=(2, 2, 2), start_frame=0) # hit -> frame 0 becomes most recent
_call(rope_sizes=(2, 2, 2), start_frame=99) # miss -> evicts now-oldest (frame 1)
surviving = {key[-2] for key in _ROTARY_POS_EMBED_CACHE}
assert 0 in surviving
assert 1 not in surviving
@@ -62,12 +62,29 @@ def resolve_inference_device_reference_folder(logger: Logger) -> str:
return device_reference_folder
def _find_reference_video(reference_folder: str, prompt: str) -> str:
def _find_reference_media(
reference_folder: str,
prompt: str,
*,
media_extension: str,
) -> str:
"""Pick a reference file whose basename contains the prompt prefix."""
prompt_prefix = prompt[:100].strip()
allowed = (media_extension.lower(), ".mp4", ".png", ".jpg", ".jpeg")
matches: list[str] = []
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4") and prompt_prefix in filename:
return os.path.join(reference_folder, filename)
raise FileNotFoundError("Reference video missing")
low = filename.lower()
if not any(low.endswith(ext) for ext in allowed):
continue
if prompt_prefix in filename:
matches.append(filename)
if not matches:
raise FileNotFoundError("Reference media missing")
preferred = media_extension.lower().lstrip(".")
for name in matches:
if name.lower().endswith(f".{preferred}"):
return os.path.join(reference_folder, name)
return os.path.join(reference_folder, matches[0])
def _remove_stale_generated_video(output_dir: str, output_video_name: str) -> None:
@@ -80,21 +97,23 @@ def _assert_similarity(
*,
logger: Logger,
output_dir: str,
output_video_name: str,
output_media_name: str,
reference_folder: str,
prompt: str,
num_inference_steps: int,
min_acceptable_ssim: float,
model_id: str,
attention_backend_name: str,
media_extension: str,
) -> None:
generated_video_path = os.path.join(output_dir, output_video_name)
generated_media_path = os.path.join(output_dir, output_media_name)
artifact_kind = "image" if media_extension.lower() in (".png", ".jpg", ".jpeg") else "video"
if not os.path.exists(reference_folder):
logger.error("Reference folder missing: %s", reference_folder)
xfail_missing_reference_in_bootstrap_mode(
generated_artifact_path=generated_video_path,
generated_artifact_path=generated_media_path,
reference_folder=reference_folder,
artifact_kind="video",
artifact_kind=artifact_kind,
)
error_msg = (
f"Reference video folder does not exist: {reference_folder}\n"
@@ -104,28 +123,32 @@ def _assert_similarity(
raise FileNotFoundError(error_msg)
try:
reference_video_path = _find_reference_video(reference_folder, prompt)
reference_media_path = _find_reference_media(
reference_folder,
prompt,
media_extension=media_extension,
)
except FileNotFoundError as error:
logger.error(
"Reference video not found for prompt: %s with backend: %s",
"Reference media not found for prompt: %s with backend: %s",
prompt,
attention_backend_name,
)
xfail_missing_reference_in_bootstrap_mode(
generated_artifact_path=generated_video_path,
generated_artifact_path=generated_media_path,
reference_folder=reference_folder,
artifact_kind="video",
artifact_kind=artifact_kind,
)
raise error
logger.info(
"Computing SSIM between %s and %s",
reference_video_path,
generated_video_path,
reference_media_path,
generated_media_path,
)
ssim_values = compute_video_ssim_torchvision(
reference_video_path,
generated_video_path,
reference_media_path,
generated_media_path,
use_ms_ssim=True,
)
@@ -136,8 +159,8 @@ def _assert_similarity(
success = write_ssim_results(
output_dir,
ssim_values,
reference_video_path,
generated_video_path,
reference_media_path,
generated_media_path,
num_inference_steps,
prompt,
)
@@ -222,6 +245,7 @@ def run_text_to_video_similarity_test(
min_acceptable_ssim: float,
init_kwargs_override: dict[str, object] | None = None,
generation_kwargs_override: dict[str, object] | None = None,
media_extension: str = ".mp4",
) -> None:
with attention_backend(attention_backend_name):
output_dir = build_generated_output_dir(
@@ -230,9 +254,9 @@ def run_text_to_video_similarity_test(
model_id,
attention_backend_name,
)
output_video_name = f"{prompt[:100].strip()}.mp4"
output_media_name = f"{prompt[:100].strip()}{media_extension}"
os.makedirs(output_dir, exist_ok=True)
_remove_stale_generated_video(output_dir, output_video_name)
_remove_stale_generated_video(output_dir, output_media_name)
params_map = select_ssim_params(
default_params_map,
@@ -274,13 +298,14 @@ def run_text_to_video_similarity_test(
_assert_similarity(
logger=logger,
output_dir=output_dir,
output_video_name=output_video_name,
output_media_name=output_media_name,
reference_folder=reference_folder,
prompt=prompt,
num_inference_steps=num_inference_steps,
min_acceptable_ssim=min_acceptable_ssim,
model_id=model_id,
attention_backend_name=attention_backend_name,
media_extension=media_extension,
)
@@ -306,9 +331,9 @@ def run_image_to_video_similarity_test(
model_id,
attention_backend_name,
)
output_video_name = f"{prompt[:100].strip()}.mp4"
output_media_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
_remove_stale_generated_video(output_dir, output_video_name)
_remove_stale_generated_video(output_dir, output_media_name)
params_map = select_ssim_params(
default_params_map,
@@ -351,11 +376,12 @@ def run_image_to_video_similarity_test(
_assert_similarity(
logger=logger,
output_dir=output_dir,
output_video_name=output_video_name,
output_media_name=output_media_name,
reference_folder=reference_folder,
prompt=prompt,
num_inference_steps=num_inference_steps,
min_acceptable_ssim=min_acceptable_ssim,
model_id=model_id,
attention_backend_name=attention_backend_name,
media_extension=".mp4",
)
@@ -0,0 +1,134 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import os
import pytest
import torch
from fastvideo.api.flux import FluxSamplingParam
from fastvideo.logger import init_logger
from fastvideo.tests.ssim.inference_similarity_utils import (
run_text_to_video_similarity_test,
)
from fastvideo.tests.ssim.reference_utils import (
get_cuda_device_name,
resolve_device_reference_folder,
)
logger = init_logger(__name__)
REQUIRED_GPUS = 1
# MS-SSIM gate (see module docstring).
FLUX_T2I_MIN_SSIM = 0.98
FLUX_MODEL_PATH = os.getenv(
"FLUX_T2I_MODEL_DIR",
"black-forest-labs/FLUX.1-dev",
)
device_reference_folder = resolve_device_reference_folder(
(
("A40", "A40"),
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
("RTX 4090", "RTX4090"),
("4090", "RTX4090"),
),
device_name=get_cuda_device_name(),
fallback_device_prefix="L40S",
logger=logger,
)
# Folder token must match Hub path with slashes → double underscore (SD3.5
# pattern in ``test_sd35_similarity.py``).
MODEL_ID = "black-forest-labs__FLUX.1-dev"
TEST_PROMPTS = [
"a photo of a cat",
]
FLUX_DEFAULT_PARAMS: dict[str, object] = {
"num_gpus": 1,
"model_path": FLUX_MODEL_PATH,
"sp_size": 1,
"tp_size": 1,
"height": 256,
"width": 256,
"num_frames": 1,
"fps": 1,
"num_inference_steps": 8,
"guidance_scale": 3.5,
"seed": 0,
}
_flux_full_defaults = FluxSamplingParam()
FLUX_FULL_QUALITY_PARAMS: dict[str, object] = {
"num_gpus": 1,
"model_path": FLUX_MODEL_PATH,
"sp_size": 1,
"tp_size": 1,
"height": _flux_full_defaults.height,
"width": _flux_full_defaults.width,
"num_frames": 1,
"fps": _flux_full_defaults.fps,
"num_inference_steps": _flux_full_defaults.num_inference_steps,
"guidance_scale": _flux_full_defaults.guidance_scale,
"seed": _flux_full_defaults.seed,
}
FLUX_MODEL_TO_PARAMS = {
MODEL_ID: FLUX_DEFAULT_PARAMS,
}
FLUX_FULL_QUALITY_MODEL_TO_PARAMS = {
MODEL_ID: FLUX_FULL_QUALITY_PARAMS,
}
@pytest.mark.skipif(
not torch.cuda.is_available(),
reason="FLUX T2I SSIM test requires CUDA",
)
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(FLUX_MODEL_TO_PARAMS.keys()))
def test_flux_t2i_similarity(
prompt: str,
attention_backend_name: str,
model_id: str,
) -> None:
is_hf_repo = "/" in FLUX_MODEL_PATH and not FLUX_MODEL_PATH.startswith("/")
if not is_hf_repo and not os.path.isdir(FLUX_MODEL_PATH):
pytest.skip(
f"FLUX weights not found at {FLUX_MODEL_PATH} "
f"(set FLUX_T2I_MODEL_DIR to override)"
)
run_text_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=FLUX_MODEL_TO_PARAMS,
full_quality_params_map=FLUX_FULL_QUALITY_MODEL_TO_PARAMS,
min_acceptable_ssim=FLUX_T2I_MIN_SSIM,
media_extension=".png",
init_kwargs_override={
"workload_type": "t2i",
"use_fsdp_inference": False,
"text_encoder_cpu_offload": False,
"vae_cpu_offload": False,
"image_encoder_cpu_offload": False,
"pin_cpu_memory": False,
},
generation_kwargs_override={
"save_video": True,
"use_embedded_guidance": True,
"true_cfg_scale": 1.0,
},
)
@@ -0,0 +1,141 @@
# SPDX-License-Identifier: Apache-2.0
"""SSIM-based regression test for GLM-Image generation."""
from __future__ import annotations
import os
from pathlib import Path
import pytest
import torch
from fastvideo.logger import init_logger
from fastvideo.tests.ssim.inference_similarity_utils import (
run_text_to_video_similarity_test,
)
from fastvideo.tests.ssim.reference_utils import (
get_cuda_device_name,
resolve_device_reference_folder,
)
logger = init_logger(__name__)
REQUIRED_GPUS = 1
REPO_ROOT = Path(__file__).resolve().parents[3]
LOCAL_WEIGHTS_DIR = Path(
os.getenv("GLM_IMAGE_LOCAL_WEIGHTS_DIR",
REPO_ROOT / "official_weights" / "glm_image"))
GLM_IMAGE_MODEL_PATH = os.getenv("GLM_IMAGE_MODEL_DIR", str(LOCAL_WEIGHTS_DIR))
device_reference_folder = resolve_device_reference_folder(
(
("A40", "A40"),
("L40S", "L40S"),
("H100", "H100"),
("H200", "H200"),
("B200", "B200"),
),
device_name=get_cuda_device_name(),
fallback_device_prefix="L40S",
logger=logger,
)
MODEL_ID = "zai-org__GLM-Image"
TEST_PROMPTS = [
"A beautiful landscape photography with rolling hills, "
"a winding river, and a vibrant sunset in the background. "
"Warm golden light, photorealistic style.",
]
GLM_IMAGE_PARAMS = {
"num_gpus": 1,
"model_path": GLM_IMAGE_MODEL_PATH,
"sp_size": 1,
"tp_size": 1,
"height": 256,
"width": 256,
"num_frames": 1,
"fps": 1,
"num_inference_steps": 4,
"guidance_scale": 1.5,
"seed": 0,
"neg_prompt": "",
}
GLM_IMAGE_FULL_QUALITY_PARAMS = {
"num_gpus": 1,
"model_path": GLM_IMAGE_MODEL_PATH,
"sp_size": 1,
"tp_size": 1,
"height": 1024,
"width": 1024,
"num_frames": 1,
"fps": 1,
"num_inference_steps": 50,
"guidance_scale": 1.5,
"seed": 0,
"neg_prompt": "",
}
GLM_IMAGE_MODEL_TO_PARAMS = {
MODEL_ID: GLM_IMAGE_PARAMS,
}
GLM_IMAGE_FULL_QUALITY_MODEL_TO_PARAMS = {
MODEL_ID: GLM_IMAGE_FULL_QUALITY_PARAMS,
}
def _has_weights() -> bool:
required = ["transformer", "vae", "text_encoder",
"vision_language_encoder", "processor", "tokenizer",
"scheduler"]
return all((LOCAL_WEIGHTS_DIR / r).exists() for r in required)
def _upstream_glm_image_available() -> bool:
try:
import transformers
except ImportError:
return False
return hasattr(transformers, "GlmImageForConditionalGeneration")
@pytest.mark.skipif(
not torch.cuda.is_available(),
reason="GLM-Image SSIM test requires CUDA",
)
@pytest.mark.skipif(
not _has_weights(),
reason=f"GLM-Image full weights not found at {LOCAL_WEIGHTS_DIR}.",
)
@pytest.mark.skipif(
not _upstream_glm_image_available(),
reason="GLM-Image needs transformers>=5.0.0rc0 (ships the AR encoder).",
)
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(GLM_IMAGE_MODEL_TO_PARAMS.keys()))
def test_glm_image_similarity(
prompt: str,
attention_backend_name: str,
model_id: str,
) -> None:
run_text_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=GLM_IMAGE_MODEL_TO_PARAMS,
full_quality_params_map=GLM_IMAGE_FULL_QUALITY_MODEL_TO_PARAMS,
min_acceptable_ssim=0.98,
init_kwargs_override={
"trust_remote_code": True,
"use_fsdp_inference": False,
},
generation_kwargs_override={
"save_video": True,
},
)
@@ -0,0 +1,231 @@
# SPDX-License-Identifier: Apache-2.0
"""AnyFlow on-policy method tests (CPU-only, no Wan instantiation).
The full AnyFlowMethod requires a real student/teacher/critic trio plus
DMD2's optimizer wiring — too heavyweight for a unit test. These tests
exercise the rollout-shape helpers and source-level invariants via
``object.__new__`` bypassing of ``__init__``.
"""
from __future__ import annotations
import inspect
from types import SimpleNamespace
from typing import Any
import pytest
import torch
from fastvideo.train.methods.distribution_matching.anyflow import AnyFlowMethod
# ---------------------------------------------------------------------------
# Helpers: build a "naked" AnyFlowMethod that skips __init__.
# ---------------------------------------------------------------------------
def _naked_method(
*,
student_sample_steps: int = 4,
use_mean_velocity: bool = True,
t_list_override: list[float] | None = None,
denoising_step_list: list[float] | None = None,
) -> AnyFlowMethod:
method = AnyFlowMethod.__new__(AnyFlowMethod)
method._student_sample_steps = int(student_sample_steps) # type: ignore[attr-defined]
method._use_mean_velocity = bool(use_mean_velocity) # type: ignore[attr-defined]
method._t_list_override = ( # type: ignore[attr-defined]
list(t_list_override) if t_list_override else None)
method._dmd_score_r = 0.0 # type: ignore[attr-defined]
method._real_score_guidance = 1.0 # type: ignore[attr-defined]
method.cuda_generator = None # type: ignore[attr-defined]
method._cfg_uncond = None # type: ignore[attr-defined]
method._denoising_step_list_cache = None # type: ignore[attr-defined]
# Stub _get_denoising_step_list — DMD2 reads method_config but for these
# focused tests we want a deterministic schedule.
raw = denoising_step_list or [999, 750, 500, 250]
cached = torch.tensor(raw, dtype=torch.long)
def _stub(self, device: torch.device) -> torch.Tensor:
return cached.to(device=device)
bound = _stub.__get__(method, AnyFlowMethod)
method._get_denoising_step_list = bound # type: ignore[assignment]
return method
# ---------------------------------------------------------------------------
# Schedule construction.
# ---------------------------------------------------------------------------
def test_get_rollout_schedule_uses_t_list_override_verbatim() -> None:
method = _naked_method(
t_list_override=[999.0, 937.0, 833.0, 624.0, 0.0])
schedule = method._get_rollout_schedule(device=torch.device("cpu"))
torch.testing.assert_close(
schedule,
torch.tensor([999.0, 937.0, 833.0, 624.0, 0.0],
dtype=torch.float32),
)
def test_get_rollout_schedule_falls_back_to_denoising_step_list() -> None:
method = _naked_method(
denoising_step_list=[999, 750, 500, 250])
schedule = method._get_rollout_schedule(device=torch.device("cpu"))
# Tail must be a 0 boundary so the final Euler step lands at t=0.
assert float(schedule[-1].item()) == 0.0
# Original step list preserved at the front.
torch.testing.assert_close(
schedule[:4],
torch.tensor([999.0, 750.0, 500.0, 250.0], dtype=torch.float32),
)
def test_get_rollout_schedule_does_not_double_append_zero() -> None:
method = _naked_method(
denoising_step_list=[999, 750, 500, 0])
schedule = method._get_rollout_schedule(device=torch.device("cpu"))
assert schedule.numel() == 4 # no extra boundary inserted.
# ---------------------------------------------------------------------------
# t_list_override validation.
# ---------------------------------------------------------------------------
def test_anyflow_method_rejects_ascending_t_list_override() -> None:
"""__init__ validates that t_list_override is descending. We test the
validation logic by stitching together a minimal cfg + role_models
path; if construction fails for an unrelated reason we still catch
the descending check via the explicit error message."""
src = inspect.getsource(AnyFlowMethod.__init__)
assert 't_list_override must be descending' in src
assert 'descending' in src
def test_anyflow_method_rejects_non_positive_student_sample_steps() -> None:
src = inspect.getsource(AnyFlowMethod.__init__)
assert 'student_sample_steps must be positive' in src
# ---------------------------------------------------------------------------
# Rollout dynamics — stubbed student.
# ---------------------------------------------------------------------------
class _SpyStudent:
"""Stand-in student that records every (t, r) pair seen during a
rollout and predicts a constant velocity field."""
def __init__(self, num_train_timesteps: int = 1000) -> None:
self.num_train_timesteps = num_train_timesteps
self.seen: list[tuple[float, float]] = []
# A single trainable parameter so callers can verify gradient flow.
self.param = torch.nn.Parameter(torch.zeros(1))
def predict_velocity_with_r(
self,
noisy: torch.Tensor,
t: torch.Tensor,
r: torch.Tensor,
batch: Any,
*,
conditional: bool = True,
cfg_uncond: Any = None,
attn_kind: str = "vsa",
) -> torch.Tensor:
del batch, conditional, cfg_uncond, attn_kind
self.seen.append((float(t.flatten()[0].item()),
float(r.flatten()[0].item())))
# Constant velocity field of magnitude param so we can backprop.
return self.param * torch.ones_like(noisy)
def _make_batch(shape: tuple[int, ...]) -> SimpleNamespace:
batch = SimpleNamespace()
batch.latents = torch.randn(*shape)
batch.dmd_latent_vis_dict = {}
return batch
def test_rollout_uses_mean_velocity_r_equals_t_next() -> None:
"""With use_mean_velocity=True, r at step i must equal t at step i+1."""
method = _naked_method(
student_sample_steps=4,
use_mean_velocity=True,
t_list_override=[999.0, 750.0, 500.0, 250.0, 0.0],
)
student = _SpyStudent()
method.student = student # type: ignore[assignment]
batch = _make_batch((1, 2, 4, 4, 4))
_ = method._student_rollout(batch, with_grad=False)
# 4 forwards = 4 (t, r) pairs.
assert len(student.seen) == 4
for i in range(3):
# r at step i must equal t at step i+1.
assert student.seen[i][1] == student.seen[i + 1][0]
def test_rollout_use_mean_velocity_false_uses_r_equal_t() -> None:
method = _naked_method(
student_sample_steps=2,
use_mean_velocity=False,
t_list_override=[999.0, 500.0, 0.0],
)
student = _SpyStudent()
method.student = student # type: ignore[assignment]
batch = _make_batch((1, 2, 4, 4, 4))
_ = method._student_rollout(batch, with_grad=False)
for t_seen, r_seen in student.seen:
assert t_seen == r_seen
def test_rollout_with_grad_true_produces_differentiable_output() -> None:
method = _naked_method(
student_sample_steps=4,
use_mean_velocity=True,
t_list_override=[999.0, 750.0, 500.0, 250.0, 0.0],
)
student = _SpyStudent()
method.student = student # type: ignore[assignment]
batch = _make_batch((1, 2, 4, 4, 4))
out = method._student_rollout(batch, with_grad=True)
assert out.requires_grad, (
"Rollout output must keep a gradient so the DMD loss can backprop "
"through the chosen step.")
out.sum().backward()
assert student.param.grad is not None
assert student.param.grad.abs().sum() > 0
def test_rollout_with_grad_false_blocks_gradient_completely() -> None:
method = _naked_method(
student_sample_steps=4,
use_mean_velocity=True,
t_list_override=[999.0, 750.0, 500.0, 250.0, 0.0],
)
student = _SpyStudent()
method.student = student # type: ignore[assignment]
batch = _make_batch((1, 2, 4, 4, 4))
out = method._student_rollout(batch, with_grad=False)
assert not out.requires_grad
def test_broadcast_grad_step_index_in_range() -> None:
method = _naked_method(student_sample_steps=4)
for _ in range(20):
idx = method._broadcast_grad_step_index(
num_steps=4, device=torch.device("cpu"))
assert 0 <= idx < 4
def test_broadcast_grad_step_index_rejects_non_positive_num_steps() -> None:
method = _naked_method()
with pytest.raises(ValueError, match="num_steps must be positive"):
method._broadcast_grad_step_index(
num_steps=0, device=torch.device("cpu"))
@@ -0,0 +1,604 @@
# SPDX-License-Identifier: Apache-2.0
"""AnyFlow pretrain method tests.
CPU-only unit tests covering:
- Config flag defaults (bit-identity preserved on legacy paths).
- ``WanTimeTextImageEmbedding`` dual-timestep forward (additive default = bit-identical
to legacy; gated mode reproduces AnyFlow's ``(1 - g) * temb + g * delta_emb`` fusion).
- ``WanTransformer3DModel.forward`` accepts ``r_timestep``.
- ``FlowMapEulerDiscreteScheduler`` numerics: ``apply_shift``, ``get_train_weight``,
``step``.
- ``(t, r)`` per-batch sampling distribution.
- Central-difference target math.
- AnyFlow HF checkpoint key remap (``remap_anyflow_keys``).
"""
from __future__ import annotations
import copy
import pytest
import torch
from fastvideo.configs.models.dits import WanVideoConfig
# ---------------------------------------------------------------------------
# Task 1: r_embedder config flags default to bit-identity preservation.
# ---------------------------------------------------------------------------
def test_wan_arch_defaults_preserve_bit_identity() -> None:
cfg = WanVideoConfig()
arch = cfg.arch_config
assert arch.r_embedder is False
assert arch.r_embedder_fusion == "additive"
assert arch.r_embedder_gate_value == 0.25
assert arch.r_embedder_deltatime_type == "r"
# ---------------------------------------------------------------------------
# Task 2: WanTimeTextImageEmbedding dual-timestep forward.
# ---------------------------------------------------------------------------
def _init_uninitialized_weights(module: torch.nn.Module, seed: int = 0) -> None:
"""FastVideo's ``ReplicatedLinear`` allocates weights with
``torch.empty`` and relies on a downstream ``load_weights`` pass to
populate them. Unit tests bypass that pass, so weights start as
uninitialized garbage (typically NaN/Inf). Manually init every
Linear / RMSNorm / LayerNorm parameter so the forward produces
deterministic finite outputs.
"""
torch.manual_seed(seed)
with torch.no_grad():
for p in module.parameters():
if p.ndim >= 2:
# Xavier-uniform scaled by inverse fan-in for stable forwards.
torch.nn.init.xavier_uniform_(p)
else:
p.zero_()
def _make_embedder(
*,
r_embedder: bool,
fusion: str = "additive",
gate: float = 0.25,
deltatime_type: str = "r",
init_seed: int = 0,
):
from fastvideo.models.dits.wanvideo import WanTimeTextImageEmbedding
emb = WanTimeTextImageEmbedding(
dim=32,
time_freq_dim=64,
text_embed_dim=16,
image_embed_dim=None,
r_embedder=r_embedder,
r_embedder_fusion=fusion,
r_embedder_gate_value=gate,
r_embedder_deltatime_type=deltatime_type,
)
_init_uninitialized_weights(emb, seed=init_seed)
emb.eval()
return emb
def test_embedder_default_path_no_delta_module() -> None:
"""When r_embedder=False, delta_embedder must not be allocated."""
emb = _make_embedder(r_embedder=False)
assert emb.delta_embedder is None
def test_embedder_default_path_is_bit_identical_to_legacy() -> None:
"""With r_embedder=False, forward output must match the legacy single-t path
(no r_timestep kwarg, no extra computation)."""
torch.manual_seed(0)
emb = _make_embedder(r_embedder=False)
t = torch.randint(0, 1000, (2,), dtype=torch.long)
txt = torch.randn(2, 4, 16)
temb_a, proj_a, _, _ = emb(t, txt)
# Calling without r_timestep again must be deterministic-equal.
temb_b, proj_b, _, _ = emb(t, txt)
torch.testing.assert_close(temb_a, temb_b)
torch.testing.assert_close(proj_a, proj_b)
assert temb_a.shape == (2, 32)
assert proj_a.shape == (2, 32 * 6)
def test_embedder_enabled_without_r_timestep_is_bit_identical_to_legacy() -> None:
"""Even with r_embedder=True, if r_timestep is None at call time the
forward must skip the delta path entirely so existing call sites that
don't pass r_timestep stay byte-equal to the legacy result."""
torch.manual_seed(0)
emb_legacy = _make_embedder(r_embedder=False)
torch.manual_seed(0)
emb_dual = _make_embedder(r_embedder=True, fusion="additive")
t = torch.randint(0, 1000, (2,), dtype=torch.long)
txt = torch.randn(2, 4, 16)
temb_legacy, proj_legacy, _, _ = emb_legacy(t, txt)
temb_dual, proj_dual, _, _ = emb_dual(t, txt) # No r_timestep.
torch.testing.assert_close(temb_legacy, temb_dual)
torch.testing.assert_close(proj_legacy, proj_dual)
def test_embedder_gated_fusion_formula() -> None:
"""Gated mode: rt_emb = (1 - g) * temb_t + g * delta_emb (with delta_input=r)."""
torch.manual_seed(0)
emb = _make_embedder(r_embedder=True, fusion="gated", gate=0.25)
t = torch.tensor([500, 500], dtype=torch.long)
r = torch.tensor([100, 100], dtype=torch.long)
txt = torch.randn(2, 4, 16)
temb_t = emb.time_embedder(t)
delta_emb = emb.delta_embedder(r)
expected = 0.75 * temb_t + 0.25 * delta_emb
rt_emb, _, _, _ = emb(t, txt, r_timestep=r)
torch.testing.assert_close(rt_emb, expected, rtol=1e-5, atol=1e-5)
def test_embedder_additive_fusion_formula() -> None:
"""Additive mode: rt_emb = temb_t + g * delta_emb."""
torch.manual_seed(0)
emb = _make_embedder(r_embedder=True, fusion="additive", gate=0.3)
t = torch.tensor([700, 700], dtype=torch.long)
r = torch.tensor([200, 200], dtype=torch.long)
txt = torch.randn(2, 4, 16)
temb_t = emb.time_embedder(t)
delta_emb = emb.delta_embedder(r)
expected = temb_t + 0.3 * delta_emb
rt_emb, _, _, _ = emb(t, txt, r_timestep=r)
torch.testing.assert_close(rt_emb, expected, rtol=1e-5, atol=1e-5)
def test_embedder_deltatime_type_t_minus_r() -> None:
"""When deltatime_type='t-r', delta_embedder consumes (t - r)."""
torch.manual_seed(0)
emb = _make_embedder(
r_embedder=True, fusion="gated", gate=0.5, deltatime_type="t-r")
t = torch.tensor([800, 800], dtype=torch.long)
r = torch.tensor([300, 300], dtype=torch.long)
txt = torch.randn(2, 4, 16)
temb_t = emb.time_embedder(t)
delta_emb = emb.delta_embedder(t - r)
expected = 0.5 * temb_t + 0.5 * delta_emb
rt_emb, _, _, _ = emb(t, txt, r_timestep=r)
torch.testing.assert_close(rt_emb, expected, rtol=1e-5, atol=1e-5)
def test_embedder_invalid_fusion_raises() -> None:
with pytest.raises(ValueError, match="r_embedder_fusion"):
_make_embedder(r_embedder=True, fusion="bogus")
def test_embedder_invalid_deltatime_type_raises() -> None:
with pytest.raises(ValueError, match="r_embedder_deltatime_type"):
_make_embedder(r_embedder=True, fusion="gated", deltatime_type="2t-r")
def test_embedder_gate_not_in_state_dict() -> None:
"""Gate is a non-persistent buffer; it must not appear in state_dict so
checkpoints stay portable across different gate hyperparameters."""
emb = _make_embedder(r_embedder=True, fusion="gated", gate=0.25)
keys = list(emb.state_dict().keys())
assert not any("_r_embedder_gate" in k for k in keys)
# ---------------------------------------------------------------------------
# Task 3: WanTransformer3DModel threads r_timestep through.
# ---------------------------------------------------------------------------
def test_wan_transformer_forward_signature_has_r_timestep() -> None:
"""The forward signature must declare r_timestep explicitly (not
swallowed by **kwargs) so callers and type checkers can see it."""
import inspect
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
sig = inspect.signature(WanTransformer3DModel.forward)
assert "r_timestep" in sig.parameters
param = sig.parameters["r_timestep"]
assert param.default is None
def test_wan_transformer_init_propagates_r_embedder_config() -> None:
"""When the arch config sets r_embedder=True the WanTransformer3DModel
constructor must instantiate the embedder with the delta path active.
We avoid full WanTransformer3DModel instantiation (which requires
distributed init) by reading the source's __init__ to confirm it
forwards the four arch flags to WanTimeTextImageEmbedding.
"""
import inspect
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
src = inspect.getsource(WanTransformer3DModel.__init__)
# All four arch config fields must be passed to WanTimeTextImageEmbedding.
assert "r_embedder=config.r_embedder" in src
assert "r_embedder_fusion=config.r_embedder_fusion" in src
assert "r_embedder_gate_value=config.r_embedder_gate_value" in src
assert "r_embedder_deltatime_type=config.r_embedder_deltatime_type" in src
def test_wan_transformer_forward_threads_r_timestep_to_embedder() -> None:
"""The forward must pass r_timestep into the embedder call (verified via
source inspection to avoid heavyweight distributed bring-up)."""
import inspect
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
src = inspect.getsource(WanTransformer3DModel.forward)
assert "r_timestep=r_timestep" in src, (
"WanTransformer3DModel.forward must forward r_timestep into "
"self.condition_embedder")
# ---------------------------------------------------------------------------
# Task 4: FlowMapEulerDiscreteScheduler numerics.
# ---------------------------------------------------------------------------
def _scheduler(*, shift: float = 1.0, n_train: int = 1000):
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
FlowMapEulerDiscreteScheduler, )
return FlowMapEulerDiscreteScheduler(
num_train_timesteps=n_train, shift=shift)
def test_flow_map_scheduler_set_timesteps_descending() -> None:
sched = _scheduler(shift=5.0)
sched.set_timesteps(num_inference_steps=4, device=torch.device("cpu"))
ts = sched.timesteps
# N inference steps → N + 1 boundary entries.
assert ts.numel() == 5
assert torch.all(ts[:-1] >= ts[1:]) # descending
assert ts[-1].item() == 0.0
assert ts[0].item() == pytest.approx(1000.0, abs=1e-3)
def test_flow_map_scheduler_set_timesteps_custom_overrides_schedule() -> None:
sched = _scheduler(shift=5.0)
custom = [999.0, 937.0, 833.0, 624.0, 0.0]
sched.set_timesteps(
num_inference_steps=4,
device=torch.device("cpu"),
custom_timesteps=custom,
)
torch.testing.assert_close(
sched.timesteps, torch.tensor(custom, dtype=torch.float32))
def test_flow_map_scheduler_custom_timesteps_must_be_descending() -> None:
sched = _scheduler(shift=5.0)
with pytest.raises(ValueError, match="descending"):
sched.set_timesteps(
num_inference_steps=4,
device=torch.device("cpu"),
custom_timesteps=[100.0, 500.0, 900.0],
)
def test_flow_map_scheduler_apply_shift_endpoints_invariant() -> None:
"""apply_shift fixes the endpoints {0, 1} and produces non-trivial
motion in the interior for shift != 1."""
sched = _scheduler(shift=5.0)
t = torch.tensor([0.0, 0.5, 1.0])
shifted = sched.apply_shift(t)
torch.testing.assert_close(
shifted, torch.tensor([0.0, 5.0 / 6.0, 1.0]), rtol=1e-6, atol=1e-6)
def test_flow_map_scheduler_apply_shift_shift_one_is_identity() -> None:
sched = _scheduler(shift=1.0)
t = torch.linspace(0.0, 1.0, 100)
torch.testing.assert_close(sched.apply_shift(t), t)
def test_flow_map_scheduler_step_one_euler_iteration_matches_formula() -> None:
"""One step: x_r = x_t - ((t - r) / N) * model_output."""
sched = _scheduler(shift=1.0)
sched.set_timesteps(num_inference_steps=4, device=torch.device("cpu"))
torch.manual_seed(0)
x_t = torch.randn(2, 4, 1, 8, 8)
v = torch.randn_like(x_t)
t = torch.tensor([750.0, 500.0])
r = torch.tensor([500.0, 250.0])
out = sched.step(v, sample=x_t, timestep=t, r_timestep=r)
expected = x_t - ((t - r) / 1000.0).view(-1, 1, 1, 1, 1) * v
torch.testing.assert_close(out, expected, rtol=1e-6, atol=1e-6)
def test_flow_map_scheduler_get_train_weight_beta08_shape_and_renorm() -> None:
"""beta08: t * sqrt(1-t), renormalized so sum equals num_train_timesteps.
The interior of the schedule must dominate the endpoints (monotone up
then monotone down)."""
sched = _scheduler()
t = torch.linspace(0.001, 0.999, 1000)
w = sched.get_train_weight(t, weight_type="beta08")
assert torch.allclose(w.sum(), torch.tensor(1000.0), rtol=1e-3)
assert torch.all(w >= 0.0)
# Endpoints smaller than the middle bump.
mid = len(w) // 2
assert w[0] < w[mid]
assert w[-1] < w[mid]
def test_flow_map_scheduler_get_train_weight_uniform_is_constant_norm() -> None:
sched = _scheduler()
t = torch.linspace(0.0, 1.0, 1000)
w = sched.get_train_weight(t, weight_type="uniform")
assert torch.allclose(w.sum(), torch.tensor(1000.0), rtol=1e-3)
# All entries equal to 1.0 after renormalization.
torch.testing.assert_close(w, torch.ones_like(w), rtol=1e-6, atol=1e-6)
def test_flow_map_scheduler_get_train_weight_accepts_absolute_units() -> None:
"""When t is provided in [0, num_train_timesteps] the helper auto-
normalizes; result must match the [0, 1] call."""
sched = _scheduler()
t_norm = torch.linspace(0.001, 0.999, 1000)
t_abs = t_norm * 1000
w_norm = sched.get_train_weight(t_norm, weight_type="beta08")
w_abs = sched.get_train_weight(t_abs, weight_type="beta08")
torch.testing.assert_close(w_norm, w_abs, rtol=1e-5, atol=1e-5)
def test_flow_map_scheduler_add_noise_matches_flow_matching_formula() -> None:
"""Linear flow-matching: x_t = (1 - sigma) * x_0 + sigma * eps, where
sigma = t / num_train_timesteps."""
sched = _scheduler()
torch.manual_seed(0)
x0 = torch.randn(2, 4, 1, 4, 4)
eps = torch.randn_like(x0)
t = torch.tensor([250.0, 750.0])
out = sched.add_noise(x0, eps, t)
sigma = (t / 1000.0).view(-1, 1, 1, 1, 1)
expected = (1.0 - sigma) * x0 + sigma * eps
torch.testing.assert_close(out, expected, rtol=1e-6, atol=1e-6)
# ---------------------------------------------------------------------------
# Task 6: (t, r) per-batch sampling.
# ---------------------------------------------------------------------------
def test_sample_pair_timesteps_partitions_batch_correctly() -> None:
"""For batch=8 with diffusion=0.5/consistency=0.25: 4 r=t, 2 r=0,
2 free entries."""
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_sample_pair_timesteps, )
torch.manual_seed(42)
t, r, is_diffusion, is_consistency = _sample_pair_timesteps(
batch_size=8,
diffusion_ratio=0.5,
consistency_ratio=0.25,
device=torch.device("cpu"),
generator=None,
)
assert t.shape == (8,)
assert r.shape == (8,)
assert int(is_diffusion.sum()) == 4
assert int(is_consistency.sum()) == 2
# The masks must be disjoint.
assert not torch.any(is_diffusion & is_consistency)
diff_idx = torch.nonzero(is_diffusion).flatten()
torch.testing.assert_close(r[diff_idx], t[diff_idx])
cons_idx = torch.nonzero(is_consistency).flatten()
torch.testing.assert_close(r[cons_idx], torch.zeros(2))
free_idx = torch.nonzero(~(is_diffusion | is_consistency)).flatten()
assert torch.all(r[free_idx] <= t[free_idx])
assert torch.all(r[free_idx] >= 0.0)
assert torch.all(t[free_idx] <= 1.0)
def test_sample_pair_timesteps_t_max_r_min_ordering() -> None:
"""For the free fraction (ratios=0), t and r come from max/min of two
uniform draws so r <= t holds."""
torch.manual_seed(0)
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_sample_pair_timesteps, )
for _ in range(50):
t, r, is_diff, is_cons = _sample_pair_timesteps(
batch_size=4,
diffusion_ratio=0.0,
consistency_ratio=0.0,
device=torch.device("cpu"),
generator=None,
)
assert torch.all(t >= r)
assert torch.all(r >= 0.0)
assert torch.all(t <= 1.0)
assert int(is_diff.sum()) == 0
assert int(is_cons.sum()) == 0
def test_sample_pair_timesteps_rejects_ratios_summing_above_one() -> None:
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_sample_pair_timesteps, )
with pytest.raises(ValueError, match="must be <= 1"):
_sample_pair_timesteps(
batch_size=8,
diffusion_ratio=0.7,
consistency_ratio=0.4,
device=torch.device("cpu"),
generator=None,
)
def test_sample_pair_timesteps_rejects_negative_ratios() -> None:
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_sample_pair_timesteps, )
with pytest.raises(ValueError, match="non-negative"):
_sample_pair_timesteps(
batch_size=8,
diffusion_ratio=-0.1,
consistency_ratio=0.25,
device=torch.device("cpu"),
generator=None,
)
# ---------------------------------------------------------------------------
# Task 7: central-difference target.
# ---------------------------------------------------------------------------
class _StubStudent:
"""Stand-in student for unit-testing the central-difference helper.
The "velocity prediction" is a closed-form function of x and t
(no actual neural network) so we can compare against the analytical
derivative.
"""
def __init__(self, alpha: float = 0.3) -> None:
self.alpha = float(alpha)
def predict_velocity_with_r(
self,
noisy: torch.Tensor,
t: torch.Tensor,
r: torch.Tensor,
batch,
*,
conditional: bool = True,
attn_kind: str = "dense",
cfg_uncond=None,
) -> torch.Tensor:
del batch, conditional, attn_kind, cfg_uncond, r
# f(x, t) = x + alpha * (t / 1000) → dF/dt = alpha / 1000.
view = [-1] + [1] * (noisy.ndim - 1)
return noisy + self.alpha * (t.view(*view).float() / 1000.0)
def test_central_difference_dF_dt_linear_function() -> None:
"""For f(x, t) = x + alpha * (t / N), the central-difference estimate
of dF/dt must be alpha / N (in absolute t-units that's alpha / N
velocity per t-unit, and our helper divides the difference by
2*delta -> exactly alpha / N)."""
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_central_difference_dF_dt, )
student = _StubStudent(alpha=0.3)
x = torch.randn(2, 1, 4, 4, 4)
latents = torch.zeros_like(x)
noise = torch.zeros_like(x) # v_pred = noise - latents = 0
t = torch.tensor([500.0, 250.0])
r = torch.tensor([100.0, 0.0])
dF = _central_difference_dF_dt(
student=student,
batch=None,
noisy=x,
latents=latents,
noise=noise,
t=t,
r=r,
delta=5.0,
num_train_timesteps=1000.0,
)
expected = torch.full_like(x, 0.3 / 1000.0)
torch.testing.assert_close(dF, expected, rtol=1e-5, atol=1e-5)
def test_central_difference_dF_dt_guidance_scaling() -> None:
"""With guidance_scale != 1, the result must be divided by guidance."""
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_central_difference_dF_dt, )
student = _StubStudent(alpha=0.6)
x = torch.zeros(1, 1, 4, 4, 4)
latents = torch.zeros_like(x)
noise = torch.zeros_like(x)
t = torch.tensor([500.0])
r = torch.tensor([100.0])
dF_g1 = _central_difference_dF_dt(
student=student, batch=None, noisy=x, latents=latents, noise=noise,
t=t, r=r, delta=5.0, num_train_timesteps=1000.0, guidance_scale=1.0)
dF_g3 = _central_difference_dF_dt(
student=student, batch=None, noisy=x, latents=latents, noise=noise,
t=t, r=r, delta=5.0, num_train_timesteps=1000.0, guidance_scale=3.0)
torch.testing.assert_close(dF_g1 / 3.0, dF_g3, rtol=1e-5, atol=1e-5)
def test_central_difference_dF_dt_rejects_zero_delta() -> None:
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
_central_difference_dF_dt, )
with pytest.raises(ValueError, match="delta must be positive"):
_central_difference_dF_dt(
student=_StubStudent(),
batch=None,
noisy=torch.zeros(1, 1, 4, 4, 4),
latents=torch.zeros(1, 1, 4, 4, 4),
noise=torch.zeros(1, 1, 4, 4, 4),
t=torch.tensor([500.0]),
r=torch.tensor([100.0]),
delta=0.0,
num_train_timesteps=1000.0,
)
# ---------------------------------------------------------------------------
# Task 9: param_names_mapping handles AnyFlow checkpoints + is a no-op on plain
# Wan checkpoints. We don't ship a separate remap_anyflow_keys helper because
# the existing param_names_mapping regex mechanism does the job (delta_embedder
# rename is a no-op when the source state dict doesn't contain those keys).
# ---------------------------------------------------------------------------
def test_param_names_mapping_includes_delta_embedder_rename() -> None:
"""The Wan arch config's param_names_mapping must rename HF AnyFlow
delta_embedder weights into FastVideo's internal mlp.fc_in/fc_out layout."""
arch = WanVideoConfig().arch_config
mapping_keys = list(arch.param_names_mapping.keys())
assert any("delta_embedder" in k for k in mapping_keys), (
"WanVideoArchConfig.param_names_mapping must include delta_embedder "
"rename so HF AnyFlow checkpoints load without a separate adapter")
def test_param_names_mapping_default_doesnt_break_plain_wan_keys() -> None:
"""The new delta_embedder regex must not match any key in a plain
pretrained Wan2.1 checkpoint (those don't have delta_embedder)."""
import re
plain_wan_keys = [
"patch_embedding.weight",
"condition_embedder.time_embedder.linear_1.weight",
"condition_embedder.time_embedder.linear_2.weight",
"condition_embedder.time_proj.weight",
"condition_embedder.text_embedder.linear_1.weight",
"blocks.0.attn1.to_q.weight",
"blocks.0.ffn.net.0.proj.weight",
]
arch = WanVideoConfig().arch_config
delta_regexes = [
k for k in arch.param_names_mapping if "delta_embedder" in k
]
assert delta_regexes, "expected at least one delta_embedder regex"
for plain in plain_wan_keys:
for rx in delta_regexes:
assert re.match(rx, plain) is None, (
f"delta_embedder regex {rx!r} unexpectedly matched plain "
f"Wan key {plain!r}")
@@ -0,0 +1,130 @@
# SPDX-License-Identifier: Apache-2.0
"""GPU smoke test for AnyFlow pretrain + on-policy.
Mirrors ``test_distill_dmd.py`` — runs the new YAML-driven training
entrypoint via ``torchrun`` for two iterations to verify end-to-end
wiring (model load, optimizer build, forward, backward, step, save).
CPU-only environments are skipped; the test is intended to fire on the
Buildkite ``/test distillation`` lane and on local boxes with at least
2 H100/H200 GPUs.
"""
from __future__ import annotations
import os
import subprocess
import sys
from pathlib import Path
import pytest
REPO_ROOT = Path(__file__).resolve().parents[4]
PRETRAIN_YAML = (
REPO_ROOT
/ "examples/train/configs/distribution_matching/wan/anyflow_pretrain_t2v.yaml")
ONPOLICY_YAML = (
REPO_ROOT
/ "examples/train/configs/distribution_matching/wan/anyflow_onpolicy_t2v.yaml")
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
def _have_enough_gpus() -> bool:
"""Return True iff at least 2 CUDA devices are visible. The smoke test
needs HSDP/FSDP with a non-trivial world size; single-GPU bring-up
races against the new framework's distributed barriers."""
try:
import torch
except Exception:
return False
if not torch.cuda.is_available():
return False
return torch.cuda.device_count() >= 2
pytestmark = pytest.mark.skipif(
not _have_enough_gpus(),
reason="AnyFlow smoke test requires >= 2 CUDA devices")
def _run_torchrun(config_path: Path, *, output_dir: Path) -> None:
if not config_path.exists():
pytest.fail(f"YAML config missing: {config_path}")
env = os.environ.copy()
env.setdefault("MASTER_ADDR", "127.0.0.1")
env.setdefault("MASTER_PORT", "29551")
env.setdefault("WANDB_MODE", "offline")
env.setdefault("TOKENIZERS_PARALLELISM", "false")
cmd = [
sys.executable,
"-m",
"torch.distributed.run",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
"--master_port", env["MASTER_PORT"],
"-m", "fastvideo.train.entrypoint.train",
"--config", str(config_path),
"--training.loop.max_train_steps", "2",
"--training.checkpoint.output_dir", str(output_dir),
"--training.distributed.num_gpus", NUM_GPUS_PER_NODE,
"--training.distributed.hsdp_shard_dim", NUM_GPUS_PER_NODE,
"--training.data.train_batch_size", "1",
]
process = subprocess.run(cmd, capture_output=True, text=True, env=env)
if process.stdout:
print("STDOUT:", process.stdout)
if process.stderr:
print("STDERR:", process.stderr)
if process.returncode != 0:
raise subprocess.CalledProcessError(
process.returncode, cmd, process.stdout, process.stderr)
def test_anyflow_pretrain_smoke(tmp_path: Path) -> None:
"""Two-iteration pretrain — exercises (t, r) sampling, central-difference
target, scale balance, optimizer step, and checkpoint save path."""
_run_torchrun(PRETRAIN_YAML, output_dir=tmp_path / "pretrain")
def test_anyflow_onpolicy_smoke(tmp_path: Path) -> None:
"""Two-iteration on-policy DMD — exercises the multi-step Euler-flow
rollout, grad-step broadcast, DMD2 alternating updates."""
# Override init_from to the public Wan2.1-T2V-1.3B-Diffusers checkpoint
# for the smoke run; param_names_mapping handles the (non-existent)
# delta_embedder rename as a no-op.
env = os.environ.copy()
env.setdefault("MASTER_ADDR", "127.0.0.1")
env.setdefault("MASTER_PORT", "29552")
env.setdefault("WANDB_MODE", "offline")
env.setdefault("TOKENIZERS_PARALLELISM", "false")
cmd = [
sys.executable,
"-m",
"torch.distributed.run",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
"--master_port", env["MASTER_PORT"],
"-m", "fastvideo.train.entrypoint.train",
"--config", str(ONPOLICY_YAML),
"--training.loop.max_train_steps", "2",
"--training.checkpoint.output_dir", str(tmp_path / "onpolicy"),
"--training.distributed.num_gpus", NUM_GPUS_PER_NODE,
"--training.distributed.hsdp_shard_dim", NUM_GPUS_PER_NODE,
"--training.data.train_batch_size", "1",
"--models.student.init_from", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--method.student_sample_steps", "2",
]
process = subprocess.run(cmd, capture_output=True, text=True, env=env)
if process.stdout:
print("STDOUT:", process.stdout)
if process.stderr:
print("STDERR:", process.stderr)
if process.returncode != 0:
raise subprocess.CalledProcessError(
process.returncode, cmd, process.stdout, process.stderr)
+185
View File
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import glob
import os
import pytest
import torch
from diffusers import FluxTransformer2DModel as HFFluxTransformer2DModel
from torch.testing import assert_close
from fastvideo.configs.models.dits.flux import FluxDiTConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
os.environ.setdefault("MASTER_ADDR", "localhost")
os.environ.setdefault("MASTER_PORT", "29517")
_REPO_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", ".."))
_DEFAULT_FLUX_TRANSFORMER = os.path.join(
_REPO_ROOT,
"official_weights",
"FLUX.1-dev",
"transformer",
)
def _flux_transformer_path() -> str:
return os.environ.get("FLUX_TRANSFORMER_PATH", _DEFAULT_FLUX_TRANSFORMER)
def _prepare_latent_image_ids(
height: int,
width: int,
device: torch.device,
dtype: torch.dtype = torch.long,
) -> torch.Tensor:
"""Match Diffusers ``FluxPipeline._prepare_latent_image_ids`` (batch omitted)."""
latent_image_ids = torch.zeros(height, width, 3)
latent_image_ids[..., 1] = latent_image_ids[..., 1] + torch.arange(height)[:, None]
latent_image_ids[..., 2] = latent_image_ids[..., 2] + torch.arange(width)[None, :]
h, w, c = latent_image_ids.shape
latent_image_ids = latent_image_ids.reshape(h * w, c)
return latent_image_ids.to(device=device, dtype=dtype)
@pytest.fixture
def torch_sdpa_attention_backend(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
requires_cuda = pytest.mark.skipif(
not torch.cuda.is_available(),
reason="FLUX DiT parity test requires CUDA",
)
requires_weights = pytest.mark.skipif(
not glob.glob(os.path.join(_flux_transformer_path(), "*.safetensors")),
reason=(
f"No safetensors under {_flux_transformer_path()} — download FLUX.1-dev "
"transformer or set FLUX_TRANSFORMER_PATH"
),
)
@requires_cuda
@requires_weights
@pytest.mark.usefixtures("distributed_setup", "torch_sdpa_attention_backend")
def test_flux_transformer_parity_vs_diffusers() -> None:
"""Single forward: FastVideo DiT vs Diffusers ``FluxTransformer2DModel``."""
device = torch.device("cuda:0")
precision = torch.bfloat16
transformer_path = _flux_transformer_path()
args = FastVideoArgs(
model_path=transformer_path,
dit_cpu_offload=False,
dit_layerwise_offload=False,
pipeline_config=PipelineConfig(dit_config=FluxDiTConfig(), dit_precision="bf16"),
)
args.device = device
generator = torch.Generator(device=device).manual_seed(0)
torch.manual_seed(0)
batch_size = 1
latent_h, latent_w = 4, 4
img_seq = latent_h * latent_w
text_len = 32
hidden_states = torch.randn(
batch_size,
img_seq,
64,
device=device,
dtype=precision,
generator=generator,
)
encoder_hidden_states = torch.randn(
batch_size,
text_len,
4096,
device=device,
dtype=precision,
generator=generator,
)
pooled_projections = torch.randn(
batch_size,
768,
device=device,
dtype=precision,
generator=generator,
)
# Diffusers pipeline passes scheduler timesteps / 1000 (float, same dtype as latents).
timestep = torch.tensor([512.0], device=device, dtype=precision) / 1000.0
guidance = torch.full((batch_size,), 3.5, device=device, dtype=torch.float32)
txt_ids = torch.zeros(text_len, 3, device=device, dtype=torch.long)
img_ids = _prepare_latent_image_ids(latent_h, latent_w, device, dtype=torch.long)
forward_batch = ForwardBatch(data_type="dummy")
# One ~12B model at a time avoids peak VRAM from holding both checkpoints.
loader = TransformerLoader()
fv_model = loader.load(transformer_path, args).to(device=device, dtype=precision)
fv_model.eval()
with (
torch.no_grad(),
torch.amp.autocast("cuda", dtype=precision),
set_forward_context(
current_timestep=512,
attn_metadata=None,
forward_batch=forward_batch,
),
):
fv_out = fv_model(
hidden_states=hidden_states.clone(),
encoder_hidden_states=encoder_hidden_states.clone(),
pooled_projections=pooled_projections.clone(),
timestep=timestep.clone(),
guidance=guidance.clone(),
txt_ids=txt_ids,
img_ids=img_ids,
return_dict=False,
)[0]
fv_out_cpu = fv_out.detach().float().cpu()
del fv_model
if torch.cuda.is_available():
torch.cuda.empty_cache()
hf_model = (
HFFluxTransformer2DModel.from_pretrained(
transformer_path,
torch_dtype=precision,
)
.to(device)
.eval()
)
with torch.no_grad(), torch.amp.autocast("cuda", dtype=precision):
hf_out = hf_model(
hidden_states=hidden_states.clone(),
encoder_hidden_states=encoder_hidden_states.clone(),
pooled_projections=pooled_projections.clone(),
timestep=timestep.clone(),
guidance=guidance.clone(),
txt_ids=txt_ids,
img_ids=img_ids,
return_dict=False,
)[0]
assert hf_out.shape == fv_out_cpu.shape
hf_cpu = hf_out.float().cpu()
abs_diff = (hf_cpu - fv_out_cpu).abs()
print(f"[FLUX DiT parity] max_diff={abs_diff.max():.4f} mean_diff={abs_diff.mean():.4f} "
f"median_diff={abs_diff.median():.4f} p99_diff="
f"{abs_diff.flatten().kthvalue(int(0.99 * abs_diff.numel())).values:.4f}")
# bfloat16 accumulation over 57 transformer layers produces tail errors up to ~0.5
# on isolated elements (median=0, mean~0.04 on L40S). atol=0.5 catches real bugs
# (wrong weights / missing layers) which produce mean_diff >> 0.1.
assert_close(hf_cpu, fv_out_cpu, atol=0.5, rtol=0.0)
+24 -5
View File
@@ -77,13 +77,32 @@ def _read_video_frames(path: str) -> torch.Tensor:
return torch.stack(frames)
def _read_image_as_single_frame_video(path: str) -> torch.Tensor:
"""Read one image as a single-frame ``(1, C, H, W)`` uint8 tensor."""
from torchvision.io import read_image
img = read_image(path)
return img.unsqueeze(0)
def _read_visual_frames(path: str) -> torch.Tensor:
"""Read a video or a single image as ``(T, C, H, W)`` uint8."""
ext = os.path.splitext(path)[1].lower()
if ext in {".png", ".jpg", ".jpeg", ".webp"}:
return _read_image_as_single_frame_video(path)
return _read_video_frames(path)
def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
"""
Compute SSIM between two videos.
Compute SSIM between two videos or single-frame image files.
Image paths (``.png``, ``.jpg``, ``.jpeg``, ``.webp``) are treated as
one-frame clips so T2I SSIM can share the same MS-SSIM path as video.
Args:
video1_path: Path to the first video.
video2_path: Path to the second video.
video1_path: Path to the first video or image.
video2_path: Path to the second video or image.
use_ms_ssim: Whether to use Multi-Scale Structural Similarity(MS-SSIM) instead of SSIM.
"""
from pytorch_msssim import ms_ssim, ssim
@@ -94,8 +113,8 @@ def compute_video_ssim_torchvision(video1_path, video2_path, use_ms_ssim=True):
if not os.path.exists(video2_path):
raise FileNotFoundError(f"Video2 not found: {video2_path}")
frames1 = _read_video_frames(video1_path)
frames2 = _read_video_frames(video2_path)
frames1 = _read_visual_frames(video1_path)
frames2 = _read_visual_frames(video2_path)
# Ensure same number of frames
min_frames = min(frames1.shape[0], frames2.shape[0])
@@ -1,5 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.train.methods.distribution_matching.anyflow import AnyFlowMethod
from fastvideo.train.methods.distribution_matching.anyflow_pretrain import (
AnyFlowPretrainMethod, )
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
from fastvideo.train.methods.distribution_matching.self_forcing import (
SelfForcingMethod, )
@@ -7,6 +10,8 @@ from fastvideo.train.methods.distribution_matching.streaming_long_tuning import
StreamingLongTuningMethod, )
__all__ = [
"AnyFlowMethod",
"AnyFlowPretrainMethod",
"DMD2Method",
"SelfForcingMethod",
"StreamingLongTuningMethod",
@@ -0,0 +1,209 @@
# SPDX-License-Identifier: Apache-2.0
"""AnyFlow on-policy distillation method.
Stage 2 of the AnyFlow two-stage recipe. Continues from a pretrained
flow-map student and refines it via distribution-matching distillation
(DMD2) where the student is rolled out for ``student_sample_steps``
Euler-flow steps from pure noise. One randomly-chosen step in the
rollout is gradient-enabled (and broadcast across ranks so every worker
agrees on which step to gradient-enable); the rest run under
``torch.no_grad``.
Inherits ``DMD2Method`` for the alternating student / critic update
machinery and the existing DMD VSD-with-fake-score loss. Overrides
``_student_rollout`` to drive the multi-step Euler-flow rollout with
``r = t_next`` (mean-velocity sampling — matches the AnyFlow paper's
``WanAnyFlowPipeline.training_rollout`` with ``use_mean_velocity=True``).
Reference: ``pipeline_wan_anyflow.py::training_rollout`` in
NVlabs/AnyFlow at commit ``549236a``.
"""
from __future__ import annotations
from typing import Any, Literal
import torch
import torch.distributed as dist
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
from fastvideo.train.utils.config import (
get_optional_float,
get_optional_int,
)
class AnyFlowMethod(DMD2Method):
"""AnyFlow on-policy distillation (multi-step rollout)."""
def __init__(
self,
*,
cfg: Any,
role_models: dict[str, Any],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
mcfg = self.method_config
student_sample_steps = get_optional_int(mcfg, "student_sample_steps", where="method.student_sample_steps")
if student_sample_steps is None:
student_sample_steps = 4
if int(student_sample_steps) <= 0:
raise ValueError("method.student_sample_steps must be positive, "
f"got {student_sample_steps}")
self._student_sample_steps = int(student_sample_steps)
use_mean_velocity_raw = mcfg.get("use_mean_velocity", True)
if not isinstance(use_mean_velocity_raw, bool):
raise ValueError("method.use_mean_velocity must be a bool, "
f"got {type(use_mean_velocity_raw).__name__}")
self._use_mean_velocity = bool(use_mean_velocity_raw)
# Optional pinned rollout schedule (descending, absolute t-units).
# Falls back to dmd_denoising_steps when absent.
raw_t_list = mcfg.get("t_list_override", None)
if raw_t_list is None:
self._t_list_override: list[float] | None = None
else:
if not isinstance(raw_t_list, list) or not raw_t_list:
raise ValueError("method.t_list_override must be a non-empty list of "
f"floats when set, got {raw_t_list!r}")
t_list = [float(x) for x in raw_t_list]
for i in range(len(t_list) - 1):
if t_list[i] < t_list[i + 1]:
raise ValueError("method.t_list_override must be descending, "
f"got {t_list!r}")
self._t_list_override = t_list
# Scoring conditioning: AnyFlow scores against r=0 for the DMD branch.
score_r_raw = mcfg.get("dmd_score_r_value", 0.0)
try:
self._dmd_score_r = float(score_r_raw)
except (TypeError, ValueError) as exc:
raise ValueError("method.dmd_score_r_value must be numeric, "
f"got {score_r_raw!r}") from exc
# Optional teacher guidance scale for the DMD loss (carry over from
# DMD2Method's behavior; default 1.0).
guidance = get_optional_float(mcfg, "real_score_guidance_scale", where="method.real_score_guidance_scale")
self._real_score_guidance = float(guidance) if guidance is not None else 1.0
# ------------------------------------------------------------------
# Rollout schedule
def _get_rollout_schedule(self, *, device: torch.device) -> torch.Tensor:
"""Build the descending timestep schedule used by the on-policy
rollout. Length is ``num_steps + 1`` so ``num_steps`` Euler steps
consume the full range.
Order of precedence:
1. ``method.t_list_override`` — used verbatim (absolute units).
2. ``method.dmd_denoising_steps`` (inherited from DMD2) appended
with a final 0 boundary if the last entry isn't already 0.
"""
if self._t_list_override is not None:
return torch.tensor(self._t_list_override, device=device, dtype=torch.float32)
steps = self._get_denoising_step_list(device).to(dtype=torch.float32)
if float(steps[-1].item()) != 0.0:
zero = torch.zeros(1, device=device, dtype=torch.float32)
steps = torch.cat([steps, zero], dim=0)
return steps
def _broadcast_grad_step_index(
self,
num_steps: int,
*,
device: torch.device,
) -> int:
"""Pick the rollout step that gets gradient enabled. In distributed
runs the choice is broadcast from rank 0 so every worker agrees."""
if num_steps <= 0:
raise ValueError("num_steps must be positive")
if dist.is_initialized() and dist.get_rank() != 0:
idx_tensor = torch.empty((1, ), dtype=torch.long, device=device)
else:
idx_tensor = torch.randint(0,
num_steps, (1, ),
device=device,
dtype=torch.long,
generator=self.cuda_generator)
if dist.is_initialized():
dist.broadcast(idx_tensor, src=0)
return int(idx_tensor.item())
# ------------------------------------------------------------------
# Rollout
def _student_rollout(
self,
batch: Any,
*,
with_grad: bool,
) -> torch.Tensor:
"""Multi-step Euler-flow rollout from pure noise.
Returns the predicted clean latent ``x_0`` after the chosen
gradient step (or the final ``x`` after the last step if
``with_grad`` is False — used by the critic path).
"""
latents = batch.latents
if latents is None or latents.ndim != 5:
raise RuntimeError("AnyFlow on-policy rollout requires TrainingBatch.latents "
"of shape [B, T, C, H, W] for shape templating")
device = latents.device
dtype = latents.dtype
schedule = self._get_rollout_schedule(device=device)
num_entries = int(schedule.numel())
num_steps = num_entries - 1
if num_steps <= 0:
raise RuntimeError("rollout schedule must have at least two entries "
f"(got {num_entries})")
if num_steps > self._student_sample_steps:
# Trim to the configured cap, keeping the last (=0) boundary.
schedule = torch.cat([schedule[:self._student_sample_steps], schedule[-1:]], dim=0)
num_steps = self._student_sample_steps
grad_step = self._broadcast_grad_step_index(num_steps, device=device) if with_grad else -1
attn_kind: Literal["dense", "vsa"] = "vsa"
n_train = float(self.student.num_train_timesteps)
x = torch.randn(latents.shape, device=device, dtype=dtype, generator=self.cuda_generator)
last_pred_x0: torch.Tensor | None = None
batch_size = int(latents.shape[0])
for i in range(num_steps):
t_cur = schedule[i].expand(batch_size)
t_next = schedule[i + 1].expand(batch_size)
r = t_next if self._use_mean_velocity else t_cur
enable_grad = bool(with_grad) and (i == grad_step)
with torch.set_grad_enabled(enable_grad):
v = self.student.predict_velocity_with_r(
x,
t_cur,
r,
batch,
conditional=True,
cfg_uncond=self._cfg_uncond,
attn_kind=attn_kind,
)
# Euler step in absolute units: x ← x - ((t_cur - t_next) / N) * v.
view = [-1] + [1] * (x.ndim - 1)
dt = ((t_cur - t_next) / n_train).view(*view)
x = x - dt * v
if enable_grad:
# We treat the rollout output (post-step) as a predicted
# clean latent — AnyFlow's last Euler step lands at t=0.
last_pred_x0 = x
if last_pred_x0 is None:
# No gradient step taken (with_grad=False path).
last_pred_x0 = x
if hasattr(batch, "dmd_latent_vis_dict"):
batch.dmd_latent_vis_dict["generator_timestep"] = (schedule[-1].detach().clone())
return last_pred_x0
@@ -0,0 +1,404 @@
# SPDX-License-Identifier: Apache-2.0
"""AnyFlow pretrain (flow-map central-difference) training method.
Stage 1 of the AnyFlow two-stage recipe. Trains a single student network
``u_θ(x_t, t, r)`` to predict the average velocity from time ``t`` back
to time ``r`` via the central-difference target
target = (eps - x_0) - ((t - r) / N) * dF/dt
where ``N = num_train_timesteps`` and ``dF/dt`` is estimated from the
student's own forward at ``(t ± δ, r)`` (with one-sided fallback near the
schedule endpoints).
Per-batch ``(t, r)`` sampling follows the AnyFlow paper:
- ``diffusion_ratio`` fraction: ``r = t`` (recovers plain flow matching).
- ``consistency_ratio`` fraction: ``r = 0`` (consistency to clean data).
- Remaining fraction: ``(t, r) = (max, min)`` of two independent uniform
draws (full reconstruction range).
Reference: ``trainer_wan_anyflow_pretrain.py`` in NVlabs/AnyFlow at
commit ``549236a``.
"""
from __future__ import annotations
from typing import Any
from collections.abc import Sequence
import torch
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.config import (
get_optional_float,
get_optional_int,
)
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
def _sample_pair_timesteps(
*,
batch_size: int,
diffusion_ratio: float,
consistency_ratio: float,
device: torch.device,
generator: torch.Generator | None,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Sample ``(t, r)`` per the AnyFlow paper.
Two uniform draws ``u1, u2 ∈ [0, 1]`` per sample, then
``t = max(u1, u2)`` and ``r = min(u1, u2)``. After the base sample,
the first ``diffusion_ratio * B`` entries get ``r = t`` (diffusion
branch, plain flow matching), and the next ``consistency_ratio * B``
entries get ``r = 0`` (consistency branch).
Returns
-------
t, r, is_diffusion, is_consistency
Each tensor has shape ``(batch_size,)``. ``t`` and ``r`` are in
``[0, 1]`` (i.e. *not* yet shifted and *not* yet in absolute
train-timestep units). ``is_diffusion`` and ``is_consistency``
are bool masks that partition a subset of the batch — entries
outside both masks are the "free" reconstruction fraction.
"""
if batch_size <= 0:
raise ValueError(f"batch_size must be positive, got {batch_size}")
if diffusion_ratio < 0.0 or consistency_ratio < 0.0:
raise ValueError("diffusion_ratio and consistency_ratio must be non-negative")
if diffusion_ratio + consistency_ratio > 1.0:
raise ValueError("diffusion_ratio + consistency_ratio must be <= 1, "
f"got {diffusion_ratio} + {consistency_ratio}")
u1 = torch.rand(batch_size, device=device, generator=generator)
u2 = torch.rand(batch_size, device=device, generator=generator)
t = torch.maximum(u1, u2)
r = torch.minimum(u1, u2)
n_diff = int(diffusion_ratio * batch_size)
n_cons = int(consistency_ratio * batch_size)
is_diffusion = torch.zeros(batch_size, dtype=torch.bool, device=device)
is_consistency = torch.zeros(batch_size, dtype=torch.bool, device=device)
is_diffusion[:n_diff] = True
is_consistency[n_diff:n_diff + n_cons] = True
# Override per the AnyFlow paper:
# - diffusion entries: r = t (plain flow matching)
# - consistency entries: r = 0 (consistency to clean data)
r = torch.where(is_diffusion, t, r)
r = torch.where(is_consistency, torch.zeros_like(r), r)
return t, r, is_diffusion, is_consistency
@torch.no_grad()
def _central_difference_dF_dt(
*,
student: Any,
batch: Any,
noisy: torch.Tensor,
latents: torch.Tensor,
noise: torch.Tensor,
t: torch.Tensor,
r: torch.Tensor,
delta: float,
num_train_timesteps: float,
attn_kind: str = "dense",
guidance_scale: float = 1.0,
) -> torch.Tensor:
"""Estimate ``dF/dt`` for the AnyFlow central-difference target.
Computes a symmetric finite difference of the velocity prediction in
*absolute train-timestep units*:
dF/dt ≈ [u_θ(x_{t+δ}, t+δ, r) - u_θ(x_{t-δ}, t-δ, r)] / (2 * δ * guidance)
The sample is also moved along the flow trajectory by the same
finite step (``v_pred * (δ / N)``) to match AnyFlow's reference
formulation in ``trainer_wan_anyflow_pretrain.py::compute_central_difference``.
Wrapped in ``torch.no_grad`` so the two extra forwards never enter
the backward graph.
"""
if delta <= 0.0:
raise ValueError(f"delta must be positive, got {delta}")
if guidance_scale <= 0.0:
raise ValueError(f"guidance_scale must be positive, got {guidance_scale}")
v_pred = noise - latents # ground-truth flow velocity
delta_x = delta / float(num_train_timesteps)
t_plus = t + delta
noisy_plus = noisy + v_pred * delta_x
f_plus = student.predict_velocity_with_r(noisy_plus, t_plus, r, batch, conditional=True, attn_kind=attn_kind)
t_minus = t - delta
noisy_minus = noisy - v_pred * delta_x
f_minus = student.predict_velocity_with_r(noisy_minus, t_minus, r, batch, conditional=True, attn_kind=attn_kind)
return (f_plus - f_minus) / (2.0 * delta * guidance_scale)
class AnyFlowPretrainMethod(TrainingMethod):
"""AnyFlow flow-map pretrain method.
Single-student training; no teacher or critic. The student must
implement ``predict_velocity_with_r(noisy, t, r, batch, ...)`` —
typically a ``WanModel`` with ``r_embedder=True`` in its arch config.
"""
def __init__(
self,
*,
cfg: Any,
role_models: dict[str, ModelBase],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
if "student" not in role_models:
raise ValueError("AnyFlowPretrainMethod requires role 'student'")
if not self.student._trainable:
raise ValueError("AnyFlowPretrainMethod requires student to be trainable")
mcfg = self.method_config
self._diffusion_ratio = float(
get_optional_float(mcfg, "diffusion_ratio", where="method.diffusion_ratio") or 0.5)
self._consistency_ratio = float(
get_optional_float(mcfg, "consistency_ratio", where="method.consistency_ratio") or 0.25)
if self._diffusion_ratio + self._consistency_ratio > 1.0:
raise ValueError("method.diffusion_ratio + method.consistency_ratio must "
f"be <= 1, got {self._diffusion_ratio} + "
f"{self._consistency_ratio}")
# δ: finite-difference step in absolute train-timestep units.
epsilon = get_optional_int(mcfg, "epsilon", where="method.epsilon")
self._fd_epsilon = float(epsilon) if epsilon is not None else 5.0
# Loss weighting scheme (uniform / gaussian / beta08).
raw_weight_type = mcfg.get("weight_type", "beta08")
if not isinstance(raw_weight_type, str):
raise ValueError("method.weight_type must be a string, got "
f"{type(raw_weight_type).__name__}")
weight_type = raw_weight_type.strip().lower()
if weight_type not in {"uniform", "gaussian", "beta08"}:
raise ValueError("method.weight_type must be one of "
"{uniform, gaussian, beta08}, "
f"got {raw_weight_type!r}")
self._weight_type = weight_type
# Guidance fused into the training target (default 1.0 = unused).
fg = get_optional_float(mcfg, "fuse_guidance_scale", where="method.fuse_guidance_scale")
self._fuse_guidance_scale = float(fg) if fg is not None else 1.0
if self._fuse_guidance_scale <= 0.0:
raise ValueError("method.fuse_guidance_scale must be positive, "
f"got {self._fuse_guidance_scale}")
# Flow-map scheduler — uses pipeline_config.flow_shift if present
# and falls back to method.shift (and finally 1.0).
shift = float(getattr(self.training_config.pipeline_config, "flow_shift", 0.0) or 0.0)
if shift <= 0.0:
shift_override = get_optional_float(mcfg, "shift", where="method.shift")
shift = float(shift_override) if shift_override is not None else 1.0
self._shift = shift
# Lazy-imported to avoid circular imports on package load.
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
FlowMapEulerDiscreteScheduler, )
self._flow_map_scheduler = FlowMapEulerDiscreteScheduler(
num_train_timesteps=int(self.student.num_train_timesteps),
shift=self._shift,
)
self.student.init_preprocessors(self.training_config)
self._init_optimizer_and_scheduler()
@property
def _optimizer_dict(self) -> dict[str, torch.optim.Optimizer]:
return {"student": self._student_optimizer}
@property
def _lr_scheduler_dict(self) -> dict[str, Any]:
return {"student": self._student_lr_scheduler}
def get_optimizers(
self,
iteration: int,
) -> Sequence[torch.optim.Optimizer]:
del iteration
return [self._student_optimizer]
def get_lr_schedulers(self, iteration: int) -> Sequence[Any]:
del iteration
return [self._student_lr_scheduler]
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
del iteration # AnyFlow pretrain has no iteration-dependent dispatch.
training_batch = self.student.prepare_batch(
batch,
generator=self.cuda_generator,
latents_source="data",
)
latents = training_batch.latents # [B, T, C, H, W] (post-permute in prepare_batch).
if latents is None or latents.ndim != 5:
raise RuntimeError("AnyFlow pretrain expects TrainingBatch.latents of shape "
"[B, T, C, H, W] after prepare_batch; got "
f"{None if latents is None else tuple(latents.shape)}")
device = latents.device
dtype = latents.dtype
batch_size = int(latents.shape[0])
# AnyFlow (t, r) sampling — overrides the timestep drawn by
# WanModel._sample_timesteps inside prepare_batch.
t_norm, r_norm, is_diffusion, is_consistency = _sample_pair_timesteps(
batch_size=batch_size,
diffusion_ratio=self._diffusion_ratio,
consistency_ratio=self._consistency_ratio,
device=device,
generator=self.cuda_generator,
)
sched = self._flow_map_scheduler
n_train = float(self.student.num_train_timesteps)
t = (sched.apply_shift(t_norm) * n_train).to(device=device, dtype=dtype)
r = (sched.apply_shift(r_norm) * n_train).to(device=device, dtype=dtype)
# Fresh noise drawn from the method's RNG; ignore the noise that
# prepare_batch attached (it pairs with the discarded timestep).
noise = torch.randn(
latents.shape,
device=device,
dtype=dtype,
generator=self.cuda_generator,
)
noisy = sched.add_noise(latents, noise, t)
# Keep training_batch coherent with the new (t, noisy): downstream
# forward_context uses these fields.
training_batch.timesteps = t
training_batch.noise = noise
training_batch.noisy_model_input = noisy.permute(0, 2, 1, 3, 4)
# Student velocity prediction at (t, r).
noise_pred = self.student.predict_velocity_with_r(
noisy,
t,
r,
training_batch,
conditional=True,
attn_kind="dense",
)
# Optional guidance distillation — fuse CFG into the training target so
# the resulting checkpoint can be sampled at guidance_scale=1.0.
if self._fuse_guidance_scale != 1.0:
with torch.no_grad():
noise_pred_uncond = self.student.predict_velocity_with_r(
noisy,
t,
r,
training_batch,
conditional=False,
attn_kind="dense",
)
g = float(self._fuse_guidance_scale)
noise_pred = (noise_pred - (1.0 - g) * noise_pred_uncond) / g
dF_dt = _central_difference_dF_dt(
student=self.student,
batch=training_batch,
noisy=noisy,
latents=latents,
noise=noise,
t=t,
r=r,
delta=self._fd_epsilon,
num_train_timesteps=n_train,
attn_kind="dense",
guidance_scale=self._fuse_guidance_scale,
)
# AnyFlow target: target = (eps - x_0) - (t - r) * dF/dt
# dF/dt is in (velocity per absolute t-unit); (t - r) is in absolute units;
# the product cancels back to velocity units, matching noise_pred.
view = [batch_size] + [1] * (latents.ndim - 1)
target = (noise - latents) - (t - r).view(*view) * dF_dt
# Per-sample squared error, then per-timestep weight, then scale-balance
# so the non-diffusion branches stay on the same magnitude as the
# diffusion branch (matches AnyFlow's stop-grad rescaling).
per_sample = torch.mean(
((noise_pred.float() - target.float())**2).reshape(batch_size, -1),
dim=-1,
)
weight = sched.get_train_weight(t, weight_type=self._weight_type)
per_sample = per_sample * weight
with torch.no_grad():
diff_mask = is_diffusion
diff_mean = per_sample[diff_mask].mean() if diff_mask.any() else per_sample.mean()
non_diff_mask = ~diff_mask
if non_diff_mask.any():
scale = diff_mean / (per_sample[non_diff_mask] + 1e-5)
else:
scale = torch.tensor(1.0, device=device, dtype=per_sample.dtype)
non_diff_idx = torch.nonzero(non_diff_mask, as_tuple=False).flatten()
if non_diff_idx.numel() > 0:
per_sample = per_sample.clone()
per_sample[non_diff_idx] = per_sample[non_diff_idx] * scale
total_loss = per_sample.mean()
loss_map = {"total_loss": total_loss}
metrics: dict[str, LogScalar] = {
"diffusion_fraction": float(is_diffusion.float().mean()),
"consistency_fraction": float(is_consistency.float().mean()),
"scale_weight_mean": float(scale.mean()) if isinstance(scale, torch.Tensor) else float(scale),
}
outputs = {
"student_ctx": (
training_batch.timesteps,
training_batch.attn_metadata,
),
}
return loss_map, outputs, metrics
def backward(
self,
loss_map: dict[str, torch.Tensor],
outputs: dict[str, Any],
*,
grad_accum_rounds: int = 1,
) -> None:
"""Route the loss backward through the student's forward_context
so attn metadata stays attached during gradient computation."""
student_ctx = outputs.get("student_ctx")
if student_ctx is None:
super().backward(loss_map, outputs, grad_accum_rounds=grad_accum_rounds)
return
self.student.backward(
loss_map["total_loss"],
student_ctx,
grad_accum_rounds=grad_accum_rounds,
)
def _init_optimizer_and_scheduler(self) -> None:
tc = self.training_config
params = [p for p in self.student.transformer.parameters() if p.requires_grad]
(
self._student_optimizer,
self._student_lr_scheduler,
) = build_optimizer_and_scheduler(
params=params,
optimizer_config=tc.optimizer,
loop_config=tc.loop,
learning_rate=float(tc.optimizer.learning_rate),
betas=tc.optimizer.betas,
scheduler_name=str(tc.optimizer.lr_scheduler),
)
+47
View File
@@ -358,6 +358,53 @@ class WanModel(ModelBase):
pred_noise = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
return pred_noise
def predict_velocity_with_r(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
r_timestep: torch.Tensor,
batch: TrainingBatch,
*,
conditional: bool,
cfg_uncond: dict[str, Any] | None = None,
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
"""AnyFlow forward: predict average velocity from ``t`` back to ``r``.
Same plumbing as :meth:`predict_noise` but injects ``r_timestep``
into the transformer kwargs. The transformer must have been
constructed with an arch config that sets ``r_embedder=True`` for
the dual-timestep branch to be active — otherwise ``r_timestep``
is silently ignored by the embedder and the forward reduces to
the single-timestep path.
"""
device_type = self.device.type
dtype = noisy_latents.dtype
if conditional:
text_dict = batch.conditional_dict
if text_dict is None:
raise RuntimeError("Missing conditional_dict in "
"TrainingBatch")
else:
text_dict = self._get_uncond_text_dict(batch, cfg_uncond=cfg_uncond)
if attn_kind == "dense":
attn_metadata = batch.attn_metadata
elif attn_kind == "vsa":
attn_metadata = batch.attn_metadata_vsa
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
):
input_kwargs = (self._build_distill_input_kwargs(noisy_latents, timestep, text_dict))
input_kwargs["r_timestep"] = r_timestep
transformer = self._get_transformer(timestep)
pred_velocity = transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
return pred_velocity
def backward(
self,
loss: torch.Tensor,
+1
View File
@@ -182,6 +182,7 @@ nav:
- Distillation:
- Data Preprocessing: distillation/data_preprocess.md
- DMD: distillation/dmd.md
- AnyFlow: distillation/anyflow.md
- Attention:
- Overview: attention/index.md
- Video Sparse Attention: attention/vsa/index.md
+4 -2
View File
@@ -21,7 +21,9 @@ dependencies = [
"requests>=2.32.2",
# Machine Learning & Transformers
"transformers>=4.57.3",
# GLM-Image's AR encoder (GlmImageForConditionalGeneration) first ships in
# transformers 5.0.0; floor bumped from >=4.57.3 to >=5.0.0 (stable, not rc).
"transformers>=5.0.0",
# <0.23: tokenizers 0.23 renamed RobertaProcessing's binding args, so
# transformers' CLIP-style tokenizer loading dies with
# "RobertaProcessing.__new__() got an unexpected keyword argument 'cls'".
@@ -227,7 +229,7 @@ skip = "./data,./wandb,apps/fastvideo_studio/package-lock.json,apps/performance_
# "tread" matches daVinci-MagiHuman's acronym "TReAD" (Token Routing and
# Early Drop). codespell lowercases ignore-words entries, so the single
# lowercase form silences all case variants.
ignore-words-list = "tread,passt"
ignore-words-list = "tread,passt,dout"
[tool.ruff]
# Allow lines to be as long as 120.
+3 -1
View File
@@ -21,7 +21,9 @@ dependencies = [
"requests>=2.32.2",
# Machine Learning & Transformers
"transformers>=4.57.3",
# GLM-Image's AR encoder (GlmImageForConditionalGeneration) first ships in
# transformers 5.0.0; floor bumped from >=4.57.3 to >=5.0.0 (stable, not rc).
"transformers>=5.0.0",
"tokenizers>=0.20.1",
"sentencepiece>=0.2.0",
"timm>=1.0.11",
@@ -0,0 +1,87 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 checkpoint strict-load verifier (no weight conversion required).
The published ``nvidia/Cosmos3-Nano`` checkpoint is diffusers-format and its
transformer weight keys map 1:1 (identity) onto FastVideo's native
``Cosmos3VFMTransformer`` parameters -- ``needs_conversion=no``. There is no
remap to apply; the checkpoint loads directly.
This utility verifies strict-load completeness (every checkpoint key has a
matching DiT parameter of the right shape, and every DiT parameter is provided
by the checkpoint) without allocating the full ~30 GB model, by reading
safetensors headers and instantiating the DiT on the ``meta`` device.
Usage:
python scripts/checkpoint_conversion/cosmos3_convert.py \
--transformer official_weights/cosmos3/transformer
"""
from __future__ import annotations
import argparse
import glob
import os
import re
import torch
from safetensors import safe_open
from fastvideo.configs.models.dits.cosmos3 import Cosmos3VideoConfig
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
def checkpoint_key_shapes(transformer_dir: str) -> dict[str, tuple[int, ...]]:
"""Read ``{key: shape}`` from a sharded safetensors transformer dir."""
shards = sorted(glob.glob(os.path.join(transformer_dir, "*.safetensors")))
if not shards:
raise FileNotFoundError(f"no .safetensors found in {transformer_dir}")
shapes: dict[str, tuple[int, ...]] = {}
for shard in shards:
with safe_open(shard, framework="pt") as handle:
for key in handle.keys():
shapes[key] = tuple(handle.get_slice(key).get_shape())
return shapes
def verify_strict_load(transformer_dir: str) -> None:
"""Raise SystemExit if the checkpoint does not strict-load into the DiT."""
ckpt = checkpoint_key_shapes(transformer_dir)
cfg = Cosmos3VideoConfig()
with torch.device("meta"):
dit = Cosmos3VFMTransformer(cfg, hf_config={})
params = {name: tuple(p.shape) for name, p in dit.named_parameters()}
buffers = {name for name, _ in dit.named_buffers()}
name_map: dict[str, str] = cfg.arch_config.param_names_mapping
def remap(key: str) -> str:
for pattern, replacement in name_map.items():
if re.match(pattern, key):
return re.sub(pattern, replacement, key)
return key
mapped = {remap(key): shape for key, shape in ckpt.items()}
unexpected = sorted(set(mapped) - set(params) - buffers)
missing = sorted(set(params) - set(mapped))
mismatched = [(k, mapped[k], params[k]) for k in (set(mapped) & set(params)) if mapped[k] != params[k]]
if unexpected or missing or mismatched:
raise SystemExit("strict-load FAILED: "
f"unexpected={unexpected[:10]} missing={missing[:10]} "
f"shape_mismatch={mismatched[:10]}")
print(f"strict-load OK: {len(ckpt)} checkpoint keys map 1:1 onto "
f"{len(params)} DiT params (identity; needs_conversion=no)")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--transformer",
default=os.path.join("official_weights", "cosmos3", "transformer"),
help="path to the checkpoint transformer/ directory",
)
args = parser.parse_args()
verify_strict_load(args.transformer)
if __name__ == "__main__":
main()
+282
View File
@@ -0,0 +1,282 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""FastVideo-side AnyFlow 14B T2V demo at NFE=4 and NFE=50.
Loads ``nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers`` into FastVideo's
``WanTransformer3DModel`` (with the ``param_names_mapping`` regex
handling the ``delta_embedder`` rename), runs the
``FlowMapEulerDiscreteScheduler`` for the requested NFE schedule, and
saves the decoded video as MP4. Matches the prompt / shift / guidance
recipe used by the parallel FastGen demo so the videos are directly
comparable.
Memory tactics (single H200, 141 GB HBM):
- Encode prompts with UMT5, free the encoder.
- Build the FastVideo Wan-14B transformer, load AnyFlow safetensor
shards via param_names_mapping translation.
- Sample at both NFEs without re-loading the transformer.
- Free the transformer, then load the Wan VAE with tiling for decode.
Configure local checkout paths via env vars (defaults assume a sibling
layout next to this repo):
ANYFLOW_LOCAL — path to ``nvidia/AnyFlow-Wan2.1-T2V-14B-Diffusers``
(default ``./anyflow-14b``)
ANYFLOW_DEMO_OUT — output directory for the rendered MP4s
(default ``./demo_videos``)
Run via::
PYTHONPATH=$PWD python scripts/demo_anyflow_14b.py
"""
from __future__ import annotations
import gc
import os
import re
import sys
import time
from pathlib import Path
import torch
from safetensors.torch import load_file
ANYFLOW_LOCAL = Path(os.environ.get("ANYFLOW_LOCAL", "./anyflow-14b")).expanduser()
OUT_DIR = Path(os.environ.get("ANYFLOW_DEMO_OUT", "./demo_videos")).expanduser()
OUT_DIR.mkdir(parents=True, exist_ok=True)
SEED = 0
DEVICE = torch.device("cuda")
DTYPE = torch.bfloat16
NUM_FRAMES = 81
HEIGHT, WIDTH = 480, 832
PROMPT = (
"CG game concept digital art, a majestic elephant with a vibrant tusk and sleek fur "
"running swiftly towards a herd of its kind. The elephant has a calm yet determined "
"expression, with its ears flapping slightly as it moves at high speed. The herd consists "
"of several other elephants of various ages and sizes, all moving in unison. The landscape "
"is vast savanna with rolling hills, tall grasses, and scattered acacia trees. The sun "
"sets behind the horizon, casting a warm golden glow over the scene. Low-angle view, focus "
"on the elephant as it accelerates towards the herd."
)
NEG_PROMPT = "blurry, low quality, distorted"
# The published nvidia/AnyFlow-* checkpoints are on-policy distilled with
# fuse_guidance_scale=3.0 baked into the weights — inference uses 1.0
# (single conditional forward, no CFG; matches AnyFlow's official demo.py).
GUIDANCE = 1.0
def banner(msg: str) -> None:
print("\n" + "=" * 80)
print(msg)
print("=" * 80)
def free_gpu() -> None:
gc.collect()
torch.cuda.empty_cache()
torch.cuda.synchronize()
used = torch.cuda.memory_allocated() / 1e9
print(f" [mem] allocated {used:.1f} GB after free")
def init_single_rank() -> None:
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29571")
os.environ.setdefault("WORLD_SIZE", "1")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("LOCAL_RANK", "0")
from fastvideo.distributed import (
init_distributed_environment,
initialize_model_parallel,
)
init_distributed_environment(world_size=1, rank=0, local_rank=0, backend="nccl")
initialize_model_parallel(
tensor_model_parallel_size=1,
sequence_model_parallel_size=1,
data_parallel_size=1,
)
def translate_keys(raw: dict, *, mapping: dict[str, str]) -> dict:
out: dict = {}
for k, v in raw.items():
new_k = k
for pat, repl in mapping.items():
new_k = re.sub(pat, repl, new_k)
out[new_k] = v
return out
def encode_prompts():
banner("(1) Encode prompts via UMT5")
from transformers import AutoTokenizer, UMT5EncoderModel
tok = AutoTokenizer.from_pretrained(str(ANYFLOW_LOCAL), subfolder="tokenizer", use_fast=False)
enc = UMT5EncoderModel.from_pretrained(
str(ANYFLOW_LOCAL), subfolder="text_encoder", torch_dtype=DTYPE,
).to(DEVICE).eval()
@torch.no_grad()
def encode_one(prompts):
out = tok(
prompts, padding="max_length", max_length=512, truncation=True,
return_attention_mask=True, return_tensors="pt")
ids = out.input_ids.to(DEVICE)
mask = out.attention_mask.to(DEVICE)
seq_lens = mask.gt(0).sum(dim=1).long()
embeds = enc(ids, mask).last_hidden_state
padded = []
for i, l in enumerate(seq_lens):
e = embeds[i, :l]
pad = torch.zeros(512 - e.size(0), e.size(1), device=DEVICE, dtype=embeds.dtype)
padded.append(torch.cat([e, pad], dim=0))
return torch.stack(padded, dim=0)
text_e = encode_one([PROMPT])
neg_e = encode_one([NEG_PROMPT])
print(f" prompts encoded: text={tuple(text_e.shape)} neg={tuple(neg_e.shape)}")
del enc, tok
free_gpu()
return text_e, neg_e
def load_transformer():
banner("(2) Build FastVideo Wan-14B + load AnyFlow weights")
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
cfg = WanVideoConfig()
arch = cfg.arch_config
# Wan2.1-T2V-14B arch (per AnyFlow checkpoint config.json).
arch.num_attention_heads = 40
arch.attention_head_dim = 128
arch.num_layers = 40
arch.ffn_dim = 13824
arch.r_embedder = True
arch.r_embedder_fusion = "gated"
arch.r_embedder_gate_value = 0.25
arch.r_embedder_deltatime_type = "r"
arch.__post_init__()
t0 = time.time()
model = WanTransformer3DModel(config=cfg, hf_config={}).to(DEVICE, dtype=DTYPE).eval()
print(f" transformer built in {time.time() - t0:.1f}s; "
f"params: {sum(p.numel() for p in model.parameters())/1e9:.2f}B")
ckpt_dir = ANYFLOW_LOCAL / "transformer"
sd: dict = {}
for shard in sorted(ckpt_dir.glob("diffusion_pytorch_model-*.safetensors")):
sd.update(load_file(str(shard), device="cpu"))
print(f" AnyFlow state dict: {len(sd)} tensors loaded")
sd = translate_keys(sd, mapping=arch.param_names_mapping)
info = model.load_state_dict(sd, strict=False)
print(f" load: missing={len(info.missing_keys)} unexpected={len(info.unexpected_keys)}")
del sd
free_gpu()
return model
@torch.no_grad()
def sample(model, text_e, neg_e, nfe: int) -> torch.Tensor:
banner(f"(3) Sample 14B NFE={nfe}")
from fastvideo.forward_context import set_forward_context
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
FlowMapEulerDiscreteScheduler, )
scheduler = FlowMapEulerDiscreteScheduler(num_train_timesteps=1000, shift=5.0)
scheduler.set_timesteps(num_inference_steps=nfe, device=DEVICE)
timesteps = scheduler.timesteps.to(DEVICE, dtype=DTYPE)
B, C = 1, 16
F = (NUM_FRAMES - 1) // 4 + 1 # temporal VAE compression = 4 (81 → 21)
H_l, W_l = HEIGHT // 8, WIDTH // 8
g = torch.Generator(device=DEVICE).manual_seed(SEED)
x = torch.randn(B, C, F, H_l, W_l, device=DEVICE, dtype=DTYPE, generator=g)
t0 = time.time()
for i, (t_cur, t_next) in enumerate(zip(timesteps[:-1], timesteps[1:])):
t_in = t_cur.expand(B).to(DTYPE)
r_in = t_next.expand(B).to(DTYPE)
with set_forward_context(current_timestep=t_in, attn_metadata=None):
flow_cond = model(
hidden_states=x, encoder_hidden_states=text_e,
timestep=t_in, r_timestep=r_in)
if GUIDANCE != 1.0:
flow_uncond = model(
hidden_states=x, encoder_hidden_states=neg_e,
timestep=t_in, r_timestep=r_in)
flow = flow_uncond + GUIDANCE * (flow_cond - flow_uncond)
else:
flow = flow_cond
x = scheduler.step(
flow, sample=x,
timestep=t_cur.repeat(B), r_timestep=t_next.repeat(B))
print(f" NFE={nfe} sample time: {time.time() - t0:.1f}s "
f"({(time.time() - t0) / nfe:.1f}s/step)")
xf = x.float()
print(f" latents mean={xf.mean().item():+.3f} std={xf.std().item():.3f} "
f"range=[{xf.min().item():+.2f}, {xf.max().item():+.2f}] "
f"finite={torch.isfinite(xf).all().item()}")
return x.detach()
@torch.no_grad()
def decode_to_mp4(latents: torch.Tensor, out_path: Path) -> tuple[Path, tuple]:
from diffusers.models.autoencoders.autoencoder_kl_wan import AutoencoderKLWan
import imageio.v3 as iio
vae = AutoencoderKLWan.from_pretrained(
str(ANYFLOW_LOCAL), subfolder="vae", torch_dtype=DTYPE,
).to(DEVICE).eval()
try:
vae.enable_tiling()
print(f" VAE tiling enabled")
except Exception:
pass
mean = torch.tensor(vae.config.latents_mean, device=DEVICE, dtype=DTYPE).view(1, -1, 1, 1, 1)
std = torch.tensor(vae.config.latents_std, device=DEVICE, dtype=DTYPE).view(1, -1, 1, 1, 1)
latents_unscaled = latents * std + mean
t0 = time.time()
frames = vae.decode(latents_unscaled, return_dict=False)[0]
print(f" VAE decode time: {time.time() - t0:.1f}s")
frames = (frames.clamp(-1, 1) + 1) / 2
frames = frames[0].permute(1, 2, 3, 0).float().cpu().numpy()
frames = (frames * 255).astype("uint8")
iio.imwrite(str(out_path), frames, fps=16, codec="libx264", quality=8)
del vae
free_gpu()
return out_path, frames.shape
def main() -> None:
torch.manual_seed(SEED)
torch.cuda.manual_seed_all(SEED)
init_single_rank()
text_e, neg_e = encode_prompts()
model = load_transformer()
latents_list = []
for nfe in [4, 50]:
lat = sample(model, text_e, neg_e, nfe=nfe)
latents_list.append((nfe, lat))
del model, text_e, neg_e
free_gpu()
for nfe, lat in latents_list:
banner(f"(4) Decode 14B NFE={nfe}")
out_path = OUT_DIR / f"fastvideo_anyflow_14b_nfe{nfe}_seed{SEED}.mp4"
path, shape = decode_to_mp4(lat, out_path)
print(f" decoded {shape}")
print(f" saved: {path} ({path.stat().st_size / 1e6:.2f} MB)")
banner("DONE")
if __name__ == "__main__":
main()
+404
View File
@@ -0,0 +1,404 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""FastVideo↔AnyFlow numerical parity verification.
Two checks, both on a single H200:
(A) Forward parity — load nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers into
FastVideo's WanTransformer3DModel (with r_embedder enabled +
param_names_mapping handling the delta_embedder rename), forward
it on identical inputs against AnyFlow's reference loader, and
compare. Expectation: rel-mean diff < 10% in bf16 (bf16 kernel noise).
(B) Any-step end-to-end sampling — run the new
FlowMapEulerDiscreteScheduler through 4 Euler-flow steps on the
same loaded weights, confirm the final latent is finite and
well-scaled.
Configure local checkout paths via env vars (defaults assume a sibling
layout next to this repo):
ANYFLOW_LOCAL — path to ``nvidia/AnyFlow-Wan2.1-T2V-1.3B-Diffusers``
(default ``./anyflow-1.3b``)
ANYFLOW_REF — path to the NVlabs/AnyFlow reference repo, used to
import its loader (default ``./anyflow-ref``)
Run via::
PYTHONPATH=$PWD python scripts/verify_anyflow_fastvideo_parity.py
"""
from __future__ import annotations
import os
import re
import sys
import time
from pathlib import Path
import torch
from safetensors.torch import load_file
ANYFLOW_LOCAL = Path(os.environ.get("ANYFLOW_LOCAL", "./anyflow-1.3b")).expanduser()
ANYFLOW_REF = Path(os.environ.get("ANYFLOW_REF", "./anyflow-ref")).expanduser()
sys.path.insert(0, str(ANYFLOW_REF))
SEED = 1234
DEVICE = torch.device("cuda")
DTYPE = torch.bfloat16
def banner(msg: str) -> None:
print("\n" + "=" * 80)
print(msg)
print("=" * 80)
# ---------------------------------------------------------------------------
# Distributed bootstrap (single-rank). Required by WanTransformer3DModel
# which calls get_sp_world_size() and uses ReplicatedLinear.
# ---------------------------------------------------------------------------
def init_single_rank() -> None:
os.environ.setdefault("MASTER_ADDR", "127.0.0.1")
os.environ.setdefault("MASTER_PORT", "29551")
os.environ.setdefault("WORLD_SIZE", "1")
os.environ.setdefault("RANK", "0")
os.environ.setdefault("LOCAL_RANK", "0")
from fastvideo.distributed import (
init_distributed_environment,
initialize_model_parallel,
)
init_distributed_environment(world_size=1, rank=0, local_rank=0,
backend="nccl")
initialize_model_parallel(
tensor_model_parallel_size=1,
sequence_model_parallel_size=1,
data_parallel_size=1,
)
print(" single-rank distributed environment + TP/SP/DP groups initialized")
# ---------------------------------------------------------------------------
# Translate AnyFlow HF safetensor keys onto FastVideo's WanTransformer3DModel
# internal layout, applying the regex from WanVideoArchConfig.param_names_mapping.
# ---------------------------------------------------------------------------
def translate_keys(
raw: dict[str, torch.Tensor],
*,
mapping: dict[str, str],
) -> dict[str, torch.Tensor]:
out: dict[str, torch.Tensor] = {}
for k, v in raw.items():
new_k = k
for pat, repl in mapping.items():
new_k = re.sub(pat, repl, new_k)
out[new_k] = v
return out
# ---------------------------------------------------------------------------
# Build FastVideo WanTransformer3DModel with AnyFlow weights loaded.
# ---------------------------------------------------------------------------
def build_fastvideo_transformer():
banner("(1) Build FastVideo WanTransformer3DModel + load AnyFlow weights")
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.models.dits.wanvideo import WanTransformer3DModel
cfg = WanVideoConfig()
arch = cfg.arch_config
# AnyFlow Wan2.1-T2V-1.3B arch dims.
arch.num_attention_heads = 12
arch.attention_head_dim = 128
arch.num_layers = 30
arch.ffn_dim = 8960
# AnyFlow dual-timestep.
arch.r_embedder = True
arch.r_embedder_fusion = "gated"
arch.r_embedder_gate_value = 0.25
arch.r_embedder_deltatime_type = "r"
arch.__post_init__()
# WanTransformer3DModel takes (config, hf_config). hf_config is
# diffusers-style — we provide a minimal dict; only fields the model
# actually reads matter.
hf_config: dict = {}
t0 = time.time()
model = WanTransformer3DModel(config=cfg, hf_config=hf_config)
model = model.to(DEVICE, dtype=DTYPE).eval()
print(f" Wan transformer built in {time.time() - t0:.1f}s; "
f"params: {sum(p.numel() for p in model.parameters())/1e9:.2f}B")
# Load AnyFlow checkpoint.
af_path = ANYFLOW_LOCAL / "transformer" / "diffusion_pytorch_model.safetensors"
af_raw = load_file(str(af_path), device="cpu")
print(f" AnyFlow checkpoint: {len(af_raw)} tensors")
translated = translate_keys(af_raw, mapping=arch.param_names_mapping)
info = model.load_state_dict(translated, strict=False)
miss, unex = info.missing_keys, info.unexpected_keys
print(f" missing_keys : {len(miss)} (first 5: {miss[:5]})")
print(f" unexpected_keys : {len(unex)} (first 5: {unex[:5]})")
return model, len(miss), len(unex)
# ---------------------------------------------------------------------------
# Build AnyFlow reference net.
# ---------------------------------------------------------------------------
def build_anyflow_reference():
banner("(2) Build AnyFlow reference loader")
from far.models import build_model
af_net = build_model("FAR_Wan_Transformer3DModel").from_pretrained(
str(ANYFLOW_LOCAL),
subfolder="transformer",
chunk_partition=None,
full_chunk_limit=0,
compressed_patch_size=[1, 4, 4],
).to(DEVICE, dtype=DTYPE).eval()
print(f" AnyFlow {type(af_net).__name__} ready")
return af_net
# ---------------------------------------------------------------------------
# Forward parity test.
# ---------------------------------------------------------------------------
def forward_compare(fv_model, af_net):
banner("(3) Forward output comparison")
B, C, F, H, W = 1, 16, 21, 60, 104
SEQ, DIM = 32, 4096
g = torch.Generator(device=DEVICE).manual_seed(SEED)
x = torch.randn(B, C, F, H, W, device=DEVICE, dtype=DTYPE, generator=g)
enc = torch.randn(B, SEQ, DIM, device=DEVICE, dtype=DTYPE, generator=g)
# AnyFlow expects per-frame (t, r) [B, F]; FastVideo T2V expects a
# single scalar per sample [B] (timestep.dim()==2 is reserved for
# Wan2.2 ti2v's per-token timestep schedule). For shared-t T2V the two
# are semantically equivalent.
t_per_frame = torch.full((B, F), 500.0, device=DEVICE, dtype=DTYPE)
r_per_frame = torch.full((B, F), 200.0, device=DEVICE, dtype=DTYPE)
t_per_sample = torch.full((B,), 500.0, device=DEVICE, dtype=DTYPE)
r_per_sample = torch.full((B,), 200.0, device=DEVICE, dtype=DTYPE)
with torch.no_grad():
# AnyFlow native: takes [B, F, C, H, W] + [B, F] (t, r).
x_af = x.permute(0, 2, 1, 3, 4).contiguous()
af_out = af_net(
x_af,
timestep=t_per_frame,
r_timestep=r_per_frame,
encoder_hidden_states=enc,
return_dict=False,
is_causal=False,
)[0]
if af_out.shape[1] != C and af_out.shape[2] == C:
af_out = af_out.permute(0, 2, 1, 3, 4).contiguous()
# FastVideo: [B, C, F, H, W] + [B] (t, r). Must wrap in
# set_forward_context so the attention layer can pick up the
# current timestep / attn_metadata.
from fastvideo.forward_context import set_forward_context
with set_forward_context(
current_timestep=t_per_sample,
attn_metadata=None,
):
fv_out = fv_model(
hidden_states=x,
encoder_hidden_states=enc,
timestep=t_per_sample,
r_timestep=r_per_sample,
)
print(f" AnyFlow out: shape={tuple(af_out.shape)} dtype={af_out.dtype}")
print(f" FastVideo : shape={tuple(fv_out.shape)} dtype={fv_out.dtype}")
if af_out.shape != fv_out.shape:
print(" ❌ shape mismatch")
return False
diff = (af_out.float() - fv_out.float()).abs()
ref = af_out.float().abs().mean().item() + 1e-12
print(f" max abs diff : {diff.max().item():.3e}")
print(f" mean abs diff: {diff.mean().item():.3e}")
print(f" rel mean diff: {diff.mean().item() / ref:.3e}")
ok = diff.mean().item() / ref < 0.10
print(f" >>> {'PASS' if ok else 'FAIL'} (target rel diff < 10%, bf16 noise)")
return ok
# ---------------------------------------------------------------------------
# Any-step sampling smoke via FlowMapEulerDiscreteScheduler.
# ---------------------------------------------------------------------------
def sample_anystep(fv_model):
banner("(4) Any-step 4-step Euler-flow sampling")
from fastvideo.models.schedulers.scheduling_flow_map_euler_discrete import (
FlowMapEulerDiscreteScheduler, )
scheduler = FlowMapEulerDiscreteScheduler(num_train_timesteps=1000, shift=5.0)
scheduler.set_timesteps(num_inference_steps=4, device=DEVICE)
timesteps = scheduler.timesteps.to(dtype=DTYPE)
B, C, F, H, W = 1, 16, 21, 60, 104
SEQ, DIM = 32, 4096
g = torch.Generator(device=DEVICE).manual_seed(SEED)
x = torch.randn(B, C, F, H, W, device=DEVICE, dtype=DTYPE, generator=g)
enc = torch.randn(B, SEQ, DIM, device=DEVICE, dtype=DTYPE, generator=g)
from fastvideo.forward_context import set_forward_context
t0 = time.time()
with torch.no_grad():
for t_cur, t_next in zip(timesteps[:-1], timesteps[1:]):
t_in = t_cur.expand(B).to(DTYPE)
r_in = t_next.expand(B).to(DTYPE)
with set_forward_context(current_timestep=t_in, attn_metadata=None):
v = fv_model(
hidden_states=x,
encoder_hidden_states=enc,
timestep=t_in,
r_timestep=r_in,
)
x = scheduler.step(
v, sample=x,
timestep=t_cur.repeat(B),
r_timestep=t_next.repeat(B),
)
elapsed = time.time() - t0
xf = x.float()
print(f" elapsed: {elapsed:.1f}s")
print(f" final latent: mean={xf.mean().item():+.3f} std={xf.std().item():.3f} "
f"range=[{xf.min().item():+.2f}, {xf.max().item():+.2f}] "
f"finite={torch.isfinite(xf).all().item()}")
ok = torch.isfinite(xf).all().item() and 0.01 < xf.std().item() < 30
print(f" >>> {'PASS' if ok else 'FAIL'}")
return ok
NUM_TRAIN_TIMESTEPS = 1000
EPSILON = 5.0 # AnyFlow paper default
@torch.no_grad()
def training_step_compare(fv_model, af_net) -> bool:
"""Inline replica of AnyFlow's train_bidirection central-difference loss
on both code paths with identical synthetic (real, noise, t, r) inputs.
Compares scalar loss + intermediate flow_pred / target tensors.
"""
banner("(5) Training-step loss comparison (central-difference target)")
from fastvideo.forward_context import set_forward_context
B, C, F, H, W = 1, 16, 21, 60, 104
SEQ, DIM = 32, 4096
g = torch.Generator(device=DEVICE).manual_seed(SEED)
real = torch.randn(B, C, F, H, W, device=DEVICE, dtype=DTYPE, generator=g)
enc = torch.randn(B, SEQ, DIM, device=DEVICE, dtype=DTYPE, generator=g)
noise = torch.randn_like(real)
t_abs_pf = torch.full((B, F), 500.0, device=DEVICE, dtype=DTYPE)
r_abs_pf = torch.full((B, F), 200.0, device=DEVICE, dtype=DTYPE)
t_abs_ps = torch.full((B,), 500.0, device=DEVICE, dtype=DTYPE)
r_abs_ps = torch.full((B,), 200.0, device=DEVICE, dtype=DTYPE)
# AnyFlow inline replica: uses [B, F, C, H, W] layout.
real_btchw = real.permute(0, 2, 1, 3, 4).contiguous()
noise_btchw = noise.permute(0, 2, 1, 3, 4).contiguous()
t_norm_pf = (t_abs_pf / NUM_TRAIN_TIMESTEPS).view(B, F, 1, 1, 1).to(DTYPE)
noisy_btchw = t_norm_pf * noise_btchw + (1 - t_norm_pf) * real_btchw
def u_func_af(x_in, t_in, r_in):
return af_net(
x_in, timestep=t_in, r_timestep=r_in,
encoder_hidden_states=enc, return_dict=False, is_causal=False)[0]
v_pred = noise_btchw - real_btchw
eps = EPSILON
F_plus = u_func_af(noisy_btchw + v_pred * (eps / NUM_TRAIN_TIMESTEPS),
t_abs_pf + eps, r_abs_pf)
F_minus = u_func_af(noisy_btchw - v_pred * (eps / NUM_TRAIN_TIMESTEPS),
t_abs_pf - eps, r_abs_pf)
dF_dt_af = (F_plus - F_minus) / (2 * eps)
target_af = ((noise_btchw - real_btchw)
- (t_abs_pf - r_abs_pf).view(B, F, 1, 1, 1) * dF_dt_af)
flow_af = u_func_af(noisy_btchw, t_abs_pf, r_abs_pf)
loss_af = (flow_af.float() - target_af.float()).pow(2).reshape(B, -1).mean(-1)
# FastVideo inline replica: uses [B, C, F, H, W] layout + [B] t/r.
real_bcfhw = real
noise_bcfhw = noise
t_norm_ps = (t_abs_ps / NUM_TRAIN_TIMESTEPS).view(B, 1, 1, 1, 1).to(DTYPE)
noisy_bcfhw = t_norm_ps * noise_bcfhw + (1 - t_norm_ps) * real_bcfhw
def u_func_fv(x_in, t_in, r_in):
with set_forward_context(current_timestep=t_in, attn_metadata=None):
return fv_model(
hidden_states=x_in,
encoder_hidden_states=enc,
timestep=t_in,
r_timestep=r_in,
)
v_pred_fv = noise_bcfhw - real_bcfhw
F_plus_fv = u_func_fv(
noisy_bcfhw + v_pred_fv * (eps / NUM_TRAIN_TIMESTEPS),
t_abs_ps + eps, r_abs_ps)
F_minus_fv = u_func_fv(
noisy_bcfhw - v_pred_fv * (eps / NUM_TRAIN_TIMESTEPS),
t_abs_ps - eps, r_abs_ps)
dF_dt_fv = (F_plus_fv - F_minus_fv) / (2 * eps)
target_fv = ((noise_bcfhw - real_bcfhw)
- (t_abs_ps - r_abs_ps).view(B, 1, 1, 1, 1) * dF_dt_fv)
flow_fv = u_func_fv(noisy_bcfhw, t_abs_ps, r_abs_ps)
loss_fv = (flow_fv.float() - target_fv.float()).pow(2).reshape(B, -1).mean(-1)
# Compare. Align AnyFlow's [B, F, C, H, W] → [B, C, F, H, W].
flow_af_aligned = flow_af.permute(0, 2, 1, 3, 4)
target_af_aligned = target_af.permute(0, 2, 1, 3, 4)
flow_diff = (flow_fv.float() - flow_af_aligned.float()).abs()
target_diff = (target_fv.float() - target_af_aligned.float()).abs()
af_loss_v = loss_af.mean().item()
fv_loss_v = loss_fv.mean().item()
abs_diff = abs(af_loss_v - fv_loss_v)
rel_diff = abs_diff / abs(af_loss_v + 1e-12)
print(f" AnyFlow loss : {af_loss_v:.6f}")
print(f" FastVideo loss: {fv_loss_v:.6f}")
print(f" abs diff : {abs_diff:.3e}")
print(f" rel diff : {rel_diff:.3e}")
print(f" flow_pred : max abs {flow_diff.max().item():.3e} "
f"mean {flow_diff.mean().item():.3e}")
print(f" target : max abs {target_diff.max().item():.3e} "
f"mean {target_diff.mean().item():.3e}")
ok = rel_diff < 0.20
print(f" >>> {'PASS' if ok else 'FAIL'} (target rel loss diff < 20%)")
return ok
def main() -> None:
torch.manual_seed(SEED)
torch.cuda.manual_seed_all(SEED)
init_single_rank()
fv_model, n_miss, n_unex = build_fastvideo_transformer()
af_net = build_anyflow_reference()
forward_ok = forward_compare(fv_model, af_net)
sample_ok = sample_anystep(fv_model)
train_ok = training_step_compare(fv_model, af_net)
banner(
f"SUMMARY: missing_keys={n_miss} unexpected_keys={n_unex} "
f"forward_parity={forward_ok} sample_smoke={sample_ok} "
f"training_parity={train_ok}"
)
sys.exit(0 if (forward_ok and sample_ok and train_ok) else 1)
if __name__ == "__main__":
main()
@@ -0,0 +1,77 @@
# Cosmos3 Audio (PR2) — Port Plan
Branch: `feat/cosmos3-audio` (stacked on `feat/cosmos3-i2v`, which has T2V/I2V/T2I).
Goal: text-to-video+sound (**t2vs**) — generate synchronized audio alongside video.
## How the framework does audio (studied 2026-06-07)
- **Sound tokenizer = AVAE** (`cosmos_framework/model/vfm/tokenizers/audio/avae.py`
+ `avae_utils/`, ~2268 lines): a 48 kHz **stereo** neural audio codec.
- checkpoint: `official_weights/cosmos3/sound_tokenizer/` (`model_type:
autoencoder_v2`, ~1.9 GB). enc=`spec_convnext` (enc_dim 192, latent_dim 128,
n_fft 64), dec=`oobleck` (dec_dim 320, strides [2,4,5,6,8]), VAE bottleneck,
`snakebeta` activations, hop_size 1920.
- interface: `encode(audio[1,C,N]) -> latent`, `decode(latent) -> audio`,
`get_latent_num_samples(N)`, `sample_rate=48000`, `audio_channels=2`,
`sound_latent_fps=25`.
- **DiT sound pathway** (`cosmos3_vfm_network.py`, 136 sound/audio refs): the MoT
has `sound2llm` / `llm2sound` / `sound_modality_embed` + `pack_sound_latents`
and joint vision+sound denoising (`preds_sound`, sound `condition_mask`, sound
noise init `cond_mask*x0 + (1-cond_mask)*noise`, velocity `pred*(1-cond_mask)`).
- FastVideo's native DiT ALREADY constructs the dormant heads
(`audio_proj_in`/`audio_proj_out`/`audio_modality_embed`, gated on
`arch.sound_gen`) for strict-load — the forward just doesn't use them yet.
- **Inference flow** (`cosmos_framework/inference/sound.py`): t2vs builds a
zero **placeholder audio** sized to the video duration (sets sound latent
length), `inject_sound_into_batch` upgrades the SequencePlan to has_sound,
the omni model denoises vision+sound jointly, then AVAE-decodes the sound
latent and `mux_audio_into_video` (PyAV, AAC) muxes it into the mp4
(`save_sound` writes a WAV).
## Components (each: native port + framework parity test, per methodology)
1. **AVAE codec** — `fastvideo/models/.../cosmos3_avae.py` + config. Port
encoder/decoder/bottleneck/snake. Parity: tiny AVAE, framework weights copied
in, bit-exact `decode` (and `encode`) on CPU/fp32. **(largest piece)**
2. **DiT sound pathway** — activate the dormant heads in `forward`; port
`pack_sound_latents` + sound token scatter/proj/modality-embed/velocity.
Parity: extend the DiT harness with sound tokens.
3. **Sound sequence packing** — extend `sequence_packing.py` with the sound
modality (positions, attn mode, condition mask). Parity vs framework
`pack_input_sequence` with sound.
4. **Pipeline (t2vs)** — placeholder audio -> joint denoise -> split ->
AVAE-decode sound -> mux into mp4 / save wav. Extend `Cosmos3DenoisingStage`
+ a sound-decode/mux stage.
5. **FastVideo AV infra** — audio in `OutputConfig` / a mux stage (check what
exists; `cosmos_framework.inference.sound.mux_audio_into_video` is the ref).
## Open decisions
- **D1 (AVAE approach)** — full native port (methodology-consistent; ~2.3k lines)
vs a documented lazy-wrapper around the framework AVAE (faster; but pulls heavy
deps and bends the "native + no-framework-at-runtime" rule). Default per
methodology: native port.
- **D2 (scope)** — t2vs (T+video+sound) first; defer audio-conditioned / v2vs.
- **D3** — confirm FastVideo can mux/emit audio (output format).
## Status
- [x] Branch forked, framework audio path studied, plan written.
- [x] D1: native port (user-chosen). D2: t2vs first.
- [x] **AVAE sound decoder (component 1) — DONE** (commit `5f81fb3d5`). Key
finding: the checkpoint is decoder-only in AutoencoderOobleck naming with
SnakeBeta + weight_g/v == FastVideo's native `OobleckVAE` decoder. Reused it
(+ `output_padding=stride%2` for the odd stride 5); `Cosmos3SoundVAE`
decoder-only wrapper; bit-exact parity vs the framework OobleckDecoder
(`test_cosmos3_avae_parity`); real 1.9 GB checkpoint strict-loads, decodes
[1,64,25] -> [1,2,48000] (1 s @ 48 kHz stereo).
- [x] **DiT sound pathway (component 2) — DONE** (commit `005d6684a`). Activated
the dormant audio heads in the forward (`_encode_sound`/`_decode_sound` mirror);
`preds_vision` + `preds_sound` bit-exact (max=mean=0.0).
- [x] **Sound sequence packing (component 3) — DONE** (commit `005d6684a`).
`Cosmos3SoundItem` + sound fields; sound shares the vision "full" split with
parallel MRoPE. Field-by-field + position_ids exact vs framework.
- [x] **t2vs pipeline + AV mux (components 4-5) — DONE** (commit `3d8355129`).
Joint [vision|sound] denoise, AVAE-decode, stereo 48 kHz AAC mux. t2vs CFG
velocity parity max=mean=0.0; real-weights run produces coherent video + real
audio (mean -10.2 dB). Example `basic_cosmos3_t2vs_new_api.py`.
**PR2 (audio/t2vs) COMPLETE** — every component bit-exact vs the framework.
+152
View File
@@ -0,0 +1,152 @@
# Cosmos3 Port Status
## Summary
- model_family: `cosmos3`
- workload_types: `T2V, I2V, T2I` supported by `WorkloadType` today; full-omni target also needs audio (AV), VLM reasoning, and action-conditioning, which require framework extensions (Q002, Q003).
- official_ref: `https://github.com/NVIDIA/cosmos-framework` — diffusers backend `diffusers_cosmos3.pipeline.Cosmos3OmniDiffusersPipeline`; HF `nvidia/Cosmos3-Nano`.
- official_ref_dir: `cosmos-framework` (symlink -> `/home/william5lin/FastVideo/cosmos-framework`, commit `003d66d4`)
- hf_weights_path: `nvidia/Cosmos3-Nano`
- local_weights_dir: `official_weights/cosmos3` (symlink -> `/home/william5lin/FastVideo/official_weights/cosmos3`, 33 GiB / 67 files)
- source_layout: `diffusers`
- local_tests_readme: `tests/local_tests/cosmos3/README.md`
## Current Phase
- phase: `FULL OMNI SUPPORTED — every modality framework-parity verified bit-exact (suite 150 passed, 0 skipped). PR1 video core (T2V/I2V/T2I) + PR2 audio (t2vs) real-weights verified on B200; PR3 action (domain-aware) + PR4 reasoning (text + vision_encoder + deepstack reasoner) bit-exact. Branch chain: feat/cosmos3-tier-a-port (T2V) -> feat/cosmos3-i2v (I2V+T2I+flow_shift) -> feat/cosmos3-audio (t2vs) -> feat/cosmos3-action -> feat/cosmos3-reasoning. Optional follow-ups: real-weights action2world (needs robot-action data) + image-conditioned-reasoning prefill wiring (vision_encoder + get_rope_index, both proven).`
- status: `in_progress`
- owner: `orchestrator`
- last_updated: `2026-06-07`
- env: `fv-cosmos3` (conda clone of fv-main; `fastvideo` editable repointed to this worktree). Run tests from the worktree cwd with this env's python.
- branch: rebased onto `origin/main` @ `1c627a3f9` (was 33 behind, merge-base 2026-05-22); now 6 commits ahead; `fastvideo` imports clean; Tier-A `13 passed, 2 skipped`.
## Component Matrix
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|---|---|---|---|---|---|---|---|---|---|
| transformer | dit | port | `diffusers_cosmos3/transformer.py:Cosmos3OmniTransformer` (model_type `qwen3_vl_text`, MoT + MRoPE) | `model_index.json: transformer`; `cosmos_framework/model/vfm/mot/cosmos3_vfm_network.py`, `omni_mot_model.py` | `fastvideo/models/dits/cosmos3.py` (branch: `Cosmos3VFMTransformer`+`Cosmos3LanguageModel` — reconcile to `Cosmos3OmniTransformer`) | skeleton | not_started | scaffold_skip | I001 |
| vae | vae | reuse | diffusers `AutoencoderKLWan` | `model_index.json: vae` | reuse Wan VAE (`fastvideo/models/vaes/`, cf. `cosmos25wanvae.py`) | not_started | passthrough? | not_started | Q001 |
| scheduler | generic | reuse (flow-coerced) | framework `FlowUniPCMultistepScheduler` (`cosmos_framework/.../fm_solvers_unipc.py`; checkpoint ships diffusers-style config) | `model_index.json: scheduler`; `cosmos_framework/.../samplers/unipc.py:UniPCSampler` | FastVideo-native `UniPCMultistepScheduler` (flow config), coerced in `initialize_pipeline` | done | n/a | framework-parity DONE (`test_cosmos3_scheduler_parity`: timesteps bit-exact, sigmas ~1e-8, trajectory <~1e-6) | I003 (resolved) |
| text_tokenizer | tokenizer | reuse | transformers `Qwen2TokenizerFast` | `model_index.json: text_tokenizer` | reuse (tokenizer = allowed third-party) | not_started | passthrough | scaffold_skip (`test_cosmos3_tokenizer_chat_template`) | - |
| vision_encoder | encoder | port | transformers `Qwen3VLVisionModel` | `model_index.json: vision_encoder` | new encoder bucket OR documented lazy-wrapper | not_started | not_started | not_started | Q002 |
| sound_tokenizer | generic/vae | port (decode) | framework AVAE `LatentAutoEncoderV2` (`avae_utils`); checkpoint is decoder-only AutoencoderOobleck-named w/ SnakeBeta | `model_index.json: sound_tokenizer` | reuse FastVideo native `OobleckVAE` decoder + `Cosmos3SoundVAE` wrapper (`models/audio/cosmos3_avae.py`) | done (decode) | n/a | DECODE bit-exact vs framework (`test_cosmos3_avae_parity`); real ckpt strict-loads | PR2 (branch feat/cosmos3-audio) |
## Conversion State
- conversion_script: `scripts/checkpoint_conversion/cosmos3_convert.py` (branch has it, 246 lines, built vs vllm-omni — repoint/verify vs diffusers checkpoint)
- converted_weights_dir: `converted_weights/cosmos3` (n/a while needs_conversion=no)
- source_layout: `diffusers`
- needs_conversion: `no` (HF already diffusers-format; verify FastVideo loaders consume directly)
- strict_load_status: `not_run`
- passthrough_components: `vae (AutoencoderKLWan), scheduler (UniPC), text_tokenizer (Qwen2)` likely passthrough
- retry_history: `none`
## Parity Commands
| Scope | Command | Last Result | Notes |
|---|---|---|---|
| Tier-A scaffold | `cd <worktree> && <fv-cosmos3 python> -m pytest tests/local_tests/cosmos3/ -q` | `13 passed, 2 skipped` (2026-06-06, post-rebase) | 2 skips: Cosmos3 tokenizer/_tokenize_prompt not yet wired on pipeline |
| component | `pytest tests/local_tests/<bucket>/test_cosmos3_<component>_parity.py -v -s` | `not_run` | after env activation + native prototypes |
| pipeline | `pytest tests/local_tests/pipelines/test_cosmos3_pipeline_parity.py -v -s` | `not_run` | |
## Open Questions
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|---|---|---|---|---|---|
| Q001 | Does Cosmos3 VAE (`AutoencoderKLWan`) match FastVideo's existing Wan VAE config/instantiation exactly (z_dim, scale factors, latents_mean/std)? | orchestrator | 3 (reuse gate) | open | |
| Q002 | `vision_encoder` (`Qwen3VLVisionModel`): native port vs documented lazy-wrapper exception? Needed for I2V/reasoning. | user/orchestrator | 3 | open | |
| Q003 | `sound_tokenizer` (`Cosmos3AVAEAudioTokenizer`) + audio output requires `WorkloadType` AV + audio regression metric. | user | 0/10 | open | full-omni scope chosen 2026-06-06; infra extensions pending |
## Issues And Blockers
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|---|---|---|---|---|---|---|---|---|
| I001 | port | transformer | high | Branch DiT (`Cosmos3VFMTransformer`+`Cosmos3LanguageModel`) built vs vllm-omni #3454; official checkpoint loads `Cosmos3OmniTransformer` (diffusers shim). Class/structure reconciliation required. | `model_index.json`; `diffusers_cosmos3/transformer.py`; branch commit `52bb65f49` | orchestrator | resolved | DiT rewritten to checkpoint layout (single `layers` dual-pathway, BaseDiT-conformant); bit-identical framework parity (3d_rope + unified_3d_mrope), commits 59a4a571c/7c4633295 |
| I002 | all | tests | medium | Tier-A conftest+tests mirror vllm-omni line-by-line (stubs, `vllm_omni...guardrails`). Must be repointed to `diffusers_cosmos3` / official structures. | `tests/local_tests/cosmos3/conftest.py` | orchestrator | open | |
| I003 | inference | scheduler | high | First real-weights T2V was all-black: checkpoint `scheduler_config.json` sets `use_karras_sigmas=true`; vendored UniPC checks karras before `use_flow_sigmas` -> diffusion (beta) sigmas -> `scheduler.step` -> NaN latents. DiT/CFG velocity was clean. The scheduler had never been parity-tested vs the framework (`test_cosmos3_denoise_cfg_parity` used diffusers UniPC on both sides). | `result_latent` NaN at denoise step 0 (v_pred clean); ffprobe 3 KB black mp4 | orchestrator | resolved | Coerce loaded config to flow setup in `initialize_pipeline`; switch pipeline+tests to native UniPC (no diffusers at runtime); add `test_cosmos3_scheduler_parity` vs framework `FlowUniPCMultistepScheduler`; repoint denoise_cfg oracle to the framework scheduler. Commit 255311cf2 |
## Escape Hatches
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|---|---|---|---|---|---|---|
| E001 | prep | dependency/env | Shared `fv-main` env has `fastvideo` editable-installed from the MAIN worktree; the cosmos3 worktree's `fastvideo` is not importable (PEP660 finder overrides PYTHONPATH), so Tier-A tests skip. How to activate the worktree's `fastvideo` for verification without disrupting ~24 other worktrees sharing the env? | Dedicated conda env for the cosmos3 worktree | resolved | Created fv-cosmos3 (clone of fv-main); repointed fastvideo editable to worktree; run from worktree cwd. Branch also rebased onto origin/main to fix stale import. |
## Decisions
| Date | Decision | Rationale | Impact |
|---|---|---|---|
| 2026-06-06 | Reference source of truth = official diffusers (`Cosmos3OmniDiffusersPipeline` + `cosmos-framework`/`diffusers-cosmos3`), not vllm-omni #3454 | Official weights now public & diffusers-format; the artifact users actually load | Repoint DiT/pipeline/conversion/tests off vllm-omni (I001, I002) |
| 2026-06-06 | Resume in worktree `/home/william5lin/FastVideo_cosmos3_port`; weights+reference symlinked (no copy) | Preserve 2,492 lines of Tier-A work; avoid 33 GB duplication | Verification needs worktree `fastvideo` active (E001) |
| 2026-06-06 | Scope = full omni (video + audio + reasoning + action) | User choice (revised from branch's original video-only scope) | Adds `vision_encoder`, `sound_tokenizer` ports + `WorkloadType` AV + audio metric |
| 2026-06-06 | Downloaded full 34.9 GB (33 GiB) `nvidia/Cosmos3-Nano` | Unblocks May-22 `PENDING` weight status (HF was 401, now public) | Real parity now possible |
| 2026-06-06 | Rebased branch onto origin/main (33 commits); resolved registry.py conflict by reconstructing from main + cosmos3 import/entry | Branch was stale; fastvideo failed to import (main removed MatrixGameI2V480PConfig) | Branch imports clean; Tier-A 13 passed/2 skipped |
| 2026-06-06 | Reference = cosmos_framework ONLY (full omni); diffusers shim dropped even for video | User directive (Phase 1 found diffusers __call__ is video-only; sound/action/reasoning live only in the framework) | Larger port; ref DiT = `Cosmos3VFMNetwork`/`Cosmos3VFMNetworkConfig` (not diffusers `Cosmos3OmniTransformer`); core model imports in fv-cosmos3 with light deps; TE only in optional dot_product_attention |
## Handoff Notes
- Prep (weights/reference/env editable installs) done in MAIN worktree; symlinked into this worktree. Env installs (`diffusers-cosmos3`, `cosmos-framework`) are in shared `fv-main`.
- Next: resolve E001 (env), then Phase 1 reference study of `diffusers_cosmos3` pipeline/transformer, then Phase 3 reuse gate (VAE/scheduler/tokenizer) + component dispatch (transformer, vision_encoder, sound_tokenizer).
- diffusers 0.36.0 imports the shim OK; checkpoint saved with 0.37.1 — watch `from_pretrained` needs (bump within FastVideo's `diffusers>=0.33.1` pin if required).
### PR1 (video core) progress — 2026-06-06
- Arch config 1:1 with checkpoint, committed `9567efdf0`.
- Framework parity-reference harness committed `dd97efda3`: `tests/local_tests/cosmos3/test_cosmos3_reference_forward.py` builds a tiny `Cosmos3VFMNetwork` on CPU/float32 (SDPA monkeypatch; flash2/3/natten are CUDA-only) and forwards `packed_seq -> {last_hidden_state, preds_vision}`. 23 tests pass in fv-cosmos3. This is the ground-truth side for DiT parity. Run: `cd <worktree> && <fv-cosmos3 py> -m pytest tests/local_tests/cosmos3/test_cosmos3_reference_forward.py -q`.
- THREE naming conventions to bridge:
1. framework-native (`Cosmos3VFMNetwork`): `language_model.model.layers.{i}.self_attn.{q,k,v,o}_proj(+ _moe_gen)`, `{q,k}_norm(+_moe_gen)`, `mlp(+_moe_gen)`, `vae2llm`/`llm2vae`, `time_embedder.mlp.{0,2}`.
2. diffusers checkpoint (on disk, what we load): `layers.{i}.self_attn.{to_q,to_k,to_v,to_out}` + `{add_q,add_k,add_v}_proj`/`to_add_out`, `{norm_q,norm_k,norm_added_q,norm_added_k}`, `mlp`/`mlp_moe_gen`, `proj_in`/`proj_out`, `time_embedder.linear_{1,2}`.
3. FastVideo DiT (our choice). Conversion maps (2)->(3); the DiT parity test copies (1)->(3).
- BaseDiT signature is `__init__(self, config: DiTConfig, hf_config: dict)`; the branch `Cosmos3VFMTransformer` uses `fastvideo_args`/SimpleNamespace and does NOT conform — rewrite to conform + match the checkpoint key surface (single `layers` dual-pathway, not split language_model/gen_layers).
- Native layers (per cosmos2_5): `ReplicatedLinear`/`MLP`/`RMSNorm` (fastvideo.layers.*), `LocalAttention`/`DistributedAttention` (fastvideo.attention), `apply_rotary_emb` (use_real_unbind_dim=-2 for Cosmos). EntryClass at module bottom; class attrs bound from config; 3D-MRoPE has no reusable util — adapt Cosmos25RotaryPosEmbed.
- NEXT: write native `fastvideo/models/dits/cosmos3.py` + fastvideo-vs-framework forward parity test (copy framework weights into the FastVideo DiT, compare outputs), then conversion script (diffusers checkpoint -> FastVideo) + strict-load, then video pipeline/packing.
### PR1 (video core) acceptance — real-weights E2E — 2026-06-07
- First real-weights T2V (`examples/inference/basic/basic_cosmos3_new_api.py`, `COSMOS3_MODEL_PATH=official_weights/cosmos3`) ran mechanically but produced an all-black 3 KB mp4. Instrumenting the denoise loop showed `v_pred` clean at step 0 but `scheduler.step` -> NaN. Root cause I003: checkpoint `scheduler_config.json` is diffusers-style (`use_karras_sigmas=true`), and the vendored UniPC checks karras before `use_flow_sigmas` -> diffusion (beta) sigmas instead of flow sigmas -> NaN. The framework actually samples with `FlowUniPCMultistepScheduler` (pure flow: `shift` + `num_train_timesteps`).
- Fix (commit `255311cf2`): coerce the loaded scheduler to the flow setup in `Cosmos3OmniDiffusersPipeline.initialize_pipeline`; use FastVideo's native UniPC (not diffusers) in pipeline + tests. Added `test_cosmos3_scheduler_parity.py` (native UniPC flow-config vs framework `FlowUniPCMultistepScheduler`: timesteps bit-exact, sigmas ~1e-8, full trajectory <~1e-6 over shift in {10,3}, steps in {4,10,35}). Repointed `test_cosmos3_denoise_cfg_parity` oracle to the framework scheduler (it previously compared diffusers-vs-diffusers, so the scheduler was never checked against the framework).
- Also wired the remaining integration glue (registry alias `Cosmos3OmniTransformer`->`Cosmos3VFMTransformer`; `text_tokenizer`->TokenizerLoader; scheduler config param-filtering; DiT `materialize_non_persistent_buffers` + compute-dtype casts; packing device-move in `to_dit_kwargs`; empty text-preprocess).
- Verified: 1280x704, 29 frames, 35 steps on a single B200 -> coherent golden-retriever-in-meadow video matching the prompt (no NaNs; per-frame pixel std ~58; visible temporal motion). Full cosmos3 suite: 95 passed, 0 skipped.
- NEXT: PR2 audio (`sound_tokenizer` AVAE) / PR3 action / PR4 reasoning. Optional: I2V/T2I real-weights spot-checks; force-push branch (needs explicit OK).
### PR1 (video core) — I2V real-weights — 2026-06-07 (branch feat/cosmos3-i2v)
- Forked `feat/cosmos3-i2v` off `feat/cosmos3-tier-a-port` (stacked, includes the T2V + scheduler fix).
- Studied the framework I2V path: `cosmos_framework.inference.vision.load_conditioning_image` (aspect-preserving resize + center crop + uint8 quantize -> `/127.5-1`) + `build_conditioned_video_batch` (frame 0 = image, remaining frames REPEAT the last conditioning frame -> static video), then VAE-encode; `condition_frame_indexes=[0]` (latent). Condition frames kept clean during sampling exactly as FastVideo already does: init noise `cond_mask*x0 + (1-cond_mask)*noise` (`omni_mot_model._prepare_inference_data`) + velocity zeroed `pred*(1-cond_mask)` each step (`_get_velocity`), no re-injection.
- Bug found + fixed (commit `bd8d604fb`): FastVideo's `_image_to_video_tensor` ZERO-filled the non-condition frames; the temporal Wan VAE (4x) makes latent frame 0 depend on several pixel frames, so zero-fill -> wrong conditioning latent. Rewrote it to repeat-fill + framework resize/crop/quantize.
- Parity: `test_cosmos3_i2v_conditioning_parity.py` vs framework `load_conditioning_image` + repeat-fill — bit-exact (max abs diff 0.0) across aspect/size/frame cases. Existing `test_cosmos3_denoise_cfg_parity` already covers the I2V cond-mask + velocity math (i2v case).
- Example: `examples/inference/basic/basic_cosmos3_i2v_new_api.py` (`InputConfig(image_path=...)`, default `assets/images/cyclist.jpg`).
- Verified on B200 (1280x704, 29f, 35 steps, real weights): output frame 0 reproduces the conditioning cyclist image; later frames show coherent forward motion down the trail following the prompt. Full suite 98 passed, 0 skipped.
- NEXT: optional T2I real-weights spot-check; then PR2 audio / PR3 action / PR4 reasoning.
### PR1 (video core) — T2I real-weights + resolution-based flow_shift — 2026-06-07 (branch feat/cosmos3-i2v)
- Studied framework T2I: tokenization uses `vlm_config.use_system_prompt` which is `false` in the checkpoint (config.json:199) — matches FastVideo's hardcoded `use_system_prompt=False` for all modes (no divergence). Canonical T2I is 960x960 (inputs/omni/t2i.json), single-frame (num_frames=1).
- Bug found + fixed (commit `604dc2637`): the stage chose `flow_shift` by task (`3.0 if is_t2i else 10.0`), but the framework picks it purely by the named resolution bucket (`OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256->3.0, 480->5.0, 720/768->10.0; model default resolution "720"). Task-based only matched T2V@720 / T2I@256 by luck; canonical T2I@960x960 is the "720" bucket -> 10.0, so `is_t2i->3.0` was wrong. Replaced with `_flow_shift_for_resolution(h,w)` (longest-side bucketing), applied to all tasks.
- Parity: `test_cosmos3_flow_shift_parity.py` checks the mapping vs framework `{VIDEO,IMAGE}_RES_SIZE_INFO` x `_RESOLUTION_SHIFT_DEFAULTS` (8B rows, 20 cases). Also hardened `_image_to_video_tensor` tensor branch to respect the [-1,1] convention (PIL path stays framework-exact).
- Example: `examples/inference/basic/basic_cosmos3_t2i_new_api.py` (num_frames=1, 960x960).
- Verified on B200 (real weights, 35 steps): coherent red-panda image matching the prompt, flow_shift=10.0. Full suite 118 passed, 0 skipped.
- Video core (T2V/I2V/T2I) is now complete and real-weights verified. NEXT: PR2 audio (sound_tokenizer AVAE + audio output) on a new stacked branch.
## Full-omni parity summary (every component, max / mean abs diff vs framework)
All run on CPU / float32 (tiny models, framework weights copied in; framework =
oracle). `tests/local_tests/cosmos3/`, suite: 150 passed, 0 skipped.
| Component / pipeline | Test | max | mean |
|---|---|---|---|
| Scheduler (UniPC flow) | test_cosmos3_scheduler_parity | timesteps 0; sigmas ~1e-8; traj <~1e-6 | ~1e-7 |
| DiT (video, unified_3d_mrope) | test_cosmos3_dit_parity_mrope | 0.0 | 0.0 |
| Sequence packing (video) | test_cosmos3_packing_parity | 0.0 (exact) | 0.0 |
| VAE (Wan2.2) | test_cosmos3_vae_parity | 0.0 | 0.0 |
| Denoise / CFG velocity | test_cosmos3_denoise_cfg_parity | <1e-6 | <1e-7 |
| flow_shift (resolution) | test_cosmos3_flow_shift_parity | exact | exact |
| I2V conditioning (static-repeat) | test_cosmos3_i2v_conditioning_parity | 0.0 | 0.0 |
| AVAE sound decoder | test_cosmos3_avae_parity | 0.0 | 0.0 |
| DiT sound pathway + packing | test_cosmos3_sound_parity | 0.0 | 0.0 |
| t2vs CFG velocity | test_cosmos3_sound_parity | 0.0 | 0.0 |
| DiT action pathway + packing | test_cosmos3_action_parity | 0.0 | 0.0 |
| action CFG velocity | test_cosmos3_action_parity | 0.0 | 0.0 |
| Reasoner prefill logits (text) | test_cosmos3_reasoning_parity | 0.0 | 0.0 |
| Reasoner greedy generation | test_cosmos3_reasoning_parity | token-exact | - |
| Deepstack reasoner forward | test_cosmos3_reasoning_parity | 0.0 | 0.0 |
| vision_encoder (Qwen3-VL ViT) | test_cosmos3_vision_encoder_parity | 0.0 | 0.0 |
Real-weights pipelines verified on B200 (`examples/inference/basic/basic_cosmos3*_new_api.py`):
T2V (1280x704), I2V (cyclist), T2I (960x960 red panda), t2vs (ocean + stereo
48kHz audio), text reasoning (greedy == framework). All coherent / prompt-matching.
+71
View File
@@ -0,0 +1,71 @@
# Cosmos3 local parity workspace
## Overview
This workspace tracks the FastVideo Cosmos3 port. Live port state, component matrix,
decisions, and blockers live in `PORT_STATUS.md`.
- **Reference (2026-06-06): official NVIDIA `cosmos-framework` diffusers backend** —
`Cosmos3OmniDiffusersPipeline` from the `diffusers-cosmos3` shim — loading the
now-public `nvidia/Cosmos3-Nano` checkpoint.
- **Scope: full omni** — T2V / I2V / T2I, audio (sound generation), VLM reasoning,
and action-conditioning.
- The original Tier-A scaffold was written against vllm-omni PR #3454 before official
weights were public; it is being repointed to the diffusers reference (see I001/I002
in `PORT_STATUS.md`).
## Reference code
Primary (official):
- Local: `cosmos-framework/` (symlink -> `/home/william5lin/FastVideo/cosmos-framework`,
commit `003d66d4`); GitHub <https://github.com/NVIDIA/cosmos-framework>
- diffusers shim `cosmos-framework/packages/diffusers-cosmos3/diffusers_cosmos3/`:
- `pipeline.py` — `Cosmos3OmniDiffusersPipeline`
- `transformer.py` — `Cosmos3OmniTransformer`
- `sequence_packing.py`
- framework model code: `cosmos_framework/model/vfm/mot/cosmos3_vfm_network.py`,
`cosmos_framework/model/vfm/omni_mot_model.py`
- Installed editable in shared `fv-main`: `diffusers-cosmos3`, `cosmos-framework`
(both `--no-deps`).
Original Tier-A reference (superseded, kept for diffing during repoint):
- vllm-omni PR #3454 <https://github.com/vllm-project/vllm-omni/pull/3454>, pinned
`8536f5b1`, checkout `/home/william5lin/cosmos3-reference`.
- The current `conftest.py` + tests still mirror this suite line-by-line.
## Weight status
DOWNLOADED (2026-06-06). `nvidia/Cosmos3-Nano` is now public and diffusers-format
(the 2026-05-22 `401` is resolved).
- Local: `official_weights/cosmos3/` (symlink -> main worktree; 33 GiB, 67 files,
`model_index.json` present)
- Source: `nvidia/Cosmos3-Nano`, default revision; `source_layout=diffusers`,
`needs_conversion=no`
- `model_index` class: `Cosmos3OmniDiffusersPipeline` (diffusers 0.37.1)
- Token: not required (public repo)
Components (from `model_index.json`): `transformer` (`Cosmos3OmniTransformer`),
`vae` (`AutoencoderKLWan`), `scheduler` (`UniPCMultistepScheduler`),
`text_tokenizer` (`Qwen2TokenizerFast`), `vision_encoder` (`Qwen3VLVisionModel`),
`sound_tokenizer` (`Cosmos3AVAEAudioTokenizer`).
## Running the Tier-A scaffold
```bash
PYTHONPATH=/home/william5lin/FastVideo_cosmos3_port \
python -m pytest tests/local_tests/cosmos3/ -q
```
NOTE: as of 2026-06-06 these report `15 skipped` because the shared `fv-main` env's
editable `fastvideo` resolves to the MAIN worktree (a PEP660 finder overrides
`PYTHONPATH`), so the worktree's cosmos3 modules are not importable. Tracked as E001
in `PORT_STATUS.md`.
## SSIM placeholder
No SSIM references seeded yet. Add SSIM coverage only after a FastVideo inference path
can load the Cosmos3 weights and generate stable T2V/I2V/T2I outputs. Audio quality
uses a separate metric (not SSIM); see `PORT_STATUS.md` Q003.
+214
View File
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared fixtures for the Cosmos3 native-pipeline local tests.
These fixtures build the FastVideo-native Cosmos3 pipeline
(``fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline.Cosmos3OmniDiffusersPipeline``)
via ``__new__`` and wire it with tiny stub components so the runtime call graph
(sequential CFG, condition-frame masking, mode dispatch) can be exercised on CPU
without real weights or ``cosmos_framework``.
The stub transformer implements the native DiT's packed-input contract
(``{"preds_vision": [[1, C, T, H, W], ...]}``) and records, per call, the first
``text_ids`` token so tests can assert the cond/uncond pass order.
"""
from __future__ import annotations
from types import SimpleNamespace
from typing import Any
import pytest
import torch
from fastvideo.models.schedulers.scheduling_unipc_multistep import (
UniPCMultistepScheduler,
)
from torch import nn
_LATENT_CHANNEL = 16
_LATENT_PATCH_SIZE = 2
_SPATIAL_FACTOR = 8
_TEMPORAL_FACTOR = 4
def pytest_configure(config: pytest.Config) -> None:
"""Register the ``local`` marker used by sibling test files."""
config.addinivalue_line(
"markers",
"local: marker for local-only parity/scaffold tests (skipped in CI)",
)
# ---------------------------------------------------------------------------
# Stub transformer: records cond/uncond call order; bounded preds_vision.
# ---------------------------------------------------------------------------
class StubCosmos3Transformer(nn.Module):
"""Records each forward's first ``text_ids`` token + returns preds_vision.
``preds_vision`` is keyed by the first text token (so the conditional and
unconditional passes return different velocities) and is zero on
conditioning frames, matching the real DiT's unpatchify output.
"""
def __init__(self, latent_channel: int = _LATENT_CHANNEL) -> None:
super().__init__()
self.latent_channel = latent_channel
self.embed_tokens = nn.Embedding(64, 8)
self.calls: list[dict[str, Any]] = []
def forward(self, **kwargs: Any) -> dict[str, Any]:
token_ids = kwargs["text_ids"]
token = int(token_ids.reshape(-1)[0].item()) if token_ids.numel() else 0
self.calls.append({"token": token, "kwargs": dict(kwargs)})
scale = 0.01 * (1.0 + (token % 7))
preds: list[torch.Tensor] = []
for latent, _shape, nfi in zip(kwargs["vision_tokens"], kwargs["vision_token_shapes"],
kwargs["vision_noisy_frame_indexes"]):
lat = latent.squeeze(0) if latent.dim() == 5 else latent # [C, T, H, W]
out = torch.zeros_like(lat)
if nfi.numel() > 0:
out[:, nfi] = scale * torch.tanh(lat[:, nfi])
preds.append(out.unsqueeze(0))
return {"preds_vision": preds}
class _StubLatentDist:
def __init__(self, latents: torch.Tensor) -> None:
self._latents = latents
def mode(self) -> torch.Tensor:
return self._latents
class StubCosmos3VAE:
"""Deterministic VAE shaped by the Wan scale factors."""
def __init__(self, z_dim: int = _LATENT_CHANNEL) -> None:
self.config = SimpleNamespace(
z_dim=z_dim,
scale_factor_temporal=_TEMPORAL_FACTOR,
scale_factor_spatial=_SPATIAL_FACTOR,
latents_mean=[0.0] * z_dim,
latents_std=[1.0] * z_dim,
)
def encode(self, video: torch.Tensor):
b, _c, t, h, w = video.shape
lt = (t - 1) // self.config.scale_factor_temporal + 1
lh = h // self.config.scale_factor_spatial
lw = w // self.config.scale_factor_spatial
return _StubLatentDist(torch.ones(b, self.config.z_dim, lt, lh, lw, dtype=video.dtype, device=video.device))
def decode(self, z: torch.Tensor):
b, _c, lt, lh, lw = z.shape
t = (lt - 1) * self.config.scale_factor_temporal + 1
h = lh * self.config.scale_factor_spatial
w = lw * self.config.scale_factor_spatial
sig = torch.nan_to_num(torch.tanh(z[:, :1, :1, :1, :1])).reshape(b, 1, 1, 1, 1)
return torch.clamp(torch.zeros(b, 3, t, h, w, dtype=z.dtype, device=z.device) + sig, -1.0, 1.0)
class StubQwen2Tokenizer:
"""Qwen2-shaped chat tokenizer stub (special tokens + chat template)."""
eos_token_id = 62
_SPECIAL = {"<|vision_start|>": 60, "<|vision_end|>": 61}
def convert_tokens_to_ids(self, token: str) -> int:
return self._SPECIAL[token]
def apply_chat_template(self, conversations, *, tokenize=True, add_generation_prompt=True, add_vision_id=False):
user = next((c["content"] for c in conversations if c["role"] == "user"), "")
n = max(1, min(8, len(user) % 8 + 1))
return [10 + (i % 40) for i in range(n)]
def make_scheduler(flow_shift: float = 10.0) -> UniPCMultistepScheduler:
return UniPCMultistepScheduler(
num_train_timesteps=1000,
solver_order=2,
prediction_type="flow_prediction",
use_flow_sigmas=True,
flow_shift=flow_shift,
)
# ---------------------------------------------------------------------------
# Pipeline factory — builds the native pipeline via __new__ + stub modules.
# ---------------------------------------------------------------------------
@pytest.fixture
def make_cosmos3_pipeline():
"""Return a factory building the native Cosmos3 pipeline wired with stubs."""
def _make(**overrides: Any):
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import ( # noqa: F401
Cosmos3OmniDiffusersPipeline, )
pipe = Cosmos3OmniDiffusersPipeline.__new__(Cosmos3OmniDiffusersPipeline)
scheduler = make_scheduler()
pipe.modules = {
"transformer": StubCosmos3Transformer(),
"vae": StubCosmos3VAE(),
"scheduler": scheduler,
"text_tokenizer": StubQwen2Tokenizer(),
}
pipe.scheduler = scheduler
pipe._base_scheduler_config = scheduler.config
pipe._current_flow_shift = float(scheduler.config.flow_shift)
pipe._engine_init_flow_shift = 10.0
for key, value in overrides.items():
setattr(pipe, key, value)
return pipe
return _make
@pytest.fixture
def make_cosmos3_stage():
"""Return a factory building a ``Cosmos3DenoisingStage`` bound to a pipeline."""
def _make(pipeline):
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
return Cosmos3DenoisingStage(
transformer=pipeline.modules["transformer"],
scheduler=pipeline.modules["scheduler"],
vae=pipeline.modules["vae"],
tokenizer=pipeline.modules["text_tokenizer"],
pipeline=pipeline,
)
return _make
def make_forward_batch(*, num_frames: int, height: int, width: int, image: Any = None, **overrides: Any):
"""Build a tiny ``ForwardBatch`` for the Cosmos3 stage."""
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
values: dict[str, Any] = dict(
data_type="video",
prompt="a calm ocean at sunrise",
negative_prompt="",
height=height,
width=width,
num_frames=num_frames,
fps=24,
num_inference_steps=2,
guidance_scale=6.0,
generator=torch.Generator("cpu").manual_seed(0),
preprocessed_image=image,
)
values.update(overrides)
return ForwardBatch(**values)
def make_fastvideo_args():
"""Build minimal ``fastvideo_args`` (only ``pipeline_config`` is read)."""
from fastvideo.configs.pipelines.cosmos3 import Cosmos3Config
cfg = Cosmos3Config()
arch = cfg.dit_config.arch_config
arch.latent_channel = _LATENT_CHANNEL
arch.latent_patch_size = _LATENT_PATCH_SIZE
arch.temporal_compression_factor = _TEMPORAL_FACTOR
arch.enable_fps_modulation = False
return SimpleNamespace(pipeline_config=cfg)
+166
View File
@@ -0,0 +1,166 @@
# Cosmos3 → FastVideo port — feedback (pitfalls, issues, difficulties)
Retrospective on porting the full **NVIDIA Cosmos3-Nano** omni world model
(video / audio / action generation + text & image reasoning) into FastVideo.
Methodology: framework-only reference, native FastVideo port, a bit-exact
framework-parity test per component, then real-weights verification. Every
modality landed bit-exact (see `PORT_STATUS.md` "Full-omni parity summary").
This doc records what bit, so the next omni/world-model port (and the `/add-model`
skill) can avoid the same traps.
---
## 1. The checkpoint's config does NOT describe the runtime — verify against the framework
The single biggest time sink. The HF checkpoint is "diffusers format", which led
to two silent traps:
- **Scheduler (caused an all-black video).** `scheduler/scheduler_config.json`
is a diffusers `UniPCMultistepScheduler` config carrying
`use_karras_sigmas=true`, `sigma_min/sigma_max`, a beta schedule, etc. But the
framework actually samples with a *flow-matching* `FlowUniPCMultistepScheduler`
(shift + num_train_timesteps only). FastVideo's vendored UniPC checks
`use_karras_sigmas` **before** `use_flow_sigmas`, so it built diffusion (beta)
sigmas → `scheduler.step` → **NaN latents → 3 KB black mp4**. The DiT/CFG
velocity was perfectly clean; only the scheduler diverged.
- Fix: coerce the loaded config to the flow setup in `initialize_pipeline`.
- **Lesson:** treat the checkpoint's generic-format config as *lossy*. Find how
the framework actually instantiates the component and match THAT, not the JSON.
- **`flow_shift` is resolution-based, not task-based.** Natural assumption:
"T2I uses a small shift, video a large one." Reality: the framework keys the
UniPC shift purely off the named resolution bucket
(`_RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256→3, 480→5, 720/768→10). The
task-based heuristic only *coincidentally* matched (T2V@720, T2I@256); canonical
T2I is 960×960 (the "720" bucket → 10), so `is_t2i→3.0` was wrong.
## 2. A "parity test" that compares two copies of the wrong thing proves nothing
The original denoise/CFG test imported **diffusers** `UniPCMultistepScheduler` and
used it on BOTH the "oracle" and "FastVideo" sides. So the scheduler was never
actually compared against the framework — which is exactly why the black-video
scheduler bug sailed through a green test suite.
- **Lesson:** the oracle side of every parity test MUST be the official framework
object, never a second instance of the unit under test. After writing a parity
test, ask: "if the framework were wrong here, would this test fail?"
## 3. Temporal-VAE conditioning: static-repeat vs zero-fill (silent corruption)
I2V/T2I condition on the input image. The framework
(`build_conditioned_video_batch`) fills frame 0 with the image and **repeats the
last conditioning frame across the whole clip** (a static video) before
VAE-encoding. The first native cut **zero-filled** the non-condition frames.
Because the Wan VAE is temporal (4× compression), latent-frame-0 (the kept-clean
condition frame) depends on several *pixel* frames — so zero-filling produced a
*wrong* conditioning latent. This is the kind of bug that doesn't crash and can
even look plausible at a glance.
- **Lesson:** when a conditioning latent feeds a temporal autoencoder, trace the
temporal receptive field; "only frame 0 matters" is false under temporal conv.
## 4. Checkpoint param names ≠ framework module structure (three naming conventions)
For the DiT there were **three** namings to bridge: framework-native
(`Cosmos3VFMNetwork`: `language_model.model.layers.*`, `vae2llm`, `q_proj_moe_gen`,
…), the diffusers checkpoint on disk (`layers.*.to_q`, `add_q_proj`, `proj_in`,
…), and the FastVideo DiT. The weight map crosses (framework)→(FastVideo) for
parity and (checkpoint)→(FastVideo) for loading.
The **sound tokenizer** was the sharpest example: the checkpoint is **decoder-only**
in diffusers `AutoencoderOobleck` naming (`decoder.conv1`, `block.N.conv_t1`,
`res_unitM`, `snake1`) — but with **`SnakeBeta`** (learned alpha *and* beta,
logscale), NOT diffusers' alpha-only `Snake1d`. So neither "use diffusers
AutoencoderOobleck" nor "port the framework `LatentAutoEncoderV2` Sequential
module" matched the on-disk keys.
- **Lesson:** dump `safetensors` keys + shapes for every sub-checkpoint *first*.
The naming reveals which existing native module (if any) already matches.
## 5. A "matching" native module can still differ on an untested config path
FastVideo already had a native `OobleckVAE` (Stable Audio) whose decoder matched
the Cosmos3 sound decoder bit-for-bit — except `OobleckDecoderBlock.conv_t1`
omitted `output_padding = stride % 2`. That omission is a **no-op for Stable
Audio's even strides** [2,4,4,8,8], so it had never mattered; Cosmos3 has an
**odd** stride (5), where the framework's `output_padding=1` makes the decode one
sample longer per odd-stride block (parity diverged 60 vs 59 samples).
- **Lesson:** reusing a native module is great, but re-run parity on the *new*
model's config — shared code can hide config-specific divergences.
## 6. Loader / registry plumbing the checkpoint format forces
- **DiT class alias.** `model_index.json` names the DiT `Cosmos3OmniTransformer`
(the diffusers shim class); the registry normalized unknown classes to a
generic `TransformersModel`. Needed an explicit registry alias
`Cosmos3OmniTransformer → Cosmos3VFMTransformer`.
- **Tokenizer module name.** `model_index.json` calls the Qwen2 tokenizer
`text_tokenizer` (not `tokenizer`); the component loader had no mapping for that
key and tried to load it as a model.
- **Scheduler config schema drift.** The vendored UniPC predates
`shift_terminal` / `sigma_min` / `sigma_max`; constructing it with the raw
checkpoint config crashes on the unexpected kwargs. Filter to the class's
`__init__` params (mirroring diffusers `from_config`).
- **Meta-device load + non-persistent buffers.** `rotary_emb.inv_freq` is derived
from `rope_theta` and is non-persistent (absent from the checkpoint), so after
the meta-device FSDP load it stays on the `meta` device → needs a
`materialize_non_persistent_buffers` hook to recompute it on the real device.
- **dtype boundaries.** Noise/VAE latents arrive fp32; the model runs bf16.
Needed explicit casts at `proj_in` and the timestep embedder (no-ops in the
fp32 parity tests, required at inference).
- **device in packing.** The packer builds ids/positions on CPU; `to_dit_kwargs`
must move every tensor to the model device before the forward.
## 7. The omni model is a Mixture-of-Transformers — modality bookkeeping is the work
The backbone is a dual-pathway MoT: **und** (causal text) + **gen** (full-attention
vision/sound/action). Once the video path worked, each extra modality was the same
*shape* of work (a proj-in + modality embed + timestep-scatter encode, a proj-out
decode, packing, a CFG-velocity slice) but with per-modality quirks:
- sound/action **share the vision "full" split** (preserving the causal+full
2-split invariant); the combined flat latent is `[vision | action | sound]` in
that order (must match the framework's per-sample concat).
- sound MRoPE uses `start_frame_offset=0` (parallel to vision); action uses
`start_frame_offset=1`; both at the vision temporal offset, tcf=1, and do NOT
advance the offset.
- action is **domain-aware** (`DomainAwareLinear`: per-embodiment weight/bias via
`nn.Embedding`, indexed by a per-token domain id).
- the unpack already zeros clean frames, so the per-step velocity masking is
defensive (but kept, to mirror the framework exactly).
- **Lesson:** build the first modality (vision) with clean seams for "a modality"
and the rest fall out; spend the care on the packing layout + MRoPE offsets,
which are the only per-modality novelties.
## 8. Reasoning reused more than expected; the encoder is just transformers
- **Text reasoning** needed *no new model code*: it's the und (causal) pathway +
`embed_tokens`/`norm`/`lm_head`, all already in the DiT. A text-only forward +
`lm_head` is token-for-token identical to the framework reasoner.
- **vision_encoder** is a stock `transformers.Qwen3VLVisionModel` (the framework
ships its own *copy* of the same class); reusing transformers' (like the Qwen2
tokenizer) is bit-exact vs the framework — re-porting a 27-layer ViT would have
been wasted effort.
- **deepstack** (image-conditioned reasoning) is the one new native piece: inject
the 3 vision-encoder deepstack features into the first 3 text layers at the
image-token positions.
- **Lesson:** before porting a big sub-model, check whether it's literally a
stock library class — and whether an existing in-repo module already implements
it (audio decoder + vision encoder were both "already there").
## 9. Running the framework as a CPU parity oracle
The framework's attention path is flash/natten (CUDA-only). Parity tests run on
CPU/float32 via an SDPA monkey-patch (`test_cosmos3_reference_forward._apply_sdpa_patches`).
A couple of framework helpers also can't import headless (`cosmos_framework.inference.args`
pulls `multistorageclient`), so a constant or two is mirrored in the test with a
cited source rather than imported.
- **Lesson:** budget for "make the oracle runnable on CPU" — monkeypatch attention,
build tiny configs, and accept a small amount of mirrored constants when a
framework module won't import in isolation.
## 10. What made it tractable
- Tiny CPU/fp32 models + copy-framework-weights-in + bit-exact compare, per
component, is a fast and decisive loop (max=mean=0.0 or it's wrong).
- A persistent `PORT_STATUS.md` (resumable state, issues, decisions) survived
several context resets.
- Stacked PR branches (one per modality) kept each parity-verified increment
reviewable and the chain bisectable.
@@ -0,0 +1,271 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 action pathway vs the framework.
Covers the action (multi-embodiment world-model) modality at the DiT level:
* **action packing** — native ``pack_cosmos3_video_sequence`` with a
``Cosmos3ActionItem`` vs framework ``pack_input_sequence`` with
``has_action``: action tokens share the vision "full" split, with ``(T,)``
shapes, a ``(T,1)`` condition mask, and 3D-MRoPE temporal positions at the
vision offset with ``start_frame_offset=1`` (parallel to vision); and
* **DiT action forward** — the dormant domain-aware ``action_proj_in`` /
``action_proj_out`` (``DomainAwareLinear``) + ``action_modality_embed`` heads,
now activated, with a per-token embodiment ``domain_id``.
Framework model + pack is the parity ORACLE (CPU/float32 via SDPA monkey-patch).
We assert the native packer matches the framework field-by-field, then that
``preds_vision`` AND ``preds_action`` match the framework forward.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_action_parity.py -q -s
"""
from __future__ import annotations
import pytest
import torch
cosmos_framework = pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
from .test_cosmos3_dit_parity import ( # noqa: E402
_fastvideo_inputs_from_packed_seq,
_framework_to_fastvideo_state_dict,
)
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
_ACTION_DIM,
_LATENT_CHANNEL,
_LATENT_PATCH_SIZE,
_RESET_SPATIAL_IDS,
_TCF,
_TEMPORAL_MODALITY_MARGIN,
_build_tiny_cosmos3_mrope,
_build_tiny_fastvideo_dit_mrope,
)
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
pytestmark = [pytest.mark.local]
_apply_sdpa_patches()
_SPECIAL_TOKENS = {"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62}
def _copy_weights_with_action(vfm, dit) -> None:
"""Copy backbone + vision weights AND the domain-aware action heads."""
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
src = dict(vfm.named_parameters())
mapped["action_proj_in.fc.weight"] = src["action2llm.fc.weight"].detach().clone()
mapped["action_proj_in.bias.weight"] = src["action2llm.bias.weight"].detach().clone()
mapped["action_proj_out.fc.weight"] = src["llm2action.fc.weight"].detach().clone()
mapped["action_proj_out.bias.weight"] = src["llm2action.bias.weight"].detach().clone()
mapped["action_modality_embed"] = src["action_modality_embed"].detach().clone()
dst = dict(dit.named_parameters())
with torch.no_grad():
for name, tensor in mapped.items():
assert name in dst, f"DiT missing param {name!r}"
assert dst[name].shape == tensor.shape, f"shape mismatch {name}"
dst[name].copy_(tensor.to(dst[name].dtype))
def _framework_pack_action(*, text_ids, vision, action, cond_vision, cond_action, domain_id, timestep,
is_image_batch):
from cosmos_framework.data.vfm.sequence_packing import (
GenerationDataClean,
SequencePlan,
pack_input_sequence,
)
gen = GenerationDataClean(
batch_size=1,
is_image_batch=is_image_batch,
x0_tokens_vision=[vision],
fps_vision=None,
num_vision_items_per_sample=[1],
x0_tokens_action=[action],
fps_action=None,
action_domain_id=[torch.tensor([domain_id], dtype=torch.long)],
)
plans = [SequencePlan(
has_text=True, has_vision=True, has_action=True,
condition_frame_indexes_vision=list(cond_vision),
condition_frame_indexes_action=list(cond_action),
)]
ps = pack_input_sequence(
sequence_plans=plans,
input_text_indexes=[list(text_ids)],
gen_data_clean=gen,
input_timesteps=torch.tensor([timestep], dtype=torch.float32),
special_tokens=_SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE,
include_end_of_generation_token=False,
position_embedding_type="unified_3d_mrope",
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TCF,
)
# The framework sets action.domain_id on the packed sequence from
# gen_data_clean inside the model (_get_velocity); mirror that for the oracle.
if ps.action is not None:
ps.action.domain_id = [torch.tensor([domain_id], dtype=torch.long)]
return ps
def _fastvideo_pack_action(*, text_ids, vision, action, cond_vision, cond_action, domain_id, timestep):
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
Cosmos3ActionItem,
Cosmos3SampleInputs,
Cosmos3VisionItem,
pack_cosmos3_video_sequence,
)
samples = [Cosmos3SampleInputs(
text_ids=list(text_ids),
vision=Cosmos3VisionItem(latent=vision, condition_frame_indexes=list(cond_vision)),
action=Cosmos3ActionItem(latent=action, condition_frame_indexes=list(cond_action), domain_id=domain_id),
timestep=float(timestep),
)]
return pack_cosmos3_video_sequence(
samples, _SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE, include_end_of_generation_token=False,
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN, reset_spatial_ids=_RESET_SPATIAL_IDS,
enable_fps_modulation=False, base_fps=24.0, temporal_compression_factor=_TCF,
)
def _fv_inputs_with_action(ps) -> dict:
kw = _fastvideo_inputs_from_packed_seq(ps)
a = ps.action
kw.update(
action_tokens=list(a.tokens),
action_token_shapes=[tuple(x) for x in a.token_shapes],
action_sequence_indexes=a.sequence_indexes,
action_timesteps=a.timesteps,
action_mse_loss_indexes=a.mse_loss_indexes,
action_noisy_frame_indexes=list(a.noisy_frame_indexes),
action_domain_id=list(a.domain_id),
)
return kw
def _diffs(a, b):
d = (a - b).abs()
return d.max().item(), d.mean().item()
# (grid_t, lh, lw, action_t, n_text, cond_vision, cond_action, domain_id)
_CASES = [
pytest.param(2, 4, 4, 6, 4, [], [], 0, id="a2v_2x2x2_act6_dom0"),
pytest.param(3, 8, 4, 9, 5, [], [], 7, id="a2v_3x4x2_act9_dom7"),
pytest.param(2, 4, 4, 5, 5, [0], [0], 3, id="ai2v_cond_act5_dom3"),
]
class TestCosmos3ActionParity:
def _build(self, num_layers=2, seed_model=42):
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers, action_gen=True)
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
_copy_weights_with_action(vfm, dit)
return vfm, dit
def _make_inputs(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom, seed=7):
torch.manual_seed(seed)
return dict(
text_ids=torch.randint(0, 60, (n_text,)).tolist(),
vision=torch.randn(1, _LATENT_CHANNEL, grid_t, lh, lw),
action=torch.randn(act_t, _ACTION_DIM), # [T, D]
cond_vision=cond_v, cond_action=cond_a, domain_id=dom, timestep=500.0,
)
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
def test_action_packing_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
ins = self._make_inputs(grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom)
fw = _framework_pack_action(is_image_batch=(grid_t == 1), **ins)
fv = _fastvideo_pack_action(**ins)
assert fv.split_lens == list(fw.split_lens), f"split_lens fv={fv.split_lens} fw={list(fw.split_lens)}"
assert fv.attn_modes == list(fw.attn_modes)
assert int(fv.sequence_length) == int(fw.sequence_length)
torch.testing.assert_close(fv.position_ids, fw.position_ids, rtol=0, atol=0)
a = fw.action
torch.testing.assert_close(fv.action_sequence_indexes, a.sequence_indexes.to(torch.long), rtol=0, atol=0)
assert fv.action_token_shapes == [tuple(x) for x in a.token_shapes]
torch.testing.assert_close(fv.action_timesteps.to(torch.float32), a.timesteps.to(torch.float32))
torch.testing.assert_close(fv.action_mse_loss_indexes, a.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
for x, y in zip(fv.action_noisy_frame_indexes, a.noisy_frame_indexes):
torch.testing.assert_close(x.to(torch.long), y.to(torch.long), rtol=0, atol=0)
print(f"\n[action_packing {grid_t}x{lh}x{lw} act={act_t} dom={dom}] position_ids + action fields exact")
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
def test_action_dit_forward_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
vfm, dit = self._build()
ins = self._make_inputs(grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom)
fw_pack = _framework_pack_action(is_image_batch=(grid_t == 1), **ins)
fv_pack = _fastvideo_pack_action(**ins)
with torch.no_grad():
fw_out = vfm(packed_seq=fw_pack)
fv_out = dit(**fv_pack.to_dit_kwargs())
fv_on_fw = dit(**_fv_inputs_with_action(fw_pack))
pv_mx, pv_mn = _diffs(fv_out["preds_vision"][0], fw_out["preds_vision"][0])
pa_mx, pa_mn = _diffs(fv_out["preds_action"][0], fw_out["preds_action"][0])
paf_mx, paf_mn = _diffs(fv_on_fw["preds_action"][0], fw_out["preds_action"][0])
print(f"\n[action_dit {grid_t}x{lh}x{lw} act={act_t} dom={dom}] "
f"preds_vision max={pv_mx:.3e} mean={pv_mn:.3e} | "
f"preds_action max={pa_mx:.3e} mean={pa_mn:.3e} | "
f"preds_action(fwpack) max={paf_mx:.3e} mean={paf_mn:.3e}")
assert fv_out["preds_action"][0].shape == fw_out["preds_action"][0].shape
torch.testing.assert_close(fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3)
torch.testing.assert_close(fv_out["preds_action"][0], fw_out["preds_action"][0], atol=1e-4, rtol=1e-3)
torch.testing.assert_close(fv_on_fw["preds_action"][0], fw_out["preds_action"][0], atol=1e-4, rtol=1e-3)
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
def test_action_cfg_velocity_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
"""Combined [vision|action] sequential-CFG velocity (action pipeline glue)
matches a framework-DiT oracle."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3ActionSpec,
Cosmos3VisionSpec,
cosmos3_get_cfg_velocity,
)
vfm, dit = self._build()
vlat_shape = (_LATENT_CHANNEL, grid_t, lh, lw)
action_shape = (act_t, _ACTION_DIM)
torch.manual_seed(3)
cond_ids = torch.randint(0, 60, (n_text,)).tolist()
uncond_ids = torch.randint(0, 60, (max(1, n_text - 1),)).tolist()
vis_numel = int(torch.tensor(vlat_shape).prod())
act_numel = int(torch.tensor(action_shape).prod())
flat = torch.randn(vis_numel + act_numel)
guidance, ts = 6.0, 500.0
def _fw_velocity(ids):
vision = flat[:vis_numel].reshape(vlat_shape).unsqueeze(0)
action = flat[vis_numel:].reshape(action_shape)
ps = _framework_pack_action(text_ids=ids, vision=vision, action=action, cond_vision=cond_v,
cond_action=cond_a, domain_id=dom, timestep=ts,
is_image_batch=(grid_t == 1))
with torch.no_grad():
out = vfm(packed_seq=ps)
pv = out["preds_vision"][0].squeeze(0) # [C,T,H,W] (zero on clean)
pa = out["preds_action"][0] # [T,D] (zero on clean)
return torch.cat([pv.reshape(-1), pa.reshape(-1)])
fw_cond, fw_uncond = _fw_velocity(cond_ids), _fw_velocity(uncond_ids)
fw_v = fw_uncond + guidance * (fw_cond - fw_uncond)
fv_v = cosmos3_get_cfg_velocity(
transformer=dit, flat_latent=flat, timestep=torch.tensor([ts]), guidance=guidance,
specs=[Cosmos3VisionSpec(shape=vlat_shape, condition_frame_indexes=list(cond_v))],
action_specs=[Cosmos3ActionSpec(shape=action_shape, condition_frame_indexes=list(cond_a), domain_id=dom)],
cond_token_ids=cond_ids, uncond_token_ids=uncond_ids,
special_tokens=_SPECIAL_TOKENS, latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN, reset_spatial_ids=_RESET_SPATIAL_IDS,
enable_fps_modulation=False, base_fps=24.0, temporal_compression_factor=_TCF,
)
assert fv_v.shape == fw_v.shape, f"shape fv={fv_v.shape} fw={fw_v.shape}"
mx, mn = _diffs(fv_v, fw_v)
print(f"\n[action_cfg_velocity {grid_t}x{lh}x{lw} act={act_t} dom={dom}] max={mx:.3e} mean={mn:.3e}")
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
@@ -0,0 +1,131 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 sound decoder vs the framework AVAE.
The Cosmos3 ``sound_tokenizer`` is an AVAE (audio VAE). Its shipped diffusers
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
(``fastvideo/models/vaes/oobleck.py``). t2vs only needs DECODE (generate sound
latents -> waveform), so this pins the decoder.
The framework decoder
(``cosmos_framework.model.vfm.tokenizers.audio.avae_utils.models.OobleckDecoder``,
``nn.Sequential`` naming, ``output_padding=stride%2`` on the transpose convs) is
the parity ORACLE. We build a tiny framework decoder, map its weights into the
FastVideo decoder (Sequential -> conv1/block.N/res_unitM/snake1/conv2), and
assert bit-exact decode. Strides include an ODD value (5, as in the real config
``[2,4,5,6,8]``) to exercise the ``output_padding`` path that diverged before.
CPU / float32. Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_avae_parity.py -q
"""
from __future__ import annotations
import pytest
import torch
# The official framework provides the parity oracle.
_fw_models = pytest.importorskip(
"cosmos_framework.model.vfm.tokenizers.audio.avae_utils.models",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
from cosmos_framework.model.vfm.tokenizers.audio.avae_utils.env import ( # noqa: E402
AttrDict,
)
from fastvideo.models.vaes.oobleck import OobleckDecoder as FvOobleckDecoder # noqa: E402
pytestmark = [pytest.mark.local]
FwOobleckDecoder = _fw_models.OobleckDecoder
def _framework_decoder(dec_dim, vocoder_input_dim, dec_c_mults, dec_strides):
"""Framework OobleckDecoder (the parity oracle), non-causal / no-antialias."""
h = AttrDict({
"vocoder_input_dim": vocoder_input_dim,
"input_channels": 1,
"stereo": True, # 2 audio channels
"dec_dim": dec_dim,
"dec_c_mults": dec_c_mults,
"dec_strides": dec_strides,
"dec_use_snake": True,
"dec_use_nearest_upsample": False,
"dec_anti_aliasing": False,
"causal": False,
"dec_use_tanh_at_final": False,
"padding_mode": "zeros",
})
return FwOobleckDecoder(h).eval()
def _framework_to_fastvideo_decoder_state(fw_decoder, num_blocks):
"""Map framework Sequential decoder weights -> FastVideo decoder names.
framework: layers.0=first conv; layers.{1..K}=OobleckDecoderBlock
(.layers.0 snake, .1 conv_t, .{2,3,4} ResidualUnit{.layers.0 snake,
.1 conv, .2 snake, .3 conv}); layers.{1+K}=final snake; layers.{2+K}=final conv.
FastVideo: conv1; block.{b}.{snake1,conv_t1,res_unit{1,2,3}.{snake1,conv1,snake2,conv2}};
snake1; conv2. Snake alpha/beta: framework [C] -> FastVideo [1,C,1].
"""
out = {}
for k, v in fw_decoder.state_dict().items():
p = k.split(".")
li = int(p[1])
if li == 0:
nk = "conv1." + ".".join(p[2:])
elif li == 1 + num_blocks:
nk = "snake1." + ".".join(p[2:])
elif li == 2 + num_blocks:
nk = "conv2." + ".".join(p[2:])
else:
b = li - 1
sub = int(p[3])
if sub == 0:
nk = f"block.{b}.snake1." + ".".join(p[4:])
elif sub == 1:
nk = f"block.{b}.conv_t1." + ".".join(p[4:])
else:
r = sub - 2 # ResidualUnit index 0..2
m = {0: "snake1", 1: "conv1", 2: "snake2", 3: "conv2"}[int(p[5])]
nk = f"block.{b}.res_unit{r + 1}.{m}." + ".".join(p[6:])
if nk.endswith(".alpha") or nk.endswith(".beta"):
v = v.reshape(1, -1, 1)
out[nk] = v
return out
# (dec_dim, vocoder_input_dim, dec_c_mults, dec_strides) — tiny; strides incl odd.
_CASES = [
pytest.param(4, 8, [1, 2], [5, 2], id="odd_stride5"),
pytest.param(6, 8, [1, 2, 4], [2, 5, 6], id="real_stride_pattern_tiny"),
pytest.param(4, 4, [1, 2], [4, 8], id="even_strides"),
]
class TestCosmos3AVAEParity:
@pytest.mark.parametrize(("dec_dim", "vin", "cmults", "strides"), _CASES)
def test_decode_matches_framework(self, dec_dim, vin, cmults, strides):
torch.manual_seed(0)
fw = _framework_decoder(dec_dim, vin, cmults, strides)
fv = FvOobleckDecoder(
channels=dec_dim,
input_channels=vin,
audio_channels=2,
upsampling_ratios=list(reversed(strides)), # framework reverses dec_strides
channel_multiples=cmults,
).eval()
state = _framework_to_fastvideo_decoder_state(fw, num_blocks=len(strides))
fv.load_state_dict(state, strict=True) # exact name + shape match
z = torch.randn(1, vin, 5)
with torch.no_grad():
a = fw(z)
b = fv(z)
assert a.shape == b.shape, f"shape: fw={a.shape} fv={b.shape}"
max_abs = (a - b).abs().max().item()
print(f"\n[avae_decode dim={dec_dim} strides={strides}] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(b, a, atol=1e-6, rtol=1e-5)
@@ -0,0 +1,342 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 denoise/CFG glue vs the framework.
The DiT forward and the sequence-packing are already framework-parity-verified
(``test_cosmos3_dit_parity*`` / ``test_cosmos3_packing_parity``). This test pins
the remaining glue that the native pipeline adds — the SEQUENTIAL classifier-free
guidance velocity and one UniPC scheduler step — against the framework math
(``diffusers_cosmos3.pipeline.Cosmos3OmniDiffusersPipeline.get_cfg_velocity`` /
``__call__``):
* for one denoise step, replicate the framework's ``get_cfg_velocity`` exactly
on top of the OFFICIAL ``Cosmos3VFMNetwork`` forward (oracle): a conditional
pass (prompt tokens) and an unconditional pass (negative-prompt tokens),
each masking the prediction on conditioning frames
(``pred * (1 - condition_mask)``), then ``v = uncond + g*(cond - uncond)``;
* run FastVideo's :func:`cosmos3_get_cfg_velocity` with the native DiT (the
framework weights copied in) + the native packer, and assert the velocity
matches the oracle;
* take one ``UniPCMultistepScheduler.step`` on each (the actual checkpoint
scheduler) and assert the stepped latent matches;
* drive :meth:`Cosmos3DenoiseEngine.denoise` for >= 2 steps and assert it
equals the manual framework step-by-step loop.
CPU / float32, via the reference SDPA monkey-patch. The official model is the
parity ORACLE.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_denoise_cfg_parity.py -q
"""
from __future__ import annotations
import pytest
import torch
# The official framework provides the parity oracle.
cosmos_framework = pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
from .test_cosmos3_dit_parity import _copy_weights # noqa: E402
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
_LATENT_CHANNEL,
_LATENT_PATCH_SIZE,
_RESET_SPATIAL_IDS,
_TCF,
_TEMPORAL_MODALITY_MARGIN,
_build_tiny_cosmos3_mrope,
_build_tiny_fastvideo_dit_mrope,
)
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
from .test_cosmos3_scheduler_parity import ( # noqa: E402
_fastvideo_scheduler,
_framework_scheduler,
)
pytestmark = [pytest.mark.local]
_apply_sdpa_patches()
# Tiny special tokens (< tiny vocab_size=64), video path appends eos + sog.
_SPECIAL_TOKENS = {
"start_of_generation": 60,
"end_of_generation": 61,
"eos_token_id": 62,
}
# Cosmos3 video flow_shift; framework scheduler is the parity oracle, FastVideo's
# vendored UniPC (flow config) is the unit under test.
_FLOW_SHIFT = 10.0
# ---------------------------------------------------------------------------
# Framework-oracle CFG velocity (replicates pipeline.get_cfg_velocity math).
# ---------------------------------------------------------------------------
def _framework_pack(*, text_ids, vision_latent, cond_frames, timestep):
from cosmos_framework.data.vfm.sequence_packing import (
GenerationDataClean,
SequencePlan,
pack_input_sequence,
)
# vision_latent is [1, C, T, H, W]; temporal dim is axis 2.
gen_data_clean = GenerationDataClean(
batch_size=1,
is_image_batch=(vision_latent.shape[2] == 1),
x0_tokens_vision=[vision_latent],
fps_vision=None,
num_vision_items_per_sample=[1],
)
plans = [SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=list(cond_frames))]
return pack_input_sequence(
sequence_plans=plans,
input_text_indexes=[list(text_ids)],
gen_data_clean=gen_data_clean,
input_timesteps=torch.tensor([timestep], dtype=torch.float32),
special_tokens=_SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE,
include_end_of_generation_token=False,
position_embedding_type="unified_3d_mrope",
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TCF,
)
def _framework_inputs(ps):
"""Framework PackedSequence -> framework Cosmos3VFMNetwork forward kwargs."""
return dict(packed_seq=ps)
def _framework_cfg_velocity(
*,
vfm,
flat_latent: torch.Tensor,
timestep: torch.Tensor,
guidance: float,
vision_shape: tuple[int, int, int, int],
cond_frames: list[int],
cond_ids: list[int],
uncond_ids: list[int],
) -> torch.Tensor:
"""Replicate the framework ``get_cfg_velocity`` on the oracle model.
Single vision item; sequential cond then uncond pass; mask condition
frames; ``v = uncond + g*(cond - uncond)``.
"""
timestep_value = float(timestep.reshape(()).item())
vision_latent = flat_latent.reshape(vision_shape) # [C, T, H, W]
def _run(text_ids: list[int]) -> torch.Tensor:
ps = _framework_pack(
text_ids=text_ids,
# The framework packer expects a 5D [1, C, T, H, W] latent.
vision_latent=vision_latent.unsqueeze(0),
cond_frames=cond_frames,
timestep=timestep_value,
)
out = vfm(**_framework_inputs(ps))
preds = out.get("preds_vision")
cond_mask = ps.vision.condition_mask[0] # [T] or [T,1,1]
if preds is None:
return torch.zeros_like(flat_latent)
pred = preds[0].squeeze(0) # [C, T, H, W]
keep = (1.0 - cond_mask.reshape(-1, 1, 1)).to(dtype=pred.dtype, device=pred.device)
velocity = pred * keep if keep.sum() > 0 else torch.zeros_like(pred)
return velocity.reshape(-1)
cond_v = _run(cond_ids)
uncond_v = _run(uncond_ids)
return uncond_v + guidance * (cond_v - uncond_v)
# ---------------------------------------------------------------------------
# Cases: T2V (no cond), I2V (cond frame 0), single-frame T2I.
# ---------------------------------------------------------------------------
_CASES = [
pytest.param(2, 4, 4, 6, [], id="t2v_2x2x2"),
pytest.param(3, 8, 4, 5, [0], id="i2v_3x4x2_cond0"),
pytest.param(1, 8, 8, 4, [], id="t2i_1x4x4"),
]
def _build_models(num_layers: int = 2, seed_model: int = 42):
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers)
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
_copy_weights(vfm, dit)
return vfm, dit
def _fastvideo_velocity(dit, *, flat_latent, timestep, guidance, vision_shape, cond_frames, cond_ids, uncond_ids):
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3VisionSpec,
cosmos3_get_cfg_velocity,
)
spec = Cosmos3VisionSpec(shape=vision_shape, condition_frame_indexes=list(cond_frames))
return cosmos3_get_cfg_velocity(
transformer=dit,
flat_latent=flat_latent,
timestep=timestep,
guidance=guidance,
specs=[spec],
cond_token_ids=cond_ids,
uncond_token_ids=uncond_ids,
special_tokens=_SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
reset_spatial_ids=_RESET_SPATIAL_IDS,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TCF,
)
class TestCosmos3DenoiseCFGParity:
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
def test_cfg_velocity_matches_framework(self, grid_t, latent_h, latent_w, n_text, cond):
vfm, dit = _build_models()
torch.manual_seed(0)
cond_ids = torch.randint(0, 60, (n_text,)).tolist()
uncond_ids = torch.randint(0, 60, (max(1, n_text - 1),)).tolist()
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
timestep = torch.tensor([[500.0]]) # framework expects [1,1]; we reshape to scalar
guidance = 6.0
fw_v = _framework_cfg_velocity(
vfm=vfm,
flat_latent=flat_latent,
timestep=timestep,
guidance=guidance,
vision_shape=vision_shape,
cond_frames=cond,
cond_ids=cond_ids,
uncond_ids=uncond_ids,
)
fv_v = _fastvideo_velocity(
dit,
flat_latent=flat_latent,
timestep=timestep,
guidance=guidance,
vision_shape=vision_shape,
cond_frames=cond,
cond_ids=cond_ids,
uncond_ids=uncond_ids,
)
assert fw_v.shape == fv_v.shape, f"shape: fw={fw_v.shape} fv={fv_v.shape}"
max_abs = (fw_v - fv_v).abs().max().item()
print(f"\n[cfg_velocity {grid_t}x{latent_h}x{latent_w} cond={cond}] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
def test_one_unipc_step_matches_framework(self):
"""CFG velocity + one UniPC step: FastVideo == framework math."""
vfm, dit = _build_models()
grid_t, latent_h, latent_w = 2, 4, 4
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
torch.manual_seed(3)
cond_ids = torch.randint(0, 60, (5,)).tolist()
uncond_ids = torch.randint(0, 60, (4,)).tolist()
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
guidance = 6.0
fw_sched = _framework_scheduler(4, _FLOW_SHIFT)
fv_sched = _fastvideo_scheduler(4, _FLOW_SHIFT)
t = fw_sched.timesteps[0]
fw_v = _framework_cfg_velocity(
vfm=vfm,
flat_latent=flat_latent,
timestep=t.reshape(1, 1),
guidance=guidance,
vision_shape=vision_shape,
cond_frames=[],
cond_ids=cond_ids,
uncond_ids=uncond_ids,
)
fw_stepped = fw_sched.step(model_output=fw_v, timestep=t, sample=flat_latent.unsqueeze(0),
return_dict=False)[0].squeeze(0)
fv_v = _fastvideo_velocity(
dit,
flat_latent=flat_latent,
timestep=t.reshape(1),
guidance=guidance,
vision_shape=vision_shape,
cond_frames=[],
cond_ids=cond_ids,
uncond_ids=uncond_ids,
)
fv_stepped = fv_sched.step(model_output=fv_v, timestep=t, sample=flat_latent.unsqueeze(0),
return_dict=False)[0].squeeze(0)
max_abs = (fw_stepped - fv_stepped).abs().max().item()
print(f"\n[unipc_step] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_stepped, fw_stepped, atol=1e-4, rtol=1e-3)
def test_full_denoise_loop_matches_framework(self):
"""Cosmos3DenoiseEngine.denoise (>= 2 steps) == framework step-by-step."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3DenoiseEngine,
Cosmos3VisionSpec,
)
vfm, dit = _build_models()
grid_t, latent_h, latent_w = 2, 4, 4
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
torch.manual_seed(5)
cond_ids = torch.randint(0, 60, (5,)).tolist()
uncond_ids = torch.randint(0, 60, (4,)).tolist()
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
guidance = 6.0
num_steps = 3
# Manual framework loop (oracle).
fw_sched = _framework_scheduler(num_steps, _FLOW_SHIFT)
fw_latent = flat_latent.clone()
for t in fw_sched.timesteps:
v = _framework_cfg_velocity(
vfm=vfm,
flat_latent=fw_latent,
timestep=t.reshape(1, 1),
guidance=guidance,
vision_shape=vision_shape,
cond_frames=[],
cond_ids=cond_ids,
uncond_ids=uncond_ids,
)
fw_latent = fw_sched.step(model_output=v, timestep=t, sample=fw_latent.unsqueeze(0),
return_dict=False)[0].squeeze(0)
# FastVideo engine loop.
fv_sched = _fastvideo_scheduler(num_steps, _FLOW_SHIFT)
engine = Cosmos3DenoiseEngine(
transformer=dit,
scheduler=fv_sched,
special_tokens=_SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
reset_spatial_ids=_RESET_SPATIAL_IDS,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TCF,
)
spec = Cosmos3VisionSpec(shape=vision_shape, condition_frame_indexes=[])
fv_latent = engine.denoise(
flat_latent=flat_latent.clone(),
timesteps=fv_sched.timesteps,
guidance=guidance,
specs=[spec],
cond_token_ids=cond_ids,
uncond_token_ids=uncond_ids,
)
assert fv_latent.shape == fw_latent.shape
max_abs = (fw_latent - fv_latent).abs().max().item()
print(f"\n[full_denoise {num_steps} steps] final latent max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_latent, fw_latent, atol=1e-4, rtol=1e-3)
@@ -0,0 +1,256 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 DiT vs official ``Cosmos3VFMNetwork``.
Builds a tiny official-framework ``Cosmos3VFMNetwork`` AND a tiny FastVideo
``Cosmos3VFMTransformer`` from the SAME tiny config, copies the framework
weights into the FastVideo DiT via an explicit framework->fastvideo name map,
runs BOTH forwards on identical deterministic inputs (CPU / float32), and
asserts ``torch.allclose`` on the vision prediction output (``preds_vision``)
and the per-token ``last_hidden_state``.
The official model is the parity ORACLE. It runs on CPU / float32 via the SDPA
monkey-patch in ``test_cosmos3_reference_forward`` (flash2/flash3/natten are
CUDA-only). The FastVideo DiT runs natively on CPU with plain SDPA.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_dit_parity.py -q
"""
from __future__ import annotations
import pytest
import torch
# The official framework provides the parity oracle.
cosmos_framework = pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
# Reuse the reference harness's tiny-model builder + SDPA monkey-patch.
from .test_cosmos3_reference_forward import ( # noqa: E402
_apply_sdpa_patches,
_build_tiny_cosmos3,
_build_tiny_packed_seq,
)
pytestmark = [pytest.mark.local]
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
_apply_sdpa_patches()
# ---------------------------------------------------------------------------
# Tiny config shared by both models (must match _build_tiny_cosmos3).
# ---------------------------------------------------------------------------
def _build_tiny_fastvideo_dit() -> "Cosmos3VFMTransformer": # noqa: F821
from fastvideo.configs.models.dits.cosmos3 import (
Cosmos3ArchConfig,
Cosmos3VideoConfig,
)
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
arch = Cosmos3ArchConfig(
hidden_size=16,
num_hidden_layers=1,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
intermediate_size=32,
vocab_size=64,
rms_norm_eps=1e-6,
attention_bias=False,
latent_patch_size=2,
latent_channel=16,
rope_theta=5_000_000.0,
mrope_section=[24, 20, 20],
position_embedding_type="3d_rope",
base_fps=24.0,
temporal_compression_factor=4,
enable_fps_modulation=False,
# Dormant heads present in the checkpoint surface (constructed for
# strict-load parity; not exercised by this video-path forward).
action_gen=True,
action_dim=64,
max_action_dim=64,
num_embodiment_domains=32,
sound_gen=True,
sound_dim=64,
)
cfg = Cosmos3VideoConfig(arch_config=arch)
model = Cosmos3VFMTransformer(cfg, hf_config={})
return model.to(torch.float32).eval()
# ---------------------------------------------------------------------------
# Framework -> FastVideo weight name map.
# ---------------------------------------------------------------------------
def _framework_to_fastvideo_state_dict(vfm, num_layers: int) -> dict[str, torch.Tensor]:
"""Translate framework param names into the FastVideo DiT param names.
Framework (Cosmos3VFMNetwork):
language_model.model.{embed_tokens,norm,norm_moe_gen}
language_model.lm_head
language_model.model.layers.{i}.self_attn.{q,k,v,o}_proj(+ _moe_gen)
language_model.model.layers.{i}.self_attn.{q,k}_norm(+ _moe_gen)
language_model.model.layers.{i}.{mlp,mlp_moe_gen}.{gate,up,down}_proj
language_model.model.layers.{i}.{input,post_attention}_layernorm(+ _moe_gen)
vae2llm / llm2vae / time_embedder.mlp.{0,2}
FastVideo (Cosmos3VFMTransformer):
embed_tokens / norm / norm_moe_gen / lm_head
layers.{i}.self_attn.{to_q,to_k,to_v,to_out} (und)
layers.{i}.self_attn.{add_q,add_k,add_v}_proj / to_add_out (gen)
layers.{i}.self_attn.{norm_q,norm_k,norm_added_q,norm_added_k}
layers.{i}.{mlp,mlp_moe_gen}.{gate,up,down}_proj
layers.{i}.{input,post_attention}_layernorm(+ _moe_gen)
proj_in / proj_out / time_embedder.linear_{1,2}
"""
src = dict(vfm.named_parameters())
out: dict[str, torch.Tensor] = {}
def take(name: str) -> torch.Tensor:
return src[name].detach().clone()
# ---- Top-level backbone ----
out["embed_tokens.weight"] = take("language_model.model.embed_tokens.weight")
out["norm.weight"] = take("language_model.model.norm.weight")
out["norm_moe_gen.weight"] = take("language_model.model.norm_moe_gen.weight")
out["lm_head.weight"] = take("language_model.lm_head.weight")
# ---- Vision adapters ----
out["proj_in.weight"] = take("vae2llm.weight")
out["proj_in.bias"] = take("vae2llm.bias")
out["proj_out.weight"] = take("llm2vae.weight")
out["proj_out.bias"] = take("llm2vae.bias")
# ---- Timestep embedder (mlp.0/mlp.2 -> linear_1/linear_2) ----
out["time_embedder.linear_1.weight"] = take("time_embedder.mlp.0.weight")
out["time_embedder.linear_1.bias"] = take("time_embedder.mlp.0.bias")
out["time_embedder.linear_2.weight"] = take("time_embedder.mlp.2.weight")
out["time_embedder.linear_2.bias"] = take("time_embedder.mlp.2.bias")
# ---- Per layer ----
und_attn = {"q_proj": "to_q", "k_proj": "to_k", "v_proj": "to_v", "o_proj": "to_out"}
gen_attn = {
"q_proj_moe_gen": "add_q_proj",
"k_proj_moe_gen": "add_k_proj",
"v_proj_moe_gen": "add_v_proj",
"o_proj_moe_gen": "to_add_out",
}
und_norm = {"q_norm": "norm_q", "k_norm": "norm_k"}
gen_norm = {"q_norm_moe_gen": "norm_added_q", "k_norm_moe_gen": "norm_added_k"}
for i in range(num_layers):
fw = f"language_model.model.layers.{i}"
fv = f"layers.{i}"
for s, d in und_attn.items():
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
for s, d in gen_attn.items():
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
for s, d in und_norm.items():
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
for s, d in gen_norm.items():
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
for mlp in ("mlp", "mlp_moe_gen"):
for proj in ("gate_proj", "up_proj", "down_proj"):
out[f"{fv}.{mlp}.{proj}.weight"] = take(f"{fw}.{mlp}.{proj}.weight")
for ln in ("input_layernorm", "input_layernorm_moe_gen", "post_attention_layernorm",
"post_attention_layernorm_moe_gen"):
out[f"{fv}.{ln}.weight"] = take(f"{fw}.{ln}.weight")
return out
def _copy_weights(vfm, dit) -> None:
"""Copy framework weights into the FastVideo DiT (shape-checked)."""
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
dst = dict(dit.named_parameters())
# Every mapped tensor must land on an existing FastVideo param with a matching shape.
for name, tensor in mapped.items():
assert name in dst, f"FastVideo DiT missing param for mapped key {name!r}"
assert dst[name].shape == tensor.shape, (f"shape mismatch for {name}: "
f"dit={tuple(dst[name].shape)} fw={tuple(tensor.shape)}")
with torch.no_grad():
for name, tensor in mapped.items():
dst[name].copy_(tensor.to(dst[name].dtype))
def _fastvideo_inputs_from_packed_seq(ps) -> dict:
"""Build the FastVideo DiT forward kwargs from a framework PackedSequence."""
v = ps.vision
return dict(
text_ids=ps.text_ids,
text_indexes=ps.text_indexes,
position_ids=ps.position_ids,
sequence_length=int(ps.sequence_length),
split_lens=list(ps.split_lens),
attn_modes=list(ps.attn_modes),
vision_tokens=list(v.tokens),
vision_token_shapes=list(v.token_shapes),
vision_sequence_indexes=v.sequence_indexes,
vision_timesteps=v.timesteps,
vision_mse_loss_indexes=v.mse_loss_indexes,
vision_noisy_frame_indexes=list(v.noisy_frame_indexes),
fps_vision=None,
)
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestCosmos3DiTParity:
def _run_both(self, seed_model: int = 42, seed_data: int = 7):
vfm = _build_tiny_cosmos3(seed=seed_model)
dit = _build_tiny_fastvideo_dit()
_copy_weights(vfm, dit)
ps = _build_tiny_packed_seq(n_text=4, seed=seed_data)
with torch.no_grad():
fw_out = vfm(packed_seq=ps)
fv_out = dit(**_fastvideo_inputs_from_packed_seq(ps))
return fw_out, fv_out
def test_weight_map_is_complete(self):
"""The framework->fastvideo map must cover EVERY FastVideo DiT parameter
that is exercised by the video path (i.e. all non-dormant params).
Dormant action/audio heads have no framework counterpart in this tiny
vision-only setup, so they are excluded from the copy; everything else
must be covered.
"""
vfm = _build_tiny_cosmos3(seed=42)
dit = _build_tiny_fastvideo_dit()
mapped = set(_framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers))
dit_params = set(n for n, _ in dit.named_parameters())
dormant = {
n
for n in dit_params
if n.startswith(("action_", "audio_"))
}
uncovered = dit_params - mapped - dormant
assert not uncovered, f"FastVideo DiT params not covered by weight map: {sorted(uncovered)}"
def test_preds_vision_parity(self):
fw_out, fv_out = self._run_both()
fw_pv = fw_out["preds_vision"][0] # [1, C, T, H, W]
fv_pv = fv_out["preds_vision"][0]
assert fw_pv.shape == fv_pv.shape, f"shape mismatch: fw={fw_pv.shape} fv={fv_pv.shape}"
max_abs = (fw_pv - fv_pv).abs().max().item()
print(f"\n[preds_vision] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_pv, fw_pv, atol=1e-4, rtol=1e-3)
def test_last_hidden_state_parity(self):
fw_out, fv_out = self._run_both()
fw_lhs = fw_out["last_hidden_state"] # [N, hidden]
fv_lhs = fv_out["last_hidden_state"]
assert fw_lhs.shape == fv_lhs.shape, f"shape mismatch: fw={fw_lhs.shape} fv={fv_lhs.shape}"
max_abs = (fw_lhs - fv_lhs).abs().max().item()
print(f"\n[last_hidden_state] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_lhs, fw_lhs, atol=1e-4, rtol=1e-3)
def test_parity_holds_across_seeds(self):
"""Re-running with a different random init still matches (not a fluke)."""
fw_out, fv_out = self._run_both(seed_model=99, seed_data=13)
torch.testing.assert_close(fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3)
@@ -0,0 +1,371 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 DiT vs ``Cosmos3VFMNetwork`` (mRoPE).
Companion to ``test_cosmos3_dit_parity.py`` (which covers ``3d_rope``). This
module exercises the rotary mode the REAL ``nvidia/Cosmos3-Nano`` checkpoint
uses: ``position_embedding_type="unified_3d_mrope"`` with the real-checkpoint
settings (``mrope_section=[24,20,20]``, ``mrope_interleaved=True``,
``rope_theta=5e6``, ``unified_3d_mrope_reset_spatial_ids=True``,
``temporal_modality_margin=15000``).
Under unified 3D mRoPE there is NO additive latent position embedding
(``latent_pos_embed is None``); all positional information rides on the
per-token 3D (T, H, W) rotary embedding. The packed-sequence ``position_ids``
are therefore shape ``[3, seq_len]``, built exactly like the framework data
packer (``cosmos_framework.data.vfm.sequence_packing``):
* text tokens broadcast one monotone id across all three axes
(``get_3d_mrope_ids_text_tokens``),
* the temporal offset is bumped by ``temporal_modality_margin`` at the
text->vision boundary,
* vision tokens lay out a (T, H, W) grid with spatial ids reset per segment
(``get_3d_mrope_ids_vae_tokens`` with ``reset_spatial_indices=True``).
Both models are built tiny from the SAME config, framework weights are copied
into the FastVideo DiT (reusing the ``3d_rope`` test's weight map — the
transformer key surface is identical across rotary modes), and BOTH forwards
run on identical deterministic CPU / float32 inputs. The official model is the
parity ORACLE (run on CPU via the SDPA monkey-patch in
``test_cosmos3_reference_forward``).
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_dit_parity_mrope.py -q
"""
from __future__ import annotations
import pytest
import torch
# The official framework provides the parity oracle.
cosmos_framework = pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
# Reuse the reference harness's SDPA monkey-patch and the 3d_rope parity
# test's weight-copy + input-builder helpers (key surface is rotary-agnostic).
from .test_cosmos3_dit_parity import ( # noqa: E402
_copy_weights,
_fastvideo_inputs_from_packed_seq,
_framework_to_fastvideo_state_dict,
)
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
pytestmark = [pytest.mark.local]
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
_apply_sdpa_patches()
# Real-checkpoint unified_3d_mrope settings (tiny model, real rope constants).
_ROPE_THETA = 5_000_000.0
_MROPE_SECTION = [24, 20, 20]
_MROPE_INTERLEAVED = True
_RESET_SPATIAL_IDS = True
_TEMPORAL_MODALITY_MARGIN = 15_000
_LATENT_PATCH_SIZE = 2
_LATENT_CHANNEL = 16
_TCF = 4 # temporal compression factor
# ---------------------------------------------------------------------------
# Tiny model builders (framework + FastVideo) with unified_3d_mrope.
# ---------------------------------------------------------------------------
_SOUND_DIM = 64
_SOUND_LATENT_FPS = 25
_ACTION_DIM = 64
_NUM_EMBODIMENT_DOMAINS = 32
def _build_tiny_cosmos3_mrope(seed: int = 42, num_layers: int = 2, sound_gen: bool = False,
action_gen: bool = False):
"""Tiny framework ``Cosmos3VFMNetwork`` with ``unified_3d_mrope``.
``rope_theta`` / ``rope_scaling`` (carrying ``mrope_section`` +
``mrope_interleaved``) are threaded through the materialized text config;
``position_embedding_type="unified_3d_mrope"`` leaves ``latent_pos_embed``
as ``None`` so positions ride solely on the 3D rotary embedding.
``sound_gen=True`` additionally builds the sound MoT heads (``sound2llm`` /
``llm2sound`` / ``sound_modality_embed``) for the t2vs parity test.
"""
from cosmos_framework.model.vfm.mot.cosmos3_vfm_network import (
Cosmos3VFMNetwork,
Cosmos3VFMNetworkConfig,
)
from cosmos_framework.model.vfm.mot.unified_mot import (
Qwen3MoTConfig,
Qwen3VLTextForCausalLM,
)
from cosmos_framework.model.vfm.vlm.qwen3_vl.configuration_qwen3_vl import Qwen3VLConfig
tiny_text_dict = dict(
model_type="qwen3_vl_text",
vocab_size=64,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=num_layers,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
rms_norm_eps=1e-6,
attention_bias=False,
attention_dropout=0.0,
rope_theta=_ROPE_THETA,
rope_scaling={
"rope_type": "default",
"mrope_section": _MROPE_SECTION,
"mrope_interleaved": _MROPE_INTERLEAVED,
},
max_position_embeddings=262144,
)
mot_cfg = Qwen3MoTConfig(
config_dict=tiny_text_dict,
qk_norm_for_text=True,
qk_norm_for_diffusion=True,
include_visual=False,
)
tiny_vlm_cfg = Qwen3VLConfig(text_config=tiny_text_dict)
sound_kwargs = dict(
sound_gen=True,
sound_dim=_SOUND_DIM,
temporal_compression_factor_sound=1,
sound_latent_fps=_SOUND_LATENT_FPS,
) if sound_gen else {}
action_kwargs = dict(
action_gen=True,
action_dim=_ACTION_DIM,
num_embodiment_domains=_NUM_EMBODIMENT_DOMAINS,
) if action_gen else {}
vfm_cfg = Cosmos3VFMNetworkConfig(
vision_gen=True,
vlm_config=tiny_vlm_cfg,
latent_patch_size=_LATENT_PATCH_SIZE,
latent_downsample_factor=8,
latent_channel_size=_LATENT_CHANNEL,
position_embedding_type="unified_3d_mrope",
max_latent_h=16,
max_latent_w=16,
max_latent_t=8,
temporal_compression_factor_vision=_TCF,
**sound_kwargs,
**action_kwargs,
)
torch.manual_seed(seed)
lm = Qwen3VLTextForCausalLM(config=mot_cfg)
vfm = Cosmos3VFMNetwork(language_model=lm, config=vfm_cfg)
# inv_freq is a non-persistent buffer; init it on CPU (mirrors from_pretrained).
vfm.language_model.model.rotary_emb.init_weights(buffer_device=None)
vfm.eval()
return vfm
def _build_tiny_fastvideo_dit_mrope(num_layers: int = 2):
from fastvideo.configs.models.dits.cosmos3 import (
Cosmos3ArchConfig,
Cosmos3VideoConfig,
)
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
arch = Cosmos3ArchConfig(
hidden_size=16,
num_hidden_layers=num_layers,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
intermediate_size=32,
vocab_size=64,
rms_norm_eps=1e-6,
attention_bias=False,
latent_patch_size=_LATENT_PATCH_SIZE,
latent_channel=_LATENT_CHANNEL,
rope_theta=_ROPE_THETA,
mrope_section=_MROPE_SECTION,
mrope_interleaved=_MROPE_INTERLEAVED,
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
position_embedding_type="unified_3d_mrope",
base_fps=24.0,
temporal_compression_factor=_TCF,
enable_fps_modulation=False,
# Dormant heads present in the checkpoint surface (constructed for
# strict-load parity; not exercised by this video-path forward).
action_gen=True,
action_dim=64,
max_action_dim=64,
num_embodiment_domains=32,
sound_gen=True,
sound_dim=64,
)
cfg = Cosmos3VideoConfig(arch_config=arch)
model = Cosmos3VFMTransformer(cfg, hf_config={})
return model.to(torch.float32).eval()
# ---------------------------------------------------------------------------
# [3, seq_len] mRoPE position-id builder (mirrors the framework data packer).
# ---------------------------------------------------------------------------
def _build_mrope_position_ids(n_text: int, grid_t: int, patch_h: int, patch_w: int) -> torch.Tensor:
"""Build ``[3, seq_len]`` (T, H, W) mRoPE ids for one text+vision sample.
Reproduces ``pack_input_sequence`` for a single causal-text + full-vision
sample: monotone text ids on all axes, ``+temporal_modality_margin`` at the
text->vision boundary, then a reset-spatial (T, H, W) vision grid.
"""
from cosmos_framework.data.vfm.sequence_packing import (
get_3d_mrope_ids_text_tokens,
get_3d_mrope_ids_vae_tokens,
)
offset: int | float = 0
text_ids, offset = get_3d_mrope_ids_text_tokens(num_tokens=n_text, temporal_offset=offset)
# End of text modality: add the boundary margin before vision.
offset += _TEMPORAL_MODALITY_MARGIN
vision_ids, offset = get_3d_mrope_ids_vae_tokens(
grid_t=grid_t,
grid_h=patch_h,
grid_w=patch_w,
temporal_offset=offset,
reset_spatial_indices=_RESET_SPATIAL_IDS,
fps=None, # integer positions (enable_fps_modulation=False)
temporal_compression_factor=_TCF,
)
return torch.cat([text_ids, vision_ids], dim=1) # [3, seq_len]
def _build_tiny_packed_seq_mrope(
*,
n_text: int = 6,
grid_t: int = 2,
latent_h: int = 4,
latent_w: int = 4,
seed: int = 7,
):
"""Minimal PackedSequence with ``[3, seq]`` mRoPE position ids.
Vision latent ``[C, grid_t, latent_h, latent_w]`` patchifies (patch=2) to a
``(grid_t, latent_h/2, latent_w/2)`` token grid; all frames are noisy.
"""
from cosmos_framework.data.vfm.sequence_packing import ModalityData, PackedSequence
patch_h = latent_h // _LATENT_PATCH_SIZE
patch_w = latent_w // _LATENT_PATCH_SIZE
n_vision = grid_t * patch_h * patch_w
total_len = n_text + n_vision
torch.manual_seed(seed)
vision_tensor = torch.randn(_LATENT_CHANNEL, grid_t, latent_h, latent_w)
text_ids = torch.randint(0, 64, (n_text,))
position_ids = _build_mrope_position_ids(n_text, grid_t, patch_h, patch_w) # [3, total_len]
noisy_frame_indexes = torch.arange(grid_t, dtype=torch.long) # all frames noisy
vision_mod = ModalityData(
sequence_indexes=torch.arange(n_text, total_len, dtype=torch.long),
timesteps=torch.full((n_vision,), 500.0),
mse_loss_indexes=torch.arange(n_text, total_len, dtype=torch.long),
token_shapes=[(grid_t, patch_h, patch_w)],
tokens=[vision_tensor],
condition_mask=[torch.zeros(grid_t, dtype=torch.long)], # 0 = noisy
noisy_frame_indexes=[noisy_frame_indexes],
)
packed_seq = PackedSequence(
sample_lens=[total_len],
split_lens=[n_text, n_vision],
attn_modes=["causal", "full"],
is_image_batch=(grid_t == 1),
sequence_length=total_len,
text_ids=text_ids,
text_indexes=torch.arange(n_text, dtype=torch.long),
position_ids=position_ids,
vision=vision_mod,
)
return packed_seq
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
# (grid_t, latent_h, latent_w): a single image, a small video, and a taller
# video, to exercise the spatial mRoPE overwrite + gen<->gen full attention.
_GRIDS = [
pytest.param(1, 8, 8, id="image_1x4x4"),
pytest.param(2, 4, 4, id="video_2x2x2"),
pytest.param(3, 8, 4, id="video_3x4x2"),
]
class TestCosmos3DiTParityMRoPE:
def _run_both(
self,
*,
grid_t: int,
latent_h: int,
latent_w: int,
seed_model: int = 42,
seed_data: int = 7,
num_layers: int = 2,
):
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers)
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
_copy_weights(vfm, dit)
ps = _build_tiny_packed_seq_mrope(
n_text=6, grid_t=grid_t, latent_h=latent_h, latent_w=latent_w, seed=seed_data
)
with torch.no_grad():
fw_out = vfm(packed_seq=ps)
fv_out = dit(**_fastvideo_inputs_from_packed_seq(ps))
return fw_out, fv_out
def test_position_ids_are_3xN_mrope(self):
"""The packed mRoPE ids are ``[3, seq_len]`` with the text->vision margin."""
ps = _build_tiny_packed_seq_mrope(n_text=6, grid_t=2, latent_h=4, latent_w=4)
pos = ps.position_ids
assert pos.ndim == 2 and pos.shape[0] == 3, f"expected [3, N], got {tuple(pos.shape)}"
assert pos.shape[1] == int(ps.sequence_length)
# Text axis is monotone 0..5 on all 3 rows; vision temporal jumps by the margin.
assert pos[0, :6].tolist() == [0, 1, 2, 3, 4, 5]
assert pos[1, :6].tolist() == [0, 1, 2, 3, 4, 5]
assert pos[2, :6].tolist() == [0, 1, 2, 3, 4, 5]
# First vision token temporal id == last_text_id (5) + margin + 1.
assert pos[0, 6].item() == 5 + _TEMPORAL_MODALITY_MARGIN + 1
# Reset spatial: first vision token H/W ids are 0.
assert pos[1, 6].item() == 0 and pos[2, 6].item() == 0
def test_no_additive_latent_pos_embed(self):
"""unified_3d_mrope must NOT build an additive latent position embedding."""
dit = _build_tiny_fastvideo_dit_mrope()
assert dit.position_embedding_type == "unified_3d_mrope"
assert dit.latent_pos_embed is None
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w"), _GRIDS)
def test_preds_vision_parity(self, grid_t, latent_h, latent_w):
fw_out, fv_out = self._run_both(grid_t=grid_t, latent_h=latent_h, latent_w=latent_w)
fw_pv = fw_out["preds_vision"][0] # [1, C, T, H, W]
fv_pv = fv_out["preds_vision"][0]
assert fw_pv.shape == fv_pv.shape, f"shape mismatch: fw={fw_pv.shape} fv={fv_pv.shape}"
max_abs = (fw_pv - fv_pv).abs().max().item()
print(f"\n[preds_vision mrope {grid_t}x{latent_h}x{latent_w}] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_pv, fw_pv, atol=1e-4, rtol=1e-3)
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w"), _GRIDS)
def test_last_hidden_state_parity(self, grid_t, latent_h, latent_w):
fw_out, fv_out = self._run_both(grid_t=grid_t, latent_h=latent_h, latent_w=latent_w)
fw_lhs = fw_out["last_hidden_state"] # [N, hidden]
fv_lhs = fv_out["last_hidden_state"]
assert fw_lhs.shape == fv_lhs.shape, f"shape mismatch: fw={fw_lhs.shape} fv={fv_lhs.shape}"
max_abs = (fw_lhs - fv_lhs).abs().max().item()
print(f"\n[last_hidden_state mrope {grid_t}x{latent_h}x{latent_w}] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(fv_lhs, fw_lhs, atol=1e-4, rtol=1e-3)
def test_parity_holds_across_seeds(self):
"""A different random init still matches bit-for-bit (not a fluke)."""
fw_out, fv_out = self._run_both(
grid_t=2, latent_h=4, latent_w=4, seed_model=99, seed_data=13
)
torch.testing.assert_close(
fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3
)
torch.testing.assert_close(
fv_out["last_hidden_state"], fw_out["last_hidden_state"], atol=1e-4, rtol=1e-3
)
@@ -0,0 +1,80 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 UniPC flow_shift vs the framework.
The framework selects the UniPC ``shift`` purely from the named resolution
bucket the (H, W) belongs to, via ``OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS``
(keyed by the VLM model size — Cosmos3-Nano uses the 8B backbone — and the
resolution string), NOT from the task (T2V/I2V/T2I share a shift at a given
resolution). FastVideo gets raw pixel ``height``/``width`` and must map back to
the same shift.
This pins ``Cosmos3DenoisingStage._flow_shift_for_resolution`` against the
framework's own tables: for every (resolution, aspect) entry in
``VIDEO_RES_SIZE_INFO`` whose resolution has an 8B shift default, the FastVideo
shift for that exact pixel size must equal the framework default.
The framework tables are the parity ORACLE.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_flow_shift_parity.py -q
"""
from __future__ import annotations
import pytest
# The official framework provides the parity oracle for the resolution->pixel
# tables. (``cosmos_framework.inference.args`` — which holds the shift constant —
# can't be imported here: it transitively requires ``multistorageclient``. The
# small shift table is mirrored verbatim below with its source location.)
_utils = pytest.importorskip(
"cosmos_framework.data.vfm.utils",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
from fastvideo.pipelines.stages.cosmos3_stages import ( # noqa: E402
Cosmos3DenoisingStage,
)
pytestmark = [pytest.mark.local]
# Cosmos3-Nano's VLM backbone is Qwen3-VL-8B (checkpoint config.json).
_MODEL_SIZE = "8B"
# Verbatim from cosmos_framework.inference.args.OmniSampleArgs
# ._RESOLUTION_SHIFT_DEFAULTS (args.py:770), restricted to the 8B rows.
_SHIFT_DEFAULTS = {
("8B", "256"): 3.0,
("8B", "480"): 5.0,
("8B", "720"): 10.0,
("8B", "768"): 10.0,
("32B", "256"): 5.0,
("32B", "480"): 5.0,
("32B", "720"): 5.0,
("32B", "768"): 5.0,
}
_VIDEO_RES = _utils.VIDEO_RES_SIZE_INFO
_IMAGE_RES = _utils.IMAGE_RES_SIZE_INFO
def _cases():
seen = set()
for resolution, by_aspect in {**_VIDEO_RES, **_IMAGE_RES}.items():
key = (_MODEL_SIZE, resolution)
if key not in _SHIFT_DEFAULTS:
continue
expected = _SHIFT_DEFAULTS[key]
for aspect, (a, b) in by_aspect.items():
cid = f"{resolution}_{aspect.replace(',', '-')}_{a}x{b}"
if cid in seen:
continue
seen.add(cid)
yield pytest.param(a, b, expected, id=cid)
class TestCosmos3FlowShiftParity:
@pytest.mark.parametrize(("dim_a", "dim_b", "expected_shift"), list(_cases()))
def test_flow_shift_matches_framework(self, dim_a, dim_b, expected_shift):
got = Cosmos3DenoisingStage._flow_shift_for_resolution(dim_a, dim_b)
assert got == expected_shift, (
f"shift for {dim_a}x{dim_b}: got {got}, framework default {expected_shift}")
@@ -0,0 +1,91 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo I2V conditioning pixel video vs the framework.
The Cosmos3 I2V path conditions on a *static repeat* of the input image. The
framework (``cosmos_framework.inference.vision``):
* ``load_conditioning_image``: aspect-preserving resize + center crop + uint8
quantization, then ``/127.5 - 1`` -> ``[3, 1, h, w]`` in [-1, 1];
* ``build_conditioned_video_batch``: frame 0 = the image, and every remaining
frame **repeats the last conditioning frame** (a static video) -> the clip
is then VAE-encoded and only the latent condition frame(s) are kept clean.
Because the VAE is temporal, zero-filling the non-condition frames (the earlier
FastVideo behavior) changes the encoded condition latent, so the repeat-fill is
correctness-critical. This pins FastVideo's
``Cosmos3DenoisingStage._image_to_video_tensor`` against the framework's
image preprocessing + repeat-fill.
CPU / float32. The framework is the parity ORACLE.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_i2v_conditioning_parity.py -q
"""
from __future__ import annotations
import numpy as np
import pytest
import torch
from PIL import Image
# The official framework provides the parity oracle.
vision = pytest.importorskip(
"cosmos_framework.inference.vision",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
from fastvideo.pipelines.stages.cosmos3_stages import ( # noqa: E402
Cosmos3DenoisingStage,
)
pytestmark = [pytest.mark.local]
def _make_image(path, h_in: int, w_in: int, seed: int = 0) -> None:
rng = np.random.default_rng(seed)
arr = rng.integers(0, 256, size=(h_in, w_in, 3), dtype=np.uint8)
Image.fromarray(arr, "RGB").save(path)
# (input H, input W, target H, target W, num_frames)
_CASES = [
pytest.param(120, 200, 256, 256, 9, id="square_from_landscape"),
pytest.param(200, 120, 704, 1280, 13, id="wide_from_portrait"),
pytest.param(256, 256, 256, 256, 5, id="same_size"),
]
class TestCosmos3I2VConditioningParity:
@pytest.mark.parametrize(("h_in", "w_in", "h", "w", "num_frames"), _CASES)
def test_conditioning_video_matches_framework(self, tmp_path, h_in, w_in, h, w, num_frames):
img_path = tmp_path / "cond.png"
_make_image(img_path, h_in, w_in)
# ---- Framework oracle ----
# load_conditioning_image -> [3, 1, h, w] in [-1, 1].
cond = vision.load_conditioning_image(img_path, target_h=h, target_w=w).float()
# Mirror build_conditioned_video_batch (vision.py lines 117-123) in fp32/CPU:
# frame 0 = image; remaining frames repeat the last conditioning frame.
t_cond = cond.shape[1]
expected = torch.zeros(1, 3, num_frames, h, w, dtype=torch.float32)
t_fill = min(t_cond, num_frames)
expected[0, :, :t_fill] = cond[:, :t_fill]
if t_fill < num_frames:
expected[0, :, t_fill:] = expected[0, :, t_fill - 1:t_fill].expand(-1, num_frames - t_fill, -1, -1)
# ---- FastVideo: same PIL image through the stage helper ----
pil = Image.open(img_path).convert("RGB")
got = Cosmos3DenoisingStage._image_to_video_tensor(
pil, num_frames, h, w, torch.device("cpu"), torch.float32)
assert got.shape == expected.shape, f"shape: got={got.shape} expected={expected.shape}"
max_abs = (got - expected).abs().max().item()
print(f"\n[i2v_cond {h}x{w} nf={num_frames}] max abs diff = {max_abs:.3e}")
torch.testing.assert_close(got, expected)
# Static repeat (not zero-fill): every frame equals frame 0, and the
# frames past frame 0 are non-zero.
assert torch.equal(got[0, :, 0], got[0, :, -1]), "non-condition frames must repeat the image"
assert got[0, :, 1:].abs().sum() > 0, "non-condition frames must not be zero-filled"
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 unified 3D mRoPE position-ID parity (Tier A scaffold).
Reference: ``vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py``
lines 113-177 (``compute_mrope_position_ids_text`` /
``compute_mrope_position_ids_vision``). The reference test asserting
these invariants lives at
``tests/diffusion/models/cosmos3/test_cosmos3_transformer.py:32-57``.
The three invariants under test:
1. Text tokens broadcast the same monotonically-increasing positions
across all three (t, h, w) axes. With ``num_tokens=3`` and
``temporal_offset=5`` the result is ``[[5,6,7], [5,6,7], [5,6,7]]``
and the next-offset is ``8``.
2. Vision tokens (no FPS modulation) flatten a ``(grid_t, grid_h, grid_w)``
position grid in t-major order. With ``(2, 2, 3)`` and offset ``10``
the resulting shape is ``(3, 12)`` and the temporal row begins
``[10]*6 + [11]*6``; next-offset is ``12``.
3. FPS-modulated vision tokens scale the temporal axis by
``base_fps / temporal_compression_factor / (fps / tcf)``. With
``fps=12``, ``base_fps=24``, ``tcf=4``, ``grid_t=2`` the first row is
``[10.0, 12.0]``.
The FastVideo side currently does NOT exist; the test is wrapped in
``try/except ImportError`` and skips. Phase 2b replaces the skip with
the real import + assertion path.
"""
from __future__ import annotations
import pytest
import torch
pytestmark = [pytest.mark.local]
def test_compute_mrope_position_ids_text_and_vision() -> None:
"""Asserts the 3 invariants of unified 3D mRoPE position-ID generation.
Once FastVideo's ``fastvideo.models.dits.cosmos3`` exports
``compute_mrope_position_ids_text`` and
``compute_mrope_position_ids_vision``, this test verifies they produce
output tensors identical to the vllm-omni reference at
transformer_cosmos3.py:113-177.
"""
try:
from fastvideo.models.dits.cosmos3 import ( # type: ignore
compute_mrope_position_ids_text,
compute_mrope_position_ids_vision,
)
except ImportError:
pytest.skip("FastVideo Cosmos3 not yet implemented (Phase 2b)")
text_ids, text_offset = compute_mrope_position_ids_text(num_tokens=3, temporal_offset=5)
assert text_ids.tolist() == [[5, 6, 7], [5, 6, 7], [5, 6, 7]]
assert text_offset == 8
vision_ids, vision_offset = compute_mrope_position_ids_vision(
2, 2, 3, temporal_offset=10, fps=None
)
assert tuple(vision_ids.shape) == (3, 12)
assert vision_ids[0].tolist() == [10] * 6 + [11] * 6
assert vision_offset == 12
modulated_ids, modulated_offset = compute_mrope_position_ids_vision(
2,
1,
1,
temporal_offset=10,
fps=12.0,
base_fps=24.0,
temporal_compression_factor=4,
)
torch.testing.assert_close(modulated_ids[0], torch.tensor([10.0, 12.0]))
assert modulated_offset == 13
@@ -0,0 +1,312 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 sequence packing vs the OFFICIAL framework.
FastVideo's native packer
(``fastvideo.pipelines.basic.cosmos3.sequence_packing.pack_cosmos3_video_sequence``)
builds the packed-sequence inputs the ``Cosmos3VFMTransformer`` consumes. This
test asserts, for the SAME logical inputs (prompt token ids, vision latents,
condition-frame indices, diffusion timestep, fps), that FastVideo's packing
matches the official ``cosmos_framework.data.vfm.sequence_packing.pack_input_sequence``
oracle field-by-field:
* ``position_ids`` (exact, ``[3, seq]``),
* ``text_ids`` / ``text_indexes``,
* ``split_lens`` / ``attn_modes`` / ``sample_lens`` / ``sequence_length``,
* vision ``sequence_indexes`` / ``token_shapes`` / ``timesteps`` /
``mse_loss_indexes`` / ``noisy_frame_indexes`` / ``condition_mask``.
Coverage spans T2V (no condition frames), I2V (condition frame 0), and T2I
(single conditioned frame), across multiple grids, plus a multi-sample batch.
Then BOTH the framework-packed and FastVideo-packed inputs are fed through the
SAME tiny FastVideo DiT (framework weights copied in as in the existing DiT
parity tests). Asserting bit-identical DiT output confirms FastVideo's own
packing drives the DiT to the same result as the framework's packing.
The official framework is the parity ORACLE; it runs on CPU / float32 via the
SDPA monkey-patch from ``test_cosmos3_reference_forward``.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_packing_parity.py -q
"""
from __future__ import annotations
import pytest
import torch
# The official framework provides the parity oracle.
cosmos_framework = pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
# Reuse the DiT parity helpers (weight copy + framework->DiT kwarg builder) and
# the mRoPE tiny-model builders (real-checkpoint rope constants).
from .test_cosmos3_dit_parity import ( # noqa: E402
_copy_weights,
_fastvideo_inputs_from_packed_seq,
)
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
_LATENT_CHANNEL,
_LATENT_PATCH_SIZE,
_MROPE_SECTION,
_RESET_SPATIAL_IDS,
_ROPE_THETA,
_TCF,
_TEMPORAL_MODALITY_MARGIN,
_build_tiny_cosmos3_mrope,
_build_tiny_fastvideo_dit_mrope,
)
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
pytestmark = [pytest.mark.local]
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
_apply_sdpa_patches()
# Tiny special-token ids (kept < tiny vocab_size=64). The video path appends
# eos + start_of_generation after the prompt tokens.
_SPECIAL_TOKENS = {
"start_of_generation": 60,
"end_of_generation": 61,
"eos_token_id": 62,
}
# ---------------------------------------------------------------------------
# Builders for the two packers from the SAME logical sample inputs.
# ---------------------------------------------------------------------------
def _make_vision(grid_t: int, latent_h: int, latent_w: int, seed: int) -> torch.Tensor:
"""Deterministic VAE latent ``[1, C, T, H, W]``."""
torch.manual_seed(seed)
return torch.randn(1, _LATENT_CHANNEL, grid_t, latent_h, latent_w)
def _framework_pack(
*,
text_ids_per_sample: list[list[int]],
visions: list[torch.Tensor],
cond_frames_per_sample: list[list[int]],
timesteps: list[float],
is_image_batch: bool,
):
from cosmos_framework.data.vfm.sequence_packing import (
GenerationDataClean,
SequencePlan,
pack_input_sequence,
)
gen_data_clean = GenerationDataClean(
batch_size=len(visions),
is_image_batch=is_image_batch,
x0_tokens_vision=list(visions),
fps_vision=None,
num_vision_items_per_sample=[1] * len(visions),
)
plans = [
SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=list(cf))
for cf in cond_frames_per_sample
]
return pack_input_sequence(
sequence_plans=plans,
input_text_indexes=[list(t) for t in text_ids_per_sample],
gen_data_clean=gen_data_clean,
input_timesteps=torch.tensor(timesteps, dtype=torch.float32),
special_tokens=_SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE,
include_end_of_generation_token=False,
position_embedding_type="unified_3d_mrope",
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TCF,
)
def _fastvideo_pack(
*,
text_ids_per_sample: list[list[int]],
visions: list[torch.Tensor],
cond_frames_per_sample: list[list[int]],
timesteps: list[float],
):
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
Cosmos3SampleInputs,
Cosmos3VisionItem,
pack_cosmos3_video_sequence,
)
samples = [
Cosmos3SampleInputs(
text_ids=list(t),
vision=Cosmos3VisionItem(latent=v, condition_frame_indexes=list(cf)),
timestep=float(ts),
)
for t, v, cf, ts in zip(text_ids_per_sample, visions, cond_frames_per_sample, timesteps)
]
return pack_cosmos3_video_sequence(
samples,
_SPECIAL_TOKENS,
latent_patch_size=_LATENT_PATCH_SIZE,
include_end_of_generation_token=False,
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
reset_spatial_ids=_RESET_SPATIAL_IDS,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TCF,
)
# ---------------------------------------------------------------------------
# Field-by-field comparison.
# ---------------------------------------------------------------------------
def _assert_packs_match(fw, fv) -> None:
"""Assert the framework PackedSequence and FastVideo pack agree field-by-field."""
# Structure.
assert fv.split_lens == list(fw.split_lens), f"split_lens: fv={fv.split_lens} fw={list(fw.split_lens)}"
assert fv.attn_modes == list(fw.attn_modes), f"attn_modes: fv={fv.attn_modes} fw={list(fw.attn_modes)}"
assert fv.sample_lens == list(fw.sample_lens), f"sample_lens: fv={fv.sample_lens} fw={list(fw.sample_lens)}"
assert int(fv.sequence_length) == int(fw.sequence_length)
# Text.
torch.testing.assert_close(fv.text_ids, fw.text_ids.to(torch.long), rtol=0, atol=0)
torch.testing.assert_close(fv.text_indexes, fw.text_indexes.to(torch.long), rtol=0, atol=0)
# position_ids: exact, [3, seq], same dtype.
assert fv.position_ids.shape == fw.position_ids.shape, (
f"position_ids shape: fv={tuple(fv.position_ids.shape)} fw={tuple(fw.position_ids.shape)}")
assert fv.position_ids.dtype == fw.position_ids.dtype, (
f"position_ids dtype: fv={fv.position_ids.dtype} fw={fw.position_ids.dtype}")
torch.testing.assert_close(fv.position_ids, fw.position_ids, rtol=0, atol=0)
# Vision.
fwv = fw.vision
torch.testing.assert_close(fv.vision_sequence_indexes, fwv.sequence_indexes.to(torch.long), rtol=0, atol=0)
assert fv.vision_token_shapes == [tuple(s) for s in fwv.token_shapes], (
f"token_shapes: fv={fv.vision_token_shapes} fw={[tuple(s) for s in fwv.token_shapes]}")
torch.testing.assert_close(fv.vision_timesteps.to(torch.float32), fwv.timesteps.to(torch.float32))
torch.testing.assert_close(fv.vision_mse_loss_indexes, fwv.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
assert len(fv.vision_noisy_frame_indexes) == len(fwv.noisy_frame_indexes)
for a, b in zip(fv.vision_noisy_frame_indexes, fwv.noisy_frame_indexes):
torch.testing.assert_close(a.to(torch.long), b.to(torch.long), rtol=0, atol=0)
assert len(fv.vision_condition_mask) == len(fwv.condition_mask)
for a, b in zip(fv.vision_condition_mask, fwv.condition_mask):
torch.testing.assert_close(a.flatten().to(torch.float32), b.flatten().to(torch.float32))
# (grid_t, latent_h, latent_w, n_text, cond_frames, id) — single-sample cases.
_CASES = [
pytest.param(1, 8, 8, 4, [], id="t2i_1x4x4"),
pytest.param(1, 4, 4, 5, [0], id="t2i_cond_1x2x2"),
pytest.param(2, 4, 4, 4, [], id="t2v_2x2x2"),
pytest.param(3, 8, 4, 6, [], id="t2v_3x4x2"),
pytest.param(2, 4, 4, 5, [0], id="i2v_2x2x2"),
pytest.param(3, 4, 4, 4, [0], id="i2v_3x2x2"),
]
class TestCosmos3PackingParity:
# -- Field-by-field packing parity -------------------------------------
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
def test_packing_fields_match_framework(self, grid_t, latent_h, latent_w, n_text, cond):
torch.manual_seed(0)
text_ids = torch.randint(0, 60, (n_text,)).tolist()
vision = _make_vision(grid_t, latent_h, latent_w, seed=123)
timestep = 500.0
fw = _framework_pack(
text_ids_per_sample=[text_ids],
visions=[vision],
cond_frames_per_sample=[cond],
timesteps=[timestep],
is_image_batch=(grid_t == 1),
)
fv = _fastvideo_pack(
text_ids_per_sample=[text_ids],
visions=[vision],
cond_frames_per_sample=[cond],
timesteps=[timestep],
)
_assert_packs_match(fw, fv)
def test_packing_fields_match_framework_multi_sample(self):
"""A batch of two samples (T2V + I2V) packs identically to the framework."""
torch.manual_seed(1)
t0 = torch.randint(0, 60, (3,)).tolist()
t1 = torch.randint(0, 60, (5,)).tolist()
v0 = _make_vision(2, 4, 4, seed=11)
v1 = _make_vision(2, 4, 4, seed=22)
kwargs = dict(
text_ids_per_sample=[t0, t1],
visions=[v0, v1],
cond_frames_per_sample=[[], [0]],
timesteps=[500.0, 250.0],
)
fw = _framework_pack(is_image_batch=False, **kwargs)
fv = _fastvideo_pack(**kwargs)
_assert_packs_match(fw, fv)
# -- End-to-end: FastVideo packing drives the DiT identically ----------
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
def test_fastvideo_packing_drives_dit_like_framework(self, grid_t, latent_h, latent_w, n_text, cond):
"""Feed BOTH the framework-packed and FastVideo-packed inputs through the
SAME FastVideo DiT (framework weights copied in); assert identical output.
"""
num_layers = 2
torch.manual_seed(0)
text_ids = torch.randint(0, 60, (n_text,)).tolist()
vision = _make_vision(grid_t, latent_h, latent_w, seed=123)
timestep = 500.0
fw_pack = _framework_pack(
text_ids_per_sample=[text_ids],
visions=[vision],
cond_frames_per_sample=[cond],
timesteps=[timestep],
is_image_batch=(grid_t == 1),
)
fv_pack = _fastvideo_pack(
text_ids_per_sample=[text_ids],
visions=[vision],
cond_frames_per_sample=[cond],
timesteps=[timestep],
)
# Guard: the two packs must agree before we trust the DiT comparison.
_assert_packs_match(fw_pack, fv_pack)
# One DiT instance, framework weights copied in (parity oracle weights).
vfm = _build_tiny_cosmos3_mrope(seed=42, num_layers=num_layers)
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
_copy_weights(vfm, dit)
with torch.no_grad():
out_fw = dit(**_fastvideo_inputs_from_packed_seq(fw_pack))
out_fv = dit(**fv_pack.to_dit_kwargs())
# last_hidden_state must be bit-identical.
lhs_fw = out_fw["last_hidden_state"]
lhs_fv = out_fv["last_hidden_state"]
assert lhs_fw.shape == lhs_fv.shape
max_abs_lhs = (lhs_fw - lhs_fv).abs().max().item()
print(f"\n[packing->dit {grid_t}x{latent_h}x{latent_w} cond={cond}] "
f"last_hidden_state max abs diff = {max_abs_lhs:.3e}")
torch.testing.assert_close(lhs_fv, lhs_fw, rtol=0, atol=0)
# preds_vision must be bit-identical when there are noisy frames to
# predict. (A fully-conditioned clip has no noisy patches, so the DiT
# emits no "preds_vision" — both packs agree the mse-loss set is empty,
# already asserted by the field-parity guard above.)
has_preds = fv_pack.vision_mse_loss_indexes.numel() > 0
assert ("preds_vision" in out_fw) == has_preds
assert ("preds_vision" in out_fv) == has_preds
if has_preds:
pv_fw = out_fw["preds_vision"][0]
pv_fv = out_fv["preds_vision"][0]
assert pv_fw.shape == pv_fv.shape
max_abs_pv = (pv_fw - pv_fv).abs().max().item()
print(f"[packing->dit {grid_t}x{latent_h}x{latent_w} cond={cond}] "
f"preds_vision max abs diff = {max_abs_pv:.3e}")
torch.testing.assert_close(pv_fv, pv_fw, rtol=0, atol=0)

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