Compare commits

...
Author SHA1 Message Date
SolitaryThinker 5b1c779041 [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-06-07 16:42:15 -07:00
SolitaryThinker c6f15de8c7 [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-06-07 16:40:46 -07:00
SolitaryThinker 245515f653 [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-06-07 16:37:55 -07:00
SolitaryThinker b24ee58f8d [docs] cosmos3 PR2 (audio/t2vs) complete; bit-exact across components 2026-06-07 16:36:45 -07:00
SolitaryThinker 62a0a0bbfb [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-06-07 16:36:45 -07:00
SolitaryThinker 0245218e34 [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-06-07 16:35:37 -07:00
SolitaryThinker 14147c57c2 [docs] cosmos3 PR2: AVAE sound decoder done; remaining audio components 2026-06-07 16:35:37 -07:00
SolitaryThinker 4fda151b7a [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-06-07 16:35:37 -07:00
SolitaryThinker 3d74fe1896 [docs] cosmos3 PR2: audio port plan (AVAE + DiT sound pathway + t2vs) 2026-06-07 16:35:37 -07:00
SolitaryThinker 06321bb364 [docs] cosmos3: record T2I verification + resolution-based flow_shift 2026-06-07 16:34:36 -07:00
SolitaryThinker 0a5d9ecbc7 [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-06-07 16:34:36 -07:00
SolitaryThinker 744ca86ced [docs] cosmos3: record I2V real-weights verification (feat/cosmos3-i2v) 2026-06-07 16:34:36 -07:00
SolitaryThinker a5be058cf1 [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-06-07 16:34:36 -07:00
SolitaryThinker 1a7e5b1665 [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-06-07 16:33:45 -07:00
SolitaryThinker 255721c1c2 [docs] cosmos3: record real-weights E2E acceptance + scheduler fix (I003) 2026-06-07 01:33:28 -07:00
SolitaryThinker 255311cf25 [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-06-07 01:31:19 -07:00
SolitaryThinker df8990807a [misc] cosmos3 PORT_STATUS: PR1 video core complete (framework parity) 2026-06-06 23:55:47 -07:00
SolitaryThinker 77bf05b2e8 [feat] cosmos3: native video pipeline + framework denoise parity 2026-06-06 23:54:32 -07:00
SolitaryThinker 3419468923 [feat] cosmos3: native MoT sequence-packing (video) + framework parity 2026-06-06 23:23:31 -07:00
SolitaryThinker 63b705340e [misc] cosmos3 PORT_STATUS: strict-load done; only pipeline remains 2026-06-06 23:11:36 -07:00
SolitaryThinker b5a8e0fb61 [feat] cosmos3: strict-load verified (identity; needs_conversion=no) 2026-06-06 23:10:39 -07:00
SolitaryThinker 30a4b61e4f [misc] cosmos3 PORT_STATUS: VAE component framework parity verified 2026-06-06 22:05:32 -07:00
SolitaryThinker d16a859cf1 [feat] cosmos3: VAE config (Wan2.2 reuse) + framework parity test 2026-06-06 22:04:25 -07:00
SolitaryThinker f009bf380a [misc] cosmos3 PORT_STATUS: DiT framework parity verified (both rope modes) 2026-06-06 21:05:44 -07:00
SolitaryThinker 7c46332955 [test] cosmos3: DiT unified_3d_mrope framework parity (real checkpoint) 2026-06-06 21:04:53 -07:00
SolitaryThinker 59a4a571cc [feat] cosmos3: native DiT + framework parity (Cosmos3VFMTransformer) 2026-06-06 20:52:36 -07:00
SolitaryThinker 2293d30eb8 [misc] cosmos3 PORT_STATUS: PR1 progress (arch config, framework parity ref) 2026-06-06 20:30:21 -07:00
SolitaryThinker dd97efda3b [test] cosmos3: official framework DiT parity reference (CPU/SDPA) 2026-06-06 20:29:08 -07:00
SolitaryThinker 90c63fd0f1 [misc] cosmos3 PORT_STATUS: framework-reference pivot + Phase 1 findings 2026-06-06 20:03:17 -07:00
SolitaryThinker 9567efdf07 [feat] cosmos3: arch config 1:1 with Cosmos3-Nano checkpoint 2026-06-06 20:02:43 -07:00
SolitaryThinker 7cb154337f [misc]: cosmos3 PORT_STATUS: post-rebase state + fv-cosmos3 env 2026-06-06 19:28:13 -07:00
SolitaryThinker 2274771b4b [misc]: cosmos3 resume: PORT_STATUS + README for official diffusers ref 2026-06-06 19:23:23 -07:00
SolitaryThinker 31f03df62b [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-06-06 19:23:23 -07:00
SolitaryThinker 69587f631e [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-06-06 19:23:23 -07:00
SolitaryThinker 0b175e414a [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-06-06 19:18:19 -07:00
SolitaryThinker cc55c10e0a [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-06-06 19:18:19 -07:00
SolitaryThinker c669b2d4b4 [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-06-06 19:18:19 -07:00
46 changed files with 8368 additions and 1 deletions
+5
View File
@@ -35,6 +35,11 @@ env
weights/
logs/
# Cosmos3 local parity assets (symlinked from main worktree)
/official_weights/
/converted_weights/
/cosmos-framework
# SSIM test outputs
fastvideo/tests/ssim/generated_videos/
**/.cache/**
@@ -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()
+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"
@@ -1,5 +1,6 @@
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.cosmos3vae import Cosmos3VAEConfig
from fastvideo.configs.models.vaes.gamecraftvae import GameCraftVAEConfig
from fastvideo.configs.models.vaes.gen3cvae import Gen3CVAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
@@ -14,6 +15,7 @@ __all__ = [
"WanVAEConfig",
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Cosmos3VAEConfig",
"Gen3CVAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
+277
View File
@@ -0,0 +1,277 @@
"""Cosmos3 (Wan2.2-TI2V-5B) VAE config and checkpoint-key mapping.
The Cosmos3 checkpoint VAE is literally ``Wan-AI/Wan2.2-TI2V-5B-Diffusers``
(diffusers ``AutoencoderKLWan``), so this config locks the Wan2.2 geometry:
residual down/up blocks, ``patch_size=2``, ``z_dim=48``, ``base_dim=160``,
``decoder_base_dim=256``, and ``scale_factor_spatial=16``. The 48-dim
``latents_mean``/``latents_std`` are taken verbatim from the Cosmos3
checkpoint's ``vae/config.json`` (identical to the canonical Wan2.2-TI2V-5B
statistics).
Mirrors the :class:`Cosmos25VAEArchConfig` pattern. ``param_names_mapping`` /
``map_official_key`` translate the *official* Wan2.2 VAE state-dict keys
(nested-residual naming, e.g. ``encoder.downsamples.{b}.downsamples.{j}`` and
``decoder.upsamples.{b}.upsamples.{j}``) into FastVideo's ``AutoencoderKLWan``
key space. The standard diffusers checkpoint already ships native FastVideo
keys, so these helpers exist for parity tooling and official ``.pth`` loading.
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig, WanVAEConfig
@dataclass
class Cosmos3VAEArchConfig(WanVAEArchConfig):
# Wan2.2-TI2V-5B geometry (differs from the Wan2.1 WanVAEArchConfig
# defaults: residual blocks, patch_size=2, z_dim=48, base_dim=160,
# decoder_base_dim=256, scale_factor_spatial=16, 12 patch channels).
_name_or_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
base_dim: int = 160
decoder_base_dim: int | None = 256
z_dim: int = 48
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
num_res_blocks: int = 2
attn_scales: tuple[float, ...] = ()
temperal_downsample: tuple[bool, ...] = (False, True, True)
dropout: float = 0.0
is_residual: bool = True
in_channels: int = 12
out_channels: int = 12
patch_size: int | None = 2
scale_factor_temporal: int = 4
scale_factor_spatial: int = 16
clip_output: bool = False
# 48-dim statistics copied verbatim from the Cosmos3 checkpoint
# (official_weights/cosmos3/vae/config.json).
latents_mean: tuple[float, ...] = (
-0.2289,
-0.0052,
-0.1323,
-0.2339,
-0.2799,
0.0174,
0.1838,
0.1557,
-0.1382,
0.0542,
0.2813,
0.0891,
0.157,
-0.0098,
0.0375,
-0.1825,
-0.2246,
-0.1207,
-0.0698,
0.5109,
0.2665,
-0.2108,
-0.2158,
0.2502,
-0.2055,
-0.0322,
0.1109,
0.1567,
-0.0729,
0.0899,
-0.2799,
-0.123,
-0.0313,
-0.1649,
0.0117,
0.0723,
-0.2839,
-0.2083,
-0.052,
0.3748,
0.0152,
0.1957,
0.1433,
-0.2944,
0.3573,
-0.0548,
-0.1681,
-0.0667,
)
latents_std: tuple[float, ...] = (
0.4765,
1.0364,
0.4514,
1.1677,
0.5313,
0.499,
0.4818,
0.5013,
0.8158,
1.0344,
0.5894,
1.0901,
0.6885,
0.6165,
0.8454,
0.4978,
0.5759,
0.3523,
0.7135,
0.6804,
0.5833,
1.4146,
0.8986,
0.5659,
0.7069,
0.5338,
0.4889,
0.4917,
0.4069,
0.4999,
0.6866,
0.4093,
0.5709,
0.6065,
0.6415,
0.4944,
0.5726,
1.2042,
0.5458,
1.6887,
0.3971,
1.06,
0.3943,
0.5537,
0.5444,
0.4089,
0.7468,
0.7744,
)
# Simple 1:1 renames. The nested-residual block remapping (encoder
# downsamples / decoder upsamples / middle / head) is handled by
# ``map_official_key()``.
param_names_mapping: dict[str, str] = field(
default_factory=lambda: {
r"^conv1\.(.*)$": r"quant_conv.\1",
r"^conv2\.(.*)$": r"post_quant_conv.\1",
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
})
@staticmethod
def map_official_key(key: str) -> str | None:
"""Map a single official Wan2.2 VAE key into FastVideo key space.
Handles the residual (Wan2.2) module layout where each down/up block
is a nested ``Sequential`` (``downsamples.{b}.downsamples.{j}`` /
``upsamples.{b}.upsamples.{j}``) rather than the flat Wan2.1 indexing.
Returns ``None`` for keys with no FastVideo counterpart.
"""
def map_residual_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^residual\.0\.gamma$", sub):
return f"{prefix}.norm1.gamma"
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv1.{m.group(1)}"
if re.match(r"^residual\.3\.gamma$", sub):
return f"{prefix}.norm2.gamma"
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv2.{m.group(1)}"
m = re.match(r"^shortcut\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv_shortcut.{m.group(1)}"
return None
def map_attn_subkey(prefix: str, sub: str) -> str | None:
if re.match(r"^norm\.gamma$", sub):
return f"{prefix}.norm.gamma"
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
if m:
return f"{prefix}.to_qkv.{m.group(1)}"
m = re.match(r"^proj\.(weight|bias)$", sub)
if m:
return f"{prefix}.proj.{m.group(1)}"
return None
def map_resample_subkey(prefix: str, sub: str) -> str | None:
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
if m:
return f"{prefix}.resample.1.{m.group(1)}"
m = re.match(r"^time_conv\.(weight|bias)$", sub)
if m:
return f"{prefix}.time_conv.{m.group(1)}"
return None
m = re.match(r"^conv1\.(weight|bias)$", key)
if m:
return f"quant_conv.{m.group(1)}"
m = re.match(r"^conv2\.(weight|bias)$", key)
if m:
return f"post_quant_conv.{m.group(1)}"
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_in.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
if m:
return f"{m.group(1)}.norm_out.gamma"
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
if m:
return f"{m.group(1)}.conv_out.{m.group(2)}"
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if m:
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
if m:
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1", m.group(2))
# Encoder: downsamples.{block}.downsamples.{j}.* (nested residual layout)
m = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
if m:
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return map_resample_subkey(f"encoder.down_blocks.{block_i}.downsampler", sub)
return map_residual_subkey(f"encoder.down_blocks.{block_i}.resnets.{res_i}", sub)
# Decoder: upsamples.{block}.upsamples.{j}.* (nested residual layout)
m = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
if m:
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return map_resample_subkey(f"decoder.up_blocks.{block_i}.upsampler", sub)
return map_residual_subkey(f"decoder.up_blocks.{block_i}.resnets.{res_i}", sub)
return None
# ``__post_init__`` (scaling_factor / shift_factor / compression ratios) is
# inherited unchanged from ``WanVAEArchConfig``.
@dataclass
class Cosmos3VAEConfig(WanVAEConfig):
"""Cosmos3 VAE config (reuses FastVideo's Wan2.2 ``AutoencoderKLWan``).
Subclasses :class:`WanVAEConfig` so the model reads the same runtime flags
(``use_feature_cache``, ``use_light_vae``, tiling) and only swaps in the
Cosmos3 = Wan2.2 ``arch_config``.
"""
arch_config: Cosmos3VAEArchConfig = field(default_factory=Cosmos3VAEArchConfig)
use_feature_cache: bool = True
use_tiling: bool = False
use_temporal_tiling: bool = False
use_parallel_tiling: bool = False
# ``__post_init__`` (blend_num_frames) is inherited from ``WanVAEConfig``.
+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
+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
+991
View File
@@ -0,0 +1,991 @@
# SPDX-License-Identifier: Apache-2.0
"""FastVideo-native Cosmos3 omni DiT (``Cosmos3VFMTransformer``).
Numerical-parity port of the official ``cosmos_framework`` ``Cosmos3VFMNetwork``
(the ``Qwen3-VL-text`` MoT backbone + the VFM vision / action / sound heads).
The module tree mirrors the published *diffusers* checkpoint layout so a
converter can strict-load with a near-identity ``param_names_mapping``:
* top level: ``embed_tokens`` / ``norm`` / ``norm_moe_gen`` / ``lm_head`` /
``proj_in`` / ``proj_out`` / ``time_embedder.linear_{1,2}`` plus the dormant
``action_*`` / ``audio_*`` heads (constructed for strict-load parity);
* per layer ``layers.{i}``: dual-pathway ``self_attn`` with understanding
(``to_{q,k,v}`` / ``to_out``) and generation (``add_{q,k,v}_proj`` /
``to_add_out``) projections + per-head QK-norms (``norm_{q,k}`` for und,
``norm_added_{q,k}`` for gen), ``mlp`` (und) and ``mlp_moe_gen`` (gen) SwiGLU
blocks, and four RMSNorms.
The forward replicates the framework contract for the video path:
``patchify -> proj_in -> (+ 3d-rope additive latent pos emb) -> scatter-add
timestep embeds onto noisy patches -> dual-pathway decoder (causal text +
full vision two-way attention, GQA, per-head QK-norm before RoPE, unified
3D-MRoPE) -> und/gen final norms -> proj_out on noisy vision patches ->
unpatchify``.
RoPE is applied with the framework's exact ``rotate_half`` math (contiguous
split-half: ``q * cos + rotate_half(q) * sin``) rather than the interleaved
``apply_rotary_emb`` helper, and attention uses plain SDPA so the whole module
runs on CPU / float32 for parity testing. No diffusers / transformers
model-class imports happen at runtime — the DiT is fully native.
This module also keeps two pure, standalone math utilities used by the Tier-A
scaffold parity tests: the unified-3D mRoPE position-ID generators
(``compute_mrope_position_ids_text`` / ``compute_mrope_position_ids_vision``)
and the batched ``patchify`` / ``unpatchify`` ``[B,C,T,H,W] <-> [B,N,p*p*C]``
roundtrip helpers.
"""
from __future__ import annotations
import math
from typing import Any
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.dits.cosmos3 import Cosmos3VideoConfig
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.visual_embedding import timestep_embedding
from fastvideo.models.dits.base import BaseDiT
EntryClass = ["Cosmos3VFMTransformer"]
# ===========================================================================
# Standalone mRoPE position-ID generators (unified 3D mRoPE)
# ===========================================================================
def compute_mrope_position_ids_text(
num_tokens: int,
temporal_offset: int,
) -> tuple[torch.Tensor, int]:
"""Generate 3D mRoPE position IDs for text tokens.
Text tokens broadcast a single monotonically-increasing position-ID
sequence across all three (t, h, w) axes.
"""
ids = torch.arange(num_tokens, dtype=torch.long) + temporal_offset
mrope_ids = ids.unsqueeze(0).expand(3, -1).contiguous()
return mrope_ids, temporal_offset + num_tokens
def compute_mrope_position_ids_vision(
grid_t: int,
grid_h: int,
grid_w: int,
temporal_offset: int | float,
fps: float | None = None,
base_fps: float = 24.0,
temporal_compression_factor: int = 4,
base_temporal_compression_factor: int | None = None,
enable_fps_modulation: bool = True,
start_frame_offset: int = 0,
) -> tuple[torch.Tensor, int]:
"""Generate 3D mRoPE position IDs for vision tokens.
Builds a ``(t, h, w)`` position grid (Qwen3-VL style, spatial indices reset
per temporal segment) flattened in t-major order. Optionally modulates the
temporal axis by ``base_fps / tcf * (1 / (fps / tcf))`` so two clips at
different FPS retain wall-clock-aligned temporal positions.
"""
fps_modulation = enable_fps_modulation and fps is not None
if fps_modulation:
assert fps is not None
tps = fps / temporal_compression_factor
effective_base_tcf = (base_temporal_compression_factor
if base_temporal_compression_factor is not None else temporal_compression_factor)
base_tps = base_fps / effective_base_tcf
frame_indices = torch.arange(grid_t, dtype=torch.float32)
t_index = (((frame_indices + start_frame_offset) / tps * base_tps + temporal_offset).view(-1, 1).expand(
-1, grid_h * grid_w).flatten())
else:
t_index = (torch.arange(grid_t, dtype=torch.long).view(-1, 1).expand(-1, grid_h * grid_w).flatten() +
int(temporal_offset) + start_frame_offset)
h_index = (torch.arange(grid_h, dtype=torch.long).view(1, -1, 1).expand(grid_t, -1, grid_w).flatten())
w_index = (torch.arange(grid_w, dtype=torch.long).view(1, 1, -1).expand(grid_t, grid_h, -1).flatten())
if fps_modulation:
mrope_ids = torch.stack([t_index, h_index.to(torch.float32), w_index.to(torch.float32)], dim=0)
else:
mrope_ids = torch.stack([t_index, h_index, w_index], dim=0)
next_offset = math.floor(mrope_ids.max().item()) + 1
return mrope_ids, next_offset
# ===========================================================================
# RoPE helpers (match cosmos_framework qwen3_vl rotate_half + apply_rotary_pos_emb)
# ===========================================================================
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
"""Contiguous split-half rotation, identical to qwen3_vl.rotate_half."""
x1 = x[..., :x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2:]
return torch.cat((-x2, x1), dim=-1)
def _apply_rotary_pos_emb(
q: torch.Tensor,
k: torch.Tensor,
cos: torch.Tensor,
sin: torch.Tensor,
unsqueeze_dim: int = 1,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Apply RoPE exactly as cosmos_framework qwen3_vl.apply_rotary_pos_emb.
``q`` / ``k`` are ``[N, heads, head_dim]`` and ``cos`` / ``sin`` are
``[N, head_dim]`` (per-token); ``unsqueeze_dim=1`` broadcasts cos/sin over
the head axis.
"""
cos = cos.unsqueeze(unsqueeze_dim)
sin = sin.unsqueeze(unsqueeze_dim)
q_embed = (q * cos) + (_rotate_half(q) * sin)
k_embed = (k * cos) + (_rotate_half(k) * sin)
return q_embed, k_embed
class Cosmos3TextRotaryEmbedding(nn.Module):
"""Unified 3D-MRoPE rotary embedding (qwen3_vl Cosmos3 flavor).
Reproduces ``Qwen3VLTextRotaryEmbedding``: ``inv_freq`` from the default
RoPE init (``1 / theta**(arange(0, head_dim, 2) / head_dim)``), interleaved
MRoPE mixing of the T/H/W frequency bands per ``mrope_section``, and a final
``cat([freqs, freqs])`` so cos/sin are ``[N, head_dim]``.
"""
def __init__(
self,
head_dim: int,
rope_theta: float,
mrope_section: list[int],
) -> None:
super().__init__()
self.head_dim = head_dim
self.rope_theta = float(rope_theta)
self.mrope_section = list(mrope_section)
self.register_buffer("inv_freq", self._compute_inv_freq(), persistent=False)
self.attention_scaling = 1.0 # default RoPE has unit attention scaling
def _compute_inv_freq(self, device: torch.device | None = None) -> torch.Tensor:
exponent = torch.arange(0, self.head_dim, 2, dtype=torch.int64, device=device).float() / self.head_dim
return 1.0 / (self.rope_theta**exponent)
def reset_inv_freq(self, device: torch.device) -> None:
"""Recompute the non-persistent ``inv_freq`` on ``device`` after a
meta-device weight load (it is derived from ``rope_theta`` and is not
part of the checkpoint)."""
self.register_buffer("inv_freq", self._compute_inv_freq(device), persistent=False)
def _apply_interleaved_mrope(self, freqs: torch.Tensor) -> torch.Tensor:
"""freqs: [3, N, head_dim//2] -> [N, head_dim//2] (interleaved T/H/W)."""
freqs_t = freqs[0].clone()
for dim, offset in enumerate((1, 2), start=1): # H, W
length = self.mrope_section[dim] * 3
idx = slice(offset, length, 3)
freqs_t[..., idx] = freqs[dim, ..., idx]
return freqs_t
@torch.no_grad()
def forward(self, position_ids: torch.Tensor, device: torch.device,
dtype: torch.dtype) -> tuple[torch.Tensor, torch.Tensor]:
"""position_ids: [N] (1D) or [3, N] (mrope) -> (cos, sin) each [N, head_dim]."""
if position_ids.ndim == 1:
position_ids = position_ids[None, :].expand(3, -1) # [3, N]
position_ids = position_ids.to(device)
inv_freq = self.inv_freq.to(device).float() # [head_dim//2]
inv_freq_expanded = inv_freq[None, :, None].expand(3, -1, 1) # [3, head_dim//2, 1]
position_ids_expanded = position_ids[:, None, :].float() # [3, 1, N]
freqs = (inv_freq_expanded @ position_ids_expanded).transpose(1, 2) # [3, N, head_dim//2]
freqs = self._apply_interleaved_mrope(freqs) # [N, head_dim//2]
emb = torch.cat((freqs, freqs), dim=-1) # [N, head_dim]
cos = (emb.cos() * self.attention_scaling).to(dtype)
sin = (emb.sin() * self.attention_scaling).to(dtype)
return cos, sin
# ===========================================================================
# 3D-RoPE additive latent position embedding (VideoRopePosition3DEmb)
# ===========================================================================
class Cosmos3VideoRopePosition3DEmb(nn.Module):
"""Additive 3D-RoPE-style latent position embedding.
Mirrors ``cosmos_framework`` ``VideoRopePosition3DEmb`` (used when
``position_embedding_type == "3d_rope"``). Produces an additive
``[N_vision, head_dim]`` tensor concatenated over per-latent (t, h, w)
grids. ``enable_fps_modulation`` is supported for completeness but the
video path keeps it disabled by default.
"""
def __init__(
self,
head_dim: int,
len_h: int,
len_w: int,
len_t: int,
base_fps: int = 24,
base_temporal_compression_factor: int = 4,
temporal_compression_factor: int = 4,
h_extrapolation_ratio: float = 1.0,
w_extrapolation_ratio: float = 1.0,
t_extrapolation_ratio: float = 1.0,
enable_fps_modulation: bool = False,
) -> None:
super().__init__()
self.base_tps = base_fps / base_temporal_compression_factor
self.temporal_compression_factor = temporal_compression_factor
self.max_h = len_h
self.max_w = len_w
self.max_t = len_t
self.enable_fps_modulation = enable_fps_modulation
dim = head_dim
dim_h = dim // 6 * 2
dim_w = dim_h
dim_t = dim - 2 * dim_h
assert dim == dim_h + dim_w + dim_t, f"bad dim: {dim} != {dim_h} + {dim_w} + {dim_t}"
self._dim_h = dim_h
self._dim_t = dim_t
self.register_buffer(
"dim_spatial_range",
torch.arange(0, dim_h, 2)[:(dim_h // 2)].float() / dim_h,
persistent=True,
)
self.register_buffer(
"dim_temporal_range",
torch.arange(0, dim_t, 2)[:(dim_t // 2)].float() / dim_t,
persistent=True,
)
self.h_ntk_factor = h_extrapolation_ratio**(dim_h / (dim_h - 2))
self.w_ntk_factor = w_extrapolation_ratio**(dim_w / (dim_w - 2))
self.t_ntk_factor = t_extrapolation_ratio**(dim_t / (dim_t - 2))
def _generate(self, t: int, h: int, w: int, device: torch.device, fps: torch.Tensor | None,
start_frame_offset: int) -> torch.Tensor:
tps = (fps / self.temporal_compression_factor) if fps is not None else None
h_theta = 10000.0 * self.h_ntk_factor
w_theta = 10000.0 * self.w_ntk_factor
t_theta = 10000.0 * self.t_ntk_factor
spatial = self.dim_spatial_range.to(device).float()
temporal = self.dim_temporal_range.to(device).float()
h_spatial_freqs = 1.0 / (h_theta**spatial)
w_spatial_freqs = 1.0 / (w_theta**spatial)
temporal_freqs = 1.0 / (t_theta**temporal)
max_needed = max(t, h, w)
seq = torch.arange(max_needed, device=device, dtype=torch.float32)
half_emb_h = torch.outer(seq[:h], h_spatial_freqs) # [h, dim_h/2]
half_emb_w = torch.outer(seq[:w], w_spatial_freqs) # [w, dim_w/2]
frame_indices = seq[:t] # [t]
if self.enable_fps_modulation and tps is not None:
scaled_time = (frame_indices + start_frame_offset) / tps[:1] * self.base_tps
half_emb_t = torch.outer(scaled_time, temporal_freqs)
else:
half_emb_t = torch.outer(frame_indices, temporal_freqs) # [t, dim_t/2]
emb_t = half_emb_t[:, None, None, :].expand(t, h, w, -1)
emb_h = half_emb_h[None, :, None, :].expand(t, h, w, -1)
emb_w = half_emb_w[None, None, :, :].expand(t, h, w, -1)
rope = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1) # [t, h, w, head_dim]
return rope.reshape(t * h * w, -1).float()
def forward(
self,
token_shapes: list[tuple[int, int, int]],
fps: torch.Tensor | None = None,
start_frame_offset: int = 0,
) -> torch.Tensor:
device = self.dim_spatial_range.device
out = []
for i, (t, h, w) in enumerate(token_shapes):
video_fps = fps[i:i + 1] if fps is not None else None
out.append(self._generate(t, h, w, device, video_fps, start_frame_offset))
return torch.cat(out, dim=0) # [N_vision, head_dim]
# ===========================================================================
# Timestep embedder (DiT-style; matches modeling_utils.TimestepEmbedder)
# ===========================================================================
class Cosmos3TimestepEmbedder(nn.Module):
"""Sinusoidal timestep -> MLP embedder. Checkpoint keys: ``linear_{1,2}``."""
def __init__(self, hidden_size: int, frequency_embedding_size: int = 256) -> None:
super().__init__()
self.frequency_embedding_size = frequency_embedding_size
self.hidden_size = hidden_size
self.linear_1 = ReplicatedLinear(frequency_embedding_size, hidden_size, bias=True)
self.act = nn.SiLU()
self.linear_2 = ReplicatedLinear(hidden_size, hidden_size, bias=True)
def forward(self, t: torch.Tensor) -> torch.Tensor:
t_freq = timestep_embedding(t, self.frequency_embedding_size)
# Sinusoidal runs in fp32 (timestep_scale makes inputs tiny); cast to the
# MLP weight dtype (bf16 at inference; no-op in the fp32 parity tests).
t_freq = t_freq.to(self.linear_1.weight.dtype)
h, _ = self.linear_1(t_freq)
h = self.act(h)
out, _ = self.linear_2(h)
return out
# ===========================================================================
# SwiGLU MLP (Qwen3-VL-text dense MLP)
# ===========================================================================
class Cosmos3MLP(nn.Module):
"""SwiGLU MLP: ``down(act(gate(x)) * up(x))``. Keys: gate/up/down_proj."""
def __init__(self, hidden_size: int, intermediate_size: int) -> None:
super().__init__()
self.gate_proj = ReplicatedLinear(hidden_size, intermediate_size, bias=False)
self.up_proj = ReplicatedLinear(hidden_size, intermediate_size, bias=False)
self.down_proj = ReplicatedLinear(intermediate_size, hidden_size, bias=False)
self.act_fn = nn.SiLU()
def forward(self, x: torch.Tensor) -> torch.Tensor:
gate, _ = self.gate_proj(x)
up, _ = self.up_proj(x)
out, _ = self.down_proj(self.act_fn(gate) * up)
return out
# ===========================================================================
# Dual-pathway packed two-way attention
# ===========================================================================
def _sdpa(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *, is_causal: bool, scale: float) -> torch.Tensor:
"""SDPA on ``[S, heads, head_dim]`` with GQA broadcast. Returns same layout."""
n_heads = q.shape[1]
n_kv = k.shape[1]
q_ = q.permute(1, 0, 2).unsqueeze(0) # [1, heads, S, head_dim]
k_ = k.permute(1, 0, 2).unsqueeze(0)
v_ = v.permute(1, 0, 2).unsqueeze(0)
if n_kv != n_heads:
k_ = k_.repeat_interleave(n_heads // n_kv, dim=1)
v_ = v_.repeat_interleave(n_heads // n_kv, dim=1)
out = F.scaled_dot_product_attention(q_, k_, v_, is_causal=is_causal, scale=scale)
return out.squeeze(0).permute(1, 0, 2) # [S, heads, head_dim]
class Cosmos3DualAttention(nn.Module):
"""Dual-pathway packed attention (understanding + generation).
Understanding (text) tokens use causal self-attention; generation (vision)
tokens use full attention where the gen query attends to ALL tokens
(und ++ gen). GQA with per-head QK-norm applied *before* RoPE. Checkpoint
keys: und ``to_{q,k,v}`` / ``to_out`` / ``norm_{q,k}``; gen
``add_{q,k,v}_proj`` / ``to_add_out`` / ``norm_added_{q,k}``.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
num_key_value_heads: int,
head_dim: int,
eps: float,
attention_bias: bool,
) -> None:
super().__init__()
self.hidden_size = hidden_size
self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
self.head_dim = head_dim
self.scaling = head_dim**-0.5
q_dim = num_attention_heads * head_dim
kv_dim = num_key_value_heads * head_dim
# Understanding pathway
self.to_q = ReplicatedLinear(hidden_size, q_dim, bias=attention_bias)
self.to_k = ReplicatedLinear(hidden_size, kv_dim, bias=attention_bias)
self.to_v = ReplicatedLinear(hidden_size, kv_dim, bias=attention_bias)
self.to_out = ReplicatedLinear(q_dim, hidden_size, bias=attention_bias)
self.norm_q = RMSNorm(head_dim, eps=eps)
self.norm_k = RMSNorm(head_dim, eps=eps)
# Generation pathway
self.add_q_proj = ReplicatedLinear(hidden_size, q_dim, bias=attention_bias)
self.add_k_proj = ReplicatedLinear(hidden_size, kv_dim, bias=attention_bias)
self.add_v_proj = ReplicatedLinear(hidden_size, kv_dim, bias=attention_bias)
self.to_add_out = ReplicatedLinear(q_dim, hidden_size, bias=attention_bias)
self.norm_added_q = RMSNorm(head_dim, eps=eps)
self.norm_added_k = RMSNorm(head_dim, eps=eps)
def forward(
self,
und_seq: torch.Tensor, # [N_und, hidden]
gen_seq: torch.Tensor, # [N_gen, hidden]
cos_und: torch.Tensor,
sin_und: torch.Tensor,
cos_gen: torch.Tensor,
sin_gen: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
n_heads = self.num_attention_heads
n_kv = self.num_key_value_heads
d = self.head_dim
q_und, _ = self.to_q(und_seq)
k_und, _ = self.to_k(und_seq)
v_und, _ = self.to_v(und_seq)
q_gen, _ = self.add_q_proj(gen_seq)
k_gen, _ = self.add_k_proj(gen_seq)
v_gen, _ = self.add_v_proj(gen_seq)
q_und = q_und.view(-1, n_heads, d)
k_und = k_und.view(-1, n_kv, d)
v_und = v_und.view(-1, n_kv, d)
q_gen = q_gen.view(-1, n_heads, d)
k_gen = k_gen.view(-1, n_kv, d)
v_gen = v_gen.view(-1, n_kv, d)
# Per-head QK-norm BEFORE RoPE
q_und = self.norm_q(q_und)
k_und = self.norm_k(k_und)
q_gen = self.norm_added_q(q_gen)
k_gen = self.norm_added_k(k_gen)
# RoPE (heads-second layout; unsqueeze cos/sin over the head axis at dim=1)
q_und, k_und = _apply_rotary_pos_emb(q_und, k_und, cos_und, sin_und, unsqueeze_dim=1)
q_gen, k_gen = _apply_rotary_pos_emb(q_gen, k_gen, cos_gen, sin_gen, unsqueeze_dim=1)
# Causal self-attention over understanding (text) tokens
und_out = _sdpa(q_und, k_und, v_und, is_causal=True, scale=self.scaling)
und_out = und_out.reshape(und_out.shape[0], n_heads * d)
# Full attention: gen query attends to ALL tokens (und ++ gen)
all_k = torch.cat([k_und, k_gen], dim=0)
all_v = torch.cat([v_und, v_gen], dim=0)
gen_out = _sdpa(q_gen, all_k, all_v, is_causal=False, scale=self.scaling)
gen_out = gen_out.reshape(gen_out.shape[0], n_heads * d)
und_out, _ = self.to_out(und_out)
gen_out, _ = self.to_add_out(gen_out)
return und_out, gen_out
# ===========================================================================
# Dual-pathway decoder layer
# ===========================================================================
class Cosmos3DecoderLayer(nn.Module):
"""MoT decoder layer: dual-pathway attention + dual SwiGLU MLP, 4 RMSNorms."""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
num_key_value_heads: int,
head_dim: int,
intermediate_size: int,
eps: float,
attention_bias: bool,
) -> None:
super().__init__()
self.self_attn = Cosmos3DualAttention(
hidden_size=hidden_size,
num_attention_heads=num_attention_heads,
num_key_value_heads=num_key_value_heads,
head_dim=head_dim,
eps=eps,
attention_bias=attention_bias,
)
self.mlp = Cosmos3MLP(hidden_size, intermediate_size)
self.mlp_moe_gen = Cosmos3MLP(hidden_size, intermediate_size)
self.input_layernorm = RMSNorm(hidden_size, eps=eps)
self.input_layernorm_moe_gen = RMSNorm(hidden_size, eps=eps)
self.post_attention_layernorm = RMSNorm(hidden_size, eps=eps)
self.post_attention_layernorm_moe_gen = RMSNorm(hidden_size, eps=eps)
def forward(
self,
und_seq: torch.Tensor,
gen_seq: torch.Tensor,
cos_und: torch.Tensor,
sin_und: torch.Tensor,
cos_gen: torch.Tensor,
sin_gen: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
# Pre-attention norm
und_norm = self.input_layernorm(und_seq)
gen_norm = self.input_layernorm_moe_gen(gen_seq)
und_attn, gen_attn = self.self_attn(und_norm, gen_norm, cos_und, sin_und, cos_gen, sin_gen)
und_res = und_seq + und_attn
gen_res = gen_seq + gen_attn
# Pre-MLP norm + dual SwiGLU
und_ln = self.post_attention_layernorm(und_res)
gen_ln = self.post_attention_layernorm_moe_gen(gen_res)
und_out = und_res + self.mlp(und_ln)
gen_out = gen_res + self.mlp_moe_gen(gen_ln)
return und_out, gen_out
# ===========================================================================
# Domain-aware linear (dormant action head; per-domain weight/bias embeddings)
# ===========================================================================
class _DomainAwareLinear(nn.Module):
"""One weight/bias pair per embodiment domain (matches framework keys)."""
def __init__(self, input_size: int, output_size: int, num_domains: int) -> None:
super().__init__()
self.input_size = int(input_size)
self.output_size = int(output_size)
self.num_domains = int(num_domains)
self.fc = nn.Embedding(self.num_domains, self.output_size * self.input_size)
self.bias = nn.Embedding(self.num_domains, self.output_size)
def forward(self, x: torch.Tensor, domain_id: torch.Tensor) -> torch.Tensor:
if domain_id.ndim == 0:
domain_id = domain_id.unsqueeze(0)
domain_id = domain_id.to(device=x.device, dtype=torch.long).reshape(-1)
weight = self.fc(domain_id).view(domain_id.shape[0], self.input_size, self.output_size)
bias = self.bias(domain_id).view(domain_id.shape[0], self.output_size)
if x.ndim == 2:
return torch.bmm(x.unsqueeze(1), weight).squeeze(1) + bias
return torch.bmm(x, weight) + bias.unsqueeze(1)
# ===========================================================================
# Full DiT
# ===========================================================================
class Cosmos3VFMTransformer(BaseDiT):
"""FastVideo-native Cosmos3 omni DiT.
Mirrors the published checkpoint's transformer key surface and the
``cosmos_framework`` ``Cosmos3VFMNetwork`` forward math for the video path.
Runs on CPU / float32 (plain SDPA + native-math RoPE).
"""
_fsdp_shard_conditions = Cosmos3VideoConfig().arch_config._fsdp_shard_conditions
_compile_conditions = Cosmos3VideoConfig().arch_config._compile_conditions
param_names_mapping = Cosmos3VideoConfig().arch_config.param_names_mapping
reverse_param_names_mapping: dict = {}
def __init__(self, config: Cosmos3VideoConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
arch = config.arch_config
self.hidden_size = arch.hidden_size
self.num_attention_heads = arch.num_attention_heads
self.num_key_value_heads = arch.num_key_value_heads
self.head_dim = arch.head_dim
self.num_hidden_layers = arch.num_hidden_layers
self.intermediate_size = arch.intermediate_size
self.vocab_size = arch.vocab_size
self.rms_norm_eps = arch.rms_norm_eps
self.attention_bias = arch.attention_bias
# VAE / patch geometry
self.latent_patch_size = arch.latent_patch_size
self.latent_channel = arch.latent_channel
# Alias kept for the standalone patchify/unpatchify helpers + scaffold tests.
self.latent_channel_size = arch.latent_channel
self.patch_latent_dim = arch.patch_latent_dim
self.num_channels_latents = arch.latent_channel
self.timestep_scale = arch.timestep_scale
self.position_embedding_type = arch.position_embedding_type
# ---- Backbone ----
self.embed_tokens = nn.Embedding(self.vocab_size, self.hidden_size)
self.layers = nn.ModuleList([
Cosmos3DecoderLayer(
hidden_size=self.hidden_size,
num_attention_heads=self.num_attention_heads,
num_key_value_heads=self.num_key_value_heads,
head_dim=self.head_dim,
intermediate_size=self.intermediate_size,
eps=self.rms_norm_eps,
attention_bias=self.attention_bias,
) for _ in range(self.num_hidden_layers)
])
self.norm = RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self.norm_moe_gen = RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self.lm_head = nn.Linear(self.hidden_size, self.vocab_size, bias=False)
# Backbone RoPE (unified 3D-MRoPE)
self.rotary_emb = Cosmos3TextRotaryEmbedding(
head_dim=self.head_dim,
rope_theta=arch.rope_theta,
mrope_section=arch.mrope_section,
)
# ---- Vision head ----
self.proj_in = ReplicatedLinear(self.patch_latent_dim, self.hidden_size, bias=True)
self.proj_out = ReplicatedLinear(self.hidden_size, self.patch_latent_dim, bias=True)
self.time_embedder = Cosmos3TimestepEmbedder(self.hidden_size)
# Additive latent position embedding only for the legacy 3d_rope variant.
self.latent_pos_embed: Cosmos3VideoRopePosition3DEmb | None = None
if self.position_embedding_type == "3d_rope":
self.latent_pos_embed = Cosmos3VideoRopePosition3DEmb(
head_dim=self.hidden_size,
len_h=getattr(arch, "max_latent_h", 32),
len_w=getattr(arch, "max_latent_w", 32),
len_t=getattr(arch, "max_latent_t", 32),
base_fps=int(arch.base_fps),
base_temporal_compression_factor=arch.temporal_compression_factor,
temporal_compression_factor=arch.temporal_compression_factor,
enable_fps_modulation=arch.enable_fps_modulation,
)
# ---- Dormant action head (constructed for strict-load parity) ----
if getattr(arch, "action_gen", False):
self.action_dim = arch.action_dim
self.num_embodiment_domains = arch.num_embodiment_domains
self.action_proj_in = _DomainAwareLinear(self.action_dim, self.hidden_size, self.num_embodiment_domains)
self.action_proj_out = _DomainAwareLinear(self.hidden_size, self.action_dim, self.num_embodiment_domains)
self.action_modality_embed = nn.Parameter(torch.zeros(self.hidden_size))
# ---- Dormant audio / sound head (constructed for strict-load parity) ----
if getattr(arch, "sound_gen", False):
self.sound_dim = arch.sound_dim
self.audio_proj_in = ReplicatedLinear(self.sound_dim, self.hidden_size, bias=True)
self.audio_proj_out = ReplicatedLinear(self.hidden_size, self.sound_dim, bias=True)
self.audio_modality_embed = nn.Parameter(torch.zeros(self.hidden_size))
self.__post_init__()
# ------------------------------------------------------------------
# Standalone batched patchify / unpatchify (scaffold-test contract)
# ------------------------------------------------------------------
def _pad_to_patch_size(self, h: int, w: int) -> tuple[int, int, int, int]:
"""Return ``(hp, wp, H_padded, W_padded)`` for ``latent_patch_size`` padding."""
p = self.latent_patch_size
h_padded = ((h + p - 1) // p) * p
w_padded = ((w + p - 1) // p) * p
return h_padded // p, w_padded // p, h_padded, w_padded
def patchify(self, latents: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor:
"""``[B, C, t, h, w] -> [B, t*hp*wp, p*p*C]`` (zero-pad h/w to patch multiples)."""
batch_size = latents.shape[0]
p = self.latent_patch_size
c = self.latent_channel_size
hp, wp, h_padded, w_padded = self._pad_to_patch_size(h, w)
if h_padded != h or w_padded != w:
latents = F.pad(latents, (0, w_padded - w, 0, h_padded - h))
x = latents.reshape(batch_size, c, t, hp, p, wp, p)
x = x.permute(0, 2, 3, 5, 4, 6, 1)
return x.reshape(batch_size, t * hp * wp, p * p * c)
def unpatchify(self, tokens: torch.Tensor, t: int, h: int, w: int) -> torch.Tensor:
"""``[B, t*hp*wp, p*p*C] -> [B, C, t, h, w]`` (crop h/w padding)."""
batch_size = tokens.shape[0]
p = self.latent_patch_size
c = self.latent_channel_size
hp, wp, h_padded, w_padded = self._pad_to_patch_size(h, w)
x = tokens.reshape(batch_size, t, hp, wp, p, p, c)
x = x.permute(0, 6, 1, 2, 4, 3, 5)
x = x.reshape(batch_size, c, t, h_padded, w_padded)
if h_padded != h or w_padded != w:
x = x[:, :, :, :h, :w]
return x
# ------------------------------------------------------------------
# Framework-faithful packed patchify / unpatchify (einsum ordering)
# ------------------------------------------------------------------
def _patchify_and_pack(
self,
tokens_vision: list[torch.Tensor],
token_shapes: list[tuple[int, int, int]],
) -> tuple[torch.Tensor, list[tuple[int, int, int]]]:
p = self.latent_patch_size
packed = []
original_shapes: list[tuple[int, int, int]] = []
for latent, (_t, _h, _w) in zip(tokens_vision, token_shapes):
latent = latent.squeeze(0) if latent.dim() == 5 else latent # [C, T, H, W]
_, t_actual, h_actual, w_actual = latent.shape
original_shapes.append((t_actual, h_actual, w_actual))
h_padded = ((h_actual + p - 1) // p) * p
w_padded = ((w_actual + p - 1) // p) * p
if h_padded != h_actual or w_padded != w_actual:
padded = torch.zeros((self.latent_channel, t_actual, h_padded, w_padded),
device=latent.device,
dtype=latent.dtype)
padded[:, :, :h_actual, :w_actual] = latent
latent = padded
h_patches = h_padded // p
w_patches = w_padded // p
latent = latent.reshape(self.latent_channel, t_actual, h_patches, p, w_patches, p)
latent = torch.einsum("cthpwq->thwpqc", latent).reshape(-1, p * p * self.latent_channel)
packed.append(latent)
return torch.cat(packed, dim=0), original_shapes
def _unpatchify_and_unpack(
self,
packed_preds: torch.Tensor,
token_shapes: list[tuple[int, int, int]],
noisy_frame_indexes: list[torch.Tensor],
original_shapes: list[tuple[int, int, int]] | None,
) -> list[torch.Tensor]:
p = self.latent_patch_size
outputs = []
start_idx = 0
for i, (t_c, _h_c, _w_c) in enumerate(token_shapes):
if original_shapes is not None:
_t_orig, h_orig, w_orig = original_shapes[i]
h_padded = ((h_orig + p - 1) // p) * p
w_padded = ((w_orig + p - 1) // p) * p
h_patches = h_padded // p
w_patches = w_padded // p
else:
h_orig, w_orig = _h_c * p, _w_c * p
h_patches, w_patches = _h_c, _w_c
nfi = noisy_frame_indexes[i]
t_n = len(nfi)
out = torch.zeros((self.latent_channel, t_c, h_orig, w_orig),
device=packed_preds.device,
dtype=packed_preds.dtype)
num_patches = t_n * h_patches * w_patches
if num_patches > 0:
end_idx = start_idx + num_patches
patches = packed_preds[start_idx:end_idx]
patches = patches.reshape(t_n, h_patches, w_patches, p, p, self.latent_channel)
latent = torch.einsum("thwpqc->cthpwq", patches)
latent = latent.reshape(self.latent_channel, t_n, h_patches * p, w_patches * p)
latent = latent[:, :, :h_orig, :w_orig]
out[:, nfi] = latent
start_idx = end_idx
outputs.append(out.unsqueeze(0)) # [1, C, T, H, W]
return outputs
def _scatter_timestep_embeds(
self,
packed_tokens: torch.Tensor,
packed_timestep_embeds: torch.Tensor,
noisy_frame_indexes: list[torch.Tensor],
token_shapes: list[tuple[int, ...]],
) -> torch.Tensor:
"""Add timestep embeds onto noisy patches (matches framework scatter_add)."""
start_noisy_index = 0
flat_idx = []
for noisy_i, shape_i in zip(noisy_frame_indexes, token_shapes):
spatial = math.prod(shape_i[1:])
spatial_idx = torch.arange(spatial, device=packed_tokens.device)
ni = (noisy_i * spatial).unsqueeze(-1).expand(-1, spatial)
ni = ni.clone() + spatial_idx + start_noisy_index
flat_idx.append(ni.flatten())
start_noisy_index += math.prod(shape_i)
flat = torch.cat(flat_idx, dim=0)
flat = flat.unsqueeze(-1).expand(-1, packed_tokens.shape[1])
return packed_tokens.scatter_add(0, flat, packed_timestep_embeds)
def materialize_non_persistent_buffers(self, device: torch.device, dtype: torch.dtype | None = None) -> None:
"""Recompute non-persistent buffers lost by the meta-device FSDP load.
Only ``rotary_emb.inv_freq`` is non-persistent (derived from
``rope_theta``, absent from the checkpoint). The FastVideo loader calls
this after ``load_model_from_full_model_state_dict``.
"""
del dtype # inv_freq stays float32 regardless of compute dtype
self.rotary_emb.reset_inv_freq(device)
# ------------------------------------------------------------------
# forward
# ------------------------------------------------------------------
def forward( # type: ignore[override]
self,
text_ids: torch.Tensor,
text_indexes: torch.Tensor,
position_ids: torch.Tensor,
sequence_length: int,
split_lens: list[int],
attn_modes: list[str],
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],
fps_vision: torch.Tensor | None = None,
sound_tokens: list[torch.Tensor] | None = None,
sound_token_shapes: list[tuple[int, int, int]] | None = None,
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] | None = None,
fps_sound: torch.Tensor | None = None,
action_tokens: list[torch.Tensor] | None = None,
action_token_shapes: list[tuple[int, ...]] | None = None,
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] | None = None,
action_domain_id: list[torch.Tensor] | None = None,
**kwargs: Any,
) -> dict[str, Any]:
"""Video-path forward returning ``{"last_hidden_state", "preds_vision"}``
(plus ``"preds_sound"`` when sound tokens are provided for t2vs).
Inputs mirror the framework ``PackedSequence`` fields for a single
text(causal)+vision(full) sample. This is the surface the parity test
and the converter exercise; a pipeline-facing wrapper composes these
from a ``PackedSequence`` builder.
"""
device = self.embed_tokens.weight.device
# 1. Text embedding scattered into the packed sequence buffer
text_embed = self.embed_tokens(text_ids) # [N_text, hidden]
target_dtype = text_embed.dtype
packed_sequence = text_embed.new_zeros((sequence_length, self.hidden_size))
packed_sequence[text_indexes] = text_embed
# 2. Vision: patchify -> proj_in -> (+ additive 3d-rope pos emb) -> timestep scatter
original_shapes: list[tuple[int, int, int]] | None = None
if vision_tokens:
packed_vision, original_shapes = self._patchify_and_pack(vision_tokens, vision_token_shapes)
# Vision latents arrive in fp32 (noise/VAE); cast to the model's
# compute dtype (bf16 at inference; no-op in the fp32 parity tests).
packed_vision, _ = self.proj_in(packed_vision.to(target_dtype))
if self.latent_pos_embed is not None:
pos_emb = self.latent_pos_embed(vision_token_shapes, fps=fps_vision).to(target_dtype)
packed_vision = packed_vision + pos_emb
if vision_mse_loss_indexes.numel() > 0:
timesteps = vision_timesteps * self.timestep_scale
ts_embeds = self.time_embedder(timesteps).to(target_dtype)
packed_vision = self._scatter_timestep_embeds(packed_vision, ts_embeds, vision_noisy_frame_indexes,
vision_token_shapes)
packed_sequence[vision_sequence_indexes] = packed_vision
# 2b. Sound (t2vs): pack [C, T] -> audio_proj_in -> + modality embed ->
# timestep scatter -> scatter into the (full) gen split. Mirrors the
# framework ``_encode_sound``; sound shares the vision "full" split.
if sound_tokens:
packed_sound = torch.cat(
[s[:, :shp[0]].permute(1, 0) for s, shp in zip(sound_tokens, sound_token_shapes)],
dim=0,
) # [total_sound_tokens, sound_dim]
packed_sound, _ = self.audio_proj_in(packed_sound.to(target_dtype))
packed_sound = packed_sound + self.audio_modality_embed
if sound_mse_loss_indexes is not None and sound_mse_loss_indexes.numel() > 0:
s_ts = sound_timesteps * self.timestep_scale
s_embeds = self.time_embedder(s_ts).to(target_dtype)
packed_sound = self._scatter_timestep_embeds(packed_sound, s_embeds, sound_noisy_frame_indexes,
sound_token_shapes)
packed_sequence[sound_sequence_indexes] = packed_sound
# 2c. Action (DomainAwareLinear): pack [T, D] per sample with a per-token
# domain id -> action_proj_in(domain) + action_modality_embed ->
# timestep scatter -> scatter into the gen split (framework
# ``_encode_action``). Action is domain-conditioned (per embodiment).
if action_tokens:
packed_action = torch.cat(
[a[:shp[0]] for a, shp in zip(action_tokens, action_token_shapes)], dim=0,
) # [total_action_tokens, action_dim]
per_token_domain = torch.cat(
[d.reshape(1).expand(shp[0]) for d, shp in zip(action_domain_id, action_token_shapes)], dim=0,
) # [total_action_tokens]
packed_action = self.action_proj_in(packed_action.to(target_dtype), per_token_domain)
packed_action = packed_action + self.action_modality_embed.view(1, -1)
if action_mse_loss_indexes is not None and action_mse_loss_indexes.numel() > 0:
a_ts = action_timesteps * self.timestep_scale
a_embeds = self.time_embedder(a_ts).to(target_dtype)
packed_action = self._scatter_timestep_embeds(packed_action, a_embeds, action_noisy_frame_indexes,
action_token_shapes)
packed_sequence[action_sequence_indexes] = packed_action
# 3. RoPE for the full packed sequence (cos/sin per token); split und/gen
cos, sin = self.rotary_emb(position_ids, device=device, dtype=target_dtype) # [N, head_dim]
und_idx, gen_idx = self._mode_indices(split_lens, attn_modes, device)
cos_und, sin_und = cos[und_idx], sin[und_idx]
cos_gen, sin_gen = cos[gen_idx], sin[gen_idx]
# 4. Dual-pathway decoder over (und = text causal, gen = vision full)
und_seq = packed_sequence[und_idx]
gen_seq = packed_sequence[gen_idx]
for layer in self.layers:
und_seq, gen_seq = layer(und_seq, gen_seq, cos_und, sin_und, cos_gen, sin_gen)
# 5. Final norms (und vs gen) and re-scatter into a joint buffer
und_seq = self.norm(und_seq)
gen_seq = self.norm_moe_gen(gen_seq)
last_hidden_state = packed_sequence.new_zeros((sequence_length, self.hidden_size))
last_hidden_state[und_idx] = und_seq
last_hidden_state[gen_idx] = gen_seq
# 6. Vision prediction: proj_out on noisy patches -> unpatchify
output: dict[str, Any] = {}
if vision_tokens and vision_mse_loss_indexes.numel() > 0:
preds, _ = self.proj_out(last_hidden_state[vision_mse_loss_indexes])
output["preds_vision"] = self._unpatchify_and_unpack(preds, vision_token_shapes,
vision_noisy_frame_indexes, original_shapes)
# 6b. Sound prediction (t2vs): audio_proj_out on noisy sound hidden
# states -> unpack to per-sample [C, T] (framework ``_decode_sound``).
if sound_tokens and sound_mse_loss_indexes is not None and sound_mse_loss_indexes.numel() > 0:
preds_sound, _ = self.audio_proj_out(last_hidden_state[sound_mse_loss_indexes])
output["preds_sound"] = self._unpack_sound(preds_sound, sound_token_shapes, sound_noisy_frame_indexes)
# 6c. Action prediction: action_proj_out(per-token domain) on noisy hidden
# states -> unpack to per-sample [T, D] (framework ``_decode_action``).
if action_tokens and action_mse_loss_indexes is not None and action_mse_loss_indexes.numel() > 0:
noisy_domain = torch.cat(
[d.reshape(1).expand(len(nfi)) for d, nfi in zip(action_domain_id, action_noisy_frame_indexes)], dim=0,
)
preds_action = self.action_proj_out(last_hidden_state[action_mse_loss_indexes], noisy_domain)
output["preds_action"] = self._unpack_action(preds_action, action_token_shapes, action_noisy_frame_indexes)
output["last_hidden_state"] = last_hidden_state
return output
def _unpack_action(
self,
packed_preds: torch.Tensor,
token_shapes: list[tuple[int, ...]],
noisy_frame_indexes: list[torch.Tensor],
) -> list[torch.Tensor]:
"""Scatter packed noisy action preds back to per-sample ``[T, D]`` (clean
frames left zero). Mirrors framework ``unpack_action``."""
outputs: list[torch.Tensor] = []
start_idx = 0
for shape, nfi in zip(token_shapes, noisy_frame_indexes):
t = shape[0]
out = torch.zeros((t, self.action_dim), device=packed_preds.device, dtype=packed_preds.dtype)
t_n = len(nfi)
if t_n > 0:
out[nfi] = packed_preds[start_idx:start_idx + t_n]
start_idx += t_n
outputs.append(out)
return outputs
def _unpack_sound(
self,
packed_preds: torch.Tensor,
token_shapes: list[tuple[int, int, int]],
noisy_frame_indexes: list[torch.Tensor],
) -> list[torch.Tensor]:
"""Scatter packed noisy sound preds back to per-sample ``[C, T]`` (clean
frames left zero). Mirrors framework ``unpack_sound_latents``."""
outputs: list[torch.Tensor] = []
start_idx = 0
for shape, nfi in zip(token_shapes, noisy_frame_indexes):
t = shape[0]
out = torch.zeros((self.sound_dim, t), device=packed_preds.device, dtype=packed_preds.dtype)
t_n = len(nfi)
if t_n > 0:
out[:, nfi] = packed_preds[start_idx:start_idx + t_n].T
start_idx += t_n
outputs.append(out)
return outputs
@staticmethod
def _mode_indices(split_lens: list[int], attn_modes: list[str],
device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
"""Token indexes for causal (und) and full (gen) splits, in pack order."""
und, gen = [], []
start = 0
for split_len, mode in zip(split_lens, attn_modes):
rng = range(start, start + split_len)
if mode == "causal":
und.extend(rng)
elif mode == "full":
gen.extend(rng)
start += split_len
return (torch.tensor(und, dtype=torch.long, device=device),
torch.tensor(gen, dtype=torch.long, device=device))
+15 -1
View File
@@ -92,6 +92,8 @@ class ComponentLoader(ABC):
"tokenizer": (TokenizerLoader, "transformers"),
"tokenizer_2": (TokenizerLoader, "transformers"),
"tokenizer_3": (TokenizerLoader, "transformers"),
# Cosmos3's model_index names its Qwen2 tokenizer "text_tokenizer".
"text_tokenizer": (TokenizerLoader, "transformers"),
"image_processor": (ImageProcessorLoader, "transformers"),
"feature_extractor": (ImageProcessorLoader, "transformers"),
"image_encoder": (ImageEncoderLoader, "transformers"),
@@ -1009,7 +1011,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
+3
View File
@@ -35,6 +35,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"),
+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,654 @@
# 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)
# ===========================================================================
# 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,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()
+20
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.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.gen3c import Gen3CConfig
@@ -555,6 +556,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,
@@ -892,6 +909,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.gamecraft.presets import (
ALL_PRESETS as GAMECRAFT_PRESETS, )
from fastvideo.pipelines.basic.gen3c.presets import (
@@ -923,6 +942,7 @@ def _register_presets() -> None:
all_preset_groups = (
COSMOS_PRESETS,
COSMOS3_PRESETS,
GAMECRAFT_PRESETS,
GEN3C_PRESETS,
HUNYUAN_PRESETS,
@@ -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()
@@ -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.
+124
View File
@@ -0,0 +1,124 @@
# 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: `PR1 (video core) + PR2 (audio/t2vs) COMPLETE, real-weights verified. Video core (T2V/I2V/T2I) + audio (AVAE decode / DiT sound pathway / sound packing / t2vs CFG velocity) all framework-parity verified bit-exact (suite 130 passed, 0 skipped). t2vs real-weights on B200 produces coherent video + real stereo 48kHz audio. Branch chain: feat/cosmos3-tier-a-port (T2V) -> feat/cosmos3-i2v (I2V+T2I) -> feat/cosmos3-audio (t2vs). Next: PR3 action / PR4 reasoning.`
- 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.
+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)
@@ -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)
@@ -0,0 +1,70 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 ``[B,C,T,H,W] <-> [B, T*hp*wp, p*p*C]`` patchify roundtrip (Tier A).
Reference: ``vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py``
lines 1009-1036 (``Cosmos3VFMTransformer.patchify`` /
``Cosmos3VFMTransformer.unpatchify``) and the reference assertion at
``tests/diffusion/models/cosmos3/test_cosmos3_transformer.py:98-101``.
Invariant: ``unpatchify(patchify(x)) == x`` for any ``x`` with shape
``(B, C, t, h, w)`` where ``h, w`` are divisible by ``latent_patch_size``.
Also exercises a non-trivial channel count (3) to ensure the
``permute([0, 2, 3, 5, 4, 6, 1])`` reordering is correct.
"""
from __future__ import annotations
import pytest
import torch
pytestmark = [pytest.mark.local]
def test_patchify_unpatchify_roundtrip() -> None:
"""Asserts that the FastVideo Cosmos3 transformer's patchify/unpatchify
pair are exact inverses for ``latent_patch_size=2``, ``latent_channel=3``.
Once FastVideo's ``fastvideo.models.dits.cosmos3.Cosmos3VFMTransformer``
lands, replace the skip with the upstream-equivalent assertion path.
"""
try:
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer # type: ignore
except ImportError:
pytest.skip("FastVideo Cosmos3 not yet implemented (Phase 2b)")
from torch import nn
model = object.__new__(Cosmos3VFMTransformer)
nn.Module.__init__(model)
model.latent_patch_size = 2
model.latent_channel_size = 3
latents = torch.arange(1 * 3 * 1 * 3 * 5, dtype=torch.float32).reshape(1, 3, 1, 3, 5)
roundtrip = model.unpatchify(model.patchify(latents, t=1, h=3, w=5), t=1, h=3, w=5)
torch.testing.assert_close(roundtrip, latents)
def test_patchify_default_patch_size() -> None:
"""Asserts shape contract for the default ``latent_patch_size=[1,2,2]``
(i.e. spatial-only patching) with a representative video latent.
With ``[B,C,T,H,W] = [1, 16, 2, 8, 8]`` and patch=2 on H/W, expected
flattened tokens = ``T * (H/2) * (W/2) = 2 * 4 * 4 = 32`` and each token
carries ``2*2*C = 4*16 = 64`` channels.
"""
try:
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer # type: ignore
except ImportError:
pytest.skip("FastVideo Cosmos3 not yet implemented (Phase 2b)")
from torch import nn
model = object.__new__(Cosmos3VFMTransformer)
nn.Module.__init__(model)
model.latent_patch_size = 2
model.latent_channel_size = 16
latents = torch.zeros(1, 16, 2, 8, 8)
tokens = model.patchify(latents, t=2, h=8, w=8)
assert tuple(tokens.shape) == (1, 32, 64)
restored = model.unpatchify(tokens, t=2, h=8, w=8)
assert tuple(restored.shape) == (1, 16, 2, 8, 8)
@@ -0,0 +1,201 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 native-pipeline call-graph contract (Tier A, no real weights).
Pins the runtime call graph of the FastVideo-native Cosmos3 pipeline against the
framework math, using the stub components from ``conftest.py`` (no real weights,
no ``cosmos_framework``). The native pipeline replaced the vllm-omni-derived
``diffuse``/``forward(req)``/``reset_cache`` skeleton with a stage-based
``Cosmos3DenoisingStage`` + ``Cosmos3DenoiseEngine`` doing SEQUENTIAL CFG.
Invariants under test:
1. SEQUENTIAL CFG order — per UniPC step, the transformer is called twice,
conditional (prompt tokens) then unconditional (negative-prompt tokens), in
that order; over N steps the call order is ``[cond, uncond] * N``.
2. CFG combination — the per-step velocity equals
``uncond + guidance * (cond - uncond)`` (verified against the stub's known
per-token output) and one UniPC step advances the latent accordingly.
3. I2V conditioning — a conditioning image is VAE-encoded and frame 0 is kept
clean: its velocity is zeroed (condition mask) so the decoded clip's
frame-0 latent equals the clean conditioning latent.
4. Mode dispatch — the stage routes T2V (num_frames>1, flow_shift=10.0,
``is_video`` tokenization) vs T2I (num_frames==1, flow_shift=3.0, image
tokenization), applying the per-mode ``flow_shift``.
"""
from __future__ import annotations
import pytest
import torch
from .conftest import make_fastvideo_args, make_forward_batch
pytestmark = [pytest.mark.local]
_LATENT_CHANNEL = 16
_LATENT_PATCH_SIZE = 2
_TEMPORAL_FACTOR = 4
_COND_TOKEN = 2
_UNCOND_TOKEN = 1
def _engine(pipeline, scheduler):
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import Cosmos3DenoiseEngine
return Cosmos3DenoiseEngine(
transformer=pipeline.modules["transformer"],
scheduler=scheduler,
special_tokens={"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62},
latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=15_000,
reset_spatial_ids=True,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TEMPORAL_FACTOR,
)
def test_sequential_cfg_calls_cond_then_uncond_each_step(make_cosmos3_pipeline) -> None:
"""Each UniPC step calls the transformer cond-then-uncond, in order."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import Cosmos3VisionSpec
pipeline = make_cosmos3_pipeline()
scheduler = pipeline.modules["scheduler"]
scheduler.set_timesteps(2, device=torch.device("cpu"))
engine = _engine(pipeline, scheduler)
shape = (_LATENT_CHANNEL, 2, 2, 2)
flat = torch.randn(int(torch.tensor(shape).prod()))
spec = Cosmos3VisionSpec(shape=shape, condition_frame_indexes=[])
engine.denoise(
flat_latent=flat,
timesteps=scheduler.timesteps,
guidance=6.0,
specs=[spec],
cond_token_ids=[_COND_TOKEN, 5, 6],
uncond_token_ids=[_UNCOND_TOKEN, 7],
)
tokens = [c["token"] for c in pipeline.modules["transformer"].calls]
# 2 steps -> 4 calls: cond, uncond, cond, uncond.
assert tokens == [_COND_TOKEN, _UNCOND_TOKEN, _COND_TOKEN, _UNCOND_TOKEN]
def test_cfg_velocity_combination_formula(make_cosmos3_pipeline) -> None:
"""The per-step velocity equals ``uncond + g*(cond - uncond)``.
The stub returns ``scale(token) * tanh(latent)`` on noisy frames, so the
expected velocity is a closed form we can check exactly.
"""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3VisionSpec,
cosmos3_get_cfg_velocity,
)
pipeline = make_cosmos3_pipeline()
transformer = pipeline.modules["transformer"]
shape = (_LATENT_CHANNEL, 2, 2, 2)
flat = torch.randn(int(torch.tensor(shape).prod()))
guidance = 6.0
v = cosmos3_get_cfg_velocity(
transformer=transformer,
flat_latent=flat,
timestep=torch.tensor([500.0]),
guidance=guidance,
specs=[Cosmos3VisionSpec(shape=shape, condition_frame_indexes=[])],
cond_token_ids=[_COND_TOKEN, 5, 6],
uncond_token_ids=[_UNCOND_TOKEN, 7],
special_tokens={"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62},
latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=15_000,
reset_spatial_ids=True,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TEMPORAL_FACTOR,
)
lat = flat.reshape(shape)
scale_cond = 0.01 * (1.0 + (_COND_TOKEN % 7))
scale_uncond = 0.01 * (1.0 + (_UNCOND_TOKEN % 7))
cond_v = (scale_cond * torch.tanh(lat)).reshape(-1)
uncond_v = (scale_uncond * torch.tanh(lat)).reshape(-1)
expected = uncond_v + guidance * (cond_v - uncond_v)
torch.testing.assert_close(v, expected, atol=1e-6, rtol=1e-5)
def test_i2v_keeps_condition_frame_clean(make_cosmos3_pipeline) -> None:
"""I2V: frame-0 velocity is zeroed so the conditioning frame stays clean."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3VisionSpec,
cosmos3_get_cfg_velocity,
)
pipeline = make_cosmos3_pipeline()
transformer = pipeline.modules["transformer"]
shape = (_LATENT_CHANNEL, 3, 2, 2) # 3 latent frames, frame 0 conditioned
flat = torch.randn(int(torch.tensor(shape).prod()))
v = cosmos3_get_cfg_velocity(
transformer=transformer,
flat_latent=flat,
timestep=torch.tensor([500.0]),
guidance=6.0,
specs=[Cosmos3VisionSpec(shape=shape, condition_frame_indexes=[0])],
cond_token_ids=[_COND_TOKEN, 5, 6],
uncond_token_ids=[_UNCOND_TOKEN, 7],
special_tokens={"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62},
latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=15_000,
reset_spatial_ids=True,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TEMPORAL_FACTOR,
)
v_grid = v.reshape(shape) # [C, T, H, W]
# Condition frame 0 velocity must be exactly zero; noisy frames non-zero.
assert torch.count_nonzero(v_grid[:, 0]) == 0
assert torch.count_nonzero(v_grid[:, 1:]) > 0
def test_stage_mode_dispatch_t2v(make_cosmos3_pipeline, make_cosmos3_stage) -> None:
"""T2V (num_frames>1): resolution-based flow_shift, video tokenization."""
pipeline = make_cosmos3_pipeline()
stage = make_cosmos3_stage(pipeline)
args = make_fastvideo_args()
batch = make_forward_batch(num_frames=5, height=16, width=16)
out = stage.forward(batch, args)
# flow_shift is resolution-based (not task-based): 16x16 -> "256" bucket -> 3.0.
# (full resolution->shift parity in test_cosmos3_flow_shift_parity.)
assert float(pipeline.scheduler.config.flow_shift) == 3.0
assert out.output is not None and out.output.dim() == 5
# T2V latent: (5-1)//4 + 1 = 2 frames; 16/8 = 2 latent h/w.
assert tuple(out.latents.shape) == (1, _LATENT_CHANNEL, 2, 2, 2)
def test_stage_mode_dispatch_t2i(make_cosmos3_pipeline, make_cosmos3_stage) -> None:
"""T2I (num_frames==1): single-frame latent; resolution-based flow_shift."""
pipeline = make_cosmos3_pipeline()
stage = make_cosmos3_stage(pipeline)
args = make_fastvideo_args()
batch = make_forward_batch(num_frames=1, height=16, width=16, guidance_scale=4.0)
out = stage.forward(batch, args)
# 16x16 -> "256" bucket -> 3.0 (resolution-based, not task-based).
assert float(pipeline.scheduler.config.flow_shift) == 3.0
assert tuple(out.latents.shape) == (1, _LATENT_CHANNEL, 1, 2, 2)
def test_stage_i2v_encodes_conditioning_image(make_cosmos3_pipeline, make_cosmos3_stage) -> None:
"""I2V stage: a conditioning image is accepted and decoded to a finite clip."""
pipeline = make_cosmos3_pipeline()
stage = make_cosmos3_stage(pipeline)
args = make_fastvideo_args()
image = torch.zeros(3, 16, 16) # [-1, 1] conditioning frame
batch = make_forward_batch(num_frames=5, height=16, width=16, image=image)
out = stage.forward(batch, args)
# flow_shift is resolution-based (not task-based): 16x16 -> 3.0.
assert float(pipeline.scheduler.config.flow_shift) == 3.0
assert out.output is not None and out.output.dim() == 5
assert torch.isfinite(out.output).all()
@@ -0,0 +1,313 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 pipeline smoke test (CPU / float32, tiny components).
Exercises the native pipeline's runtime path end-to-end on tiny stub-or-real
components, with NO real weights and NO cosmos_framework dependency:
* a stub transformer implementing the native DiT's packed-input ->
``{"preds_vision": [...]}`` contract with bounded, deterministic output
(the REAL DiT forward / unpatchify is exhaustively bit-identical-tested in
``test_cosmos3_dit_parity*`` and ``test_cosmos3_denoise_cfg_parity``; an
untrained real DiT emits unbounded velocities that overflow UniPC, so the
smoke uses a stub to keep finiteness deterministic);
* a stub VAE exposing the ``AutoencoderKLWan`` surface
(``config.latents_mean/std/scale_factor_spatial``, ``encode().mode()``,
``decode()``) used by the encode/normalize + decode/denormalize bridges;
* a stub Qwen2-shaped tokenizer (chat template + special tokens).
Two paths are covered:
1. ``Cosmos3DenoiseEngine.denoise`` for >= 2 UniPC steps over a tiny T2V
latent, asserting a finite final latent of the right shape, plus the VAE
decode + ``(1 + x)/2`` clamp producing a finite ``[B, 3, T, H, W]`` video;
2. the real ``Cosmos3DenoisingStage.forward`` (full tokenize -> noise ->
denoise -> decode wiring) driven through a ``__new__``-built pipeline +
``ForwardBatch``, for both T2V and I2V.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_pipeline_smoke.py -q
"""
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,
)
pytestmark = [pytest.mark.local]
_LATENT_CHANNEL = 16
_LATENT_PATCH_SIZE = 2
_SPATIAL_FACTOR = 8
_TEMPORAL_FACTOR = 4
# ---------------------------------------------------------------------------
# Stub transformer implementing the native DiT packed-input contract.
# ---------------------------------------------------------------------------
class StubCosmos3Transformer(torch.nn.Module):
"""Bounded stand-in for ``Cosmos3VFMTransformer.forward``.
Consumes the same packed kwargs (``vision_token_shapes`` /
``vision_noisy_frame_indexes`` / ``vision_tokens`` / ``text_ids``) and
returns ``{"preds_vision": [[1, C, T, H, W], ...]}`` with predictions only on
noisy frames (zeros on conditioning frames), matching the real DiT's
``_unpatchify_and_unpack`` output structure. The prediction is a small
``tanh`` of the input latent, scaled by the first text id so the cond and
uncond passes differ (exercising the CFG combination).
"""
def __init__(self, latent_channel: int = _LATENT_CHANNEL) -> None:
super().__init__()
self.latent_channel = latent_channel
# A real attribute the stage reads for device/dtype.
self.embed_tokens = torch.nn.Embedding(64, 8)
def forward(self, **kwargs: Any) -> dict[str, Any]:
token_ids = kwargs["text_ids"]
token = float(token_ids.reshape(-1)[0].item()) if token_ids.numel() else 1.0
scale = 0.01 * (1.0 + (token % 7))
token_shapes = kwargs["vision_token_shapes"]
noisy = kwargs["vision_noisy_frame_indexes"]
tokens = kwargs["vision_tokens"]
preds: list[torch.Tensor] = []
for latent, (t, _h, _w), nfi in zip(tokens, token_shapes, noisy):
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)) # [1, C, T, H, W]
return {"preds_vision": preds}
# ---------------------------------------------------------------------------
# Stub VAE matching the AutoencoderKLWan surface used by the bridges.
# ---------------------------------------------------------------------------
class _StubLatentDist:
def __init__(self, latents: torch.Tensor) -> None:
self._latents = latents
def mode(self) -> torch.Tensor:
return self._latents
class StubCosmos3VAE:
"""Minimal VAE: deterministic encode/decode shaped by 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
latents = torch.ones(b, self.config.z_dim, lt, lh, lw, dtype=video.dtype, device=video.device)
return _StubLatentDist(latents)
def decode(self, z: torch.Tensor):
# Upsample latents back to pixel dims; clamp like AutoencoderKLWan.
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
# Bounded function of z so the output reflects (and stays finite with)
# the latent: tanh maps any finite z to [-1, 1]; nan_to_num guards
# against non-finite latents from an untrained denoise.
z_signal = torch.nan_to_num(torch.tanh(z[:, :1, :1, :1, :1]))
out = torch.zeros(b, 3, t, h, w, dtype=z.dtype, device=z.device) + z_signal.reshape(b, 1, 1, 1, 1)
return torch.clamp(out, -1.0, 1.0)
# ---------------------------------------------------------------------------
# Stub Qwen2-shaped tokenizer (chat template + special tokens).
# ---------------------------------------------------------------------------
class StubQwen2Tokenizer:
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):
# Deterministic small token ids from the user message length.
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() -> UniPCMultistepScheduler:
return UniPCMultistepScheduler(
num_train_timesteps=1000,
solver_order=2,
prediction_type="flow_prediction",
use_flow_sigmas=True,
flow_shift=10.0,
)
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestCosmos3PipelineSmoke:
def test_denoise_engine_and_decode_finite(self):
"""Engine.denoise (>= 2 steps) + VAE decode -> finite output, right shape."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3DenoiseEngine,
Cosmos3VisionSpec,
_VaeNorm,
cosmos3_vae_decode,
)
dit = StubCosmos3Transformer()
vae = StubCosmos3VAE()
scheduler = _make_scheduler()
scheduler.set_timesteps(2, device=torch.device("cpu"))
latent_shape = (_LATENT_CHANNEL, 2, 2, 2) # [C, T, H, W]
torch.manual_seed(1)
flat = torch.randn(int(torch.tensor(latent_shape).prod()))
engine = Cosmos3DenoiseEngine(
transformer=dit,
scheduler=scheduler,
special_tokens={"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62},
latent_patch_size=_LATENT_PATCH_SIZE,
temporal_modality_margin=15_000,
reset_spatial_ids=True,
enable_fps_modulation=False,
base_fps=24.0,
temporal_compression_factor=_TEMPORAL_FACTOR,
)
spec = Cosmos3VisionSpec(shape=latent_shape, condition_frame_indexes=[])
out_flat = engine.denoise(
flat_latent=flat,
timesteps=scheduler.timesteps,
guidance=6.0,
specs=[spec],
cond_token_ids=[10, 11, 12],
uncond_token_ids=[13, 14],
)
assert out_flat.shape == flat.shape
assert torch.isfinite(out_flat).all()
norm = _VaeNorm.from_vae(vae, torch.float32)
result_latent = out_flat.reshape(latent_shape).unsqueeze(0)
decoded = cosmos3_vae_decode(vae, result_latent, norm)
video = ((1.0 + decoded) / 2.0).clamp(0.0, 1.0)
assert video.dim() == 5 and video.shape[1] == 3
assert torch.isfinite(video).all()
assert float(video.min()) >= 0.0 and float(video.max()) <= 1.0
def _make_pipeline(self):
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import 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
return pipe
def _make_args(self):
from fastvideo.configs.pipelines.cosmos3 import Cosmos3Config
cfg = Cosmos3Config()
# Shrink the DiT arch config to the tiny smoke geometry.
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)
def _make_batch(self, *, num_frames, height, width, image=None):
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
return ForwardBatch(
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,
)
def test_stage_forward_t2v_finite(self):
"""The real Cosmos3DenoisingStage.forward runs T2V end-to-end."""
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
pipe = self._make_pipeline()
args = self._make_args()
stage = Cosmos3DenoisingStage(
transformer=pipe.modules["transformer"],
scheduler=pipe.modules["scheduler"],
vae=pipe.modules["vae"],
tokenizer=pipe.modules["text_tokenizer"],
pipeline=pipe,
)
# 5 frames -> latent_t = (5-1)//4 + 1 = 2; 16x16 px -> 2x2 latent.
batch = self._make_batch(num_frames=5, height=16, width=16)
out = stage.forward(batch, args)
assert out.output is not None
assert out.output.dim() == 5 and out.output.shape[1] == 3
assert torch.isfinite(out.output).all()
assert torch.isfinite(out.latents).all()
def test_stage_forward_i2v_keeps_condition_frame(self):
"""I2V: a conditioning image is VAE-encoded and frame 0 stays clean."""
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
pipe = self._make_pipeline()
args = self._make_args()
stage = Cosmos3DenoisingStage(
transformer=pipe.modules["transformer"],
scheduler=pipe.modules["scheduler"],
vae=pipe.modules["vae"],
tokenizer=pipe.modules["text_tokenizer"],
pipeline=pipe,
)
# Conditioning image as a [3, H, W] tensor in [-1, 1].
image = torch.zeros(3, 16, 16)
batch = self._make_batch(num_frames=5, height=16, width=16, image=image)
out = stage.forward(batch, args)
assert out.output is not None
assert out.output.dim() == 5 and out.output.shape[1] == 3
assert torch.isfinite(out.output).all()
def test_tokenize_caption_special_tokens(self):
"""Pipeline.tokenize_caption uses the Qwen2 chat template + ids."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import cosmos3_special_tokens
pipe = self._make_pipeline()
ids = pipe.tokenize_caption("hello world", is_video=True, use_system_prompt=False)
assert isinstance(ids, list) and len(ids) > 0
special = cosmos3_special_tokens(pipe.get_module("text_tokenizer"))
assert special == {"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62}
@@ -0,0 +1,480 @@
# SPDX-License-Identifier: Apache-2.0
"""
Numerical-parity reference test for the official Cosmos3 DiT (Cosmos3VFMNetwork).
Runs a tiny deterministic forward of the OFFICIAL framework model on CPU / float32
using a torch SDPA monkey-patch (flash2/flash3/natten are CUDA-only; SDPA works on CPU).
The test exercises the full forward contract:
packed_seq -> vfm(packed_seq) -> {last_hidden_state, preds_vision}
and is used as the "ground truth" side of any FastVideo parity check.
Environment requirements
------------------------
- cosmos_framework must be installed (editable) in the active interpreter.
The canonical env is: /home/william5lin/miniconda3/envs/fv-cosmos3/bin/python
- No transformer_engine / natten / GPU required.
- PYTHONSAFEPATH=1 is recommended to avoid cwd import shadowing.
Run:
PYTHONSAFEPATH=1 pytest fastvideo/tests/layers/test_cosmos3_reference_forward.py -v
"""
import math
import sys
import pytest
import torch
import torch.nn.functional as F
# ---------------------------------------------------------------------------
# Skip guard: cosmos_framework may not be installed in the default dev env.
# ---------------------------------------------------------------------------
cosmos_framework = pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
pytestmark = [pytest.mark.local]
# ---------------------------------------------------------------------------
# SDPA attention monkey-patch
# ---------------------------------------------------------------------------
# The imaginaire attention backend (flash2/flash3) requires CUDA *and*
# float16/bfloat16. For CPU/float32 parity testing we replace it with a
# simple SDPA implementation that handles both standard and varlen packed
# formats (cumulative_seqlen_{Q,KV}).
# ---------------------------------------------------------------------------
def _sdpa_attention(
query,
key,
value,
is_causal=False,
causal_type=None,
scale=None,
seqlens_Q=None,
seqlens_KV=None,
cumulative_seqlen_Q=None,
cumulative_seqlen_KV=None,
max_seqlen_Q=None,
max_seqlen_KV=None,
backend=None,
return_lse=False,
backend_kwargs=None,
deterministic=False,
):
"""Minimal SDPA wrapper that mirrors the imaginaire attention signature."""
B, Sq, H, D = query.shape
Hkv = key.shape[2]
attn_scale = scale if scale is not None else D**-0.5
if cumulative_seqlen_Q is not None:
# Varlen packed layout: B==1, tokens from different samples are concatenated.
oq = cumulative_seqlen_Q.cpu().tolist()
okv = cumulative_seqlen_KV.cpu().tolist()
outs = []
for i in range(len(oq) - 1):
qi = query[0, oq[i] : oq[i + 1]].unsqueeze(0).permute(0, 2, 1, 3) # [1,H,S,D]
ki = key[0, okv[i] : okv[i + 1]].unsqueeze(0).permute(0, 2, 1, 3)
vi = value[0, okv[i] : okv[i + 1]].unsqueeze(0).permute(0, 2, 1, 3)
if Hkv != H:
ki = ki.repeat_interleave(H // Hkv, dim=1)
vi = vi.repeat_interleave(H // Hkv, dim=1)
oi = F.scaled_dot_product_attention(qi, ki, vi, scale=attn_scale, is_causal=is_causal)
outs.append(oi.permute(0, 2, 1, 3)) # [1,S,H,D]
out = torch.cat(outs, dim=1) # [1,S_total,H,D]
else:
q = query.permute(0, 2, 1, 3)
k = key.permute(0, 2, 1, 3)
v = value.permute(0, 2, 1, 3)
if Hkv != H:
k = k.repeat_interleave(H // Hkv, dim=1)
v = v.repeat_interleave(H // Hkv, dim=1)
out = F.scaled_dot_product_attention(q, k, v, scale=attn_scale, is_causal=is_causal)
out = out.permute(0, 2, 1, 3) # [B,S,H,D]
if return_lse:
lse = torch.zeros(B, Sq, H, 1, dtype=query.dtype, device=query.device)
return out, lse
return out
def _sdpa_merge_attentions(outputs, lse_tensors, torch_compile=False):
"""Log-sum-exp weighted merge of two attention outputs."""
if len(outputs) == 1:
return outputs[0], lse_tensors[0]
o1, lse1 = outputs[0], lse_tensors[0]
o2, lse2 = outputs[1], lse_tensors[1]
m = torch.maximum(lse1, lse2)
w1 = torch.exp(lse1 - m)
w2 = torch.exp(lse2 - m)
ws = w1 + w2
return (o1 * w1 + o2 * w2) / ws, m + torch.log(ws)
def _apply_sdpa_patches():
"""Monkey-patch every attention reference in cosmos_framework to use SDPA."""
import cosmos_framework.model.attention as attn_pkg
import cosmos_framework.model.attention.frontend as attn_frontend
import cosmos_framework.model.vfm.mot.attention as vfm_attn
import cosmos_framework.model.vfm.mot.unified_mot as mot_module
attn_frontend.attention = _sdpa_attention
attn_pkg.attention = _sdpa_attention
mot_module.imaginaire_attention = _sdpa_attention
vfm_attn.attention = _sdpa_attention
attn_frontend.merge_attentions = _sdpa_merge_attentions
attn_pkg.merge_attentions = _sdpa_merge_attentions
vfm_attn.merge_attentions = _sdpa_merge_attentions
# Apply patches at import time (before any model is constructed).
_apply_sdpa_patches()
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _build_tiny_cosmos3(seed: int = 42):
"""Construct a tiny Cosmos3VFMNetwork on CPU / float32.
Architecture:
hidden_size=16, intermediate_size=32, num_hidden_layers=1,
num_attention_heads=2, num_key_value_heads=2, head_dim=8,
vocab_size=64, latent_channel_size=16, latent_patch_size=2,
max_latent_{h,w,t}=8,8,4
"""
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=1,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
rms_norm_eps=1e-6,
attention_bias=False,
attention_dropout=0.0,
)
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)
vfm_cfg = Cosmos3VFMNetworkConfig(
vision_gen=True,
vlm_config=tiny_vlm_cfg,
latent_patch_size=2,
latent_downsample_factor=8,
latent_channel_size=16,
position_embedding_type="3d_rope",
max_latent_h=8,
max_latent_w=8,
max_latent_t=4,
)
torch.manual_seed(seed)
lm = Qwen3VLTextForCausalLM(config=mot_cfg)
vfm = Cosmos3VFMNetwork(language_model=lm, config=vfm_cfg)
vfm.eval()
return vfm
def _build_tiny_packed_seq(*, n_text: int = 4, seed: int = 7):
"""Build a minimal PackedSequence: 4 text tokens + 1 vision patch.
Vision: C=16, T=1, H=2, W=2 → after patch_size=2: 1*1*1 = 1 patch.
All vision frames are noisy (timestep=500).
"""
from cosmos_framework.data.vfm.sequence_packing import ModalityData, PackedSequence
torch.manual_seed(seed)
vision_tensor = torch.randn(16, 1, 2, 2) # [C=16, T=1, H=2, W=2]
text_ids = torch.randint(0, 64, (n_text,))
n_vision = 1 # 1 patch after patchify
total_len = n_text + n_vision
vision_mod = ModalityData(
sequence_indexes=torch.arange(n_text, total_len, dtype=torch.long),
timesteps=torch.tensor([500.0]), # one noisy frame
mse_loss_indexes=torch.arange(n_text, total_len, dtype=torch.long),
token_shapes=[(1, 1, 1)], # (t_patches, h_patches, w_patches) = (1,1,1)
tokens=[vision_tensor],
condition_mask=[torch.zeros(1, dtype=torch.long)], # 0=noisy
noisy_frame_indexes=[torch.tensor([0])],
)
packed_seq = PackedSequence(
sample_lens=[total_len],
split_lens=[n_text, n_vision],
attn_modes=["causal", "full"],
is_image_batch=True,
sequence_length=total_len,
text_ids=text_ids,
text_indexes=torch.arange(n_text, dtype=torch.long),
position_ids=torch.arange(total_len),
vision=vision_mod,
)
return packed_seq
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestCosmos3ReferenceConfig:
"""Step 1: verify config construction and field enumeration."""
def test_qwen3vl_text_config_fields(self):
from cosmos_framework.model.vfm.vlm.qwen3_vl.configuration_qwen3_vl import Qwen3VLTextConfig
cfg = Qwen3VLTextConfig(
vocab_size=64,
hidden_size=16,
intermediate_size=32,
num_hidden_layers=1,
num_attention_heads=2,
num_key_value_heads=2,
head_dim=8,
)
assert cfg.vocab_size == 64
assert cfg.hidden_size == 16
assert cfg.num_attention_heads == 2
assert cfg.num_key_value_heads == 2
assert cfg.head_dim == 8
def test_cosmos3_vfm_network_config(self):
from cosmos_framework.model.vfm.mot.cosmos3_vfm_network import Cosmos3VFMNetworkConfig
from cosmos_framework.model.vfm.vlm.qwen3_vl.configuration_qwen3_vl import Qwen3VLConfig
tiny_vlm = Qwen3VLConfig(text_config=dict(vocab_size=64, hidden_size=16))
cfg = Cosmos3VFMNetworkConfig(
vision_gen=True,
vlm_config=tiny_vlm,
latent_patch_size=2,
latent_downsample_factor=8,
latent_channel_size=16,
position_embedding_type="3d_rope",
max_latent_h=8,
max_latent_w=8,
max_latent_t=4,
)
assert cfg.vision_gen is True
assert cfg.latent_patch_size == 2
assert cfg.latent_channel_size == 16
class TestCosmos3ReferenceInstantiation:
"""Step 2: verify parameter key pattern and module tree."""
@pytest.fixture(scope="class")
def tiny_vfm(self):
return _build_tiny_cosmos3(seed=42)
def test_instantiation_succeeds(self, tiny_vfm):
assert tiny_vfm is not None
def test_param_key_pattern_attention(self, tiny_vfm):
"""Verify understanding and generation attention projections exist."""
param_keys = {n for n, _ in tiny_vfm.named_parameters()}
layer = "language_model.model.layers.0.self_attn"
# Understanding pathway
for proj in ("q_proj", "k_proj", "v_proj", "o_proj"):
assert f"{layer}.{proj}.weight" in param_keys, f"Missing {layer}.{proj}.weight"
# Generation pathway (moe_gen suffix)
for proj in ("q_proj_moe_gen", "k_proj_moe_gen", "v_proj_moe_gen", "o_proj_moe_gen"):
assert f"{layer}.{proj}.weight" in param_keys, f"Missing {layer}.{proj}.weight"
# QK norms
for norm in ("q_norm", "k_norm", "q_norm_moe_gen", "k_norm_moe_gen"):
assert f"{layer}.{norm}.weight" in param_keys, f"Missing {layer}.{norm}.weight"
def test_param_key_pattern_mlp(self, tiny_vfm):
param_keys = {n for n, _ in tiny_vfm.named_parameters()}
for pathway in ("mlp", "mlp_moe_gen"):
base = f"language_model.model.layers.0.{pathway}"
for proj in ("gate_proj", "up_proj", "down_proj"):
assert f"{base}.{proj}.weight" in param_keys
def test_param_key_pattern_layernorms(self, tiny_vfm):
param_keys = {n for n, _ in tiny_vfm.named_parameters()}
layer = "language_model.model.layers.0"
for ln in (
"input_layernorm",
"input_layernorm_moe_gen",
"post_attention_layernorm",
"post_attention_layernorm_moe_gen",
):
assert f"{layer}.{ln}.weight" in param_keys
def test_param_key_pattern_toplevel(self, tiny_vfm):
param_keys = {n for n, _ in tiny_vfm.named_parameters()}
# LM submodules
assert "language_model.model.embed_tokens.weight" in param_keys
assert "language_model.model.norm.weight" in param_keys
assert "language_model.model.norm_moe_gen.weight" in param_keys
assert "language_model.lm_head.weight" in param_keys
# VFM vision head
assert "vae2llm.weight" in param_keys
assert "vae2llm.bias" in param_keys
assert "llm2vae.weight" in param_keys
assert "llm2vae.bias" in param_keys
# Timestep embedder
assert "time_embedder.mlp.0.weight" in param_keys
assert "time_embedder.mlp.2.weight" in param_keys
def test_expected_param_count(self, tiny_vfm):
"""Sanity-check total param count for the tiny model."""
n_params = sum(p.numel() for p in tiny_vfm.parameters())
# Rough bound: tiny model should be under 50 K params
assert n_params < 50_000, f"Unexpected param count: {n_params}"
def test_dtype_is_float32(self, tiny_vfm):
for name, p in tiny_vfm.named_parameters():
if "inv_freq" in name:
continue # inv_freq stays float32 always
assert p.dtype == torch.float32, f"{name} has dtype {p.dtype}"
class TestCosmos3ReferenceForward:
"""Step 3 + 4: verify forward contract and determinism."""
@pytest.fixture(scope="class")
def tiny_vfm(self):
return _build_tiny_cosmos3(seed=42)
@pytest.fixture(scope="class")
def packed_seq(self):
return _build_tiny_packed_seq(n_text=4, seed=7)
@pytest.fixture(scope="class")
def fwd_output(self, tiny_vfm, packed_seq):
with torch.no_grad():
return tiny_vfm(packed_seq=packed_seq)
def test_forward_returns_dict(self, fwd_output):
assert isinstance(fwd_output, dict)
def test_last_hidden_state_shape(self, fwd_output):
# 4 text + 1 vision patch = 5 total tokens; hidden_size=16
lhs = fwd_output["last_hidden_state"]
assert lhs.shape == torch.Size([5, 16]), f"Got {lhs.shape}"
def test_last_hidden_state_finite(self, fwd_output):
assert torch.isfinite(fwd_output["last_hidden_state"]).all()
def test_preds_vision_present(self, fwd_output):
assert "preds_vision" in fwd_output
def test_preds_vision_shape(self, fwd_output):
# latent_channel=16, T=1, H=2, W=2 → [1, 16, 1, 2, 2]
pv = fwd_output["preds_vision"][0]
assert pv.shape == torch.Size([1, 16, 1, 2, 2]), f"Got {pv.shape}"
def test_preds_vision_finite(self, fwd_output):
assert torch.isfinite(fwd_output["preds_vision"][0]).all()
def test_forward_deterministic_same_seed(self):
"""Two models with identical seed should produce identical output."""
ps = _build_tiny_packed_seq(n_text=4, seed=7)
vfm1 = _build_tiny_cosmos3(seed=42)
vfm2 = _build_tiny_cosmos3(seed=42)
with torch.no_grad():
out1 = vfm1(packed_seq=ps)
out2 = vfm2(packed_seq=ps)
assert torch.allclose(out1["last_hidden_state"], out2["last_hidden_state"])
assert torch.allclose(out1["preds_vision"][0], out2["preds_vision"][0])
def test_forward_repeatable_same_model(self, tiny_vfm, packed_seq):
"""Same model, same input → identical output on two calls."""
with torch.no_grad():
out1 = tiny_vfm(packed_seq=packed_seq)
out2 = tiny_vfm(packed_seq=packed_seq)
assert torch.allclose(out1["last_hidden_state"], out2["last_hidden_state"])
def test_different_seed_gives_different_output(self):
"""Different model seeds should give different outputs."""
ps = _build_tiny_packed_seq(n_text=4, seed=7)
vfm1 = _build_tiny_cosmos3(seed=42)
vfm2 = _build_tiny_cosmos3(seed=99)
with torch.no_grad():
out1 = vfm1(packed_seq=ps)
out2 = vfm2(packed_seq=ps)
# With very high probability random init → different outputs
assert not torch.allclose(out1["last_hidden_state"], out2["last_hidden_state"])
def test_float32_dtype_preserved(self, fwd_output):
assert fwd_output["last_hidden_state"].dtype == torch.float32
def test_cpu_device(self, fwd_output):
assert fwd_output["last_hidden_state"].device.type == "cpu"
class TestCosmos3ReasonerForward:
"""Optional: reasoner (und-only) pathway via standard [B,T,H] layout."""
def test_reasoner_forward_shape(self):
"""reasoner_forward runs the und tower only; no PackedSequence needed."""
vfm = _build_tiny_cosmos3(seed=42)
lm = vfm.language_model
input_ids = torch.randint(0, 64, (1, 6))
with torch.no_grad():
out = lm.model.reasoner_forward(input_ids=input_ids, cache=None)
# [B=1, T=6, hidden_size=16]
assert out.shape == torch.Size([1, 6, 16])
def test_reasoner_forward_finite(self):
vfm = _build_tiny_cosmos3(seed=42)
lm = vfm.language_model
input_ids = torch.randint(0, 64, (1, 6))
with torch.no_grad():
out = lm.model.reasoner_forward(input_ids=input_ids, cache=None)
assert torch.isfinite(out).all()
def test_reasoner_forward_deterministic(self):
vfm = _build_tiny_cosmos3(seed=42)
lm = vfm.language_model
input_ids = torch.randint(0, 64, (1, 6))
with torch.no_grad():
out1 = lm.model.reasoner_forward(input_ids=input_ids, cache=None)
out2 = lm.model.reasoner_forward(input_ids=input_ids, cache=None)
assert torch.allclose(out1, out2)
# ---------------------------------------------------------------------------
# Convenience: print a concise summary when run directly.
# ---------------------------------------------------------------------------
if __name__ == "__main__":
print("Building tiny Cosmos3VFMNetwork...")
vfm = _build_tiny_cosmos3(seed=42)
print("\nParameter keys:")
for name, p in sorted(vfm.named_parameters()):
print(f" {name}: {tuple(p.shape)}")
print("\nRunning forward...")
ps = _build_tiny_packed_seq(n_text=4, seed=7)
with torch.no_grad():
out = vfm(packed_seq=ps)
lhs = out["last_hidden_state"]
pv = out["preds_vision"][0]
print(f"\nlast_hidden_state: {lhs.shape}, finite={torch.isfinite(lhs).all().item()}, mean={lhs.mean():.4f}")
print(f"preds_vision[0]: {pv.shape}, finite={torch.isfinite(pv).all().item()}, mean={pv.mean():.4f}")
# Determinism
with torch.no_grad():
out2 = vfm(packed_seq=ps)
print(f"\nDeterministic repeat: {torch.allclose(lhs, out2['last_hidden_state'])}")
@@ -0,0 +1,105 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 scheduler default + per-request override parity (Tier A scaffold).
Reference:
* ``vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py:275-307`` —
initial UniPCMultistepScheduler load (preserves solver_order,
timestep_spacing, beta_schedule, sigma bounds, flow_shift) and
one-time override at engine-init if ``od_config.flow_shift`` is set.
* ``vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py:498-512`` —
``_set_flow_shift(target_shift)``: rebuild the scheduler via
``UniPCMultistepScheduler.from_config(base_config, flow_shift=target)``
when the requested target differs from the current shift.
* ``vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py:1069-1110`` —
per-request mode defaults: T2I uses ``shift=3.0``; T2V/I2V use the
engine-init shift (typically 1.0); ``flow_shift`` may be overridden
per request via ``sampling_params.extra_args["flow_shift"]``.
The invariant under test: for the same RNG seed and the same number of
inference steps, the scheduler's ``timesteps`` tensor must be identical
whenever the ``flow_shift`` is identical, and must change deterministically
when ``flow_shift`` is overridden via ``_set_flow_shift``.
"""
from __future__ import annotations
import pytest
import torch
pytestmark = [pytest.mark.local]
def test_t2i_default_flow_shift_is_3() -> None:
"""Asserts that T2I requests rebuild the scheduler at ``flow_shift=3.0``.
Cross-check: pipeline_cosmos3.py:1073-1080 sets
``default_flow_shift = 3.0`` for T2I, and
pipeline_cosmos3.py:1110 calls ``self._set_flow_shift(flow_shift_target)``
which rebuilds the scheduler via ``UniPCMultistepScheduler.from_config(
base_config, flow_shift=3.0)``.
"""
try:
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import ( # type: ignore
Cosmos3OmniDiffusersPipeline,
)
except ImportError:
pytest.skip("FastVideo Cosmos3 pipeline not yet implemented (Phase 2b)")
pipeline = Cosmos3OmniDiffusersPipeline.__new__(Cosmos3OmniDiffusersPipeline)
if not hasattr(pipeline, "_set_flow_shift") or not hasattr(pipeline, "scheduler"):
pytest.skip("FastVideo Cosmos3 scheduler/_set_flow_shift not yet wired")
pipeline._set_flow_shift(3.0)
assert float(pipeline.scheduler.config.flow_shift) == 3.0
def test_t2v_default_flow_shift_is_engine_init() -> None:
"""Asserts T2V/I2V use the engine-init shift (e.g. 1.0), NOT a fixed default.
Cross-check: pipeline_cosmos3.py:1091 sets
``default_flow_shift = self._engine_init_flow_shift`` for T2V/I2V
(NOT ``None`` — passing ``None`` would leak a prior T2I rebuild
forward).
"""
try:
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import ( # type: ignore
Cosmos3OmniDiffusersPipeline,
)
except ImportError:
pytest.skip("FastVideo Cosmos3 pipeline not yet implemented (Phase 2b)")
pipeline = Cosmos3OmniDiffusersPipeline.__new__(Cosmos3OmniDiffusersPipeline)
if not hasattr(pipeline, "_engine_init_flow_shift") or not hasattr(pipeline, "_set_flow_shift"):
pytest.skip("FastVideo Cosmos3 _engine_init_flow_shift not yet wired")
init_shift = float(pipeline._engine_init_flow_shift)
pipeline._set_flow_shift(init_shift)
assert float(pipeline.scheduler.config.flow_shift) == init_shift
def test_scheduler_timesteps_deterministic_under_seed() -> None:
"""Asserts that ``scheduler.set_timesteps(N)`` is deterministic given the
same N and the same flow_shift.
The UniPC scheduler's timestep sequence does not depend on a torch
seed (it's a closed-form function of N + scheduler config), so
invoking ``set_timesteps`` twice should produce identical tensors.
"""
try:
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import ( # type: ignore
Cosmos3OmniDiffusersPipeline,
)
except ImportError:
pytest.skip("FastVideo Cosmos3 pipeline not yet implemented (Phase 2b)")
pipeline = Cosmos3OmniDiffusersPipeline.__new__(Cosmos3OmniDiffusersPipeline)
if not hasattr(pipeline, "scheduler") or not hasattr(pipeline, "_set_flow_shift"):
pytest.skip("FastVideo Cosmos3 scheduler not yet wired")
pipeline._set_flow_shift(3.0)
pipeline.scheduler.set_timesteps(35, device=torch.device("cpu"))
seq_a = pipeline.scheduler.timesteps.clone()
pipeline.scheduler.set_timesteps(35, device=torch.device("cpu"))
seq_b = pipeline.scheduler.timesteps.clone()
torch.testing.assert_close(seq_a, seq_b)
@@ -0,0 +1,124 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo's UniPC scheduler vs the framework's.
The Cosmos3 video sampler is the framework's flow-matching UniPC
(``cosmos_framework.model.vfm.diffusion.samplers.fm_solvers_unipc.FlowUniPCMultistepScheduler``),
driven by ``UniPCSampler`` with config ``num_train_timesteps=1000``,
``use_dynamic_shifting=False`` and a per-mode ``shift`` (10.0 for T2V/I2V,
3.0 for T2I). FastVideo reuses its vendored
``UniPCMultistepScheduler`` configured for pure flow matching
(``use_flow_sigmas=True``, ``prediction_type="flow_prediction"``,
``predict_x0=True``, ``solver_type="bh2"``, ``solver_order=2``,
``final_sigmas_type="zero"``) with ``flow_shift`` set to the same shift.
This pins the scheduler — the one Cosmos3 component whose earlier test
(``test_cosmos3_denoise_cfg_parity``) compared diffusers-vs-diffusers rather
than against the framework oracle — by:
* asserting the discrete ``timesteps`` and ``sigmas`` match the framework;
* running a full multi-step UniPC trajectory with a fixed sequence of
pseudo-random "velocity" model outputs (identical on both sides, so the
DiT is factored out) and asserting every intermediate + final latent
matches the framework bit-for-bit.
CPU / float32. The framework scheduler is the parity ORACLE.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_scheduler_parity.py -q
"""
from __future__ import annotations
import pytest
import torch
# The official framework provides the parity oracle.
fm_unipc = pytest.importorskip(
"cosmos_framework.model.vfm.diffusion.samplers.fm_solvers_unipc",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
FlowUniPCMultistepScheduler = fm_unipc.FlowUniPCMultistepScheduler
from fastvideo.models.schedulers.scheduling_unipc_multistep import ( # noqa: E402
UniPCMultistepScheduler,
)
pytestmark = [pytest.mark.local]
def _framework_scheduler(num_steps: int, shift: float) -> FlowUniPCMultistepScheduler:
"""Exactly how ``UniPCSampler`` builds + primes its scheduler."""
sched = FlowUniPCMultistepScheduler(
num_train_timesteps=1000,
shift=1.0,
use_dynamic_shifting=False,
)
sched.set_timesteps(num_steps, device=torch.device("cpu"), shift=shift)
return sched
def _fastvideo_scheduler(num_steps: int, shift: float) -> UniPCMultistepScheduler:
"""FastVideo's vendored UniPC configured for the framework's flow setup."""
sched = UniPCMultistepScheduler(
num_train_timesteps=1000,
solver_order=2,
prediction_type="flow_prediction",
use_flow_sigmas=True,
predict_x0=True,
solver_type="bh2",
final_sigmas_type="zero",
flow_shift=shift,
)
sched.set_timesteps(num_steps, device=torch.device("cpu"))
return sched
_SHIFTS = [pytest.param(10.0, id="shift10_video"), pytest.param(3.0, id="shift3_t2i")]
_STEPS = [pytest.param(4, id="4steps"), pytest.param(10, id="10steps"), pytest.param(35, id="35steps")]
class TestCosmos3SchedulerParity:
@pytest.mark.parametrize("shift", _SHIFTS)
@pytest.mark.parametrize("num_steps", _STEPS)
def test_timesteps_and_sigmas_match_framework(self, num_steps, shift):
fw = _framework_scheduler(num_steps, shift)
fv = _fastvideo_scheduler(num_steps, shift)
t_max = (fw.timesteps.float() - fv.timesteps.float()).abs().max().item()
s_max = (fw.sigmas.float() - fv.sigmas.float()).abs().max().item()
print(f"\n[sched n={num_steps} shift={shift}] timesteps max diff={t_max:.3e} "
f"sigmas max diff={s_max:.3e}")
assert fw.timesteps.shape == fv.timesteps.shape
assert fw.sigmas.shape == fv.sigmas.shape
torch.testing.assert_close(fv.timesteps, fw.timesteps)
torch.testing.assert_close(fv.sigmas, fw.sigmas)
@pytest.mark.parametrize("shift", _SHIFTS)
@pytest.mark.parametrize("num_steps", _STEPS)
def test_full_trajectory_matches_framework(self, num_steps, shift):
# A small latent so order-2 einsum paths exercise; batch axis as the
# samplers expect ([B, C, T, H, W]).
shape = (1, 4, 2, 3, 3)
torch.manual_seed(123)
init = torch.randn(shape, dtype=torch.float32)
# One pseudo-random "velocity" per step, identical on both sides.
velocities = [torch.randn(shape, dtype=torch.float32) for _ in range(num_steps)]
fw = _framework_scheduler(num_steps, shift)
fv = _fastvideo_scheduler(num_steps, shift)
fw_lat = init.clone()
fv_lat = init.clone()
worst = 0.0
for i, t in enumerate(fw.timesteps):
v = velocities[i]
fw_lat = fw.step(model_output=v, timestep=t, sample=fw_lat, return_dict=False)[0]
fv_lat = fv.step(model_output=v, timestep=t, sample=fv_lat, return_dict=False)[0]
step_max = (fw_lat - fv_lat).abs().max().item()
worst = max(worst, step_max)
assert not torch.isnan(fv_lat).any(), f"FastVideo latent NaN at step {i}"
mean_abs = (fw_lat - fv_lat).abs().mean().item()
print(f"\n[traj n={num_steps} shift={shift}] worst step max diff={worst:.3e} "
f"final mean diff={mean_abs:.3e}")
torch.testing.assert_close(fv_lat, fw_lat, atol=1e-5, rtol=1e-4)
@@ -0,0 +1,288 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 sound (t2vs) pathway vs the framework.
Covers the audio generation pathway end-to-end at the DiT level:
* **sound packing** — FastVideo's native packer
(``pack_cosmos3_video_sequence`` with a ``Cosmos3SoundItem``) vs the
framework ``pack_input_sequence`` with ``has_sound``: sound tokens share the
vision "full" split, with ``(T,1,1)`` shapes, a ``(T,1)`` condition mask,
and 3D-MRoPE temporal positions starting at the vision temporal offset
(parallel to vision); and
* **DiT sound forward** — the dormant ``audio_proj_in`` / ``audio_proj_out`` /
``audio_modality_embed`` heads, now activated (framework ``sound2llm`` /
``llm2sound`` / ``sound_modality_embed``).
Both tiny models are built sound-enabled from the SAME config, framework weights
(incl. the sound heads) are copied into the FastVideo DiT, and the FRAMEWORK
model + framework pack is the parity ORACLE (CPU/float32 via the SDPA
monkey-patch). We assert the native packer matches the framework field-by-field,
then that ``preds_vision`` AND ``preds_sound`` match the framework forward.
Run:
cd <worktree> && <fv-cosmos3 python> -m pytest \
tests/local_tests/cosmos3/test_cosmos3_sound_parity.py -q -s
"""
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 ( # noqa: E402
_fastvideo_inputs_from_packed_seq,
_framework_to_fastvideo_state_dict,
)
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
_LATENT_CHANNEL,
_LATENT_PATCH_SIZE,
_RESET_SPATIAL_IDS,
_SOUND_DIM,
_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_sound(vfm, dit) -> None:
"""Copy backbone + vision weights AND the sound MoT heads into the DiT."""
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
src = dict(vfm.named_parameters())
mapped["audio_proj_in.weight"] = src["sound2llm.weight"].detach().clone()
mapped["audio_proj_in.bias"] = src["sound2llm.bias"].detach().clone()
mapped["audio_proj_out.weight"] = src["llm2sound.weight"].detach().clone()
mapped["audio_proj_out.bias"] = src["llm2sound.bias"].detach().clone()
mapped["audio_modality_embed"] = src["sound_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_sound(*, text_ids, vision, sound, cond_vision, cond_sound, 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_sound=[sound],
fps_sound=None,
)
plans = [SequencePlan(
has_text=True, has_vision=True, has_sound=True,
condition_frame_indexes_vision=list(cond_vision),
condition_frame_indexes_sound=list(cond_sound),
)]
return 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,
)
def _fastvideo_pack_sound(*, text_ids, vision, sound, cond_vision, cond_sound, timestep):
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
Cosmos3SampleInputs,
Cosmos3SoundItem,
Cosmos3VisionItem,
pack_cosmos3_video_sequence,
)
samples = [Cosmos3SampleInputs(
text_ids=list(text_ids),
vision=Cosmos3VisionItem(latent=vision, condition_frame_indexes=list(cond_vision)),
sound=Cosmos3SoundItem(latent=sound, condition_frame_indexes=list(cond_sound)),
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_sound(ps) -> dict:
"""Framework PackedSequence (with sound) -> native DiT forward kwargs."""
kw = _fastvideo_inputs_from_packed_seq(ps)
s = ps.sound
kw.update(
sound_tokens=list(s.tokens),
sound_token_shapes=[tuple(x) for x in s.token_shapes],
sound_sequence_indexes=s.sequence_indexes,
sound_timesteps=s.timesteps,
sound_mse_loss_indexes=s.mse_loss_indexes,
sound_noisy_frame_indexes=list(s.noisy_frame_indexes),
fps_sound=None,
)
return kw
def _diffs(a: torch.Tensor, b: torch.Tensor) -> tuple[float, float]:
d = (a - b).abs()
return d.max().item(), d.mean().item()
# (grid_t, latent_h, latent_w, sound_t, n_text, cond_vision, cond_sound)
_CASES = [
pytest.param(2, 4, 4, 5, 4, [], [], id="t2vs_2x2x2_snd5"),
pytest.param(3, 8, 4, 8, 5, [], [], id="t2vs_3x4x2_snd8"),
pytest.param(2, 4, 4, 6, 5, [0], [], id="i2vs_cond_snd6"),
]
class TestCosmos3SoundParity:
def _build(self, num_layers=2, seed_model=42):
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers, sound_gen=True)
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
_copy_weights_with_sound(vfm, dit)
return vfm, dit
def _make_inputs(self, grid_t, latent_h, latent_w, sound_t, n_text, cond_v, cond_s, seed=7):
torch.manual_seed(seed)
vision = torch.randn(1, _LATENT_CHANNEL, grid_t, latent_h, latent_w)
sound = torch.randn(_SOUND_DIM, sound_t) # [C, T]
text_ids = torch.randint(0, 60, (n_text,)).tolist()
return dict(text_ids=text_ids, vision=vision, sound=sound,
cond_vision=cond_v, cond_sound=cond_s, timestep=500.0)
@pytest.mark.parametrize(("grid_t", "lh", "lw", "snd_t", "n_text", "cond_v", "cond_s"), _CASES)
def test_sound_packing_matches_framework(self, grid_t, lh, lw, snd_t, n_text, cond_v, cond_s):
ins = self._make_inputs(grid_t, lh, lw, snd_t, n_text, cond_v, cond_s)
fw = _framework_pack_sound(is_image_batch=(grid_t == 1), **ins)
fv = _fastvideo_pack_sound(**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) # [3, seq], exact
# Sound fields.
s = fw.sound
torch.testing.assert_close(fv.sound_sequence_indexes, s.sequence_indexes.to(torch.long), rtol=0, atol=0)
assert fv.sound_token_shapes == [tuple(x) for x in s.token_shapes]
torch.testing.assert_close(fv.sound_timesteps.to(torch.float32), s.timesteps.to(torch.float32))
torch.testing.assert_close(fv.sound_mse_loss_indexes, s.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
for a, b in zip(fv.sound_noisy_frame_indexes, s.noisy_frame_indexes):
torch.testing.assert_close(a.to(torch.long), b.to(torch.long), rtol=0, atol=0)
for a, b in zip(fv.sound_condition_mask, s.condition_mask):
torch.testing.assert_close(a.flatten().to(torch.float32), b.flatten().to(torch.float32))
print(f"\n[sound_packing {grid_t}x{lh}x{lw} snd={snd_t}] position_ids + sound fields exact")
@pytest.mark.parametrize(("grid_t", "lh", "lw", "snd_t", "n_text", "cond_v", "cond_s"), _CASES)
def test_sound_dit_forward_matches_framework(self, grid_t, lh, lw, snd_t, n_text, cond_v, cond_s):
vfm, dit = self._build()
ins = self._make_inputs(grid_t, lh, lw, snd_t, n_text, cond_v, cond_s)
fw_pack = _framework_pack_sound(is_image_batch=(grid_t == 1), **ins)
fv_pack = _fastvideo_pack_sound(**ins)
with torch.no_grad():
fw_out = vfm(packed_seq=fw_pack) # framework model + framework pack (oracle)
fv_out = dit(**fv_pack.to_dit_kwargs()) # native model + native pack
# Also run the native DiT on the framework pack to isolate the forward.
fv_on_fw = dit(**_fv_inputs_with_sound(fw_pack))
# preds_vision parity.
pv_max, pv_mean = _diffs(fv_out["preds_vision"][0], fw_out["preds_vision"][0])
# preds_sound parity.
ps_max, ps_mean = _diffs(fv_out["preds_sound"][0], fw_out["preds_sound"][0])
# native-DiT-on-framework-pack (forward only) parity.
psf_max, psf_mean = _diffs(fv_on_fw["preds_sound"][0], fw_out["preds_sound"][0])
print(f"\n[sound_dit {grid_t}x{lh}x{lw} snd={snd_t}] "
f"preds_vision max={pv_max:.3e} mean={pv_mean:.3e} | "
f"preds_sound max={ps_max:.3e} mean={ps_mean:.3e} | "
f"preds_sound(fwpack) max={psf_max:.3e} mean={psf_mean:.3e}")
assert fv_out["preds_sound"][0].shape == fw_out["preds_sound"][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_sound"][0], fw_out["preds_sound"][0], atol=1e-4, rtol=1e-3)
torch.testing.assert_close(fv_on_fw["preds_sound"][0], fw_out["preds_sound"][0], atol=1e-4, rtol=1e-3)
@pytest.mark.parametrize(("grid_t", "lh", "lw", "snd_t", "n_text", "cond_v", "cond_s"), _CASES)
def test_t2vs_cfg_velocity_matches_framework(self, grid_t, lh, lw, snd_t, n_text, cond_v, cond_s):
"""The combined [vision|sound] sequential-CFG velocity (the t2vs denoise
step's pipeline glue) matches a framework-DiT oracle."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
Cosmos3SoundSpec,
Cosmos3VisionSpec,
cosmos3_get_cfg_velocity,
)
vfm, dit = self._build()
vision_shape = (_LATENT_CHANNEL, grid_t, lh // _LATENT_PATCH_SIZE, lw // _LATENT_PATCH_SIZE)
# _make_inputs builds the vision LATENT [1,C,T,H,W]; here we drive the
# combined flat latent directly, so use the un-patchified latent shape.
vlat_shape = (_LATENT_CHANNEL, grid_t, lh, lw)
sound_shape = (_SOUND_DIM, snd_t)
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())
snd_numel = int(torch.tensor(sound_shape).prod())
flat = torch.randn(vis_numel + snd_numel)
guidance, ts = 6.0, 500.0
def _fw_velocity(ids):
vision = flat[:vis_numel].reshape(vlat_shape).unsqueeze(0) # [1,C,T,H,W]
sound = flat[vis_numel:].reshape(sound_shape) # [C,T]
ps = _framework_pack_sound(text_ids=ids, vision=vision, sound=sound,
cond_vision=cond_v, cond_sound=cond_s,
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)
psd = out["preds_sound"][0] # [C,T] (zero on clean)
return torch.cat([pv.reshape(-1), psd.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))],
sound_specs=[Cosmos3SoundSpec(shape=sound_shape, condition_frame_indexes=list(cond_s))],
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[t2vs_cfg_velocity {grid_t}x{lh}x{lw} snd={snd_t}] max={mx:.3e} mean={mn:.3e}")
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
@@ -0,0 +1,169 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 DiT param-name / checkpoint-key-surface contract.
The FastVideo native Cosmos3 DiT (``fastvideo.models.dits.cosmos3.
Cosmos3VFMTransformer``) deliberately mirrors the *published diffusers
checkpoint* transformer key surface so the converter at
``scripts/checkpoint_conversion/cosmos3_convert.py`` can strict-load with a
near-identity ``param_names_mapping``.
This pins the FastVideo side of that contract: a tiny 1-layer DiT must expose
exactly the published checkpoint's parameter names. The checkpoint key surface
(validated against ``nvidia/Cosmos3-Nano``; ``{i}`` ranges over the layers):
Top level:
embed_tokens.weight, norm.weight, norm_moe_gen.weight, lm_head.weight,
proj_in.{weight,bias}, proj_out.{weight,bias},
time_embedder.linear_1.{weight,bias}, time_embedder.linear_2.{weight,bias}
Dormant heads (present for strict-load):
action_proj_in.fc.weight, action_proj_in.bias.weight,
action_proj_out.fc.weight, action_proj_out.bias.weight,
action_modality_embed,
audio_proj_in.{weight,bias}, audio_proj_out.{weight,bias},
audio_modality_embed
Per layer ``layers.{i}``:
self_attn.{to_q,to_k,to_v,to_out,add_q_proj,add_k_proj,add_v_proj,
to_add_out,norm_q,norm_k,norm_added_q,norm_added_k}.weight,
mlp.{gate_proj,up_proj,down_proj}.weight,
mlp_moe_gen.{gate_proj,up_proj,down_proj}.weight,
{input_layernorm,input_layernorm_moe_gen,post_attention_layernorm,
post_attention_layernorm_moe_gen}.weight
The earlier scaffold pinned the dead vllm-omni layout
(``language_model.layers`` / ``gen_layers`` / ``cross_attention`` /
``vae2llm`` / ``llm2vae``); that structure is gone — the native DiT is a single
dual-pathway ``layers`` ModuleList matching the diffusers checkpoint.
"""
from __future__ import annotations
import pytest
pytestmark = [pytest.mark.local]
def _build_tiny_dit():
"""Construct a tiny 1-layer FastVideo Cosmos3 DiT, or skip if unavailable."""
try:
from fastvideo.configs.models.dits.cosmos3 import (
Cosmos3ArchConfig,
Cosmos3VideoConfig,
)
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
except ImportError: # pragma: no cover - import guard
pytest.skip("FastVideo Cosmos3 DiT not importable in this environment")
import torch
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,
latent_patch_size=2,
latent_channel=16,
position_embedding_type="3d_rope",
enable_fps_modulation=False,
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)
return Cosmos3VFMTransformer(cfg, hf_config={}).to(torch.float32)
def test_fastvideo_cosmos3_dit_module_tree_param_names() -> None:
"""The native DiT param-name set must equal the published checkpoint surface."""
model = _build_tiny_dit()
names = {name for name, _ in model.named_parameters()}
expected_top = {
"embed_tokens.weight",
"norm.weight",
"norm_moe_gen.weight",
"lm_head.weight",
"proj_in.weight",
"proj_in.bias",
"proj_out.weight",
"proj_out.bias",
"time_embedder.linear_1.weight",
"time_embedder.linear_1.bias",
"time_embedder.linear_2.weight",
"time_embedder.linear_2.bias",
}
expected_dormant = {
"action_proj_in.fc.weight",
"action_proj_in.bias.weight",
"action_proj_out.fc.weight",
"action_proj_out.bias.weight",
"action_modality_embed",
"audio_proj_in.weight",
"audio_proj_in.bias",
"audio_proj_out.weight",
"audio_proj_out.bias",
"audio_modality_embed",
}
layer_suffixes = {
"self_attn.to_q.weight",
"self_attn.to_k.weight",
"self_attn.to_v.weight",
"self_attn.to_out.weight",
"self_attn.add_q_proj.weight",
"self_attn.add_k_proj.weight",
"self_attn.add_v_proj.weight",
"self_attn.to_add_out.weight",
"self_attn.norm_q.weight",
"self_attn.norm_k.weight",
"self_attn.norm_added_q.weight",
"self_attn.norm_added_k.weight",
"mlp.gate_proj.weight",
"mlp.up_proj.weight",
"mlp.down_proj.weight",
"mlp_moe_gen.gate_proj.weight",
"mlp_moe_gen.up_proj.weight",
"mlp_moe_gen.down_proj.weight",
"input_layernorm.weight",
"input_layernorm_moe_gen.weight",
"post_attention_layernorm.weight",
"post_attention_layernorm_moe_gen.weight",
}
expected_layer = {f"layers.0.{s}" for s in layer_suffixes}
expected = expected_top | expected_dormant | expected_layer
assert names == expected, (f"Cosmos3 DiT param surface mismatch.\n"
f" missing: {sorted(expected - names)}\n"
f" unexpected: {sorted(names - expected)}")
def test_fastvideo_cosmos3_dit_no_dead_vllm_omni_layout() -> None:
"""The dead vllm-omni layout (split language_model/gen_layers/cross_attention,
vae2llm/llm2vae) must NOT appear in the native DiT param tree."""
model = _build_tiny_dit()
names = {name for name, _ in model.named_parameters()}
dead_fragments = (
"language_model.",
"gen_layers.",
"cross_attention.",
"vae2llm",
"llm2vae",
)
offenders = sorted(n for n in names if any(frag in n for frag in dead_fragments))
assert not offenders, f"native DiT still exposes dead vllm-omni param names: {offenders}"
def test_fastvideo_cosmos3_dit_per_layer_counts() -> None:
"""Per-layer block count: 22 weights/layer; full 1-layer model = 44 params.
(12 attention + 6 dual-MLP + 4 layernorm per layer; + 12 top-level
+ 10 dormant-head params.)
"""
model = _build_tiny_dit()
names = [name for name, _ in model.named_parameters()]
per_layer = [n for n in names if n.startswith("layers.0.")]
assert len(per_layer) == 22, f"expected 22 per-layer params, got {len(per_layer)}"
assert len(names) == 44, f"expected 44 total params for tiny 1-layer DiT, got {len(names)}"
@@ -0,0 +1,80 @@
# SPDX-License-Identifier: Apache-2.0
"""Strict-load completeness: real Cosmos3-Nano transformer <-> FastVideo DiT.
The published ``nvidia/Cosmos3-Nano`` checkpoint is diffusers-format
(``needs_conversion=no``). This test verifies that EVERY transformer weight key
in the checkpoint maps 1:1 (via the DiT's ``param_names_mapping``) onto a
FastVideo ``Cosmos3VFMTransformer`` parameter of matching shape, and that no DiT
parameter is left unfilled -- i.e. a ``strict=True`` load will succeed.
It runs on the ``meta`` device (no 30 GB allocation, no GPU) by reading only the
safetensors headers. Skips cleanly if the checkpoint is not present locally.
"""
from __future__ import annotations
import glob
import os
import re
import pytest
import torch
pytestmark = [pytest.mark.local]
_CKPT_DIR = os.path.join("official_weights", "cosmos3", "transformer")
def _checkpoint_key_shapes() -> dict[str, tuple[int, ...]]:
from safetensors import safe_open
shards = sorted(glob.glob(os.path.join(_CKPT_DIR, "*.safetensors")))
if not shards:
pytest.skip(f"Cosmos3 transformer checkpoint not present: {_CKPT_DIR}")
out: dict[str, tuple[int, ...]] = {}
for shard in shards:
with safe_open(shard, framework="pt") as f:
for k in f.keys():
out[k] = tuple(f.get_slice(k).get_shape())
return out
def _meta_dit():
from fastvideo.configs.models.dits.cosmos3 import Cosmos3VideoConfig
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
cfg = Cosmos3VideoConfig()
with torch.device("meta"):
dit = Cosmos3VFMTransformer(cfg, hf_config={})
return dit, cfg
def _apply_mapping(key: str, pmap: dict[str, str]) -> str:
for pat, repl in pmap.items():
if re.match(pat, key):
return re.sub(pat, repl, key)
return key
def test_strict_load_completeness():
ckpt = _checkpoint_key_shapes()
dit, cfg = _meta_dit()
dit_params = {n: tuple(p.shape) for n, p in dit.named_parameters()}
dit_buffers = {n: tuple(b.shape) for n, b in dit.named_buffers()}
pmap = cfg.arch_config.param_names_mapping
mapped = {_apply_mapping(k, pmap): v for k, v in ckpt.items()}
ckpt_names = set(mapped)
param_names = set(dit_params)
buffer_names = set(dit_buffers)
# Every checkpoint key must land on a DiT parameter (non-persistent buffers excepted).
unexpected = sorted(ckpt_names - param_names - buffer_names)
assert not unexpected, f"checkpoint keys with no DiT param: {unexpected[:20]}"
# Every DiT parameter must be provided by the checkpoint (true strict load).
missing = sorted(param_names - ckpt_names)
assert not missing, f"DiT params not provided by checkpoint: {missing[:20]}"
# Shapes must match exactly.
mism = [(k, mapped[k], dit_params[k]) for k in (ckpt_names & param_names) if mapped[k] != dit_params[k]]
assert not mism, f"shape mismatches: {mism[:10]}"
assert len(ckpt) == len(dit_params), f"key count mismatch: ckpt={len(ckpt)} dit={len(dit_params)}"
@@ -0,0 +1,83 @@
# SPDX-License-Identifier: Apache-2.0
"""Cosmos3 prompt tokenization contract — chat template + special tokens.
The native pipeline tokenizes via the module-level helpers
``cosmos3_special_tokens`` / ``cosmos3_tokenize_caption``
(``fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline``), which wrap a Qwen2
chat-template tokenizer:
* special tokens: ``start_of_generation=<|vision_start|>``,
``end_of_generation=<|vision_end|>``, ``eos_token_id=tokenizer.eos_token_id``
(the framework ``llm_special_tokens``);
* ``tokenize_caption`` applies the chat template with
``add_generation_prompt=True`` / ``add_vision_id=False`` and an optional
image/video system prompt.
These contract checks run against the conftest Qwen2-shaped stub tokenizer (no
real weights). The byte-for-byte real-token-id check (``eos == 151645``,
``<|vision_start|> == 151652``) needs the real ``nvidia/Cosmos3-Nano``
``text_tokenizer`` and is skipped cleanly when it is unavailable.
"""
from __future__ import annotations
import pytest
from .conftest import StubQwen2Tokenizer
pytestmark = [pytest.mark.local]
def test_native_special_tokens_resolution() -> None:
"""``cosmos3_special_tokens`` resolves the three generation special tokens."""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import cosmos3_special_tokens
special = cosmos3_special_tokens(StubQwen2Tokenizer())
assert set(special) == {"start_of_generation", "end_of_generation", "eos_token_id"}
assert special["start_of_generation"] == StubQwen2Tokenizer().convert_tokens_to_ids("<|vision_start|>")
assert special["end_of_generation"] == StubQwen2Tokenizer().convert_tokens_to_ids("<|vision_end|>")
assert special["eos_token_id"] == StubQwen2Tokenizer.eos_token_id
def test_native_tokenize_caption_uses_chat_template() -> None:
"""``cosmos3_tokenize_caption`` returns a non-empty token-id list.
The video / image system-prompt variants tokenize independently (the chat
template prepends a role=system turn when ``use_system_prompt`` is set), and
the result is always a plain list of ints.
"""
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import cosmos3_tokenize_caption
tok = StubQwen2Tokenizer()
ids_video = cosmos3_tokenize_caption(tok, "a robot dances", is_video=True, use_system_prompt=False)
ids_image = cosmos3_tokenize_caption(tok, "a robot", is_video=False, use_system_prompt=True)
assert isinstance(ids_video, list) and all(isinstance(i, int) for i in ids_video) and ids_video
assert isinstance(ids_image, list) and ids_image
def test_cosmos3_special_token_ids_real_weights() -> None:
"""Byte-for-byte Qwen2 special-token ids (needs the real text_tokenizer).
Asserts ``eos_token_id == 151645`` and
``convert_tokens_to_ids('<|vision_start|>') == 151652`` on the real
``nvidia/Cosmos3-Nano`` Qwen2 tokenizer. Skipped cleanly when the real
tokenizer is not loadable in this environment.
"""
try:
from transformers import AutoTokenizer
except ImportError:
pytest.skip("transformers not available")
import os
candidate_paths = [
os.path.join(os.environ.get("COSMOS3_WEIGHTS_DIR", ""), "text_tokenizer"),
"official_weights/cosmos3/text_tokenizer",
]
tok_path = next((p for p in candidate_paths if p and os.path.isdir(p)), None)
if tok_path is None:
pytest.skip("real nvidia/Cosmos3-Nano text_tokenizer not available "
"(set COSMOS3_WEIGHTS_DIR or provide official_weights/cosmos3)")
tokenizer = AutoTokenizer.from_pretrained(tok_path)
assert tokenizer.eos_token_id == 151645
assert tokenizer.convert_tokens_to_ids("<|vision_start|>") == 151652
@@ -0,0 +1,392 @@
# SPDX-License-Identifier: Apache-2.0
"""Numerical-parity test: FastVideo Cosmos3 (Wan2.2) VAE vs OFFICIAL framework VAE.
The Cosmos3 checkpoint VAE is literally ``Wan-AI/Wan2.2-TI2V-5B-Diffusers``
(diffusers ``AutoencoderKLWan``). FastVideo reuses its native ``AutoencoderKLWan``
(``fastvideo/models/vaes/wanvae.py``) with the Wan2.2-residual geometry locked
in ``Cosmos3VAEConfig`` (``fastvideo/configs/models/vaes/cosmos3vae.py``).
Parity oracle: the OFFICIAL framework VAE
``cosmos_framework.model.vfm.tokenizers.wan2pt2_vae_4x16x16.WanVAE_``
(CausalConv3d / ResidualBlock / Encoder3d / Decoder3d), which is the same
architecture loaded by ``Cosmos3-Nano.yaml`` via ``Wan2pt2VAEInterface``.
Approach (preferred per the porting plan): build a *tiny* FastVideo
``AutoencoderKLWan`` and a *tiny* framework ``WanVAE_`` with matching small
Wan2.2 geometry, copy weights via an explicit name map, then compare ENCODE
and DECODE of a small deterministic video on CPU/float32.
Why tiny weight-copy (not real weights): it runs on CPU in <1s, needs no GPU
and no 33 GiB checkpoint, and exercises the *full* encoder + decoder conv
stack. The module structures are isomorphic (verified: 0 unmapped / 0 missing /
0 extra / 0 shape mismatches), so the copy is exact and the comparison is
meaningful end-to-end. A real-weights cross-check is included but skips cleanly
when the checkpoint / diffusers are unavailable.
Normalization handling
----------------------
- The framework ``WanVAE_.encode(x, scale)`` applies ``(mu - mean) * inv_std``
internally; ``decode(z, scale, ...)`` inverts it. FastVideo's ``encode`` /
``decode`` operate on *raw* (un-normalized) latents. To compare the conv
stacks directly we pass an identity scale ``(mean=0, inv_std=1)`` to the
framework so both sides see the same raw latent space.
- The framework ``WanVAE_.decode`` does NOT clamp its output (clamping happens
in the outer ``Wan2pt2VAEInterface.decode`` wrapper), whereas FastVideo's
``AutoencoderKLWan.decode`` ends with ``torch.clamp(out, -1, 1)``. We
therefore clamp the framework decode to ``[-1, 1]`` before comparing — the
only intended behavioral difference between the two paths.
Run:
PYTHONSAFEPATH=1 pytest tests/local_tests/cosmos3/test_cosmos3_vae_parity.py -v
"""
from __future__ import annotations
import re
import sys
import types
import pytest
import torch
pytestmark = [pytest.mark.local]
# ---------------------------------------------------------------------------
# Import the framework VAE module.
#
# ``wan2pt2_vae_4x16x16`` imports ``cosmos_framework.utils.easy_io`` at module
# scope, which pulls in optional cloud-storage backends (boto3 /
# multistorageclient) that are not installed in the CPU test env. We only need
# the nn.Modules, not checkpoint I/O, so we stub ``easy_io`` before import.
# ---------------------------------------------------------------------------
def _import_framework_vae():
pytest.importorskip(
"cosmos_framework",
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
)
if "cosmos_framework.utils.easy_io.easy_io" not in sys.modules:
pkg = types.ModuleType("cosmos_framework.utils.easy_io")
pkg.__path__ = [] # type: ignore[attr-defined]
sys.modules.setdefault("cosmos_framework.utils.easy_io", pkg)
eio = types.ModuleType("cosmos_framework.utils.easy_io.easy_io")
eio.easy_io = types.SimpleNamespace( # type: ignore[attr-defined]
load=lambda *a, **k: (_ for _ in ()).throw(
RuntimeError("easy_io is stubbed for the CPU parity test")))
sys.modules["cosmos_framework.utils.easy_io.easy_io"] = eio
try:
import cosmos_framework.model.vfm.tokenizers.wan2pt2_vae_4x16x16 as fw_vae
except Exception as exc: # pragma: no cover - env-dependent
pytest.skip(f"framework wan2pt2 VAE not importable: {exc!r}")
return fw_vae
# ---------------------------------------------------------------------------
# Tiny matching Wan2.2 geometry (small dims, real structure).
# ---------------------------------------------------------------------------
TINY_DIM = 8
TINY_DEC_DIM = 12
TINY_ZDIM = 4
TINY_DIM_MULT = (1, 2, 4, 4)
TINY_NUM_RES_BLOCKS = 2
TINY_TDOWN = (False, True, True)
def _build_framework_vae(fw_vae, seed: int = 0):
torch.manual_seed(seed)
model = fw_vae.WanVAE_(
dim=TINY_DIM,
dec_dim=TINY_DEC_DIM,
z_dim=TINY_ZDIM,
dim_mult=list(TINY_DIM_MULT),
num_res_blocks=TINY_NUM_RES_BLOCKS,
attn_scales=[],
temperal_downsample=list(TINY_TDOWN),
dropout=0.0,
temporal_window=4,
)
model.eval()
return model
def _build_fastvideo_vae():
from fastvideo.configs.models.vaes.cosmos3vae import (
Cosmos3VAEArchConfig,
Cosmos3VAEConfig,
)
from fastvideo.models.vaes.wanvae import AutoencoderKLWan
# Start from the locked Cosmos3 arch then shrink the geometry; keep
# is_residual / patch_size / channels exactly as the real config.
arch = Cosmos3VAEArchConfig(
base_dim=TINY_DIM,
decoder_base_dim=TINY_DEC_DIM,
z_dim=TINY_ZDIM,
dim_mult=TINY_DIM_MULT,
num_res_blocks=TINY_NUM_RES_BLOCKS,
temperal_downsample=TINY_TDOWN,
latents_mean=tuple([0.0] * TINY_ZDIM),
latents_std=tuple([1.0] * TINY_ZDIM),
)
cfg = Cosmos3VAEConfig(arch_config=arch)
cfg.use_feature_cache = True
cfg.load_encoder = True
cfg.load_decoder = True
model = AutoencoderKLWan(cfg)
model.eval()
return model
# ---------------------------------------------------------------------------
# Explicit framework -> FastVideo state-dict key map (residual Wan2.2 layout).
# This mirrors Cosmos3VAEArchConfig.map_official_key but is kept inline so the
# test is self-documenting and independent of the production helper.
# ---------------------------------------------------------------------------
def _map_residual_sub(prefix: str, sub: str) -> str | None:
if sub == "residual.0.gamma":
return f"{prefix}.norm1.gamma"
m = re.match(r"residual\.2\.(weight|bias)$", sub)
if m:
return f"{prefix}.conv1.{m.group(1)}"
if sub == "residual.3.gamma":
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_resample_sub(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
def _map_fw_to_fv(key: str) -> str | 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_sub(f"{m.group(1)}.mid_block.resnets.0", m.group(2))
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if m:
# attention subkeys are identically named
return 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_sub(f"{m.group(1)}.mid_block.resnets.1", m.group(2))
m = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
if m:
b, j, sub = int(m.group(1)), int(m.group(2)), m.group(3)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return _map_resample_sub(f"encoder.down_blocks.{b}.downsampler", sub)
return _map_residual_sub(f"encoder.down_blocks.{b}.resnets.{j}", sub)
m = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
if m:
b, j, sub = int(m.group(1)), int(m.group(2)), m.group(3)
if sub.startswith("resample.") or sub.startswith("time_conv."):
return _map_resample_sub(f"decoder.up_blocks.{b}.upsampler", sub)
return _map_residual_sub(f"decoder.up_blocks.{b}.resnets.{j}", sub)
return None
def _copy_weights(fw_model, fv_model) -> None:
"""Copy framework weights into the FastVideo model via the explicit map.
Asserts an exact 1:1 mapping (no unmapped source keys, no uncovered target
keys) so the parity comparison cannot be silently weakened by a partial
copy.
"""
fw_sd = fw_model.state_dict()
fv_sd = fv_model.state_dict()
new_sd: dict[str, torch.Tensor] = {}
unmapped = []
for k, v in fw_sd.items():
nk = _map_fw_to_fv(k)
if nk is None:
unmapped.append(k)
else:
new_sd[nk] = v
assert not unmapped, f"unmapped framework keys: {unmapped[:10]}"
missing = set(fv_sd) - set(new_sd)
extra = set(new_sd) - set(fv_sd)
assert not missing, f"FastVideo keys not produced by map: {sorted(missing)[:10]}"
assert not extra, f"mapped keys absent in FastVideo: {sorted(extra)[:10]}"
shape_bad = [(k, tuple(new_sd[k].shape), tuple(fv_sd[k].shape))
for k in new_sd if new_sd[k].shape != fv_sd[k].shape]
assert not shape_bad, f"shape mismatches: {shape_bad[:10]}"
missing_keys, unexpected_keys = fv_model.load_state_dict(new_sd, strict=True)
assert not missing_keys and not unexpected_keys
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture(scope="module")
def vae_pair():
fw_vae = _import_framework_vae()
fw_model = _build_framework_vae(fw_vae, seed=0)
fv_model = _build_fastvideo_vae()
_copy_weights(fw_model, fv_model)
return fw_model, fv_model
@pytest.fixture(scope="module")
def tiny_video() -> torch.Tensor:
# Wan VAE temporal constraint: T == 1 or (T - 1) % 4 == 0.
# Spatial dims must be divisible by scale_factor_spatial=16 (after the
# internal 2x patchify the encoder still needs H/2, W/2 divisible by 8).
torch.manual_seed(123)
return torch.randn(1, 3, 5, 32, 32, dtype=torch.float32)
def _fv_encode_mu(fv_model, video: torch.Tensor) -> torch.Tensor:
out = fv_model.encode(video)
dist = out.latent_dist if hasattr(out, "latent_dist") else out
return dist.mode()
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestCosmos3VAEParityTinyWeightCopy:
"""Bit-exact parity between FastVideo and framework Wan2.2 VAE (tiny copy)."""
def test_key_map_is_one_to_one(self, vae_pair):
# _copy_weights already asserts this; re-run on a fresh pair to make the
# invariant an explicit, named test.
fw_model, fv_model = vae_pair
fw_sd = fw_model.state_dict()
mapped = {}
unmapped = []
for k in fw_sd:
nk = _map_fw_to_fv(k)
(mapped.setdefault(nk, k) if nk is not None else unmapped.append(k))
assert not unmapped
assert set(mapped) == set(fv_model.state_dict())
def test_encode_parity(self, vae_pair, tiny_video):
fw_model, fv_model = vae_pair
zeros = torch.zeros(TINY_ZDIM)
ones = torch.ones(TINY_ZDIM)
with torch.no_grad():
fw_mu = fw_model.encode(tiny_video, scale=(zeros, ones))
fv_mu = _fv_encode_mu(fv_model, tiny_video)
assert fw_mu.shape == fv_mu.shape, f"{fw_mu.shape} vs {fv_mu.shape}"
max_abs = (fw_mu - fv_mu).abs().max().item()
print(f"\n[ENCODE] max abs diff = {max_abs:.3e} shape={tuple(fw_mu.shape)}")
# Bit-exact: identical weights + identical (deterministic) conv stack.
torch.testing.assert_close(fv_mu, fw_mu, rtol=0.0, atol=1e-6)
def test_decode_parity(self, vae_pair, tiny_video):
fw_model, fv_model = vae_pair
zeros = torch.zeros(TINY_ZDIM)
ones = torch.ones(TINY_ZDIM)
with torch.no_grad():
# Shared raw latent (framework encode with identity scale).
z = fw_model.encode(tiny_video, scale=(zeros, ones))
fw_dec = fw_model.decode(z, scale=(zeros, ones), clear_decoder_cache=True)
fv_dec = fv_model.decode(z)
assert fw_dec.shape == fv_dec.shape, f"{fw_dec.shape} vs {fv_dec.shape}"
# FastVideo clamps to [-1, 1]; the framework WanVAE_.decode does not
# (its outer interface wrapper does). Clamp the framework output to the
# same range — the only intended behavioral difference.
fw_dec_clamped = fw_dec.clamp(-1.0, 1.0)
max_abs = (fw_dec_clamped - fv_dec).abs().max().item()
max_abs_raw = (fw_dec - fv_dec).abs().max().item()
print(f"\n[DECODE] max abs diff (clamp-matched) = {max_abs:.3e} "
f"shape={tuple(fw_dec.shape)} (raw, pre-clamp diff = {max_abs_raw:.3e})")
torch.testing.assert_close(fv_dec, fw_dec_clamped, rtol=0.0, atol=1e-6)
def test_roundtrip_finite(self, vae_pair, tiny_video):
fw_model, fv_model = vae_pair
zeros = torch.zeros(TINY_ZDIM)
ones = torch.ones(TINY_ZDIM)
with torch.no_grad():
z = fw_model.encode(tiny_video, scale=(zeros, ones))
fv_dec = fv_model.decode(z)
assert torch.isfinite(fv_dec).all()
assert fv_dec.min() >= -1.0 - 1e-6 and fv_dec.max() <= 1.0 + 1e-6
class TestCosmos3VAEConfigLock:
"""The Cosmos3 VAE config must encode the Wan2.2-TI2V-5B geometry."""
def test_config_matches_checkpoint_geometry(self):
from fastvideo.configs.models.vaes import Cosmos3VAEConfig
cfg = Cosmos3VAEConfig()
arch = cfg.arch_config
assert arch._name_or_path == "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
assert arch.base_dim == 160
assert arch.decoder_base_dim == 256
assert arch.z_dim == 48
assert tuple(arch.dim_mult) == (1, 2, 4, 4)
assert arch.num_res_blocks == 2
assert arch.in_channels == 12
assert arch.out_channels == 12
assert arch.patch_size == 2
assert arch.scale_factor_temporal == 4
assert arch.scale_factor_spatial == 16
assert arch.is_residual is True
assert arch.clip_output is False
assert tuple(arch.temperal_downsample) == (False, True, True)
assert len(arch.latents_mean) == 48
assert len(arch.latents_std) == 48
# attribute delegation through ModelConfig.__getattr__
assert cfg.z_dim == 48
assert cfg.is_residual is True
def test_config_latents_match_checkpoint_json(self):
"""latents_mean/std must equal the Cosmos3 checkpoint values when the
checkpoint config.json is available."""
import json
import math
import os
ckpt_path = os.path.join("official_weights", "cosmos3", "vae", "config.json")
if not os.path.exists(ckpt_path):
pytest.skip(f"checkpoint config not available: {ckpt_path}")
from fastvideo.configs.models.vaes import Cosmos3VAEConfig
with open(ckpt_path) as f:
ckpt = json.load(f)
arch = Cosmos3VAEConfig().arch_config
for field_name in ("latents_mean", "latents_std"):
mine = list(getattr(arch, field_name))
theirs = ckpt[field_name]
assert len(mine) == len(theirs) == 48
for a, b in zip(mine, theirs):
assert math.isclose(a, b, rel_tol=0.0, abs_tol=1e-7), (
f"{field_name}: {a} != {b}")
# NOTE: A real-weights cross-check vs diffusers AutoencoderKLWan was intentionally
# omitted — the parity oracle for this port is the official cosmos_framework only.
# Real-weight validation is covered framework-side during checkpoint conversion.