Compare commits
42
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c59f568d56 | ||
|
|
fdea9e9898 | ||
|
|
9b19098434 | ||
|
|
24e2e73454 | ||
|
|
e5bd10678f | ||
|
|
5c749a6911 | ||
|
|
5623c1ba1b | ||
|
|
c72fc2de2c | ||
|
|
41e2d3ee8b | ||
|
|
1ec268cf80 | ||
|
|
637d0bf943 | ||
|
|
417e653604 | ||
|
|
9de2ab7505 | ||
|
|
7b4f757fc2 | ||
|
|
35fe3db250 | ||
|
|
2114d19ced | ||
|
|
895cef22a3 | ||
|
|
284c50fc96 | ||
|
|
94d9dda5cb | ||
|
|
f661bbf104 | ||
|
|
a2a204c25d | ||
|
|
4533447e58 | ||
|
|
c94d625844 | ||
|
|
7427b9d2d3 | ||
|
|
9dab81f180 | ||
|
|
840b43b4c3 | ||
|
|
ace34f866a | ||
|
|
b5f90a6b77 | ||
|
|
ed8f118acd | ||
|
|
b379ea96b4 | ||
|
|
6771f4ec2e | ||
|
|
b48eafe515 | ||
|
|
1d199050af | ||
|
|
8b65f7ce6c | ||
|
|
2d2657821e | ||
|
|
ded49ab1c8 | ||
|
|
10f5372f42 | ||
|
|
ed4125e7d8 | ||
|
|
26e78f8fb1 | ||
|
|
12f94fd6e7 | ||
|
|
facd035b97 | ||
|
|
c3c6c0cb2d |
@@ -37,6 +37,11 @@ logs/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# Cosmos3 local parity assets (symlinked from main worktree)
|
||||
/official_weights/
|
||||
/converted_weights/
|
||||
/cosmos-framework
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
**/.cache/**
|
||||
|
||||
@@ -0,0 +1,81 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
InputConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — image-to-video (I2V) path through
|
||||
# FastVideo's native Cosmos3 pipeline. The input image conditions latent frame 0
|
||||
# (kept clean during denoising); the rest of the clip is generated to follow it.
|
||||
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
|
||||
# ``official_weights/cosmos3``) to skip the Hugging Face download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3_i2v"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
image_path = os.environ.get("COSMOS3_IMAGE_PATH", "assets/images/cyclist.jpg")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A mountain biker rides forward along the sunlit forest trail, wheels "
|
||||
"kicking up dust as trees and dappled light sweep past, smooth cinematic "
|
||||
"tracking shot from behind."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
inputs=InputConfig(image_path=image_path),
|
||||
sampling=SamplingConfig(
|
||||
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
|
||||
# overridable via env for quick smoke runs.
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,77 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — this example exercises the
|
||||
# text-to-video (T2V) path through FastVideo's native Cosmos3 pipeline.
|
||||
# Point COSMOS3_MODEL_PATH at a local diffusers checkpoint (e.g.
|
||||
# ``official_weights/cosmos3``) to skip the Hugging Face download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A golden retriever puppy runs across a sunlit meadow toward the camera, "
|
||||
"ears flopping and wildflowers swaying in the breeze. Shallow depth of "
|
||||
"field, warm afternoon light, smooth cinematic tracking shot."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
# cosmos3_nano native defaults are 704x1280, 189 frames, 35 steps;
|
||||
# overridable via env for quick smoke runs.
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,78 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
# NVIDIA Cosmos3-Nano omni world model — text-to-image (T2I) path through
|
||||
# FastVideo's native Cosmos3 pipeline. T2I is the single-frame case
|
||||
# (num_frames=1); the canonical Cosmos3 T2I resolution is 960x960 (the model's
|
||||
# "720" bucket, UniPC flow_shift=10.0). Point COSMOS3_MODEL_PATH at a local
|
||||
# diffusers checkpoint (e.g. ``official_weights/cosmos3``) to skip the download.
|
||||
OUTPUT_PATH = "video_samples_cosmos3_t2i"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(
|
||||
text_encoder=True,
|
||||
pin_cpu_memory=True,
|
||||
dit=False,
|
||||
vae=False,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start_time
|
||||
|
||||
prompt = (
|
||||
"A photograph of a red panda sitting on a mossy log in a misty bamboo "
|
||||
"forest, soft golden morning light filtering through the leaves, shallow "
|
||||
"depth of field, crisp fur detail, serene atmosphere."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
# T2I is single-frame; canonical Cosmos3 T2I is 960x960. Overridable
|
||||
# via env for quick smoke runs.
|
||||
num_frames=1,
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "960")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "960")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
|
||||
start_time = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start_time
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate image: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,67 @@
|
||||
import os
|
||||
import time
|
||||
|
||||
# t2vs (text -> video + sound). The Cosmos3 denoise stage generates a joint
|
||||
# [vision | sound] latent and AVAE-decodes the sound to a waveform muxed into the
|
||||
# mp4. The joint-sound path is gated on COSMOS3_T2VS (set here for the example).
|
||||
os.environ.setdefault("COSMOS3_T2VS", "1")
|
||||
|
||||
from fastvideo import VideoGenerator # noqa: E402
|
||||
from fastvideo.api import ( # noqa: E402
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
OUTPUT_PATH = "video_samples_cosmos3_t2vs"
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.environ.get("COSMOS3_MODEL_PATH", "nvidia/Cosmos3-Nano")
|
||||
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_name,
|
||||
engine=EngineConfig(
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
offload=OffloadConfig(text_encoder=True, pin_cpu_memory=True, dit=False, vae=False),
|
||||
),
|
||||
)
|
||||
|
||||
load_start = time.perf_counter()
|
||||
generator = VideoGenerator.from_config(generator_config)
|
||||
load_time = time.perf_counter() - load_start
|
||||
|
||||
prompt = (
|
||||
"Ocean waves crash against a rocky shore at sunset, white foam spraying "
|
||||
"into the air as seagulls wheel overhead. Golden light, cinematic wide "
|
||||
"shot, the rhythmic roar of the surf."
|
||||
)
|
||||
request = GenerationRequest(
|
||||
prompt=prompt,
|
||||
sampling=SamplingConfig(
|
||||
num_frames=int(os.environ.get("COSMOS3_NUM_FRAMES", "189")),
|
||||
height=int(os.environ.get("COSMOS3_HEIGHT", "704")),
|
||||
width=int(os.environ.get("COSMOS3_WIDTH", "1280")),
|
||||
num_inference_steps=int(os.environ.get("COSMOS3_STEPS", "35")),
|
||||
guidance_scale=6.0,
|
||||
fps=24,
|
||||
seed=1024,
|
||||
),
|
||||
output=OutputConfig(output_path=OUTPUT_PATH, save_video=True, return_frames=False),
|
||||
)
|
||||
|
||||
start = time.perf_counter()
|
||||
result = generator.generate(request)
|
||||
gen_time = time.perf_counter() - start
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video+sound: {gen_time} seconds")
|
||||
print(f"Output written to: {result.video_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,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.glm_image import GlmImageVAEConfig
|
||||
@@ -16,6 +17,7 @@ __all__ = [
|
||||
"WanVAEConfig",
|
||||
"CosmosVAEConfig",
|
||||
"Cosmos25VAEConfig",
|
||||
"Cosmos3VAEConfig",
|
||||
"Gen3CVAEConfig",
|
||||
"Hunyuan15VAEConfig",
|
||||
"LTX2VAEConfig",
|
||||
|
||||
@@ -0,0 +1,277 @@
|
||||
"""Cosmos3 (Wan2.2-TI2V-5B) VAE config and checkpoint-key mapping.
|
||||
|
||||
The Cosmos3 checkpoint VAE is literally ``Wan-AI/Wan2.2-TI2V-5B-Diffusers``
|
||||
(diffusers ``AutoencoderKLWan``), so this config locks the Wan2.2 geometry:
|
||||
residual down/up blocks, ``patch_size=2``, ``z_dim=48``, ``base_dim=160``,
|
||||
``decoder_base_dim=256``, and ``scale_factor_spatial=16``. The 48-dim
|
||||
``latents_mean``/``latents_std`` are taken verbatim from the Cosmos3
|
||||
checkpoint's ``vae/config.json`` (identical to the canonical Wan2.2-TI2V-5B
|
||||
statistics).
|
||||
|
||||
Mirrors the :class:`Cosmos25VAEArchConfig` pattern. ``param_names_mapping`` /
|
||||
``map_official_key`` translate the *official* Wan2.2 VAE state-dict keys
|
||||
(nested-residual naming, e.g. ``encoder.downsamples.{b}.downsamples.{j}`` and
|
||||
``decoder.upsamples.{b}.upsamples.{j}``) into FastVideo's ``AutoencoderKLWan``
|
||||
key space. The standard diffusers checkpoint already ships native FastVideo
|
||||
keys, so these helpers exist for parity tooling and official ``.pth`` loading.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.vaes.wanvae import WanVAEArchConfig, WanVAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3VAEArchConfig(WanVAEArchConfig):
|
||||
# Wan2.2-TI2V-5B geometry (differs from the Wan2.1 WanVAEArchConfig
|
||||
# defaults: residual blocks, patch_size=2, z_dim=48, base_dim=160,
|
||||
# decoder_base_dim=256, scale_factor_spatial=16, 12 patch channels).
|
||||
_name_or_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
base_dim: int = 160
|
||||
decoder_base_dim: int | None = 256
|
||||
z_dim: int = 48
|
||||
dim_mult: tuple[int, ...] = (1, 2, 4, 4)
|
||||
num_res_blocks: int = 2
|
||||
attn_scales: tuple[float, ...] = ()
|
||||
temperal_downsample: tuple[bool, ...] = (False, True, True)
|
||||
dropout: float = 0.0
|
||||
is_residual: bool = True
|
||||
in_channels: int = 12
|
||||
out_channels: int = 12
|
||||
patch_size: int | None = 2
|
||||
scale_factor_temporal: int = 4
|
||||
scale_factor_spatial: int = 16
|
||||
clip_output: bool = False
|
||||
|
||||
# 48-dim statistics copied verbatim from the Cosmos3 checkpoint
|
||||
# (official_weights/cosmos3/vae/config.json).
|
||||
latents_mean: tuple[float, ...] = (
|
||||
-0.2289,
|
||||
-0.0052,
|
||||
-0.1323,
|
||||
-0.2339,
|
||||
-0.2799,
|
||||
0.0174,
|
||||
0.1838,
|
||||
0.1557,
|
||||
-0.1382,
|
||||
0.0542,
|
||||
0.2813,
|
||||
0.0891,
|
||||
0.157,
|
||||
-0.0098,
|
||||
0.0375,
|
||||
-0.1825,
|
||||
-0.2246,
|
||||
-0.1207,
|
||||
-0.0698,
|
||||
0.5109,
|
||||
0.2665,
|
||||
-0.2108,
|
||||
-0.2158,
|
||||
0.2502,
|
||||
-0.2055,
|
||||
-0.0322,
|
||||
0.1109,
|
||||
0.1567,
|
||||
-0.0729,
|
||||
0.0899,
|
||||
-0.2799,
|
||||
-0.123,
|
||||
-0.0313,
|
||||
-0.1649,
|
||||
0.0117,
|
||||
0.0723,
|
||||
-0.2839,
|
||||
-0.2083,
|
||||
-0.052,
|
||||
0.3748,
|
||||
0.0152,
|
||||
0.1957,
|
||||
0.1433,
|
||||
-0.2944,
|
||||
0.3573,
|
||||
-0.0548,
|
||||
-0.1681,
|
||||
-0.0667,
|
||||
)
|
||||
latents_std: tuple[float, ...] = (
|
||||
0.4765,
|
||||
1.0364,
|
||||
0.4514,
|
||||
1.1677,
|
||||
0.5313,
|
||||
0.499,
|
||||
0.4818,
|
||||
0.5013,
|
||||
0.8158,
|
||||
1.0344,
|
||||
0.5894,
|
||||
1.0901,
|
||||
0.6885,
|
||||
0.6165,
|
||||
0.8454,
|
||||
0.4978,
|
||||
0.5759,
|
||||
0.3523,
|
||||
0.7135,
|
||||
0.6804,
|
||||
0.5833,
|
||||
1.4146,
|
||||
0.8986,
|
||||
0.5659,
|
||||
0.7069,
|
||||
0.5338,
|
||||
0.4889,
|
||||
0.4917,
|
||||
0.4069,
|
||||
0.4999,
|
||||
0.6866,
|
||||
0.4093,
|
||||
0.5709,
|
||||
0.6065,
|
||||
0.6415,
|
||||
0.4944,
|
||||
0.5726,
|
||||
1.2042,
|
||||
0.5458,
|
||||
1.6887,
|
||||
0.3971,
|
||||
1.06,
|
||||
0.3943,
|
||||
0.5537,
|
||||
0.5444,
|
||||
0.4089,
|
||||
0.7468,
|
||||
0.7744,
|
||||
)
|
||||
|
||||
# Simple 1:1 renames. The nested-residual block remapping (encoder
|
||||
# downsamples / decoder upsamples / middle / head) is handled by
|
||||
# ``map_official_key()``.
|
||||
param_names_mapping: dict[str, str] = field(
|
||||
default_factory=lambda: {
|
||||
r"^conv1\.(.*)$": r"quant_conv.\1",
|
||||
r"^conv2\.(.*)$": r"post_quant_conv.\1",
|
||||
r"^encoder\.conv1\.(.*)$": r"encoder.conv_in.\1",
|
||||
r"^decoder\.conv1\.(.*)$": r"decoder.conv_in.\1",
|
||||
r"^encoder\.head\.0\.gamma$": r"encoder.norm_out.gamma",
|
||||
r"^encoder\.head\.2\.(.*)$": r"encoder.conv_out.\1",
|
||||
r"^decoder\.head\.0\.gamma$": r"decoder.norm_out.gamma",
|
||||
r"^decoder\.head\.2\.(.*)$": r"decoder.conv_out.\1",
|
||||
})
|
||||
|
||||
@staticmethod
|
||||
def map_official_key(key: str) -> str | None:
|
||||
"""Map a single official Wan2.2 VAE key into FastVideo key space.
|
||||
|
||||
Handles the residual (Wan2.2) module layout where each down/up block
|
||||
is a nested ``Sequential`` (``downsamples.{b}.downsamples.{j}`` /
|
||||
``upsamples.{b}.upsamples.{j}``) rather than the flat Wan2.1 indexing.
|
||||
Returns ``None`` for keys with no FastVideo counterpart.
|
||||
"""
|
||||
|
||||
def map_residual_subkey(prefix: str, sub: str) -> str | None:
|
||||
if re.match(r"^residual\.0\.gamma$", sub):
|
||||
return f"{prefix}.norm1.gamma"
|
||||
m = re.match(r"^residual\.2\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv1.{m.group(1)}"
|
||||
if re.match(r"^residual\.3\.gamma$", sub):
|
||||
return f"{prefix}.norm2.gamma"
|
||||
m = re.match(r"^residual\.6\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv2.{m.group(1)}"
|
||||
m = re.match(r"^shortcut\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.conv_shortcut.{m.group(1)}"
|
||||
return None
|
||||
|
||||
def map_attn_subkey(prefix: str, sub: str) -> str | None:
|
||||
if re.match(r"^norm\.gamma$", sub):
|
||||
return f"{prefix}.norm.gamma"
|
||||
m = re.match(r"^to_qkv\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.to_qkv.{m.group(1)}"
|
||||
m = re.match(r"^proj\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.proj.{m.group(1)}"
|
||||
return None
|
||||
|
||||
def map_resample_subkey(prefix: str, sub: str) -> str | None:
|
||||
m = re.match(r"^resample\.1\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.resample.1.{m.group(1)}"
|
||||
m = re.match(r"^time_conv\.(weight|bias)$", sub)
|
||||
if m:
|
||||
return f"{prefix}.time_conv.{m.group(1)}"
|
||||
return None
|
||||
|
||||
m = re.match(r"^conv1\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"quant_conv.{m.group(1)}"
|
||||
m = re.match(r"^conv2\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"post_quant_conv.{m.group(1)}"
|
||||
m = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.conv_in.{m.group(2)}"
|
||||
m = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.norm_out.gamma"
|
||||
m = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
|
||||
if m:
|
||||
return f"{m.group(1)}.conv_out.{m.group(2)}"
|
||||
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
|
||||
if m:
|
||||
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.0", m.group(2))
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
|
||||
if m:
|
||||
return map_attn_subkey(f"{m.group(1)}.mid_block.attentions.0", m.group(2))
|
||||
m = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
|
||||
if m:
|
||||
return map_residual_subkey(f"{m.group(1)}.mid_block.resnets.1", m.group(2))
|
||||
|
||||
# Encoder: downsamples.{block}.downsamples.{j}.* (nested residual layout)
|
||||
m = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
|
||||
if m:
|
||||
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
|
||||
if sub.startswith("resample.") or sub.startswith("time_conv."):
|
||||
return map_resample_subkey(f"encoder.down_blocks.{block_i}.downsampler", sub)
|
||||
return map_residual_subkey(f"encoder.down_blocks.{block_i}.resnets.{res_i}", sub)
|
||||
|
||||
# Decoder: upsamples.{block}.upsamples.{j}.* (nested residual layout)
|
||||
m = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
|
||||
if m:
|
||||
block_i, res_i, sub = int(m.group(1)), int(m.group(2)), m.group(3)
|
||||
if sub.startswith("resample.") or sub.startswith("time_conv."):
|
||||
return map_resample_subkey(f"decoder.up_blocks.{block_i}.upsampler", sub)
|
||||
return map_residual_subkey(f"decoder.up_blocks.{block_i}.resnets.{res_i}", sub)
|
||||
|
||||
return None
|
||||
|
||||
# ``__post_init__`` (scaling_factor / shift_factor / compression ratios) is
|
||||
# inherited unchanged from ``WanVAEArchConfig``.
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3VAEConfig(WanVAEConfig):
|
||||
"""Cosmos3 VAE config (reuses FastVideo's Wan2.2 ``AutoencoderKLWan``).
|
||||
|
||||
Subclasses :class:`WanVAEConfig` so the model reads the same runtime flags
|
||||
(``use_feature_cache``, ``use_light_vae``, tiling) and only swaps in the
|
||||
Cosmos3 = Wan2.2 ``arch_config``.
|
||||
"""
|
||||
|
||||
arch_config: Cosmos3VAEArchConfig = field(default_factory=Cosmos3VAEArchConfig)
|
||||
|
||||
use_feature_cache: bool = True
|
||||
use_tiling: bool = False
|
||||
use_temporal_tiling: bool = False
|
||||
use_parallel_tiling: bool = False
|
||||
|
||||
# ``__post_init__`` (blend_num_frames) is inherited from ``WanVAEConfig``.
|
||||
@@ -0,0 +1,69 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 pipeline configuration.
|
||||
|
||||
Reference of record: the official ``cosmos-framework`` / ``nvidia/Cosmos3-Nano``
|
||||
checkpoint (``model_index.json``). Cosmos3 is structurally different from
|
||||
Cosmos 2.5:
|
||||
|
||||
- Dual-pathway (UND + GEN) DiT lives entirely inside ``Cosmos3VFMTransformer``
|
||||
(``Cosmos3VideoConfig``).
|
||||
- No separate text encoder — the Qwen3-VL-text backbone is inside the DiT, so
|
||||
``text_encoder_configs`` is the empty tuple. The Qwen2 tokenizer is loaded as
|
||||
the ``text_tokenizer`` checkpoint module by the component loader.
|
||||
- VAE is Wan2.2 ``AutoencoderKLWan`` (z_dim=48, scale_factor_spatial=16),
|
||||
configured by ``Cosmos3VAEConfig`` (the checkpoint's exact latents_mean/std).
|
||||
- Scheduler is FastVideo-native ``UniPCMultistepScheduler`` configured for
|
||||
pure flow matching (flow_prediction, use_flow_sigmas), equivalent to the
|
||||
framework's ``FlowUniPCMultistepScheduler``. The checkpoint's diffusers-style
|
||||
scheduler config (karras/sigma_min/max) is coerced to the flow setup in
|
||||
``Cosmos3OmniDiffusersPipeline.initialize_pipeline``.
|
||||
- T2I default ``flow_shift`` is 3.0 (set per-request by ``_set_flow_shift``);
|
||||
T2V/I2V use the engine-init default of 1.0 baked into this config.
|
||||
"""
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.configs.models.dits.cosmos3 import (Cosmos3ArchConfig, Cosmos3VideoConfig)
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.vaes import Cosmos3VAEConfig # Wan2.2 AutoencoderKLWan
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3Config(PipelineConfig):
|
||||
"""Configuration for the Cosmos3 video generation pipeline (T2V/I2V/T2I).
|
||||
|
||||
Wires the framework-parity-verified Cosmos3 components: the native
|
||||
``Cosmos3VideoConfig`` DiT, the Wan2.2 ``Cosmos3VAEConfig`` VAE, the Qwen2
|
||||
tokenizer (loaded as ``text_tokenizer``), and the UniPC scheduler.
|
||||
"""
|
||||
|
||||
dit_config: DiTConfig = field(default_factory=lambda: Cosmos3VideoConfig(arch_config=Cosmos3ArchConfig()))
|
||||
|
||||
vae_config: VAEConfig = field(default_factory=Cosmos3VAEConfig)
|
||||
|
||||
# No separate text encoder: the Qwen3-VL-text backbone lives inside the DiT
|
||||
# and the pipeline tokenizes in Cosmos3DenoisingStage, so all three
|
||||
# text-encoder lists are empty (the generic text-encode stage is not used).
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(default_factory=tuple)
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=tuple)
|
||||
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor], ...] = field(default_factory=tuple)
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "bf16"
|
||||
text_encoder_precisions: tuple[str, ...] = field(default_factory=tuple)
|
||||
|
||||
embedded_cfg_scale: float = 0.0
|
||||
# T2V/I2V engine-init flow_shift (framework text2video/image2video default);
|
||||
# T2I overrides to 3.0 per request via Cosmos3DenoisingStage._set_flow_shift.
|
||||
flow_shift: float = 10.0
|
||||
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_decoder = True
|
||||
@@ -0,0 +1,133 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 sound tokenizer (AVAE) — decode path.
|
||||
|
||||
The Cosmos3 ``sound_tokenizer`` is an AVAE (audio VAE). Its shipped diffusers
|
||||
checkpoint is **decoder-only** (``decoder.*``) in ``AutoencoderOobleck`` naming
|
||||
with SnakeBeta activations and ``weight_g``/``weight_v`` weight-norm — exactly
|
||||
FastVideo's native :class:`~fastvideo.models.vaes.oobleck.OobleckDecoder`
|
||||
(verified bit-exact vs the framework in ``test_cosmos3_avae_parity``). Text-to-
|
||||
video+sound (t2vs) only needs DECODE: the DiT generates the sound latent and
|
||||
this module decodes it to a waveform, so only the decoder is ported (the
|
||||
SpectrogramConvNeXt encoder is not exported in the checkpoint).
|
||||
|
||||
Mirrors the framework ``AVAEModel.decode``: run the Oobleck decoder, then clamp
|
||||
to [-1, 1]. The VAE bottleneck's decode is the identity (the DiT already emits
|
||||
the post-bottleneck latent), so there is no bottleneck step here.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.oobleck import OobleckDecoder
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SoundVAEArchConfig:
|
||||
"""Cosmos3 AVAE decoder constants (from ``sound_tokenizer/config.json``)."""
|
||||
|
||||
dec_dim: int = 320 # decoder base channels
|
||||
vocoder_input_dim: int = 64 # latent channels in
|
||||
dec_c_mults: list[int] = field(default_factory=lambda: [1, 2, 4, 8, 16])
|
||||
dec_strides: list[int] = field(default_factory=lambda: [2, 4, 5, 6, 8])
|
||||
audio_channels: int = 2 # stereo
|
||||
sampling_rate: int = 48000
|
||||
|
||||
@property
|
||||
def hop_size(self) -> int:
|
||||
return int(np.prod(self.dec_strides)) # 1920
|
||||
|
||||
|
||||
class Cosmos3SoundVAE(nn.Module):
|
||||
"""Decoder-only Cosmos3 AVAE: latent ``[B, z, T]`` -> waveform ``[B, C, N]``."""
|
||||
|
||||
def __init__(self, arch: Cosmos3SoundVAEArchConfig | None = None) -> None:
|
||||
super().__init__()
|
||||
self.arch = arch or Cosmos3SoundVAEArchConfig()
|
||||
self.decoder = OobleckDecoder(
|
||||
channels=self.arch.dec_dim,
|
||||
input_channels=self.arch.vocoder_input_dim,
|
||||
audio_channels=self.arch.audio_channels,
|
||||
# The framework builds decoder blocks from ``reversed(dec_strides)``
|
||||
# (deepest first), so block strides are e.g. [8,6,5,4,2].
|
||||
upsampling_ratios=list(reversed(self.arch.dec_strides)),
|
||||
channel_multiples=list(self.arch.dec_c_mults),
|
||||
)
|
||||
|
||||
@property
|
||||
def sample_rate(self) -> int:
|
||||
return self.arch.sampling_rate
|
||||
|
||||
@property
|
||||
def audio_channels(self) -> int:
|
||||
return self.arch.audio_channels
|
||||
|
||||
@property
|
||||
def hop_size(self) -> int:
|
||||
return self.arch.hop_size
|
||||
|
||||
def get_latent_num_samples(self, num_audio_samples: int) -> int:
|
||||
"""Latent length for a given audio length (``AVAEInterface``: ``N // hop``)."""
|
||||
return int(num_audio_samples) // self.arch.hop_size
|
||||
|
||||
@torch.no_grad()
|
||||
def decode(self, latent: torch.Tensor) -> torch.Tensor:
|
||||
"""Decode normalized latent ``[B, z, T]`` to waveform ``[B, C, N]`` in [-1, 1].
|
||||
|
||||
Matches ``AVAEModel.decode``: Oobleck decoder then clamp to [-1, 1] (the
|
||||
VAE bottleneck decode is identity).
|
||||
"""
|
||||
audio = self.decoder(latent) # [B, C, N]
|
||||
return audio.clamp(-1.0, 1.0)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(
|
||||
cls,
|
||||
model_path: str,
|
||||
*,
|
||||
torch_dtype: torch.dtype | None = None,
|
||||
) -> "Cosmos3SoundVAE":
|
||||
"""Build + load the decoder from a ``sound_tokenizer`` directory.
|
||||
|
||||
Reads ``config.json`` (``dec_dim`` / ``vocoder_input_dim`` /
|
||||
``dec_c_mults`` / ``dec_strides`` / ``sampling_rate`` / ``stereo``) and
|
||||
loads the ``decoder.*`` weights (the checkpoint is decoder-only).
|
||||
"""
|
||||
from safetensors.torch import load_file
|
||||
|
||||
cfg_path = os.path.join(model_path, "config.json")
|
||||
with open(cfg_path) as f:
|
||||
cfg = json.load(f)
|
||||
arch = Cosmos3SoundVAEArchConfig(
|
||||
dec_dim=int(cfg["dec_dim"]),
|
||||
vocoder_input_dim=int(cfg["vocoder_input_dim"]),
|
||||
dec_c_mults=list(cfg["dec_c_mults"]),
|
||||
dec_strides=list(cfg["dec_strides"]),
|
||||
audio_channels=2 if cfg.get("stereo", True) else 1,
|
||||
sampling_rate=int(cfg.get("sampling_rate", 48000)),
|
||||
)
|
||||
model = cls(arch)
|
||||
|
||||
weights_path = os.path.join(model_path, "diffusion_pytorch_model.safetensors")
|
||||
state = load_file(weights_path)
|
||||
# Decoder-only checkpoint: strip the ``decoder.`` prefix.
|
||||
dec_state = {k[len("decoder."):]: v for k, v in state.items() if k.startswith("decoder.")}
|
||||
model.decoder.load_state_dict(dec_state, strict=True)
|
||||
logger.info("Loaded Cosmos3 sound AVAE decoder (%d params) from %s",
|
||||
sum(p.numel() for p in model.parameters()), model_path)
|
||||
|
||||
if torch_dtype is not None:
|
||||
model = model.to(dtype=torch_dtype)
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
EntryClass = Cosmos3SoundVAE
|
||||
File diff suppressed because it is too large
Load Diff
@@ -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"),
|
||||
@@ -1137,7 +1139,19 @@ class SchedulerLoader(ComponentLoader):
|
||||
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
# Diffusers checkpoints can carry newer scheduler config keys than the
|
||||
# vendored scheduler accepts (e.g. shift_terminal / sigma_min / sigma_max
|
||||
# from a newer diffusers release). Filter to the class's __init__ params,
|
||||
# mirroring diffusers' ``from_config``, so loading is robust to schema
|
||||
# drift instead of crashing on an unexpected kwarg.
|
||||
import inspect
|
||||
valid_params = set(inspect.signature(scheduler_cls.__init__).parameters)
|
||||
filtered_config = {k: v for k, v in config.items() if k in valid_params}
|
||||
dropped = sorted(set(config) - set(filtered_config))
|
||||
if dropped:
|
||||
logger.warning("Scheduler %s: dropping unsupported config keys %s", class_name, dropped)
|
||||
|
||||
scheduler = scheduler_cls(**filtered_config)
|
||||
if fastvideo_args.pipeline_config.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
|
||||
return scheduler
|
||||
|
||||
@@ -37,6 +37,9 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
|
||||
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
|
||||
# Cosmos3-Nano's checkpoint model_index names the DiT "Cosmos3OmniTransformer";
|
||||
# map that HF class name to FastVideo's native Cosmos3VFMTransformer.
|
||||
"Cosmos3OmniTransformer": ("dits", "cosmos3", "Cosmos3VFMTransformer"),
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
|
||||
|
||||
@@ -94,6 +94,9 @@ class OobleckDecoderBlock(nn.Module):
|
||||
input_dim, output_dim,
|
||||
kernel_size=2 * stride, stride=stride,
|
||||
padding=math.ceil(stride / 2),
|
||||
# Clean L*stride upsample for both parities; a no-op (0) for even
|
||||
# strides (Stable Audio), needed for odd strides (Cosmos3: 5).
|
||||
output_padding=stride % 2,
|
||||
))
|
||||
self.res_unit1 = OobleckResidualUnit(output_dim, dilation=1)
|
||||
self.res_unit2 = OobleckResidualUnit(output_dim, dilation=3)
|
||||
|
||||
@@ -0,0 +1,708 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-native Cosmos3 video pipeline (T2V / I2V / T2I).
|
||||
|
||||
This replaces the earlier vllm-omni-derived skeleton with a native, stage-based
|
||||
:class:`ComposedPipelineBase` pipeline that wires the framework-parity-verified
|
||||
Cosmos3 components:
|
||||
|
||||
* tokenizer: Qwen2 ``Qwen2TokenizerFast`` + chat template (the only allowed
|
||||
third-party model-adjacent dependency; tokenizers are explicitly permitted),
|
||||
* VAE: FastVideo-native ``AutoencoderKLWan`` (Wan2.2) via ``Cosmos3VAEConfig``;
|
||||
encode normalizes ``(mu - mean) * inv_std`` and decode denormalizes + clamps,
|
||||
* sequence-packing: :func:`pack_cosmos3_video_sequence` (native, parity-tested),
|
||||
* DiT: ``Cosmos3VFMTransformer`` (native, bit-identical to the framework),
|
||||
* scheduler: FastVideo-native ``UniPCMultistepScheduler`` configured for pure
|
||||
flow matching (``flow_prediction`` + ``use_flow_sigmas``), numerically
|
||||
equivalent to the framework's ``FlowUniPCMultistepScheduler`` (parity-tested
|
||||
in ``test_cosmos3_scheduler_parity``).
|
||||
|
||||
The denoise/CFG glue is a faithful port of the framework's
|
||||
``Cosmos3OmniDiffusersPipeline`` math (mirrored in the framework-equivalent
|
||||
``diffusers_cosmos3.pipeline``): per UniPC timestep, run a SEQUENTIAL conditional
|
||||
then unconditional pass (each repacks the sequence with the prompt / negative
|
||||
prompt token ids, forwards the DiT, and zeros the prediction on conditioning
|
||||
frames), then combine ``v = uncond + guidance * (cond - uncond)`` and take one
|
||||
``scheduler.step(model_output=v, timestep, sample=latent)``. ``timestep_scale``
|
||||
is applied to the per-token timesteps *inside* the DiT (its ``forward`` already
|
||||
multiplies ``vision_timesteps * timestep_scale`` before the time embedder), so
|
||||
the loop passes raw scheduler timesteps to the packer.
|
||||
|
||||
The pure denoise math lives in :class:`Cosmos3DenoiseEngine` and the free
|
||||
function :func:`cosmos3_get_cfg_velocity` so it can be unit-/parity-tested
|
||||
directly against the framework oracle without constructing the full pipeline.
|
||||
|
||||
No diffusers/transformers *model* classes are imported at runtime here; only the
|
||||
Qwen2 tokenizer (loaded by the component loader) and the UniPC scheduler are
|
||||
third-party, both explicitly allowed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_unipc_multistep import (
|
||||
UniPCMultistepScheduler, )
|
||||
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
|
||||
Cosmos3ActionItem,
|
||||
Cosmos3SampleInputs,
|
||||
Cosmos3SoundItem,
|
||||
Cosmos3VisionItem,
|
||||
pack_cosmos3_video_sequence,
|
||||
)
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# System prompts, verbatim from the framework (diffusers_cosmos3.pipeline).
|
||||
_SYSTEM_PROMPT_IMAGE = "You are a helpful assistant who will generate images from a give prompt."
|
||||
_SYSTEM_PROMPT_VIDEO = "You are a helpful assistant who will generate videos from a give prompt."
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Special-token resolution (Qwen2 chat tokenizer)
|
||||
# ===========================================================================
|
||||
def cosmos3_special_tokens(tokenizer: Any) -> dict[str, int]:
|
||||
"""Resolve the Cosmos3 generation special tokens from a Qwen2 tokenizer.
|
||||
|
||||
Mirrors the framework's ``llm_special_tokens``:
|
||||
``start_of_generation=<|vision_start|>``, ``end_of_generation=<|vision_end|>``,
|
||||
``eos_token_id=tokenizer.eos_token_id``.
|
||||
"""
|
||||
return {
|
||||
"start_of_generation": int(tokenizer.convert_tokens_to_ids("<|vision_start|>")),
|
||||
"end_of_generation": int(tokenizer.convert_tokens_to_ids("<|vision_end|>")),
|
||||
"eos_token_id": int(tokenizer.eos_token_id),
|
||||
}
|
||||
|
||||
|
||||
def cosmos3_tokenize_caption(
|
||||
tokenizer: Any,
|
||||
caption: str,
|
||||
*,
|
||||
is_video: bool = False,
|
||||
use_system_prompt: bool = False,
|
||||
) -> list[int]:
|
||||
"""Tokenize a caption with the Qwen2 chat template (framework-faithful).
|
||||
|
||||
Optionally prepends an image/video system prompt; always adds the
|
||||
generation prompt and disables ``add_vision_id`` (matching the framework's
|
||||
``tokenize_caption``).
|
||||
"""
|
||||
conversations: list[dict[str, str]] = []
|
||||
if use_system_prompt:
|
||||
conversations.append({
|
||||
"role": "system",
|
||||
"content": _SYSTEM_PROMPT_VIDEO if is_video else _SYSTEM_PROMPT_IMAGE,
|
||||
})
|
||||
conversations.append({"role": "user", "content": caption})
|
||||
token_ids = tokenizer.apply_chat_template(
|
||||
conversations,
|
||||
tokenize=True,
|
||||
add_generation_prompt=True,
|
||||
add_vision_id=False,
|
||||
)
|
||||
return list(token_ids)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Reasoning (VLM text generation) — und (causal) pathway + lm_head
|
||||
# ===========================================================================
|
||||
def cosmos3_generate_reasoner_text(
|
||||
transformer: Any,
|
||||
input_ids: list[int],
|
||||
max_new_tokens: int,
|
||||
*,
|
||||
eos_token_id: int | list[int] | None = None,
|
||||
) -> list[int]:
|
||||
"""Greedy text reasoning via the und (causal) backbone + ``lm_head``.
|
||||
|
||||
Mirrors the framework ``generate_reasoner_text`` (text-only prefill, greedy):
|
||||
only the und-pathway weights (no ``_moe_gen``) + ``embed_tokens`` / ``norm`` /
|
||||
``lm_head`` participate; the generation pathway and the VFM multimodal
|
||||
embedders are bypassed (no vision/sound/action tokens). Token-for-token
|
||||
identical to the framework reasoner (``test_cosmos3_reasoning_parity``).
|
||||
|
||||
Re-prefills each step (no KV cache) — correctness-first; a KV-cache fast path
|
||||
is a later optimization. Returns the newly generated token ids.
|
||||
"""
|
||||
device = next(transformer.parameters()).device
|
||||
ids = [int(x) for x in input_ids]
|
||||
eos: set[int] = set()
|
||||
if eos_token_id is not None:
|
||||
eos = {int(eos_token_id)} if isinstance(eos_token_id, int) else {int(x) for x in eos_token_id}
|
||||
|
||||
new_tokens: list[int] = []
|
||||
for _ in range(int(max_new_tokens)):
|
||||
n = len(ids)
|
||||
pos = torch.arange(n).unsqueeze(0).expand(3, -1).contiguous().to(device)
|
||||
out = transformer(
|
||||
text_ids=torch.tensor(ids, device=device, dtype=torch.long),
|
||||
text_indexes=torch.arange(n, device=device),
|
||||
position_ids=pos,
|
||||
sequence_length=n,
|
||||
split_lens=[n],
|
||||
attn_modes=["causal"],
|
||||
vision_tokens=[],
|
||||
vision_token_shapes=[],
|
||||
vision_sequence_indexes=torch.empty(0, dtype=torch.long, device=device),
|
||||
vision_timesteps=torch.empty(0, device=device),
|
||||
vision_mse_loss_indexes=torch.empty(0, dtype=torch.long, device=device),
|
||||
vision_noisy_frame_indexes=[],
|
||||
)
|
||||
logits = transformer.lm_head(out["last_hidden_state"][n - 1]) # [vocab]
|
||||
nxt = int(logits.argmax().item())
|
||||
ids.append(nxt)
|
||||
new_tokens.append(nxt)
|
||||
if nxt in eos:
|
||||
break
|
||||
return new_tokens
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# VAE encode/decode bridge (normalize / denormalize, matching the framework)
|
||||
# ===========================================================================
|
||||
@dataclass
|
||||
class _VaeNorm:
|
||||
"""Cached ``mean`` / ``inv_std`` for VAE (de)normalization."""
|
||||
|
||||
mean: torch.Tensor # [z_dim]
|
||||
inv_std: torch.Tensor # [z_dim]
|
||||
|
||||
@classmethod
|
||||
def from_vae(cls, vae: Any, dtype: torch.dtype) -> _VaeNorm:
|
||||
mean = torch.tensor(list(vae.config.latents_mean), dtype=dtype)
|
||||
std = torch.tensor(list(vae.config.latents_std), dtype=dtype)
|
||||
return cls(mean=mean, inv_std=1.0 / std)
|
||||
|
||||
|
||||
def cosmos3_vae_encode(vae: Any, video: torch.Tensor, norm: _VaeNorm) -> torch.Tensor:
|
||||
"""Encode ``[B, 3, T, H, W]`` pixels in [-1, 1] to NORMALIZED latents.
|
||||
|
||||
Matches the framework ``DiffusersWan22VAE.encode``: take the posterior mode
|
||||
and apply ``(mu - mean) * inv_std``. FastVideo's ``AutoencoderKLWan.encode``
|
||||
returns a ``DiagonalGaussianDistribution``; we read ``.mode()``.
|
||||
"""
|
||||
in_dtype = video.dtype
|
||||
device = video.device
|
||||
mean = norm.mean.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
inv_std = norm.inv_std.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
raw_mu = vae.encode(video).mode()
|
||||
return ((raw_mu - mean) * inv_std).to(in_dtype)
|
||||
|
||||
|
||||
def cosmos3_vae_decode(vae: Any, latents: torch.Tensor, norm: _VaeNorm) -> torch.Tensor:
|
||||
"""Decode NORMALIZED latents ``[B, z, T, H, W]`` to pixels ``[B, 3, T, H, W]``.
|
||||
|
||||
Inverts the normalization (``z / inv_std + mean``) then calls
|
||||
``vae.decode`` (which already clamps to [-1, 1]).
|
||||
"""
|
||||
in_dtype = latents.dtype
|
||||
device = latents.device
|
||||
mean = norm.mean.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
inv_std = norm.inv_std.to(device=device, dtype=in_dtype).view(1, -1, 1, 1, 1)
|
||||
z_raw = latents / inv_std + mean
|
||||
out = vae.decode(z_raw)
|
||||
if isinstance(out, tuple):
|
||||
out = out[0]
|
||||
if hasattr(out, "sample"):
|
||||
out = out.sample
|
||||
return out.to(in_dtype)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Per-vision-item packing geometry
|
||||
# ===========================================================================
|
||||
@dataclass
|
||||
class Cosmos3VisionSpec:
|
||||
"""Geometry + conditioning for one vision item in a denoise run.
|
||||
|
||||
Args:
|
||||
condition_frame_indexes: Latent-frame indices kept clean.
|
||||
shape: ``(C, T, H, W)`` of the latent for this item.
|
||||
"""
|
||||
|
||||
shape: tuple[int, int, int, int]
|
||||
condition_frame_indexes: list[int]
|
||||
|
||||
@property
|
||||
def numel(self) -> int:
|
||||
return int(math.prod(self.shape))
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Pure denoise/CFG math (parity oracle target)
|
||||
# ===========================================================================
|
||||
def _split_flat_latent(flat: torch.Tensor, specs: list[Any]) -> list[torch.Tensor]:
|
||||
"""Split a flat vector into per-item tensors via each spec's ``numel``/``shape``.
|
||||
|
||||
Shared by vision (``[C, T, H, W]``), sound (``[C, T]``), and action
|
||||
(``[T, D]``) specs — every spec exposes ``numel`` and ``shape``.
|
||||
"""
|
||||
out: list[torch.Tensor] = []
|
||||
offset = 0
|
||||
for spec in specs:
|
||||
out.append(flat[offset:offset + spec.numel].reshape(spec.shape))
|
||||
offset += spec.numel
|
||||
return out
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SoundSpec:
|
||||
"""Geometry + conditioning for one sound item in a denoise run.
|
||||
|
||||
Args:
|
||||
shape: ``(C, T)`` of the sound latent (channels, temporal frames).
|
||||
condition_frame_indexes: Latent-frame indices kept clean (``[]`` for t2vs).
|
||||
fps: Sound latent FPS (``sound_latent_fps``); used iff fps modulation is on.
|
||||
"""
|
||||
|
||||
shape: tuple[int, int]
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
fps: float | None = None
|
||||
|
||||
@property
|
||||
def numel(self) -> int:
|
||||
return int(math.prod(self.shape))
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3ActionSpec:
|
||||
"""Geometry + conditioning for one action item in a denoise run.
|
||||
|
||||
Args:
|
||||
shape: ``(T, action_dim)`` of the action latent.
|
||||
condition_frame_indexes: Frame indices kept clean (conditioning actions).
|
||||
domain_id: Embodiment domain id for the domain-aware action projection.
|
||||
fps: Action FPS; used iff fps modulation is on.
|
||||
"""
|
||||
|
||||
shape: tuple[int, int]
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
domain_id: int = 0
|
||||
fps: float | None = None
|
||||
|
||||
@property
|
||||
def numel(self) -> int:
|
||||
return int(math.prod(self.shape))
|
||||
|
||||
|
||||
def cosmos3_get_cfg_velocity(
|
||||
*,
|
||||
transformer: Any,
|
||||
flat_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
guidance: float,
|
||||
specs: list[Cosmos3VisionSpec],
|
||||
cond_token_ids: list[int],
|
||||
uncond_token_ids: list[int],
|
||||
special_tokens: dict[str, int],
|
||||
latent_patch_size: int,
|
||||
temporal_modality_margin: int,
|
||||
reset_spatial_ids: bool,
|
||||
enable_fps_modulation: bool,
|
||||
base_fps: float,
|
||||
temporal_compression_factor: int,
|
||||
include_end_of_generation_token: bool = False,
|
||||
fps_per_item: list[float] | None = None,
|
||||
normalize_cfg: bool = False,
|
||||
sound_specs: list[Cosmos3SoundSpec] | None = None,
|
||||
sound_fps_per_item: list[float] | None = None,
|
||||
action_specs: list[Cosmos3ActionSpec] | None = None,
|
||||
action_fps_per_item: list[float] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Sequential-CFG velocity for one denoise step (framework math).
|
||||
|
||||
Replicates the framework ``get_cfg_velocity``:
|
||||
|
||||
1. split ``flat_latent`` into per-vision-item ``[C, T, H, W]`` latents,
|
||||
2. run a conditional pass (prompt tokens) and an unconditional pass
|
||||
(negative-prompt tokens); each repacks via
|
||||
:func:`pack_cosmos3_video_sequence`, forwards the DiT to obtain
|
||||
``preds_vision`` (a list of ``[1, C, T, H, W]`` unpatchified noisy-frame
|
||||
predictions), and zeros the prediction on conditioning frames
|
||||
(``pred * (1 - condition_mask)``),
|
||||
3. combine ``v = uncond + guidance * (cond - uncond)`` (optionally
|
||||
norm-rescaled), returned flattened to match ``flat_latent``.
|
||||
|
||||
``timestep`` is a scalar tensor (raw scheduler timestep); ``timestep_scale``
|
||||
is applied inside the DiT, so it is passed through unscaled here.
|
||||
"""
|
||||
assert timestep.numel() == 1, "timestep must be a scalar"
|
||||
timestep_value = float(timestep.reshape(()).item())
|
||||
|
||||
# Combined flat layout: [all vision | all action | all sound], matching the
|
||||
# framework per-sample concat order ([vision_i | action_i | sound_i]); single
|
||||
# sample here.
|
||||
vision_total = sum(spec.numel for spec in specs)
|
||||
action_total = sum(spec.numel for spec in action_specs) if action_specs else 0
|
||||
noise_x_vision = _split_flat_latent(flat_latent[:vision_total], specs)
|
||||
noise_x_action = (_split_flat_latent(flat_latent[vision_total:vision_total +
|
||||
action_total], action_specs) if action_specs else None)
|
||||
noise_x_sound = (_split_flat_latent(flat_latent[vision_total +
|
||||
action_total:], sound_specs) if sound_specs else None)
|
||||
device = next(transformer.parameters()).device
|
||||
|
||||
def _run(token_ids: list[int]) -> torch.Tensor:
|
||||
sound_items: list[Cosmos3SoundItem] = []
|
||||
if sound_specs is not None and noise_x_sound is not None:
|
||||
sound_items = [
|
||||
Cosmos3SoundItem(
|
||||
latent=noise_x_sound[i],
|
||||
condition_frame_indexes=list(ss.condition_frame_indexes),
|
||||
fps=(sound_fps_per_item[i] if sound_fps_per_item is not None else None),
|
||||
) for i, ss in enumerate(sound_specs)
|
||||
]
|
||||
action_items: list[Cosmos3ActionItem] = []
|
||||
if action_specs is not None and noise_x_action is not None:
|
||||
action_items = [
|
||||
Cosmos3ActionItem(
|
||||
latent=noise_x_action[i],
|
||||
condition_frame_indexes=list(asp.condition_frame_indexes),
|
||||
domain_id=asp.domain_id,
|
||||
fps=(action_fps_per_item[i] if action_fps_per_item is not None else None),
|
||||
) for i, asp in enumerate(action_specs)
|
||||
]
|
||||
samples = [
|
||||
Cosmos3SampleInputs(
|
||||
text_ids=list(token_ids),
|
||||
vision=Cosmos3VisionItem(
|
||||
latent=latent,
|
||||
condition_frame_indexes=list(spec.condition_frame_indexes),
|
||||
fps=(fps_per_item[i] if fps_per_item is not None else None),
|
||||
),
|
||||
sound=(sound_items[i] if i < len(sound_items) else None),
|
||||
action=(action_items[i] if i < len(action_items) else None),
|
||||
timestep=timestep_value,
|
||||
) for i, (latent, spec) in enumerate(zip(noise_x_vision, specs, strict=False))
|
||||
]
|
||||
packed = pack_cosmos3_video_sequence(
|
||||
samples,
|
||||
special_tokens,
|
||||
latent_patch_size=latent_patch_size,
|
||||
include_end_of_generation_token=include_end_of_generation_token,
|
||||
temporal_modality_margin=temporal_modality_margin,
|
||||
reset_spatial_ids=reset_spatial_ids,
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=temporal_compression_factor,
|
||||
)
|
||||
out = transformer(**packed.to_dit_kwargs(device=device))
|
||||
|
||||
# Vision velocity: zero on conditioning frames, per item, flattened.
|
||||
vision_vel = torch.zeros(vision_total, device=flat_latent.device, dtype=flat_latent.dtype)
|
||||
preds = out.get("preds_vision")
|
||||
if preds is not None:
|
||||
items: list[torch.Tensor] = []
|
||||
for pred, cond_mask in zip(preds, packed.vision_condition_mask, strict=False):
|
||||
pred = pred.squeeze(0) if pred.dim() == 5 else pred # [C, T, H, W]
|
||||
keep = (1.0 - cond_mask).to(dtype=pred.dtype, device=pred.device) # [T,1,1]
|
||||
items.append(pred * keep if keep.sum() > 0 else torch.zeros_like(pred))
|
||||
vision_vel = torch.cat([v.reshape(-1) for v in items]).to(flat_latent.dtype)
|
||||
|
||||
parts = [vision_vel]
|
||||
|
||||
if action_specs:
|
||||
# Action velocity: preds_action are per-item [T, D], already zero on
|
||||
# clean frames; zero on cond frames defensively.
|
||||
action_vel = torch.zeros(action_total, device=flat_latent.device, dtype=flat_latent.dtype)
|
||||
preds_a = out.get("preds_action")
|
||||
if preds_a is not None:
|
||||
a_items: list[torch.Tensor] = []
|
||||
for pred, cond_mask in zip(preds_a, packed.action_condition_mask, strict=False):
|
||||
pred = pred.squeeze(0) if pred.dim() == 3 else pred # [T, D]
|
||||
keep = (1.0 - cond_mask).reshape(-1, 1).to(dtype=pred.dtype, device=pred.device) # [T, 1]
|
||||
a_items.append(pred * keep)
|
||||
action_vel = torch.cat([v.reshape(-1) for v in a_items]).to(flat_latent.dtype)
|
||||
parts.append(action_vel)
|
||||
|
||||
if sound_specs:
|
||||
# Sound velocity: preds_sound are per-item [C, T], already zero on clean
|
||||
# frames (unpack fills only noisy frames); zero on cond frames defensively.
|
||||
sound_total = sum(spec.numel for spec in sound_specs)
|
||||
sound_vel = torch.zeros(sound_total, device=flat_latent.device, dtype=flat_latent.dtype)
|
||||
preds_s = out.get("preds_sound")
|
||||
if preds_s is not None:
|
||||
s_items: list[torch.Tensor] = []
|
||||
for pred, cond_mask in zip(preds_s, packed.sound_condition_mask, strict=False):
|
||||
pred = pred.squeeze(0) if pred.dim() == 3 else pred # [C, T]
|
||||
keep = (1.0 - cond_mask).reshape(1, -1).to(dtype=pred.dtype, device=pred.device) # [1, T]
|
||||
s_items.append(pred * keep)
|
||||
sound_vel = torch.cat([v.reshape(-1) for v in s_items]).to(flat_latent.dtype)
|
||||
parts.append(sound_vel)
|
||||
|
||||
return vision_vel if len(parts) == 1 else torch.cat(parts)
|
||||
|
||||
cond_v = _run(cond_token_ids)
|
||||
uncond_v = _run(uncond_token_ids)
|
||||
v_pred = uncond_v + guidance * (cond_v - uncond_v)
|
||||
if normalize_cfg:
|
||||
scale = (torch.norm(cond_v) / (torch.norm(v_pred) + 1e-8)).clamp(min=0.0, max=1.0)
|
||||
v_pred = v_pred * scale
|
||||
return v_pred
|
||||
|
||||
|
||||
class Cosmos3DenoiseEngine:
|
||||
"""Stateless denoise driver tying CFG velocity to UniPC stepping.
|
||||
|
||||
Holds the transformer + scheduler + packing constants and runs the full
|
||||
UniPC denoise loop. Kept separate from the pipeline so it can be exercised
|
||||
in isolation (smoke + parity tests) with stub or real components.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
transformer: Any,
|
||||
scheduler: Any,
|
||||
special_tokens: dict[str, int],
|
||||
latent_patch_size: int,
|
||||
temporal_modality_margin: int,
|
||||
reset_spatial_ids: bool,
|
||||
enable_fps_modulation: bool,
|
||||
base_fps: float,
|
||||
temporal_compression_factor: int,
|
||||
include_end_of_generation_token: bool = False,
|
||||
) -> None:
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
self.special_tokens = special_tokens
|
||||
self.latent_patch_size = latent_patch_size
|
||||
self.temporal_modality_margin = temporal_modality_margin
|
||||
self.reset_spatial_ids = reset_spatial_ids
|
||||
self.enable_fps_modulation = enable_fps_modulation
|
||||
self.base_fps = base_fps
|
||||
self.temporal_compression_factor = temporal_compression_factor
|
||||
self.include_end_of_generation_token = include_end_of_generation_token
|
||||
|
||||
def velocity(
|
||||
self,
|
||||
*,
|
||||
flat_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
guidance: float,
|
||||
specs: list[Cosmos3VisionSpec],
|
||||
cond_token_ids: list[int],
|
||||
uncond_token_ids: list[int],
|
||||
fps_per_item: list[float] | None = None,
|
||||
sound_specs: list[Cosmos3SoundSpec] | None = None,
|
||||
sound_fps_per_item: list[float] | None = None,
|
||||
action_specs: list[Cosmos3ActionSpec] | None = None,
|
||||
action_fps_per_item: list[float] | None = None,
|
||||
) -> torch.Tensor:
|
||||
return cosmos3_get_cfg_velocity(
|
||||
transformer=self.transformer,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
specs=specs,
|
||||
cond_token_ids=cond_token_ids,
|
||||
uncond_token_ids=uncond_token_ids,
|
||||
special_tokens=self.special_tokens,
|
||||
latent_patch_size=self.latent_patch_size,
|
||||
temporal_modality_margin=self.temporal_modality_margin,
|
||||
reset_spatial_ids=self.reset_spatial_ids,
|
||||
enable_fps_modulation=self.enable_fps_modulation,
|
||||
base_fps=self.base_fps,
|
||||
temporal_compression_factor=self.temporal_compression_factor,
|
||||
include_end_of_generation_token=self.include_end_of_generation_token,
|
||||
fps_per_item=fps_per_item,
|
||||
sound_specs=sound_specs,
|
||||
sound_fps_per_item=sound_fps_per_item,
|
||||
action_specs=action_specs,
|
||||
action_fps_per_item=action_fps_per_item,
|
||||
)
|
||||
|
||||
def denoise(
|
||||
self,
|
||||
*,
|
||||
flat_latent: torch.Tensor,
|
||||
timesteps: torch.Tensor,
|
||||
guidance: float,
|
||||
specs: list[Cosmos3VisionSpec],
|
||||
cond_token_ids: list[int],
|
||||
uncond_token_ids: list[int],
|
||||
fps_per_item: list[float] | None = None,
|
||||
progress_bar: Any | None = None,
|
||||
sound_specs: list[Cosmos3SoundSpec] | None = None,
|
||||
sound_fps_per_item: list[float] | None = None,
|
||||
action_specs: list[Cosmos3ActionSpec] | None = None,
|
||||
action_fps_per_item: list[float] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run the full UniPC denoise loop, returning the final flat latent.
|
||||
|
||||
For each timestep: compute the sequential-CFG velocity, then
|
||||
``scheduler.step(model_output=v, timestep, sample=latent.unsqueeze(0))``
|
||||
(the framework steps with a leading batch axis), squeezing back to flat.
|
||||
For t2vs the flat latent is ``[vision | sound]`` and the velocity covers
|
||||
both; the scheduler steps the combined vector jointly.
|
||||
"""
|
||||
latent = flat_latent
|
||||
iterator = progress_bar(timesteps) if progress_bar is not None else timesteps
|
||||
for t in iterator:
|
||||
v_pred = self.velocity(
|
||||
flat_latent=latent,
|
||||
timestep=t.reshape(1),
|
||||
guidance=guidance,
|
||||
specs=specs,
|
||||
cond_token_ids=cond_token_ids,
|
||||
uncond_token_ids=uncond_token_ids,
|
||||
fps_per_item=fps_per_item,
|
||||
sound_specs=sound_specs,
|
||||
sound_fps_per_item=sound_fps_per_item,
|
||||
action_specs=action_specs,
|
||||
action_fps_per_item=action_fps_per_item,
|
||||
)
|
||||
stepped = self.scheduler.step(
|
||||
model_output=v_pred,
|
||||
timestep=t,
|
||||
sample=latent.unsqueeze(0),
|
||||
return_dict=False,
|
||||
)[0]
|
||||
latent = stepped.squeeze(0)
|
||||
return latent
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Pipeline (ComposedPipelineBase)
|
||||
# ===========================================================================
|
||||
class Cosmos3OmniDiffusersPipeline(ComposedPipelineBase):
|
||||
"""Cosmos3 video generation pipeline (T2V / I2V / T2I).
|
||||
|
||||
Stage-based ``ComposedPipelineBase`` pipeline. The required modules
|
||||
(``transformer`` / ``vae`` / ``scheduler`` / ``text_tokenizer``) are loaded
|
||||
from the ``nvidia/Cosmos3-Nano`` checkpoint by the component loader. The
|
||||
class name matches the checkpoint ``model_index.json`` ``_class_name`` so
|
||||
the registry resolves it directly.
|
||||
|
||||
The denoise/CFG/VAE math is delegated to module-level helpers
|
||||
(:func:`cosmos3_get_cfg_velocity`, :class:`Cosmos3DenoiseEngine`,
|
||||
:func:`cosmos3_vae_encode` / :func:`cosmos3_vae_decode`) which are
|
||||
framework-parity tested in ``tests/local_tests/cosmos3``.
|
||||
"""
|
||||
|
||||
is_video_pipeline = True
|
||||
# ``vision_encoder`` / ``sound_tokenizer`` ship in the checkpoint but the
|
||||
# video path does not need them; they are intentionally omitted here.
|
||||
_required_config_modules = ["text_tokenizer", "vae", "transformer", "scheduler"]
|
||||
|
||||
# Engine-init flow_shift (T2V/I2V); T2I overrides to 3.0 per request.
|
||||
_engine_init_flow_shift: float = 1.0
|
||||
# Class-attribute defaults so ``__new__``-based unit tests can read these
|
||||
# before ``initialize_pipeline`` runs.
|
||||
scheduler: Any = None
|
||||
_base_scheduler_config: Any = None
|
||||
_current_flow_shift: float | None = None
|
||||
|
||||
@staticmethod
|
||||
def _flow_scheduler_config(config: Any) -> dict[str, Any]:
|
||||
"""Coerce a loaded UniPC config to the framework's flow-matching setup.
|
||||
|
||||
The checkpoint ``scheduler_config.json`` carries diffusers-style fields
|
||||
(``use_karras_sigmas=True``, ``sigma_min``/``sigma_max``, beta schedule)
|
||||
that do not describe the framework sampler. The framework uses
|
||||
``FlowUniPCMultistepScheduler`` (pure flow matching: ``shift`` +
|
||||
``num_train_timesteps`` only). FastVideo's vendored UniPC checks
|
||||
``use_karras_sigmas`` *before* ``use_flow_sigmas``, so leaving karras on
|
||||
builds diffusion-style sigmas and the denoise diverges to NaN. Force the
|
||||
flow config here (parity-verified in ``test_cosmos3_scheduler_parity``).
|
||||
"""
|
||||
cfg = dict(config)
|
||||
cfg.update(
|
||||
use_karras_sigmas=False,
|
||||
use_exponential_sigmas=False,
|
||||
use_beta_sigmas=False,
|
||||
use_flow_sigmas=True,
|
||||
prediction_type="flow_prediction",
|
||||
predict_x0=True,
|
||||
final_sigmas_type="zero",
|
||||
)
|
||||
return cfg
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Bind the loaded scheduler + snapshot its config so per-request
|
||||
flow_shift rebuilds are cheap and the engine-init shift is applied."""
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
engine_shift = getattr(pipeline_config, "flow_shift", None)
|
||||
if engine_shift is not None:
|
||||
self._engine_init_flow_shift = float(engine_shift)
|
||||
scheduler = self.get_module("scheduler")
|
||||
if scheduler is not None:
|
||||
# Rebuild from a flow-coerced config so the runtime scheduler matches
|
||||
# the framework sampler (the loaded checkpoint config is diffusers-style).
|
||||
flow_config = self._flow_scheduler_config(scheduler.config)
|
||||
self.scheduler = UniPCMultistepScheduler.from_config(flow_config)
|
||||
if isinstance(self.modules, dict):
|
||||
self.modules["scheduler"] = self.scheduler
|
||||
self._base_scheduler_config = self.scheduler.config
|
||||
self._current_flow_shift = float(getattr(self.scheduler.config, "flow_shift", 1.0))
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Wire the Cosmos3 stages.
|
||||
|
||||
The whole text->latent->denoise->decode flow is custom (sequential CFG
|
||||
with per-pass repacking), so a single :class:`Cosmos3DenoisingStage`
|
||||
owns it. ``InputValidationStage`` runs first for the standard checks.
|
||||
"""
|
||||
from fastvideo.pipelines.stages import InputValidationStage
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import Cosmos3DenoisingStage
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=Cosmos3DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae"),
|
||||
tokenizer=self.get_module("text_tokenizer"),
|
||||
pipeline=self,
|
||||
),
|
||||
)
|
||||
|
||||
# -- Scheduler control --------------------------------------------------
|
||||
|
||||
def _set_flow_shift(self, target_shift: float) -> None:
|
||||
"""Set UniPC ``flow_shift`` to ``target_shift``.
|
||||
|
||||
Lazily builds a default UniPC scheduler when called before
|
||||
``initialize_pipeline`` (e.g. the ``__new__``-based scheduler-parity
|
||||
tests); otherwise rebuilds from the snapshotted base config only when
|
||||
the target differs from the current shift.
|
||||
"""
|
||||
target = float(target_shift)
|
||||
base_config = self._base_scheduler_config
|
||||
if base_config is None:
|
||||
self.scheduler = UniPCMultistepScheduler(
|
||||
num_train_timesteps=1000,
|
||||
solver_order=2,
|
||||
prediction_type="flow_prediction",
|
||||
use_flow_sigmas=True,
|
||||
flow_shift=target,
|
||||
)
|
||||
self._base_scheduler_config = self.scheduler.config
|
||||
self._current_flow_shift = target
|
||||
return
|
||||
current = self._current_flow_shift
|
||||
if current is not None and target == float(current):
|
||||
return
|
||||
self.scheduler = UniPCMultistepScheduler.from_config(base_config, flow_shift=target)
|
||||
if isinstance(self.modules, dict):
|
||||
self.modules["scheduler"] = self.scheduler
|
||||
self._current_flow_shift = target
|
||||
|
||||
# -- Tokenization -------------------------------------------------------
|
||||
|
||||
def tokenize_caption(self, caption: str, *, is_video: bool = False, use_system_prompt: bool = False) -> list[int]:
|
||||
return cosmos3_tokenize_caption(self.get_module("text_tokenizer"),
|
||||
caption,
|
||||
is_video=is_video,
|
||||
use_system_prompt=use_system_prompt)
|
||||
|
||||
|
||||
# Entry point for the pipeline registry. The class name matches the checkpoint
|
||||
# ``model_index.json`` ``_class_name`` so ``resolve_pipeline_cls`` finds it.
|
||||
EntryClass = Cosmos3OmniDiffusersPipeline
|
||||
@@ -0,0 +1,85 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 (Cosmos3-Nano) inference presets.
|
||||
|
||||
Defaults track the official ``cosmos-framework`` ``sample_args`` for the video
|
||||
paths (``text2video`` / ``image2video``: guidance=6.0, num_steps=35, shift=10.0,
|
||||
fps=24, num_frames=189) and ``text2image`` (guidance=4.0, num_steps=50,
|
||||
shift=3.0). The default resolution is 16:9 at a VAE-aligned 704x1280 (spatial
|
||||
compression 16 -> 44x80 latent grid).
|
||||
"""
|
||||
from fastvideo.api.presets import InferencePreset, PresetStageSpec
|
||||
|
||||
_DENOISE_STAGE = PresetStageSpec(
|
||||
name="denoise",
|
||||
kind="denoising",
|
||||
description="Cosmos3 sequential-CFG UniPC denoising pass",
|
||||
allowed_overrides=frozenset({
|
||||
"num_inference_steps",
|
||||
"guidance_scale",
|
||||
}),
|
||||
)
|
||||
|
||||
# Framework video negative prompt (Cosmos quality prompt).
|
||||
COSMOS3_VIDEO_NEGATIVE_PROMPT = (
|
||||
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
|
||||
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
|
||||
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, "
|
||||
"jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, "
|
||||
"fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
|
||||
"Overall, the video is of poor quality.")
|
||||
|
||||
COSMOS3_NANO = InferencePreset(
|
||||
name="cosmos3_nano",
|
||||
version=1,
|
||||
model_family="cosmos3",
|
||||
description="Cosmos3-Nano text-to-video",
|
||||
workload_type="t2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 189,
|
||||
"fps": 24,
|
||||
"guidance_scale": 6.0,
|
||||
"num_inference_steps": 35,
|
||||
"negative_prompt": COSMOS3_VIDEO_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
COSMOS3_NANO_I2V = InferencePreset(
|
||||
name="cosmos3_nano_i2v",
|
||||
version=1,
|
||||
model_family="cosmos3",
|
||||
description="Cosmos3-Nano image-to-video",
|
||||
workload_type="i2v",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 189,
|
||||
"fps": 24,
|
||||
"guidance_scale": 6.0,
|
||||
"num_inference_steps": 35,
|
||||
"negative_prompt": COSMOS3_VIDEO_NEGATIVE_PROMPT,
|
||||
},
|
||||
)
|
||||
|
||||
COSMOS3_NANO_T2I = InferencePreset(
|
||||
name="cosmos3_nano_t2i",
|
||||
version=1,
|
||||
model_family="cosmos3",
|
||||
description="Cosmos3-Nano text-to-image",
|
||||
workload_type="t2i",
|
||||
stage_schemas=(_DENOISE_STAGE, ),
|
||||
defaults={
|
||||
"height": 1024,
|
||||
"width": 1024,
|
||||
"num_frames": 1,
|
||||
"fps": 24,
|
||||
"guidance_scale": 4.0,
|
||||
"num_inference_steps": 50,
|
||||
"negative_prompt": "",
|
||||
},
|
||||
)
|
||||
|
||||
ALL_PRESETS = (COSMOS3_NANO, COSMOS3_NANO_I2V, COSMOS3_NANO_T2I)
|
||||
@@ -0,0 +1,549 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""FastVideo-native Cosmos3 sequence packing (video subset).
|
||||
|
||||
Numerical-parity port of the official ``cosmos_framework`` data packer
|
||||
(``cosmos_framework.data.vfm.sequence_packing.pack_input_sequence``) restricted
|
||||
to the VIDEO generation path that the FastVideo Cosmos3 DiT consumes (T2V / I2V
|
||||
/ T2I). It builds, per sample, two splits:
|
||||
|
||||
* a ``causal`` text split (prompt token ids, plus the trailing ``eos`` and
|
||||
``start_of_generation`` markers the framework appends when a generation
|
||||
modality follows), and
|
||||
* a ``full`` vision split (VAE latent patch tokens).
|
||||
|
||||
The 3D-MRoPE position ids ``[3, seq]`` are produced exactly like the framework:
|
||||
text tokens broadcast a single monotone id across the (t, h, w) axes, the
|
||||
temporal offset is bumped by ``temporal_modality_margin`` at the text->vision
|
||||
boundary, and vision tokens lay out a (T, H, W) grid with spatial ids reset per
|
||||
segment. Condition frames (I2V cond frame 0, T2I single conditioned frame, ...)
|
||||
are kept in the packed sequence and rope grid but excluded from the MSE-loss /
|
||||
timestep bookkeeping, mirroring the framework.
|
||||
|
||||
The output ``Cosmos3PackedSequence`` maps 1:1 onto the
|
||||
``Cosmos3VFMTransformer.forward`` kwargs via :meth:`to_dit_kwargs`. This module
|
||||
is pure torch/python; it imports no diffusers/transformers model classes.
|
||||
|
||||
Reference of record: ``cosmos_framework`` (NVIDIA), the parity oracle used by
|
||||
``tests/local_tests/cosmos3/test_cosmos3_packing_parity.py``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.dits.cosmos3 import (
|
||||
compute_mrope_position_ids_text,
|
||||
compute_mrope_position_ids_vision,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Cosmos3VisionItem",
|
||||
"Cosmos3SampleInputs",
|
||||
"Cosmos3PackedSequence",
|
||||
"pack_cosmos3_video_sequence",
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inputs
|
||||
# ---------------------------------------------------------------------------
|
||||
@dataclass
|
||||
class Cosmos3VisionItem:
|
||||
"""One vision latent for a sample.
|
||||
|
||||
Args:
|
||||
latent: VAE latent ``[C, T, H, W]`` (a leading batch axis of size 1 is
|
||||
accepted and squeezed).
|
||||
condition_frame_indexes: Latent-frame indices that are *conditioned*
|
||||
(clean) rather than noisy. ``[]`` for T2V, ``[0]`` for I2V, and the
|
||||
single conditioned frame for T2I.
|
||||
fps: Frames-per-second for this clip; only used when
|
||||
``enable_fps_modulation`` is set.
|
||||
"""
|
||||
|
||||
latent: torch.Tensor
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SoundItem:
|
||||
"""One sound latent for a sample (t2vs).
|
||||
|
||||
Args:
|
||||
latent: AVAE sound latent ``[C, T]`` (channels, temporal frames).
|
||||
condition_frame_indexes: Latent-frame indices that are *conditioned*
|
||||
(clean). ``[]`` for t2vs (all frames generated).
|
||||
fps: Sound latent FPS (``sound_latent_fps``, e.g. 25); only used when
|
||||
``enable_fps_modulation`` is set.
|
||||
"""
|
||||
|
||||
latent: torch.Tensor
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3ActionItem:
|
||||
"""One action latent for a sample (action-conditioned world model).
|
||||
|
||||
Args:
|
||||
latent: Action latent ``[T, action_dim]`` (per-frame action vectors).
|
||||
condition_frame_indexes: Frame indices kept clean (conditioning actions).
|
||||
domain_id: Embodiment domain id (scalar / ``[1]``) for the
|
||||
domain-aware action projection.
|
||||
fps: Action FPS; only used when ``enable_fps_modulation`` is set.
|
||||
"""
|
||||
|
||||
latent: torch.Tensor
|
||||
condition_frame_indexes: list[int] = field(default_factory=list)
|
||||
domain_id: int = 0
|
||||
fps: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cosmos3SampleInputs:
|
||||
"""Per-sample packing inputs (text prompt + vision item, +sound, +action)."""
|
||||
|
||||
text_ids: list[int]
|
||||
vision: Cosmos3VisionItem
|
||||
timestep: float
|
||||
sound: Cosmos3SoundItem | None = None
|
||||
action: Cosmos3ActionItem | None = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Output
|
||||
# ---------------------------------------------------------------------------
|
||||
@dataclass
|
||||
class Cosmos3PackedSequence:
|
||||
"""Packed-sequence inputs consumed by ``Cosmos3VFMTransformer.forward``.
|
||||
|
||||
Field names mirror the framework ``PackedSequence`` (+ its ``vision``
|
||||
``ModalityData``) so the parity test can compare field-by-field.
|
||||
"""
|
||||
|
||||
# Sequence structure.
|
||||
sample_lens: list[int]
|
||||
split_lens: list[int]
|
||||
attn_modes: list[str]
|
||||
sequence_length: int
|
||||
is_image_batch: bool
|
||||
|
||||
# Text modality.
|
||||
text_ids: torch.Tensor
|
||||
text_indexes: torch.Tensor
|
||||
position_ids: torch.Tensor # [3, sequence_length]
|
||||
|
||||
# Vision modality.
|
||||
vision_tokens: list[torch.Tensor]
|
||||
vision_token_shapes: list[tuple[int, int, int]]
|
||||
vision_sequence_indexes: torch.Tensor
|
||||
vision_timesteps: torch.Tensor
|
||||
vision_mse_loss_indexes: torch.Tensor
|
||||
vision_noisy_frame_indexes: list[torch.Tensor]
|
||||
vision_condition_mask: list[torch.Tensor]
|
||||
fps_vision: torch.Tensor | None = None
|
||||
|
||||
# Sound modality (t2vs); empty/None when no sound.
|
||||
sound_tokens: list[torch.Tensor] = field(default_factory=list)
|
||||
sound_token_shapes: list[tuple[int, int, int]] = field(default_factory=list)
|
||||
sound_sequence_indexes: torch.Tensor | None = None
|
||||
sound_timesteps: torch.Tensor | None = None
|
||||
sound_mse_loss_indexes: torch.Tensor | None = None
|
||||
sound_noisy_frame_indexes: list[torch.Tensor] = field(default_factory=list)
|
||||
sound_condition_mask: list[torch.Tensor] = field(default_factory=list)
|
||||
fps_sound: torch.Tensor | None = None
|
||||
|
||||
# Action modality (action-conditioned world model); empty/None when no action.
|
||||
action_tokens: list[torch.Tensor] = field(default_factory=list)
|
||||
action_token_shapes: list[tuple[int, ...]] = field(default_factory=list)
|
||||
action_sequence_indexes: torch.Tensor | None = None
|
||||
action_timesteps: torch.Tensor | None = None
|
||||
action_mse_loss_indexes: torch.Tensor | None = None
|
||||
action_noisy_frame_indexes: list[torch.Tensor] = field(default_factory=list)
|
||||
action_condition_mask: list[torch.Tensor] = field(default_factory=list)
|
||||
action_domain_id: list[torch.Tensor] = field(default_factory=list)
|
||||
|
||||
def to_dit_kwargs(self, device: torch.device | str | None = None) -> dict[str, Any]:
|
||||
"""Return the kwargs dict for ``Cosmos3VFMTransformer.forward``.
|
||||
|
||||
Packing is device-agnostic (ids/indexes/position-ids are built on CPU).
|
||||
When ``device`` is given, every tensor input is moved to it so the DiT
|
||||
forward runs on a single device (e.g. the model's GPU at inference).
|
||||
"""
|
||||
|
||||
def _mv(x: Any) -> Any:
|
||||
return x.to(device) if (device is not None and torch.is_tensor(x)) else x
|
||||
|
||||
return dict(
|
||||
text_ids=_mv(self.text_ids),
|
||||
text_indexes=_mv(self.text_indexes),
|
||||
position_ids=_mv(self.position_ids),
|
||||
sequence_length=int(self.sequence_length),
|
||||
split_lens=list(self.split_lens),
|
||||
attn_modes=list(self.attn_modes),
|
||||
vision_tokens=[_mv(t) for t in self.vision_tokens],
|
||||
vision_token_shapes=list(self.vision_token_shapes),
|
||||
vision_sequence_indexes=_mv(self.vision_sequence_indexes),
|
||||
vision_timesteps=_mv(self.vision_timesteps),
|
||||
vision_mse_loss_indexes=_mv(self.vision_mse_loss_indexes),
|
||||
vision_noisy_frame_indexes=[_mv(t) for t in self.vision_noisy_frame_indexes],
|
||||
fps_vision=self.fps_vision,
|
||||
sound_tokens=[_mv(t) for t in self.sound_tokens],
|
||||
sound_token_shapes=list(self.sound_token_shapes),
|
||||
sound_sequence_indexes=_mv(self.sound_sequence_indexes),
|
||||
sound_timesteps=_mv(self.sound_timesteps),
|
||||
sound_mse_loss_indexes=_mv(self.sound_mse_loss_indexes),
|
||||
sound_noisy_frame_indexes=[_mv(t) for t in self.sound_noisy_frame_indexes],
|
||||
fps_sound=_mv(self.fps_sound),
|
||||
action_tokens=[_mv(t) for t in self.action_tokens],
|
||||
action_token_shapes=list(self.action_token_shapes),
|
||||
action_sequence_indexes=_mv(self.action_sequence_indexes),
|
||||
action_timesteps=_mv(self.action_timesteps),
|
||||
action_mse_loss_indexes=_mv(self.action_mse_loss_indexes),
|
||||
action_noisy_frame_indexes=[_mv(t) for t in self.action_noisy_frame_indexes],
|
||||
action_domain_id=[_mv(t) for t in self.action_domain_id],
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Packing
|
||||
# ---------------------------------------------------------------------------
|
||||
def pack_cosmos3_video_sequence(
|
||||
samples: list[Cosmos3SampleInputs],
|
||||
special_tokens: dict[str, int],
|
||||
*,
|
||||
latent_patch_size: int = 2,
|
||||
include_end_of_generation_token: bool = False,
|
||||
temporal_modality_margin: int = 15_000,
|
||||
reset_spatial_ids: bool = True,
|
||||
enable_fps_modulation: bool = False,
|
||||
base_fps: float = 24.0,
|
||||
temporal_compression_factor: int = 4,
|
||||
initial_mrope_temporal_offset: int | float = 0,
|
||||
) -> Cosmos3PackedSequence:
|
||||
"""Pack prompts + vision latents into the Cosmos3 DiT packed-sequence inputs.
|
||||
|
||||
Video subset of ``cosmos_framework`` ``pack_input_sequence`` under
|
||||
``unified_3d_mrope``: each sample is ``[causal text, full vision]``.
|
||||
|
||||
Args:
|
||||
samples: Per-sample text prompt token ids + vision item + timestep.
|
||||
special_tokens: Must contain ``eos_token_id`` and
|
||||
``start_of_generation`` (and ``end_of_generation`` if
|
||||
``include_end_of_generation_token``). ``bos_token_id`` is honored if
|
||||
present (prepended) to match the framework.
|
||||
latent_patch_size: Latent patch size used by the DiT.
|
||||
include_end_of_generation_token: Append the framework's end-of-generation
|
||||
marker after the vision split.
|
||||
temporal_modality_margin: Temporal-offset bump applied at the
|
||||
text->vision boundary (``unified_3d_mrope_temporal_modality_margin``).
|
||||
reset_spatial_ids: Reset vision spatial ids to 0 per segment.
|
||||
enable_fps_modulation: Use float, fps-scaled temporal positions.
|
||||
base_fps: Base FPS used when ``enable_fps_modulation``.
|
||||
temporal_compression_factor: VAE temporal compression factor.
|
||||
initial_mrope_temporal_offset: Per-sample starting temporal offset.
|
||||
|
||||
Returns:
|
||||
A :class:`Cosmos3PackedSequence`.
|
||||
"""
|
||||
assert "eos_token_id" in special_tokens, "special_tokens must contain eos_token_id"
|
||||
assert "start_of_generation" in special_tokens, "special_tokens must contain start_of_generation"
|
||||
if latent_patch_size < 1:
|
||||
raise ValueError(f"latent_patch_size must be >= 1, got {latent_patch_size}")
|
||||
|
||||
# Build-time accumulators (concatenated across samples).
|
||||
sample_lens: list[int] = []
|
||||
split_lens: list[int] = []
|
||||
attn_modes: list[str] = []
|
||||
|
||||
text_ids: list[int] = []
|
||||
text_indexes: list[int] = []
|
||||
position_id_blocks: list[torch.Tensor] = [] # each [3, n]
|
||||
|
||||
vision_tokens: list[torch.Tensor] = []
|
||||
vision_token_shapes: list[tuple[int, int, int]] = []
|
||||
vision_sequence_indexes: list[int] = []
|
||||
vision_timesteps: list[float] = []
|
||||
vision_mse_loss_indexes: list[int] = []
|
||||
vision_noisy_frame_indexes: list[torch.Tensor] = []
|
||||
vision_condition_mask: list[torch.Tensor] = []
|
||||
fps_values: list[float] = []
|
||||
|
||||
sound_tokens: list[torch.Tensor] = []
|
||||
sound_token_shapes: list[tuple[int, int, int]] = []
|
||||
sound_sequence_indexes: list[int] = []
|
||||
sound_timesteps: list[float] = []
|
||||
sound_mse_loss_indexes: list[int] = []
|
||||
sound_noisy_frame_indexes: list[torch.Tensor] = []
|
||||
sound_condition_mask: list[torch.Tensor] = []
|
||||
sound_fps_values: list[float] = []
|
||||
|
||||
action_tokens: list[torch.Tensor] = []
|
||||
action_token_shapes: list[tuple[int, ...]] = []
|
||||
action_sequence_indexes: list[int] = []
|
||||
action_timesteps: list[float] = []
|
||||
action_mse_loss_indexes: list[int] = []
|
||||
action_noisy_frame_indexes: list[torch.Tensor] = []
|
||||
action_condition_mask: list[torch.Tensor] = []
|
||||
action_domain_id: list[torch.Tensor] = []
|
||||
|
||||
curr = 0 # running position in the packed sequence
|
||||
is_image_batch = True
|
||||
|
||||
for sample in samples:
|
||||
temporal_offset: int | float = initial_mrope_temporal_offset
|
||||
sample_len = 0
|
||||
|
||||
# ---- 1. Text split (causal) ----
|
||||
if "bos_token_id" in special_tokens:
|
||||
shifted_text_ids = [special_tokens["bos_token_id"], *sample.text_ids]
|
||||
else:
|
||||
shifted_text_ids = list(sample.text_ids)
|
||||
# The video path always has a following generation modality, so the
|
||||
# framework appends eos + start_of_generation.
|
||||
shifted_text_ids = [*shifted_text_ids, special_tokens["eos_token_id"], special_tokens["start_of_generation"]]
|
||||
text_split_len = len(shifted_text_ids)
|
||||
|
||||
text_ids.extend(shifted_text_ids)
|
||||
text_indexes.extend(range(curr, curr + text_split_len))
|
||||
|
||||
text_mrope, temporal_offset = compute_mrope_position_ids_text(
|
||||
num_tokens=text_split_len,
|
||||
temporal_offset=int(temporal_offset),
|
||||
)
|
||||
position_id_blocks.append(text_mrope)
|
||||
|
||||
attn_modes.append("causal")
|
||||
split_lens.append(text_split_len)
|
||||
curr += text_split_len
|
||||
sample_len += text_split_len
|
||||
|
||||
# End of text modality: bump temporal offset before vision.
|
||||
temporal_offset += temporal_modality_margin
|
||||
# Sound shares the vision temporal start (parallel temporal positions).
|
||||
vision_start_temporal_offset = temporal_offset
|
||||
|
||||
# ---- 2. Vision split (full) ----
|
||||
latent = sample.vision.latent
|
||||
latent = latent.squeeze(0) if latent.dim() == 5 else latent # [C, T, H, W]
|
||||
_c, latent_t, latent_h, latent_w = latent.shape
|
||||
patch_h = math.ceil(latent_h / latent_patch_size)
|
||||
patch_w = math.ceil(latent_w / latent_patch_size)
|
||||
num_vision_tokens = latent_t * patch_h * patch_w
|
||||
|
||||
vision_tokens.append(sample.vision.latent)
|
||||
vision_token_shapes.append((latent_t, patch_h, patch_w))
|
||||
vision_sequence_indexes.extend(range(curr, curr + num_vision_tokens))
|
||||
|
||||
condition_set = {idx for idx in sample.vision.condition_frame_indexes if 0 <= idx < latent_t}
|
||||
cond_mask = torch.zeros((latent_t, 1, 1), device=latent.device, dtype=latent.dtype)
|
||||
for frame_idx in condition_set:
|
||||
cond_mask[frame_idx, 0, 0] = 1.0
|
||||
vision_condition_mask.append(cond_mask)
|
||||
|
||||
noisy_frames = torch.tensor(
|
||||
[idx for idx in range(latent_t) if idx not in condition_set],
|
||||
device=latent.device,
|
||||
dtype=torch.long,
|
||||
)
|
||||
vision_noisy_frame_indexes.append(noisy_frames)
|
||||
|
||||
# MSE-loss indices + per-token timesteps cover only the noisy frames.
|
||||
frame_token_stride = patch_h * patch_w
|
||||
for frame_idx in range(latent_t):
|
||||
if frame_idx in condition_set:
|
||||
continue
|
||||
frame_start = curr + frame_idx * frame_token_stride
|
||||
vision_mse_loss_indexes.extend(range(frame_start, frame_start + frame_token_stride))
|
||||
vision_timesteps.extend([float(sample.timestep)] * frame_token_stride)
|
||||
|
||||
vision_fps = sample.vision.fps if enable_fps_modulation else None
|
||||
if vision_fps is not None:
|
||||
fps_values.append(float(vision_fps))
|
||||
vision_mrope, temporal_offset = compute_mrope_position_ids_vision(
|
||||
grid_t=latent_t,
|
||||
grid_h=patch_h,
|
||||
grid_w=patch_w,
|
||||
temporal_offset=temporal_offset,
|
||||
fps=vision_fps,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=temporal_compression_factor,
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
)
|
||||
position_id_blocks.append(vision_mrope)
|
||||
|
||||
curr += num_vision_tokens
|
||||
sample_len += num_vision_tokens
|
||||
|
||||
# ---- 2a2. Action split: shares the vision "full" split ----
|
||||
# Mirrors framework ``_pack_action_tokens``: action latent [T, D] -> T
|
||||
# tokens (token shape (T,)), domain-aware, 3D-MRoPE at the vision temporal
|
||||
# offset with ``start_frame_offset=1`` (parallel to vision; tcf=1; does
|
||||
# not advance the offset).
|
||||
action_split_len = 0
|
||||
if sample.action is not None:
|
||||
action_latent = sample.action.latent # [T, D]
|
||||
action_t = int(action_latent.shape[0])
|
||||
action_split_len = action_t
|
||||
|
||||
action_tokens.append(action_latent)
|
||||
action_token_shapes.append((action_t, ))
|
||||
action_sequence_indexes.extend(range(curr, curr + action_t))
|
||||
action_domain_id.append(torch.tensor([int(sample.action.domain_id)], dtype=torch.long))
|
||||
|
||||
a_cond_set = {idx for idx in sample.action.condition_frame_indexes if 0 <= idx < action_t}
|
||||
a_cond_mask = torch.zeros((action_t, 1), device=action_latent.device, dtype=action_latent.dtype)
|
||||
for fi in a_cond_set:
|
||||
a_cond_mask[fi, 0] = 1.0
|
||||
action_condition_mask.append(a_cond_mask)
|
||||
|
||||
a_noisy = torch.tensor([idx for idx in range(action_t) if idx not in a_cond_set],
|
||||
device=action_latent.device,
|
||||
dtype=torch.long)
|
||||
action_noisy_frame_indexes.append(a_noisy)
|
||||
|
||||
for fi in range(action_t):
|
||||
if fi in a_cond_set:
|
||||
continue
|
||||
action_mse_loss_indexes.append(curr + fi)
|
||||
action_timesteps.append(float(sample.timestep))
|
||||
|
||||
action_fps = sample.action.fps if enable_fps_modulation else None
|
||||
action_mrope, _ = compute_mrope_position_ids_vision(
|
||||
grid_t=action_t,
|
||||
grid_h=1,
|
||||
grid_w=1,
|
||||
temporal_offset=vision_start_temporal_offset,
|
||||
fps=action_fps,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=1, # action is at frame rate
|
||||
base_temporal_compression_factor=temporal_compression_factor,
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
start_frame_offset=1,
|
||||
)
|
||||
position_id_blocks.append(action_mrope)
|
||||
curr += action_t
|
||||
sample_len += action_t
|
||||
|
||||
# ---- 2b. Sound split (t2vs): shares the vision "full" split ----
|
||||
# Mirrors framework ``_pack_sound_tokens``: sound latent [C, T] -> T
|
||||
# tokens (token shape (T,1,1)), packed right after vision, with 3D-MRoPE
|
||||
# temporal positions starting at the vision temporal offset (parallel to
|
||||
# vision, start_frame_offset=0, tcf=1) and NOT advancing it.
|
||||
sound_split_len = 0
|
||||
if sample.sound is not None:
|
||||
sound_latent = sample.sound.latent
|
||||
sound_latent = sound_latent.squeeze(0) if sound_latent.dim() == 3 else sound_latent # [C, T]
|
||||
_sc, sound_t = sound_latent.shape
|
||||
sound_split_len = sound_t
|
||||
|
||||
sound_tokens.append(sound_latent)
|
||||
sound_token_shapes.append((sound_t, 1, 1))
|
||||
sound_sequence_indexes.extend(range(curr, curr + sound_t))
|
||||
|
||||
s_cond_set = {idx for idx in sample.sound.condition_frame_indexes if 0 <= idx < sound_t}
|
||||
s_cond_mask = torch.zeros((sound_t, 1), device=sound_latent.device, dtype=sound_latent.dtype)
|
||||
for fi in s_cond_set:
|
||||
s_cond_mask[fi, 0] = 1.0
|
||||
sound_condition_mask.append(s_cond_mask)
|
||||
|
||||
s_noisy = torch.tensor([idx for idx in range(sound_t) if idx not in s_cond_set],
|
||||
device=sound_latent.device,
|
||||
dtype=torch.long)
|
||||
sound_noisy_frame_indexes.append(s_noisy)
|
||||
|
||||
for fi in range(sound_t):
|
||||
if fi in s_cond_set:
|
||||
continue
|
||||
sound_mse_loss_indexes.append(curr + fi) # 1 token per sound frame
|
||||
sound_timesteps.append(float(sample.timestep))
|
||||
|
||||
sound_fps = sample.sound.fps if enable_fps_modulation else None
|
||||
if sound_fps is not None:
|
||||
sound_fps_values.append(float(sound_fps))
|
||||
sound_mrope, _ = compute_mrope_position_ids_vision(
|
||||
grid_t=sound_t,
|
||||
grid_h=1,
|
||||
grid_w=1,
|
||||
temporal_offset=vision_start_temporal_offset,
|
||||
fps=sound_fps,
|
||||
base_fps=base_fps,
|
||||
temporal_compression_factor=1, # sound latent already at sound_latent_fps
|
||||
enable_fps_modulation=enable_fps_modulation,
|
||||
start_frame_offset=0,
|
||||
)
|
||||
position_id_blocks.append(sound_mrope)
|
||||
curr += sound_t
|
||||
sample_len += sound_t
|
||||
|
||||
# ---- 3. Optional end-of-generation marker ----
|
||||
eov_len = 0
|
||||
if include_end_of_generation_token:
|
||||
assert "end_of_generation" in special_tokens, ("special_tokens must contain end_of_generation when "
|
||||
"include_end_of_generation_token=True")
|
||||
text_ids.append(special_tokens["end_of_generation"])
|
||||
text_indexes.append(curr)
|
||||
eov_dtype = torch.float32 if enable_fps_modulation else torch.long
|
||||
eov_ids = torch.full((3, 1), temporal_offset, dtype=eov_dtype)
|
||||
position_id_blocks.append(eov_ids)
|
||||
temporal_offset += 1
|
||||
curr += 1
|
||||
eov_len = 1
|
||||
sample_len += 1
|
||||
|
||||
# Vision + action + sound + any trailing eov marker share one "full" split.
|
||||
attn_modes.append("full")
|
||||
split_lens.append(num_vision_tokens + action_split_len + sound_split_len + eov_len)
|
||||
sample_lens.append(sample_len)
|
||||
|
||||
if latent_t != 1:
|
||||
is_image_batch = False
|
||||
|
||||
sequence_length = sum(sample_lens)
|
||||
|
||||
# position_ids: float iff any block is float (fps modulation path).
|
||||
any_float = any(b.dtype.is_floating_point for b in position_id_blocks)
|
||||
if any_float:
|
||||
position_id_blocks = [b.to(torch.float32) for b in position_id_blocks]
|
||||
position_ids = torch.cat(position_id_blocks, dim=1) # [3, sequence_length]
|
||||
|
||||
timesteps_dtype = torch.float32
|
||||
return Cosmos3PackedSequence(
|
||||
sample_lens=sample_lens,
|
||||
split_lens=split_lens,
|
||||
attn_modes=attn_modes,
|
||||
sequence_length=sequence_length,
|
||||
is_image_batch=is_image_batch,
|
||||
text_ids=torch.tensor(text_ids, dtype=torch.long),
|
||||
text_indexes=torch.tensor(text_indexes, dtype=torch.long),
|
||||
position_ids=position_ids,
|
||||
vision_tokens=vision_tokens,
|
||||
vision_token_shapes=vision_token_shapes,
|
||||
vision_sequence_indexes=torch.tensor(vision_sequence_indexes, dtype=torch.long),
|
||||
vision_timesteps=torch.tensor(vision_timesteps, dtype=timesteps_dtype),
|
||||
vision_mse_loss_indexes=torch.tensor(vision_mse_loss_indexes, dtype=torch.long),
|
||||
vision_noisy_frame_indexes=vision_noisy_frame_indexes,
|
||||
vision_condition_mask=vision_condition_mask,
|
||||
fps_vision=(torch.tensor(fps_values, dtype=torch.float32) if fps_values else None),
|
||||
sound_tokens=sound_tokens,
|
||||
sound_token_shapes=sound_token_shapes,
|
||||
sound_sequence_indexes=(torch.tensor(sound_sequence_indexes, dtype=torch.long) if sound_tokens else None),
|
||||
sound_timesteps=(torch.tensor(sound_timesteps, dtype=timesteps_dtype) if sound_tokens else None),
|
||||
sound_mse_loss_indexes=(torch.tensor(sound_mse_loss_indexes, dtype=torch.long) if sound_tokens else None),
|
||||
sound_noisy_frame_indexes=sound_noisy_frame_indexes,
|
||||
sound_condition_mask=sound_condition_mask,
|
||||
fps_sound=(torch.tensor(sound_fps_values, dtype=torch.float32) if sound_fps_values else None),
|
||||
action_tokens=action_tokens,
|
||||
action_token_shapes=action_token_shapes,
|
||||
action_sequence_indexes=(torch.tensor(action_sequence_indexes, dtype=torch.long) if action_tokens else None),
|
||||
action_timesteps=(torch.tensor(action_timesteps, dtype=timesteps_dtype) if action_tokens else None),
|
||||
action_mse_loss_indexes=(torch.tensor(action_mse_loss_indexes, dtype=torch.long) if action_tokens else None),
|
||||
action_noisy_frame_indexes=action_noisy_frame_indexes,
|
||||
action_condition_mask=action_condition_mask,
|
||||
action_domain_id=action_domain_id,
|
||||
)
|
||||
@@ -0,0 +1,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,6 +20,7 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
|
||||
Cosmos25Config,
|
||||
Cosmos25_14BConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.cosmos3 import Cosmos3Config
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
@@ -772,6 +773,22 @@ def _register_configs() -> None:
|
||||
default_preset="gen3c_cosmos_7b",
|
||||
)
|
||||
|
||||
# Cosmos 3 (must register before Cosmos 2.5 and generic Cosmos detectors
|
||||
# so the cosmos3 path-detection takes precedence)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=Cosmos3Config,
|
||||
workload_types=(WorkloadType.T2V, ),
|
||||
hf_model_paths=[
|
||||
"nvidia/Cosmos3-Nano",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "cosmos3" in path.lower() or "cosmos-3" in path.lower(),
|
||||
],
|
||||
model_family="cosmos3",
|
||||
default_preset="cosmos3_nano",
|
||||
)
|
||||
|
||||
# Cosmos 2.5 (2B)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
@@ -1194,6 +1211,8 @@ def _register_presets() -> None:
|
||||
from fastvideo.api.presets import register_preset
|
||||
from fastvideo.pipelines.basic.cosmos.presets import (
|
||||
ALL_PRESETS as COSMOS_PRESETS, )
|
||||
from fastvideo.pipelines.basic.cosmos3.presets import (
|
||||
ALL_PRESETS as COSMOS3_PRESETS, )
|
||||
from fastvideo.pipelines.basic.dreamx_world.presets import (
|
||||
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.gamecraft.presets import (
|
||||
@@ -1231,6 +1250,7 @@ def _register_presets() -> None:
|
||||
|
||||
all_preset_groups = (
|
||||
COSMOS_PRESETS,
|
||||
COSMOS3_PRESETS,
|
||||
DREAMX_WORLD_PRESETS,
|
||||
FLUX2_PRESETS,
|
||||
GAMECRAFT_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.
|
||||
@@ -0,0 +1,152 @@
|
||||
# Cosmos3 Port Status
|
||||
|
||||
## Summary
|
||||
|
||||
- model_family: `cosmos3`
|
||||
- workload_types: `T2V, I2V, T2I` supported by `WorkloadType` today; full-omni target also needs audio (AV), VLM reasoning, and action-conditioning, which require framework extensions (Q002, Q003).
|
||||
- official_ref: `https://github.com/NVIDIA/cosmos-framework` — diffusers backend `diffusers_cosmos3.pipeline.Cosmos3OmniDiffusersPipeline`; HF `nvidia/Cosmos3-Nano`.
|
||||
- official_ref_dir: `cosmos-framework` (symlink -> `/home/william5lin/FastVideo/cosmos-framework`, commit `003d66d4`)
|
||||
- hf_weights_path: `nvidia/Cosmos3-Nano`
|
||||
- local_weights_dir: `official_weights/cosmos3` (symlink -> `/home/william5lin/FastVideo/official_weights/cosmos3`, 33 GiB / 67 files)
|
||||
- source_layout: `diffusers`
|
||||
- local_tests_readme: `tests/local_tests/cosmos3/README.md`
|
||||
|
||||
## Current Phase
|
||||
|
||||
- phase: `FULL OMNI SUPPORTED — every modality framework-parity verified bit-exact (suite 150 passed, 0 skipped). PR1 video core (T2V/I2V/T2I) + PR2 audio (t2vs) real-weights verified on B200; PR3 action (domain-aware) + PR4 reasoning (text + vision_encoder + deepstack reasoner) bit-exact. Branch chain: feat/cosmos3-tier-a-port (T2V) -> feat/cosmos3-i2v (I2V+T2I+flow_shift) -> feat/cosmos3-audio (t2vs) -> feat/cosmos3-action -> feat/cosmos3-reasoning. Optional follow-ups: real-weights action2world (needs robot-action data) + image-conditioned-reasoning prefill wiring (vision_encoder + get_rope_index, both proven).`
|
||||
- status: `in_progress`
|
||||
- owner: `orchestrator`
|
||||
- last_updated: `2026-06-07`
|
||||
- env: `fv-cosmos3` (conda clone of fv-main; `fastvideo` editable repointed to this worktree). Run tests from the worktree cwd with this env's python.
|
||||
- branch: rebased onto `origin/main` @ `1c627a3f9` (was 33 behind, merge-base 2026-05-22); now 6 commits ahead; `fastvideo` imports clean; Tier-A `13 passed, 2 skipped`.
|
||||
|
||||
## Component Matrix
|
||||
|
||||
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|
||||
|---|---|---|---|---|---|---|---|---|---|
|
||||
| transformer | dit | port | `diffusers_cosmos3/transformer.py:Cosmos3OmniTransformer` (model_type `qwen3_vl_text`, MoT + MRoPE) | `model_index.json: transformer`; `cosmos_framework/model/vfm/mot/cosmos3_vfm_network.py`, `omni_mot_model.py` | `fastvideo/models/dits/cosmos3.py` (branch: `Cosmos3VFMTransformer`+`Cosmos3LanguageModel` — reconcile to `Cosmos3OmniTransformer`) | skeleton | not_started | scaffold_skip | I001 |
|
||||
| vae | vae | reuse | diffusers `AutoencoderKLWan` | `model_index.json: vae` | reuse Wan VAE (`fastvideo/models/vaes/`, cf. `cosmos25wanvae.py`) | not_started | passthrough? | not_started | Q001 |
|
||||
| scheduler | generic | reuse (flow-coerced) | framework `FlowUniPCMultistepScheduler` (`cosmos_framework/.../fm_solvers_unipc.py`; checkpoint ships diffusers-style config) | `model_index.json: scheduler`; `cosmos_framework/.../samplers/unipc.py:UniPCSampler` | FastVideo-native `UniPCMultistepScheduler` (flow config), coerced in `initialize_pipeline` | done | n/a | framework-parity DONE (`test_cosmos3_scheduler_parity`: timesteps bit-exact, sigmas ~1e-8, trajectory <~1e-6) | I003 (resolved) |
|
||||
| text_tokenizer | tokenizer | reuse | transformers `Qwen2TokenizerFast` | `model_index.json: text_tokenizer` | reuse (tokenizer = allowed third-party) | not_started | passthrough | scaffold_skip (`test_cosmos3_tokenizer_chat_template`) | - |
|
||||
| vision_encoder | encoder | port | transformers `Qwen3VLVisionModel` | `model_index.json: vision_encoder` | new encoder bucket OR documented lazy-wrapper | not_started | not_started | not_started | Q002 |
|
||||
| sound_tokenizer | generic/vae | port (decode) | framework AVAE `LatentAutoEncoderV2` (`avae_utils`); checkpoint is decoder-only AutoencoderOobleck-named w/ SnakeBeta | `model_index.json: sound_tokenizer` | reuse FastVideo native `OobleckVAE` decoder + `Cosmos3SoundVAE` wrapper (`models/audio/cosmos3_avae.py`) | done (decode) | n/a | DECODE bit-exact vs framework (`test_cosmos3_avae_parity`); real ckpt strict-loads | PR2 (branch feat/cosmos3-audio) |
|
||||
|
||||
## Conversion State
|
||||
|
||||
- conversion_script: `scripts/checkpoint_conversion/cosmos3_convert.py` (branch has it, 246 lines, built vs vllm-omni — repoint/verify vs diffusers checkpoint)
|
||||
- converted_weights_dir: `converted_weights/cosmos3` (n/a while needs_conversion=no)
|
||||
- source_layout: `diffusers`
|
||||
- needs_conversion: `no` (HF already diffusers-format; verify FastVideo loaders consume directly)
|
||||
- strict_load_status: `not_run`
|
||||
- passthrough_components: `vae (AutoencoderKLWan), scheduler (UniPC), text_tokenizer (Qwen2)` likely passthrough
|
||||
- retry_history: `none`
|
||||
|
||||
## Parity Commands
|
||||
|
||||
| Scope | Command | Last Result | Notes |
|
||||
|---|---|---|---|
|
||||
| Tier-A scaffold | `cd <worktree> && <fv-cosmos3 python> -m pytest tests/local_tests/cosmos3/ -q` | `13 passed, 2 skipped` (2026-06-06, post-rebase) | 2 skips: Cosmos3 tokenizer/_tokenize_prompt not yet wired on pipeline |
|
||||
| component | `pytest tests/local_tests/<bucket>/test_cosmos3_<component>_parity.py -v -s` | `not_run` | after env activation + native prototypes |
|
||||
| pipeline | `pytest tests/local_tests/pipelines/test_cosmos3_pipeline_parity.py -v -s` | `not_run` | |
|
||||
|
||||
## Open Questions
|
||||
|
||||
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|
||||
|---|---|---|---|---|---|
|
||||
| Q001 | Does Cosmos3 VAE (`AutoencoderKLWan`) match FastVideo's existing Wan VAE config/instantiation exactly (z_dim, scale factors, latents_mean/std)? | orchestrator | 3 (reuse gate) | open | |
|
||||
| Q002 | `vision_encoder` (`Qwen3VLVisionModel`): native port vs documented lazy-wrapper exception? Needed for I2V/reasoning. | user/orchestrator | 3 | open | |
|
||||
| Q003 | `sound_tokenizer` (`Cosmos3AVAEAudioTokenizer`) + audio output requires `WorkloadType` AV + audio regression metric. | user | 0/10 | open | full-omni scope chosen 2026-06-06; infra extensions pending |
|
||||
|
||||
## Issues And Blockers
|
||||
|
||||
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|---|---|
|
||||
| I001 | port | transformer | high | Branch DiT (`Cosmos3VFMTransformer`+`Cosmos3LanguageModel`) built vs vllm-omni #3454; official checkpoint loads `Cosmos3OmniTransformer` (diffusers shim). Class/structure reconciliation required. | `model_index.json`; `diffusers_cosmos3/transformer.py`; branch commit `52bb65f49` | orchestrator | resolved | DiT rewritten to checkpoint layout (single `layers` dual-pathway, BaseDiT-conformant); bit-identical framework parity (3d_rope + unified_3d_mrope), commits 59a4a571c/7c4633295 |
|
||||
| I002 | all | tests | medium | Tier-A conftest+tests mirror vllm-omni line-by-line (stubs, `vllm_omni...guardrails`). Must be repointed to `diffusers_cosmos3` / official structures. | `tests/local_tests/cosmos3/conftest.py` | orchestrator | open | |
|
||||
| I003 | inference | scheduler | high | First real-weights T2V was all-black: checkpoint `scheduler_config.json` sets `use_karras_sigmas=true`; vendored UniPC checks karras before `use_flow_sigmas` -> diffusion (beta) sigmas -> `scheduler.step` -> NaN latents. DiT/CFG velocity was clean. The scheduler had never been parity-tested vs the framework (`test_cosmos3_denoise_cfg_parity` used diffusers UniPC on both sides). | `result_latent` NaN at denoise step 0 (v_pred clean); ffprobe 3 KB black mp4 | orchestrator | resolved | Coerce loaded config to flow setup in `initialize_pipeline`; switch pipeline+tests to native UniPC (no diffusers at runtime); add `test_cosmos3_scheduler_parity` vs framework `FlowUniPCMultistepScheduler`; repoint denoise_cfg oracle to the framework scheduler. Commit 255311cf2 |
|
||||
|
||||
## Escape Hatches
|
||||
|
||||
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|
|
||||
| E001 | prep | dependency/env | Shared `fv-main` env has `fastvideo` editable-installed from the MAIN worktree; the cosmos3 worktree's `fastvideo` is not importable (PEP660 finder overrides PYTHONPATH), so Tier-A tests skip. How to activate the worktree's `fastvideo` for verification without disrupting ~24 other worktrees sharing the env? | Dedicated conda env for the cosmos3 worktree | resolved | Created fv-cosmos3 (clone of fv-main); repointed fastvideo editable to worktree; run from worktree cwd. Branch also rebased onto origin/main to fix stale import. |
|
||||
|
||||
## Decisions
|
||||
|
||||
| Date | Decision | Rationale | Impact |
|
||||
|---|---|---|---|
|
||||
| 2026-06-06 | Reference source of truth = official diffusers (`Cosmos3OmniDiffusersPipeline` + `cosmos-framework`/`diffusers-cosmos3`), not vllm-omni #3454 | Official weights now public & diffusers-format; the artifact users actually load | Repoint DiT/pipeline/conversion/tests off vllm-omni (I001, I002) |
|
||||
| 2026-06-06 | Resume in worktree `/home/william5lin/FastVideo_cosmos3_port`; weights+reference symlinked (no copy) | Preserve 2,492 lines of Tier-A work; avoid 33 GB duplication | Verification needs worktree `fastvideo` active (E001) |
|
||||
| 2026-06-06 | Scope = full omni (video + audio + reasoning + action) | User choice (revised from branch's original video-only scope) | Adds `vision_encoder`, `sound_tokenizer` ports + `WorkloadType` AV + audio metric |
|
||||
| 2026-06-06 | Downloaded full 34.9 GB (33 GiB) `nvidia/Cosmos3-Nano` | Unblocks May-22 `PENDING` weight status (HF was 401, now public) | Real parity now possible |
|
||||
| 2026-06-06 | Rebased branch onto origin/main (33 commits); resolved registry.py conflict by reconstructing from main + cosmos3 import/entry | Branch was stale; fastvideo failed to import (main removed MatrixGameI2V480PConfig) | Branch imports clean; Tier-A 13 passed/2 skipped |
|
||||
| 2026-06-06 | Reference = cosmos_framework ONLY (full omni); diffusers shim dropped even for video | User directive (Phase 1 found diffusers __call__ is video-only; sound/action/reasoning live only in the framework) | Larger port; ref DiT = `Cosmos3VFMNetwork`/`Cosmos3VFMNetworkConfig` (not diffusers `Cosmos3OmniTransformer`); core model imports in fv-cosmos3 with light deps; TE only in optional dot_product_attention |
|
||||
|
||||
## Handoff Notes
|
||||
|
||||
- Prep (weights/reference/env editable installs) done in MAIN worktree; symlinked into this worktree. Env installs (`diffusers-cosmos3`, `cosmos-framework`) are in shared `fv-main`.
|
||||
- Next: resolve E001 (env), then Phase 1 reference study of `diffusers_cosmos3` pipeline/transformer, then Phase 3 reuse gate (VAE/scheduler/tokenizer) + component dispatch (transformer, vision_encoder, sound_tokenizer).
|
||||
- diffusers 0.36.0 imports the shim OK; checkpoint saved with 0.37.1 — watch `from_pretrained` needs (bump within FastVideo's `diffusers>=0.33.1` pin if required).
|
||||
|
||||
### PR1 (video core) progress — 2026-06-06
|
||||
- Arch config 1:1 with checkpoint, committed `9567efdf0`.
|
||||
- Framework parity-reference harness committed `dd97efda3`: `tests/local_tests/cosmos3/test_cosmos3_reference_forward.py` builds a tiny `Cosmos3VFMNetwork` on CPU/float32 (SDPA monkeypatch; flash2/3/natten are CUDA-only) and forwards `packed_seq -> {last_hidden_state, preds_vision}`. 23 tests pass in fv-cosmos3. This is the ground-truth side for DiT parity. Run: `cd <worktree> && <fv-cosmos3 py> -m pytest tests/local_tests/cosmos3/test_cosmos3_reference_forward.py -q`.
|
||||
- THREE naming conventions to bridge:
|
||||
1. framework-native (`Cosmos3VFMNetwork`): `language_model.model.layers.{i}.self_attn.{q,k,v,o}_proj(+ _moe_gen)`, `{q,k}_norm(+_moe_gen)`, `mlp(+_moe_gen)`, `vae2llm`/`llm2vae`, `time_embedder.mlp.{0,2}`.
|
||||
2. diffusers checkpoint (on disk, what we load): `layers.{i}.self_attn.{to_q,to_k,to_v,to_out}` + `{add_q,add_k,add_v}_proj`/`to_add_out`, `{norm_q,norm_k,norm_added_q,norm_added_k}`, `mlp`/`mlp_moe_gen`, `proj_in`/`proj_out`, `time_embedder.linear_{1,2}`.
|
||||
3. FastVideo DiT (our choice). Conversion maps (2)->(3); the DiT parity test copies (1)->(3).
|
||||
- BaseDiT signature is `__init__(self, config: DiTConfig, hf_config: dict)`; the branch `Cosmos3VFMTransformer` uses `fastvideo_args`/SimpleNamespace and does NOT conform — rewrite to conform + match the checkpoint key surface (single `layers` dual-pathway, not split language_model/gen_layers).
|
||||
- Native layers (per cosmos2_5): `ReplicatedLinear`/`MLP`/`RMSNorm` (fastvideo.layers.*), `LocalAttention`/`DistributedAttention` (fastvideo.attention), `apply_rotary_emb` (use_real_unbind_dim=-2 for Cosmos). EntryClass at module bottom; class attrs bound from config; 3D-MRoPE has no reusable util — adapt Cosmos25RotaryPosEmbed.
|
||||
- NEXT: write native `fastvideo/models/dits/cosmos3.py` + fastvideo-vs-framework forward parity test (copy framework weights into the FastVideo DiT, compare outputs), then conversion script (diffusers checkpoint -> FastVideo) + strict-load, then video pipeline/packing.
|
||||
|
||||
### PR1 (video core) acceptance — real-weights E2E — 2026-06-07
|
||||
- First real-weights T2V (`examples/inference/basic/basic_cosmos3_new_api.py`, `COSMOS3_MODEL_PATH=official_weights/cosmos3`) ran mechanically but produced an all-black 3 KB mp4. Instrumenting the denoise loop showed `v_pred` clean at step 0 but `scheduler.step` -> NaN. Root cause I003: checkpoint `scheduler_config.json` is diffusers-style (`use_karras_sigmas=true`), and the vendored UniPC checks karras before `use_flow_sigmas` -> diffusion (beta) sigmas instead of flow sigmas -> NaN. The framework actually samples with `FlowUniPCMultistepScheduler` (pure flow: `shift` + `num_train_timesteps`).
|
||||
- Fix (commit `255311cf2`): coerce the loaded scheduler to the flow setup in `Cosmos3OmniDiffusersPipeline.initialize_pipeline`; use FastVideo's native UniPC (not diffusers) in pipeline + tests. Added `test_cosmos3_scheduler_parity.py` (native UniPC flow-config vs framework `FlowUniPCMultistepScheduler`: timesteps bit-exact, sigmas ~1e-8, full trajectory <~1e-6 over shift in {10,3}, steps in {4,10,35}). Repointed `test_cosmos3_denoise_cfg_parity` oracle to the framework scheduler (it previously compared diffusers-vs-diffusers, so the scheduler was never checked against the framework).
|
||||
- Also wired the remaining integration glue (registry alias `Cosmos3OmniTransformer`->`Cosmos3VFMTransformer`; `text_tokenizer`->TokenizerLoader; scheduler config param-filtering; DiT `materialize_non_persistent_buffers` + compute-dtype casts; packing device-move in `to_dit_kwargs`; empty text-preprocess).
|
||||
- Verified: 1280x704, 29 frames, 35 steps on a single B200 -> coherent golden-retriever-in-meadow video matching the prompt (no NaNs; per-frame pixel std ~58; visible temporal motion). Full cosmos3 suite: 95 passed, 0 skipped.
|
||||
- NEXT: PR2 audio (`sound_tokenizer` AVAE) / PR3 action / PR4 reasoning. Optional: I2V/T2I real-weights spot-checks; force-push branch (needs explicit OK).
|
||||
|
||||
### PR1 (video core) — I2V real-weights — 2026-06-07 (branch feat/cosmos3-i2v)
|
||||
- Forked `feat/cosmos3-i2v` off `feat/cosmos3-tier-a-port` (stacked, includes the T2V + scheduler fix).
|
||||
- Studied the framework I2V path: `cosmos_framework.inference.vision.load_conditioning_image` (aspect-preserving resize + center crop + uint8 quantize -> `/127.5-1`) + `build_conditioned_video_batch` (frame 0 = image, remaining frames REPEAT the last conditioning frame -> static video), then VAE-encode; `condition_frame_indexes=[0]` (latent). Condition frames kept clean during sampling exactly as FastVideo already does: init noise `cond_mask*x0 + (1-cond_mask)*noise` (`omni_mot_model._prepare_inference_data`) + velocity zeroed `pred*(1-cond_mask)` each step (`_get_velocity`), no re-injection.
|
||||
- Bug found + fixed (commit `bd8d604fb`): FastVideo's `_image_to_video_tensor` ZERO-filled the non-condition frames; the temporal Wan VAE (4x) makes latent frame 0 depend on several pixel frames, so zero-fill -> wrong conditioning latent. Rewrote it to repeat-fill + framework resize/crop/quantize.
|
||||
- Parity: `test_cosmos3_i2v_conditioning_parity.py` vs framework `load_conditioning_image` + repeat-fill — bit-exact (max abs diff 0.0) across aspect/size/frame cases. Existing `test_cosmos3_denoise_cfg_parity` already covers the I2V cond-mask + velocity math (i2v case).
|
||||
- Example: `examples/inference/basic/basic_cosmos3_i2v_new_api.py` (`InputConfig(image_path=...)`, default `assets/images/cyclist.jpg`).
|
||||
- Verified on B200 (1280x704, 29f, 35 steps, real weights): output frame 0 reproduces the conditioning cyclist image; later frames show coherent forward motion down the trail following the prompt. Full suite 98 passed, 0 skipped.
|
||||
- NEXT: optional T2I real-weights spot-check; then PR2 audio / PR3 action / PR4 reasoning.
|
||||
|
||||
### PR1 (video core) — T2I real-weights + resolution-based flow_shift — 2026-06-07 (branch feat/cosmos3-i2v)
|
||||
- Studied framework T2I: tokenization uses `vlm_config.use_system_prompt` which is `false` in the checkpoint (config.json:199) — matches FastVideo's hardcoded `use_system_prompt=False` for all modes (no divergence). Canonical T2I is 960x960 (inputs/omni/t2i.json), single-frame (num_frames=1).
|
||||
- Bug found + fixed (commit `604dc2637`): the stage chose `flow_shift` by task (`3.0 if is_t2i else 10.0`), but the framework picks it purely by the named resolution bucket (`OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256->3.0, 480->5.0, 720/768->10.0; model default resolution "720"). Task-based only matched T2V@720 / T2I@256 by luck; canonical T2I@960x960 is the "720" bucket -> 10.0, so `is_t2i->3.0` was wrong. Replaced with `_flow_shift_for_resolution(h,w)` (longest-side bucketing), applied to all tasks.
|
||||
- Parity: `test_cosmos3_flow_shift_parity.py` checks the mapping vs framework `{VIDEO,IMAGE}_RES_SIZE_INFO` x `_RESOLUTION_SHIFT_DEFAULTS` (8B rows, 20 cases). Also hardened `_image_to_video_tensor` tensor branch to respect the [-1,1] convention (PIL path stays framework-exact).
|
||||
- Example: `examples/inference/basic/basic_cosmos3_t2i_new_api.py` (num_frames=1, 960x960).
|
||||
- Verified on B200 (real weights, 35 steps): coherent red-panda image matching the prompt, flow_shift=10.0. Full suite 118 passed, 0 skipped.
|
||||
- Video core (T2V/I2V/T2I) is now complete and real-weights verified. NEXT: PR2 audio (sound_tokenizer AVAE + audio output) on a new stacked branch.
|
||||
|
||||
## Full-omni parity summary (every component, max / mean abs diff vs framework)
|
||||
|
||||
All run on CPU / float32 (tiny models, framework weights copied in; framework =
|
||||
oracle). `tests/local_tests/cosmos3/`, suite: 150 passed, 0 skipped.
|
||||
|
||||
| Component / pipeline | Test | max | mean |
|
||||
|---|---|---|---|
|
||||
| Scheduler (UniPC flow) | test_cosmos3_scheduler_parity | timesteps 0; sigmas ~1e-8; traj <~1e-6 | ~1e-7 |
|
||||
| DiT (video, unified_3d_mrope) | test_cosmos3_dit_parity_mrope | 0.0 | 0.0 |
|
||||
| Sequence packing (video) | test_cosmos3_packing_parity | 0.0 (exact) | 0.0 |
|
||||
| VAE (Wan2.2) | test_cosmos3_vae_parity | 0.0 | 0.0 |
|
||||
| Denoise / CFG velocity | test_cosmos3_denoise_cfg_parity | <1e-6 | <1e-7 |
|
||||
| flow_shift (resolution) | test_cosmos3_flow_shift_parity | exact | exact |
|
||||
| I2V conditioning (static-repeat) | test_cosmos3_i2v_conditioning_parity | 0.0 | 0.0 |
|
||||
| AVAE sound decoder | test_cosmos3_avae_parity | 0.0 | 0.0 |
|
||||
| DiT sound pathway + packing | test_cosmos3_sound_parity | 0.0 | 0.0 |
|
||||
| t2vs CFG velocity | test_cosmos3_sound_parity | 0.0 | 0.0 |
|
||||
| DiT action pathway + packing | test_cosmos3_action_parity | 0.0 | 0.0 |
|
||||
| action CFG velocity | test_cosmos3_action_parity | 0.0 | 0.0 |
|
||||
| Reasoner prefill logits (text) | test_cosmos3_reasoning_parity | 0.0 | 0.0 |
|
||||
| Reasoner greedy generation | test_cosmos3_reasoning_parity | token-exact | - |
|
||||
| Deepstack reasoner forward | test_cosmos3_reasoning_parity | 0.0 | 0.0 |
|
||||
| vision_encoder (Qwen3-VL ViT) | test_cosmos3_vision_encoder_parity | 0.0 | 0.0 |
|
||||
|
||||
Real-weights pipelines verified on B200 (`examples/inference/basic/basic_cosmos3*_new_api.py`):
|
||||
T2V (1280x704), I2V (cyclist), T2I (960x960 red panda), t2vs (ocean + stereo
|
||||
48kHz audio), text reasoning (greedy == framework). All coherent / prompt-matching.
|
||||
@@ -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.
|
||||
@@ -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,166 @@
|
||||
# Cosmos3 → FastVideo port — feedback (pitfalls, issues, difficulties)
|
||||
|
||||
Retrospective on porting the full **NVIDIA Cosmos3-Nano** omni world model
|
||||
(video / audio / action generation + text & image reasoning) into FastVideo.
|
||||
Methodology: framework-only reference, native FastVideo port, a bit-exact
|
||||
framework-parity test per component, then real-weights verification. Every
|
||||
modality landed bit-exact (see `PORT_STATUS.md` "Full-omni parity summary").
|
||||
|
||||
This doc records what bit, so the next omni/world-model port (and the `/add-model`
|
||||
skill) can avoid the same traps.
|
||||
|
||||
---
|
||||
|
||||
## 1. The checkpoint's config does NOT describe the runtime — verify against the framework
|
||||
|
||||
The single biggest time sink. The HF checkpoint is "diffusers format", which led
|
||||
to two silent traps:
|
||||
|
||||
- **Scheduler (caused an all-black video).** `scheduler/scheduler_config.json`
|
||||
is a diffusers `UniPCMultistepScheduler` config carrying
|
||||
`use_karras_sigmas=true`, `sigma_min/sigma_max`, a beta schedule, etc. But the
|
||||
framework actually samples with a *flow-matching* `FlowUniPCMultistepScheduler`
|
||||
(shift + num_train_timesteps only). FastVideo's vendored UniPC checks
|
||||
`use_karras_sigmas` **before** `use_flow_sigmas`, so it built diffusion (beta)
|
||||
sigmas → `scheduler.step` → **NaN latents → 3 KB black mp4**. The DiT/CFG
|
||||
velocity was perfectly clean; only the scheduler diverged.
|
||||
- Fix: coerce the loaded config to the flow setup in `initialize_pipeline`.
|
||||
- **Lesson:** treat the checkpoint's generic-format config as *lossy*. Find how
|
||||
the framework actually instantiates the component and match THAT, not the JSON.
|
||||
|
||||
- **`flow_shift` is resolution-based, not task-based.** Natural assumption:
|
||||
"T2I uses a small shift, video a large one." Reality: the framework keys the
|
||||
UniPC shift purely off the named resolution bucket
|
||||
(`_RESOLUTION_SHIFT_DEFAULTS`, 8B backbone: 256→3, 480→5, 720/768→10). The
|
||||
task-based heuristic only *coincidentally* matched (T2V@720, T2I@256); canonical
|
||||
T2I is 960×960 (the "720" bucket → 10), so `is_t2i→3.0` was wrong.
|
||||
|
||||
## 2. A "parity test" that compares two copies of the wrong thing proves nothing
|
||||
|
||||
The original denoise/CFG test imported **diffusers** `UniPCMultistepScheduler` and
|
||||
used it on BOTH the "oracle" and "FastVideo" sides. So the scheduler was never
|
||||
actually compared against the framework — which is exactly why the black-video
|
||||
scheduler bug sailed through a green test suite.
|
||||
- **Lesson:** the oracle side of every parity test MUST be the official framework
|
||||
object, never a second instance of the unit under test. After writing a parity
|
||||
test, ask: "if the framework were wrong here, would this test fail?"
|
||||
|
||||
## 3. Temporal-VAE conditioning: static-repeat vs zero-fill (silent corruption)
|
||||
|
||||
I2V/T2I condition on the input image. The framework
|
||||
(`build_conditioned_video_batch`) fills frame 0 with the image and **repeats the
|
||||
last conditioning frame across the whole clip** (a static video) before
|
||||
VAE-encoding. The first native cut **zero-filled** the non-condition frames.
|
||||
Because the Wan VAE is temporal (4× compression), latent-frame-0 (the kept-clean
|
||||
condition frame) depends on several *pixel* frames — so zero-filling produced a
|
||||
*wrong* conditioning latent. This is the kind of bug that doesn't crash and can
|
||||
even look plausible at a glance.
|
||||
- **Lesson:** when a conditioning latent feeds a temporal autoencoder, trace the
|
||||
temporal receptive field; "only frame 0 matters" is false under temporal conv.
|
||||
|
||||
## 4. Checkpoint param names ≠ framework module structure (three naming conventions)
|
||||
|
||||
For the DiT there were **three** namings to bridge: framework-native
|
||||
(`Cosmos3VFMNetwork`: `language_model.model.layers.*`, `vae2llm`, `q_proj_moe_gen`,
|
||||
…), the diffusers checkpoint on disk (`layers.*.to_q`, `add_q_proj`, `proj_in`,
|
||||
…), and the FastVideo DiT. The weight map crosses (framework)→(FastVideo) for
|
||||
parity and (checkpoint)→(FastVideo) for loading.
|
||||
|
||||
The **sound tokenizer** was the sharpest example: the checkpoint is **decoder-only**
|
||||
in diffusers `AutoencoderOobleck` naming (`decoder.conv1`, `block.N.conv_t1`,
|
||||
`res_unitM`, `snake1`) — but with **`SnakeBeta`** (learned alpha *and* beta,
|
||||
logscale), NOT diffusers' alpha-only `Snake1d`. So neither "use diffusers
|
||||
AutoencoderOobleck" nor "port the framework `LatentAutoEncoderV2` Sequential
|
||||
module" matched the on-disk keys.
|
||||
- **Lesson:** dump `safetensors` keys + shapes for every sub-checkpoint *first*.
|
||||
The naming reveals which existing native module (if any) already matches.
|
||||
|
||||
## 5. A "matching" native module can still differ on an untested config path
|
||||
|
||||
FastVideo already had a native `OobleckVAE` (Stable Audio) whose decoder matched
|
||||
the Cosmos3 sound decoder bit-for-bit — except `OobleckDecoderBlock.conv_t1`
|
||||
omitted `output_padding = stride % 2`. That omission is a **no-op for Stable
|
||||
Audio's even strides** [2,4,4,8,8], so it had never mattered; Cosmos3 has an
|
||||
**odd** stride (5), where the framework's `output_padding=1` makes the decode one
|
||||
sample longer per odd-stride block (parity diverged 60 vs 59 samples).
|
||||
- **Lesson:** reusing a native module is great, but re-run parity on the *new*
|
||||
model's config — shared code can hide config-specific divergences.
|
||||
|
||||
## 6. Loader / registry plumbing the checkpoint format forces
|
||||
|
||||
- **DiT class alias.** `model_index.json` names the DiT `Cosmos3OmniTransformer`
|
||||
(the diffusers shim class); the registry normalized unknown classes to a
|
||||
generic `TransformersModel`. Needed an explicit registry alias
|
||||
`Cosmos3OmniTransformer → Cosmos3VFMTransformer`.
|
||||
- **Tokenizer module name.** `model_index.json` calls the Qwen2 tokenizer
|
||||
`text_tokenizer` (not `tokenizer`); the component loader had no mapping for that
|
||||
key and tried to load it as a model.
|
||||
- **Scheduler config schema drift.** The vendored UniPC predates
|
||||
`shift_terminal` / `sigma_min` / `sigma_max`; constructing it with the raw
|
||||
checkpoint config crashes on the unexpected kwargs. Filter to the class's
|
||||
`__init__` params (mirroring diffusers `from_config`).
|
||||
- **Meta-device load + non-persistent buffers.** `rotary_emb.inv_freq` is derived
|
||||
from `rope_theta` and is non-persistent (absent from the checkpoint), so after
|
||||
the meta-device FSDP load it stays on the `meta` device → needs a
|
||||
`materialize_non_persistent_buffers` hook to recompute it on the real device.
|
||||
- **dtype boundaries.** Noise/VAE latents arrive fp32; the model runs bf16.
|
||||
Needed explicit casts at `proj_in` and the timestep embedder (no-ops in the
|
||||
fp32 parity tests, required at inference).
|
||||
- **device in packing.** The packer builds ids/positions on CPU; `to_dit_kwargs`
|
||||
must move every tensor to the model device before the forward.
|
||||
|
||||
## 7. The omni model is a Mixture-of-Transformers — modality bookkeeping is the work
|
||||
|
||||
The backbone is a dual-pathway MoT: **und** (causal text) + **gen** (full-attention
|
||||
vision/sound/action). Once the video path worked, each extra modality was the same
|
||||
*shape* of work (a proj-in + modality embed + timestep-scatter encode, a proj-out
|
||||
decode, packing, a CFG-velocity slice) but with per-modality quirks:
|
||||
- sound/action **share the vision "full" split** (preserving the causal+full
|
||||
2-split invariant); the combined flat latent is `[vision | action | sound]` in
|
||||
that order (must match the framework's per-sample concat).
|
||||
- sound MRoPE uses `start_frame_offset=0` (parallel to vision); action uses
|
||||
`start_frame_offset=1`; both at the vision temporal offset, tcf=1, and do NOT
|
||||
advance the offset.
|
||||
- action is **domain-aware** (`DomainAwareLinear`: per-embodiment weight/bias via
|
||||
`nn.Embedding`, indexed by a per-token domain id).
|
||||
- the unpack already zeros clean frames, so the per-step velocity masking is
|
||||
defensive (but kept, to mirror the framework exactly).
|
||||
- **Lesson:** build the first modality (vision) with clean seams for "a modality"
|
||||
and the rest fall out; spend the care on the packing layout + MRoPE offsets,
|
||||
which are the only per-modality novelties.
|
||||
|
||||
## 8. Reasoning reused more than expected; the encoder is just transformers
|
||||
|
||||
- **Text reasoning** needed *no new model code*: it's the und (causal) pathway +
|
||||
`embed_tokens`/`norm`/`lm_head`, all already in the DiT. A text-only forward +
|
||||
`lm_head` is token-for-token identical to the framework reasoner.
|
||||
- **vision_encoder** is a stock `transformers.Qwen3VLVisionModel` (the framework
|
||||
ships its own *copy* of the same class); reusing transformers' (like the Qwen2
|
||||
tokenizer) is bit-exact vs the framework — re-porting a 27-layer ViT would have
|
||||
been wasted effort.
|
||||
- **deepstack** (image-conditioned reasoning) is the one new native piece: inject
|
||||
the 3 vision-encoder deepstack features into the first 3 text layers at the
|
||||
image-token positions.
|
||||
- **Lesson:** before porting a big sub-model, check whether it's literally a
|
||||
stock library class — and whether an existing in-repo module already implements
|
||||
it (audio decoder + vision encoder were both "already there").
|
||||
|
||||
## 9. Running the framework as a CPU parity oracle
|
||||
|
||||
The framework's attention path is flash/natten (CUDA-only). Parity tests run on
|
||||
CPU/float32 via an SDPA monkey-patch (`test_cosmos3_reference_forward._apply_sdpa_patches`).
|
||||
A couple of framework helpers also can't import headless (`cosmos_framework.inference.args`
|
||||
pulls `multistorageclient`), so a constant or two is mirrored in the test with a
|
||||
cited source rather than imported.
|
||||
- **Lesson:** budget for "make the oracle runnable on CPU" — monkeypatch attention,
|
||||
build tiny configs, and accept a small amount of mirrored constants when a
|
||||
framework module won't import in isolation.
|
||||
|
||||
## 10. What made it tractable
|
||||
|
||||
- Tiny CPU/fp32 models + copy-framework-weights-in + bit-exact compare, per
|
||||
component, is a fast and decisive loop (max=mean=0.0 or it's wrong).
|
||||
- A persistent `PORT_STATUS.md` (resumable state, issues, decisions) survived
|
||||
several context resets.
|
||||
- Stacked PR branches (one per modality) kept each parity-verified increment
|
||||
reviewable and the chain bisectable.
|
||||
@@ -0,0 +1,271 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 action pathway vs the framework.
|
||||
|
||||
Covers the action (multi-embodiment world-model) modality at the DiT level:
|
||||
|
||||
* **action packing** — native ``pack_cosmos3_video_sequence`` with a
|
||||
``Cosmos3ActionItem`` vs framework ``pack_input_sequence`` with
|
||||
``has_action``: action tokens share the vision "full" split, with ``(T,)``
|
||||
shapes, a ``(T,1)`` condition mask, and 3D-MRoPE temporal positions at the
|
||||
vision offset with ``start_frame_offset=1`` (parallel to vision); and
|
||||
* **DiT action forward** — the dormant domain-aware ``action_proj_in`` /
|
||||
``action_proj_out`` (``DomainAwareLinear``) + ``action_modality_embed`` heads,
|
||||
now activated, with a per-token embodiment ``domain_id``.
|
||||
|
||||
Framework model + pack is the parity ORACLE (CPU/float32 via SDPA monkey-patch).
|
||||
We assert the native packer matches the framework field-by-field, then that
|
||||
``preds_vision`` AND ``preds_action`` match the framework forward.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_action_parity.py -q -s
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from .test_cosmos3_dit_parity import ( # noqa: E402
|
||||
_fastvideo_inputs_from_packed_seq,
|
||||
_framework_to_fastvideo_state_dict,
|
||||
)
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_ACTION_DIM,
|
||||
_LATENT_CHANNEL,
|
||||
_LATENT_PATCH_SIZE,
|
||||
_RESET_SPATIAL_IDS,
|
||||
_TCF,
|
||||
_TEMPORAL_MODALITY_MARGIN,
|
||||
_build_tiny_cosmos3_mrope,
|
||||
_build_tiny_fastvideo_dit_mrope,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
_apply_sdpa_patches()
|
||||
|
||||
_SPECIAL_TOKENS = {"start_of_generation": 60, "end_of_generation": 61, "eos_token_id": 62}
|
||||
|
||||
|
||||
def _copy_weights_with_action(vfm, dit) -> None:
|
||||
"""Copy backbone + vision weights AND the domain-aware action heads."""
|
||||
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
|
||||
src = dict(vfm.named_parameters())
|
||||
mapped["action_proj_in.fc.weight"] = src["action2llm.fc.weight"].detach().clone()
|
||||
mapped["action_proj_in.bias.weight"] = src["action2llm.bias.weight"].detach().clone()
|
||||
mapped["action_proj_out.fc.weight"] = src["llm2action.fc.weight"].detach().clone()
|
||||
mapped["action_proj_out.bias.weight"] = src["llm2action.bias.weight"].detach().clone()
|
||||
mapped["action_modality_embed"] = src["action_modality_embed"].detach().clone()
|
||||
dst = dict(dit.named_parameters())
|
||||
with torch.no_grad():
|
||||
for name, tensor in mapped.items():
|
||||
assert name in dst, f"DiT missing param {name!r}"
|
||||
assert dst[name].shape == tensor.shape, f"shape mismatch {name}"
|
||||
dst[name].copy_(tensor.to(dst[name].dtype))
|
||||
|
||||
|
||||
def _framework_pack_action(*, text_ids, vision, action, cond_vision, cond_action, domain_id, timestep,
|
||||
is_image_batch):
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
GenerationDataClean,
|
||||
SequencePlan,
|
||||
pack_input_sequence,
|
||||
)
|
||||
|
||||
gen = GenerationDataClean(
|
||||
batch_size=1,
|
||||
is_image_batch=is_image_batch,
|
||||
x0_tokens_vision=[vision],
|
||||
fps_vision=None,
|
||||
num_vision_items_per_sample=[1],
|
||||
x0_tokens_action=[action],
|
||||
fps_action=None,
|
||||
action_domain_id=[torch.tensor([domain_id], dtype=torch.long)],
|
||||
)
|
||||
plans = [SequencePlan(
|
||||
has_text=True, has_vision=True, has_action=True,
|
||||
condition_frame_indexes_vision=list(cond_vision),
|
||||
condition_frame_indexes_action=list(cond_action),
|
||||
)]
|
||||
ps = pack_input_sequence(
|
||||
sequence_plans=plans,
|
||||
input_text_indexes=[list(text_ids)],
|
||||
gen_data_clean=gen,
|
||||
input_timesteps=torch.tensor([timestep], dtype=torch.float32),
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
# The framework sets action.domain_id on the packed sequence from
|
||||
# gen_data_clean inside the model (_get_velocity); mirror that for the oracle.
|
||||
if ps.action is not None:
|
||||
ps.action.domain_id = [torch.tensor([domain_id], dtype=torch.long)]
|
||||
return ps
|
||||
|
||||
|
||||
def _fastvideo_pack_action(*, text_ids, vision, action, cond_vision, cond_action, domain_id, timestep):
|
||||
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
|
||||
Cosmos3ActionItem,
|
||||
Cosmos3SampleInputs,
|
||||
Cosmos3VisionItem,
|
||||
pack_cosmos3_video_sequence,
|
||||
)
|
||||
|
||||
samples = [Cosmos3SampleInputs(
|
||||
text_ids=list(text_ids),
|
||||
vision=Cosmos3VisionItem(latent=vision, condition_frame_indexes=list(cond_vision)),
|
||||
action=Cosmos3ActionItem(latent=action, condition_frame_indexes=list(cond_action), domain_id=domain_id),
|
||||
timestep=float(timestep),
|
||||
)]
|
||||
return pack_cosmos3_video_sequence(
|
||||
samples, _SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE, include_end_of_generation_token=False,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN, reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False, base_fps=24.0, temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
def _fv_inputs_with_action(ps) -> dict:
|
||||
kw = _fastvideo_inputs_from_packed_seq(ps)
|
||||
a = ps.action
|
||||
kw.update(
|
||||
action_tokens=list(a.tokens),
|
||||
action_token_shapes=[tuple(x) for x in a.token_shapes],
|
||||
action_sequence_indexes=a.sequence_indexes,
|
||||
action_timesteps=a.timesteps,
|
||||
action_mse_loss_indexes=a.mse_loss_indexes,
|
||||
action_noisy_frame_indexes=list(a.noisy_frame_indexes),
|
||||
action_domain_id=list(a.domain_id),
|
||||
)
|
||||
return kw
|
||||
|
||||
|
||||
def _diffs(a, b):
|
||||
d = (a - b).abs()
|
||||
return d.max().item(), d.mean().item()
|
||||
|
||||
|
||||
# (grid_t, lh, lw, action_t, n_text, cond_vision, cond_action, domain_id)
|
||||
_CASES = [
|
||||
pytest.param(2, 4, 4, 6, 4, [], [], 0, id="a2v_2x2x2_act6_dom0"),
|
||||
pytest.param(3, 8, 4, 9, 5, [], [], 7, id="a2v_3x4x2_act9_dom7"),
|
||||
pytest.param(2, 4, 4, 5, 5, [0], [0], 3, id="ai2v_cond_act5_dom3"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3ActionParity:
|
||||
|
||||
def _build(self, num_layers=2, seed_model=42):
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers, action_gen=True)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights_with_action(vfm, dit)
|
||||
return vfm, dit
|
||||
|
||||
def _make_inputs(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom, seed=7):
|
||||
torch.manual_seed(seed)
|
||||
return dict(
|
||||
text_ids=torch.randint(0, 60, (n_text,)).tolist(),
|
||||
vision=torch.randn(1, _LATENT_CHANNEL, grid_t, lh, lw),
|
||||
action=torch.randn(act_t, _ACTION_DIM), # [T, D]
|
||||
cond_vision=cond_v, cond_action=cond_a, domain_id=dom, timestep=500.0,
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
|
||||
def test_action_packing_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
|
||||
ins = self._make_inputs(grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom)
|
||||
fw = _framework_pack_action(is_image_batch=(grid_t == 1), **ins)
|
||||
fv = _fastvideo_pack_action(**ins)
|
||||
assert fv.split_lens == list(fw.split_lens), f"split_lens fv={fv.split_lens} fw={list(fw.split_lens)}"
|
||||
assert fv.attn_modes == list(fw.attn_modes)
|
||||
assert int(fv.sequence_length) == int(fw.sequence_length)
|
||||
torch.testing.assert_close(fv.position_ids, fw.position_ids, rtol=0, atol=0)
|
||||
a = fw.action
|
||||
torch.testing.assert_close(fv.action_sequence_indexes, a.sequence_indexes.to(torch.long), rtol=0, atol=0)
|
||||
assert fv.action_token_shapes == [tuple(x) for x in a.token_shapes]
|
||||
torch.testing.assert_close(fv.action_timesteps.to(torch.float32), a.timesteps.to(torch.float32))
|
||||
torch.testing.assert_close(fv.action_mse_loss_indexes, a.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
|
||||
for x, y in zip(fv.action_noisy_frame_indexes, a.noisy_frame_indexes):
|
||||
torch.testing.assert_close(x.to(torch.long), y.to(torch.long), rtol=0, atol=0)
|
||||
print(f"\n[action_packing {grid_t}x{lh}x{lw} act={act_t} dom={dom}] position_ids + action fields exact")
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
|
||||
def test_action_dit_forward_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
|
||||
vfm, dit = self._build()
|
||||
ins = self._make_inputs(grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom)
|
||||
fw_pack = _framework_pack_action(is_image_batch=(grid_t == 1), **ins)
|
||||
fv_pack = _fastvideo_pack_action(**ins)
|
||||
with torch.no_grad():
|
||||
fw_out = vfm(packed_seq=fw_pack)
|
||||
fv_out = dit(**fv_pack.to_dit_kwargs())
|
||||
fv_on_fw = dit(**_fv_inputs_with_action(fw_pack))
|
||||
pv_mx, pv_mn = _diffs(fv_out["preds_vision"][0], fw_out["preds_vision"][0])
|
||||
pa_mx, pa_mn = _diffs(fv_out["preds_action"][0], fw_out["preds_action"][0])
|
||||
paf_mx, paf_mn = _diffs(fv_on_fw["preds_action"][0], fw_out["preds_action"][0])
|
||||
print(f"\n[action_dit {grid_t}x{lh}x{lw} act={act_t} dom={dom}] "
|
||||
f"preds_vision max={pv_mx:.3e} mean={pv_mn:.3e} | "
|
||||
f"preds_action max={pa_mx:.3e} mean={pa_mn:.3e} | "
|
||||
f"preds_action(fwpack) max={paf_mx:.3e} mean={paf_mn:.3e}")
|
||||
assert fv_out["preds_action"][0].shape == fw_out["preds_action"][0].shape
|
||||
torch.testing.assert_close(fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3)
|
||||
torch.testing.assert_close(fv_out["preds_action"][0], fw_out["preds_action"][0], atol=1e-4, rtol=1e-3)
|
||||
torch.testing.assert_close(fv_on_fw["preds_action"][0], fw_out["preds_action"][0], atol=1e-4, rtol=1e-3)
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "lh", "lw", "act_t", "n_text", "cond_v", "cond_a", "dom"), _CASES)
|
||||
def test_action_cfg_velocity_matches_framework(self, grid_t, lh, lw, act_t, n_text, cond_v, cond_a, dom):
|
||||
"""Combined [vision|action] sequential-CFG velocity (action pipeline glue)
|
||||
matches a framework-DiT oracle."""
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3ActionSpec,
|
||||
Cosmos3VisionSpec,
|
||||
cosmos3_get_cfg_velocity,
|
||||
)
|
||||
|
||||
vfm, dit = self._build()
|
||||
vlat_shape = (_LATENT_CHANNEL, grid_t, lh, lw)
|
||||
action_shape = (act_t, _ACTION_DIM)
|
||||
torch.manual_seed(3)
|
||||
cond_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (max(1, n_text - 1),)).tolist()
|
||||
vis_numel = int(torch.tensor(vlat_shape).prod())
|
||||
act_numel = int(torch.tensor(action_shape).prod())
|
||||
flat = torch.randn(vis_numel + act_numel)
|
||||
guidance, ts = 6.0, 500.0
|
||||
|
||||
def _fw_velocity(ids):
|
||||
vision = flat[:vis_numel].reshape(vlat_shape).unsqueeze(0)
|
||||
action = flat[vis_numel:].reshape(action_shape)
|
||||
ps = _framework_pack_action(text_ids=ids, vision=vision, action=action, cond_vision=cond_v,
|
||||
cond_action=cond_a, domain_id=dom, timestep=ts,
|
||||
is_image_batch=(grid_t == 1))
|
||||
with torch.no_grad():
|
||||
out = vfm(packed_seq=ps)
|
||||
pv = out["preds_vision"][0].squeeze(0) # [C,T,H,W] (zero on clean)
|
||||
pa = out["preds_action"][0] # [T,D] (zero on clean)
|
||||
return torch.cat([pv.reshape(-1), pa.reshape(-1)])
|
||||
|
||||
fw_cond, fw_uncond = _fw_velocity(cond_ids), _fw_velocity(uncond_ids)
|
||||
fw_v = fw_uncond + guidance * (fw_cond - fw_uncond)
|
||||
fv_v = cosmos3_get_cfg_velocity(
|
||||
transformer=dit, flat_latent=flat, timestep=torch.tensor([ts]), guidance=guidance,
|
||||
specs=[Cosmos3VisionSpec(shape=vlat_shape, condition_frame_indexes=list(cond_v))],
|
||||
action_specs=[Cosmos3ActionSpec(shape=action_shape, condition_frame_indexes=list(cond_a), domain_id=dom)],
|
||||
cond_token_ids=cond_ids, uncond_token_ids=uncond_ids,
|
||||
special_tokens=_SPECIAL_TOKENS, latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN, reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False, base_fps=24.0, temporal_compression_factor=_TCF,
|
||||
)
|
||||
assert fv_v.shape == fw_v.shape, f"shape fv={fv_v.shape} fw={fw_v.shape}"
|
||||
mx, mn = _diffs(fv_v, fw_v)
|
||||
print(f"\n[action_cfg_velocity {grid_t}x{lh}x{lw} act={act_t} dom={dom}] max={mx:.3e} mean={mn:.3e}")
|
||||
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
|
||||
@@ -0,0 +1,131 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 sound decoder vs the framework AVAE.
|
||||
|
||||
The Cosmos3 ``sound_tokenizer`` is an AVAE (audio VAE). Its shipped diffusers
|
||||
checkpoint is **decoder-only** (``decoder.*``; the SpectrogramConvNeXt encoder is
|
||||
not exported) in diffusers ``AutoencoderOobleck`` naming, but with **SnakeBeta**
|
||||
activations (alpha+beta, logscale) and ``weight_g``/``weight_v`` weight-norm —
|
||||
i.e. exactly FastVideo's existing native ``OobleckVAE`` decoder
|
||||
(``fastvideo/models/vaes/oobleck.py``). t2vs only needs DECODE (generate sound
|
||||
latents -> waveform), so this pins the decoder.
|
||||
|
||||
The framework decoder
|
||||
(``cosmos_framework.model.vfm.tokenizers.audio.avae_utils.models.OobleckDecoder``,
|
||||
``nn.Sequential`` naming, ``output_padding=stride%2`` on the transpose convs) is
|
||||
the parity ORACLE. We build a tiny framework decoder, map its weights into the
|
||||
FastVideo decoder (Sequential -> conv1/block.N/res_unitM/snake1/conv2), and
|
||||
assert bit-exact decode. Strides include an ODD value (5, as in the real config
|
||||
``[2,4,5,6,8]``) to exercise the ``output_padding`` path that diverged before.
|
||||
|
||||
CPU / float32. Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_avae_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
_fw_models = pytest.importorskip(
|
||||
"cosmos_framework.model.vfm.tokenizers.audio.avae_utils.models",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
from cosmos_framework.model.vfm.tokenizers.audio.avae_utils.env import ( # noqa: E402
|
||||
AttrDict,
|
||||
)
|
||||
|
||||
from fastvideo.models.vaes.oobleck import OobleckDecoder as FvOobleckDecoder # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
FwOobleckDecoder = _fw_models.OobleckDecoder
|
||||
|
||||
|
||||
def _framework_decoder(dec_dim, vocoder_input_dim, dec_c_mults, dec_strides):
|
||||
"""Framework OobleckDecoder (the parity oracle), non-causal / no-antialias."""
|
||||
h = AttrDict({
|
||||
"vocoder_input_dim": vocoder_input_dim,
|
||||
"input_channels": 1,
|
||||
"stereo": True, # 2 audio channels
|
||||
"dec_dim": dec_dim,
|
||||
"dec_c_mults": dec_c_mults,
|
||||
"dec_strides": dec_strides,
|
||||
"dec_use_snake": True,
|
||||
"dec_use_nearest_upsample": False,
|
||||
"dec_anti_aliasing": False,
|
||||
"causal": False,
|
||||
"dec_use_tanh_at_final": False,
|
||||
"padding_mode": "zeros",
|
||||
})
|
||||
return FwOobleckDecoder(h).eval()
|
||||
|
||||
|
||||
def _framework_to_fastvideo_decoder_state(fw_decoder, num_blocks):
|
||||
"""Map framework Sequential decoder weights -> FastVideo decoder names.
|
||||
|
||||
framework: layers.0=first conv; layers.{1..K}=OobleckDecoderBlock
|
||||
(.layers.0 snake, .1 conv_t, .{2,3,4} ResidualUnit{.layers.0 snake,
|
||||
.1 conv, .2 snake, .3 conv}); layers.{1+K}=final snake; layers.{2+K}=final conv.
|
||||
FastVideo: conv1; block.{b}.{snake1,conv_t1,res_unit{1,2,3}.{snake1,conv1,snake2,conv2}};
|
||||
snake1; conv2. Snake alpha/beta: framework [C] -> FastVideo [1,C,1].
|
||||
"""
|
||||
out = {}
|
||||
for k, v in fw_decoder.state_dict().items():
|
||||
p = k.split(".")
|
||||
li = int(p[1])
|
||||
if li == 0:
|
||||
nk = "conv1." + ".".join(p[2:])
|
||||
elif li == 1 + num_blocks:
|
||||
nk = "snake1." + ".".join(p[2:])
|
||||
elif li == 2 + num_blocks:
|
||||
nk = "conv2." + ".".join(p[2:])
|
||||
else:
|
||||
b = li - 1
|
||||
sub = int(p[3])
|
||||
if sub == 0:
|
||||
nk = f"block.{b}.snake1." + ".".join(p[4:])
|
||||
elif sub == 1:
|
||||
nk = f"block.{b}.conv_t1." + ".".join(p[4:])
|
||||
else:
|
||||
r = sub - 2 # ResidualUnit index 0..2
|
||||
m = {0: "snake1", 1: "conv1", 2: "snake2", 3: "conv2"}[int(p[5])]
|
||||
nk = f"block.{b}.res_unit{r + 1}.{m}." + ".".join(p[6:])
|
||||
if nk.endswith(".alpha") or nk.endswith(".beta"):
|
||||
v = v.reshape(1, -1, 1)
|
||||
out[nk] = v
|
||||
return out
|
||||
|
||||
|
||||
# (dec_dim, vocoder_input_dim, dec_c_mults, dec_strides) — tiny; strides incl odd.
|
||||
_CASES = [
|
||||
pytest.param(4, 8, [1, 2], [5, 2], id="odd_stride5"),
|
||||
pytest.param(6, 8, [1, 2, 4], [2, 5, 6], id="real_stride_pattern_tiny"),
|
||||
pytest.param(4, 4, [1, 2], [4, 8], id="even_strides"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3AVAEParity:
|
||||
|
||||
@pytest.mark.parametrize(("dec_dim", "vin", "cmults", "strides"), _CASES)
|
||||
def test_decode_matches_framework(self, dec_dim, vin, cmults, strides):
|
||||
torch.manual_seed(0)
|
||||
fw = _framework_decoder(dec_dim, vin, cmults, strides)
|
||||
fv = FvOobleckDecoder(
|
||||
channels=dec_dim,
|
||||
input_channels=vin,
|
||||
audio_channels=2,
|
||||
upsampling_ratios=list(reversed(strides)), # framework reverses dec_strides
|
||||
channel_multiples=cmults,
|
||||
).eval()
|
||||
state = _framework_to_fastvideo_decoder_state(fw, num_blocks=len(strides))
|
||||
fv.load_state_dict(state, strict=True) # exact name + shape match
|
||||
|
||||
z = torch.randn(1, vin, 5)
|
||||
with torch.no_grad():
|
||||
a = fw(z)
|
||||
b = fv(z)
|
||||
assert a.shape == b.shape, f"shape: fw={a.shape} fv={b.shape}"
|
||||
max_abs = (a - b).abs().max().item()
|
||||
print(f"\n[avae_decode dim={dec_dim} strides={strides}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(b, a, atol=1e-6, rtol=1e-5)
|
||||
@@ -0,0 +1,342 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 denoise/CFG glue vs the framework.
|
||||
|
||||
The DiT forward and the sequence-packing are already framework-parity-verified
|
||||
(``test_cosmos3_dit_parity*`` / ``test_cosmos3_packing_parity``). This test pins
|
||||
the remaining glue that the native pipeline adds — the SEQUENTIAL classifier-free
|
||||
guidance velocity and one UniPC scheduler step — against the framework math
|
||||
(``diffusers_cosmos3.pipeline.Cosmos3OmniDiffusersPipeline.get_cfg_velocity`` /
|
||||
``__call__``):
|
||||
|
||||
* for one denoise step, replicate the framework's ``get_cfg_velocity`` exactly
|
||||
on top of the OFFICIAL ``Cosmos3VFMNetwork`` forward (oracle): a conditional
|
||||
pass (prompt tokens) and an unconditional pass (negative-prompt tokens),
|
||||
each masking the prediction on conditioning frames
|
||||
(``pred * (1 - condition_mask)``), then ``v = uncond + g*(cond - uncond)``;
|
||||
* run FastVideo's :func:`cosmos3_get_cfg_velocity` with the native DiT (the
|
||||
framework weights copied in) + the native packer, and assert the velocity
|
||||
matches the oracle;
|
||||
* take one ``UniPCMultistepScheduler.step`` on each (the actual checkpoint
|
||||
scheduler) and assert the stepped latent matches;
|
||||
* drive :meth:`Cosmos3DenoiseEngine.denoise` for >= 2 steps and assert it
|
||||
equals the manual framework step-by-step loop.
|
||||
|
||||
CPU / float32, via the reference SDPA monkey-patch. The official model is the
|
||||
parity ORACLE.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_denoise_cfg_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from .test_cosmos3_dit_parity import _copy_weights # noqa: E402
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_LATENT_CHANNEL,
|
||||
_LATENT_PATCH_SIZE,
|
||||
_RESET_SPATIAL_IDS,
|
||||
_TCF,
|
||||
_TEMPORAL_MODALITY_MARGIN,
|
||||
_build_tiny_cosmos3_mrope,
|
||||
_build_tiny_fastvideo_dit_mrope,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
from .test_cosmos3_scheduler_parity import ( # noqa: E402
|
||||
_fastvideo_scheduler,
|
||||
_framework_scheduler,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
_apply_sdpa_patches()
|
||||
|
||||
# Tiny special tokens (< tiny vocab_size=64), video path appends eos + sog.
|
||||
_SPECIAL_TOKENS = {
|
||||
"start_of_generation": 60,
|
||||
"end_of_generation": 61,
|
||||
"eos_token_id": 62,
|
||||
}
|
||||
|
||||
# Cosmos3 video flow_shift; framework scheduler is the parity oracle, FastVideo's
|
||||
# vendored UniPC (flow config) is the unit under test.
|
||||
_FLOW_SHIFT = 10.0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Framework-oracle CFG velocity (replicates pipeline.get_cfg_velocity math).
|
||||
# ---------------------------------------------------------------------------
|
||||
def _framework_pack(*, text_ids, vision_latent, cond_frames, timestep):
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
GenerationDataClean,
|
||||
SequencePlan,
|
||||
pack_input_sequence,
|
||||
)
|
||||
|
||||
# vision_latent is [1, C, T, H, W]; temporal dim is axis 2.
|
||||
gen_data_clean = GenerationDataClean(
|
||||
batch_size=1,
|
||||
is_image_batch=(vision_latent.shape[2] == 1),
|
||||
x0_tokens_vision=[vision_latent],
|
||||
fps_vision=None,
|
||||
num_vision_items_per_sample=[1],
|
||||
)
|
||||
plans = [SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=list(cond_frames))]
|
||||
return pack_input_sequence(
|
||||
sequence_plans=plans,
|
||||
input_text_indexes=[list(text_ids)],
|
||||
gen_data_clean=gen_data_clean,
|
||||
input_timesteps=torch.tensor([timestep], dtype=torch.float32),
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
def _framework_inputs(ps):
|
||||
"""Framework PackedSequence -> framework Cosmos3VFMNetwork forward kwargs."""
|
||||
return dict(packed_seq=ps)
|
||||
|
||||
|
||||
def _framework_cfg_velocity(
|
||||
*,
|
||||
vfm,
|
||||
flat_latent: torch.Tensor,
|
||||
timestep: torch.Tensor,
|
||||
guidance: float,
|
||||
vision_shape: tuple[int, int, int, int],
|
||||
cond_frames: list[int],
|
||||
cond_ids: list[int],
|
||||
uncond_ids: list[int],
|
||||
) -> torch.Tensor:
|
||||
"""Replicate the framework ``get_cfg_velocity`` on the oracle model.
|
||||
|
||||
Single vision item; sequential cond then uncond pass; mask condition
|
||||
frames; ``v = uncond + g*(cond - uncond)``.
|
||||
"""
|
||||
timestep_value = float(timestep.reshape(()).item())
|
||||
vision_latent = flat_latent.reshape(vision_shape) # [C, T, H, W]
|
||||
|
||||
def _run(text_ids: list[int]) -> torch.Tensor:
|
||||
ps = _framework_pack(
|
||||
text_ids=text_ids,
|
||||
# The framework packer expects a 5D [1, C, T, H, W] latent.
|
||||
vision_latent=vision_latent.unsqueeze(0),
|
||||
cond_frames=cond_frames,
|
||||
timestep=timestep_value,
|
||||
)
|
||||
out = vfm(**_framework_inputs(ps))
|
||||
preds = out.get("preds_vision")
|
||||
cond_mask = ps.vision.condition_mask[0] # [T] or [T,1,1]
|
||||
if preds is None:
|
||||
return torch.zeros_like(flat_latent)
|
||||
pred = preds[0].squeeze(0) # [C, T, H, W]
|
||||
keep = (1.0 - cond_mask.reshape(-1, 1, 1)).to(dtype=pred.dtype, device=pred.device)
|
||||
velocity = pred * keep if keep.sum() > 0 else torch.zeros_like(pred)
|
||||
return velocity.reshape(-1)
|
||||
|
||||
cond_v = _run(cond_ids)
|
||||
uncond_v = _run(uncond_ids)
|
||||
return uncond_v + guidance * (cond_v - uncond_v)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Cases: T2V (no cond), I2V (cond frame 0), single-frame T2I.
|
||||
# ---------------------------------------------------------------------------
|
||||
_CASES = [
|
||||
pytest.param(2, 4, 4, 6, [], id="t2v_2x2x2"),
|
||||
pytest.param(3, 8, 4, 5, [0], id="i2v_3x4x2_cond0"),
|
||||
pytest.param(1, 8, 8, 4, [], id="t2i_1x4x4"),
|
||||
]
|
||||
|
||||
|
||||
def _build_models(num_layers: int = 2, seed_model: int = 42):
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights(vfm, dit)
|
||||
return vfm, dit
|
||||
|
||||
|
||||
def _fastvideo_velocity(dit, *, flat_latent, timestep, guidance, vision_shape, cond_frames, cond_ids, uncond_ids):
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3VisionSpec,
|
||||
cosmos3_get_cfg_velocity,
|
||||
)
|
||||
|
||||
spec = Cosmos3VisionSpec(shape=vision_shape, condition_frame_indexes=list(cond_frames))
|
||||
return cosmos3_get_cfg_velocity(
|
||||
transformer=dit,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
specs=[spec],
|
||||
cond_token_ids=cond_ids,
|
||||
uncond_token_ids=uncond_ids,
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
class TestCosmos3DenoiseCFGParity:
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
|
||||
def test_cfg_velocity_matches_framework(self, grid_t, latent_h, latent_w, n_text, cond):
|
||||
vfm, dit = _build_models()
|
||||
torch.manual_seed(0)
|
||||
cond_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (max(1, n_text - 1),)).tolist()
|
||||
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
|
||||
timestep = torch.tensor([[500.0]]) # framework expects [1,1]; we reshape to scalar
|
||||
guidance = 6.0
|
||||
|
||||
fw_v = _framework_cfg_velocity(
|
||||
vfm=vfm,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=cond,
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fv_v = _fastvideo_velocity(
|
||||
dit,
|
||||
flat_latent=flat_latent,
|
||||
timestep=timestep,
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=cond,
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
assert fw_v.shape == fv_v.shape, f"shape: fw={fw_v.shape} fv={fv_v.shape}"
|
||||
max_abs = (fw_v - fv_v).abs().max().item()
|
||||
print(f"\n[cfg_velocity {grid_t}x{latent_h}x{latent_w} cond={cond}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_v, fw_v, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_one_unipc_step_matches_framework(self):
|
||||
"""CFG velocity + one UniPC step: FastVideo == framework math."""
|
||||
vfm, dit = _build_models()
|
||||
grid_t, latent_h, latent_w = 2, 4, 4
|
||||
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
torch.manual_seed(3)
|
||||
cond_ids = torch.randint(0, 60, (5,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (4,)).tolist()
|
||||
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
|
||||
guidance = 6.0
|
||||
|
||||
fw_sched = _framework_scheduler(4, _FLOW_SHIFT)
|
||||
fv_sched = _fastvideo_scheduler(4, _FLOW_SHIFT)
|
||||
t = fw_sched.timesteps[0]
|
||||
|
||||
fw_v = _framework_cfg_velocity(
|
||||
vfm=vfm,
|
||||
flat_latent=flat_latent,
|
||||
timestep=t.reshape(1, 1),
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=[],
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fw_stepped = fw_sched.step(model_output=fw_v, timestep=t, sample=flat_latent.unsqueeze(0),
|
||||
return_dict=False)[0].squeeze(0)
|
||||
|
||||
fv_v = _fastvideo_velocity(
|
||||
dit,
|
||||
flat_latent=flat_latent,
|
||||
timestep=t.reshape(1),
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=[],
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fv_stepped = fv_sched.step(model_output=fv_v, timestep=t, sample=flat_latent.unsqueeze(0),
|
||||
return_dict=False)[0].squeeze(0)
|
||||
|
||||
max_abs = (fw_stepped - fv_stepped).abs().max().item()
|
||||
print(f"\n[unipc_step] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_stepped, fw_stepped, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_full_denoise_loop_matches_framework(self):
|
||||
"""Cosmos3DenoiseEngine.denoise (>= 2 steps) == framework step-by-step."""
|
||||
from fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import (
|
||||
Cosmos3DenoiseEngine,
|
||||
Cosmos3VisionSpec,
|
||||
)
|
||||
|
||||
vfm, dit = _build_models()
|
||||
grid_t, latent_h, latent_w = 2, 4, 4
|
||||
vision_shape = (_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
torch.manual_seed(5)
|
||||
cond_ids = torch.randint(0, 60, (5,)).tolist()
|
||||
uncond_ids = torch.randint(0, 60, (4,)).tolist()
|
||||
flat_latent = torch.randn(int(torch.tensor(vision_shape).prod()))
|
||||
guidance = 6.0
|
||||
num_steps = 3
|
||||
|
||||
# Manual framework loop (oracle).
|
||||
fw_sched = _framework_scheduler(num_steps, _FLOW_SHIFT)
|
||||
fw_latent = flat_latent.clone()
|
||||
for t in fw_sched.timesteps:
|
||||
v = _framework_cfg_velocity(
|
||||
vfm=vfm,
|
||||
flat_latent=fw_latent,
|
||||
timestep=t.reshape(1, 1),
|
||||
guidance=guidance,
|
||||
vision_shape=vision_shape,
|
||||
cond_frames=[],
|
||||
cond_ids=cond_ids,
|
||||
uncond_ids=uncond_ids,
|
||||
)
|
||||
fw_latent = fw_sched.step(model_output=v, timestep=t, sample=fw_latent.unsqueeze(0),
|
||||
return_dict=False)[0].squeeze(0)
|
||||
|
||||
# FastVideo engine loop.
|
||||
fv_sched = _fastvideo_scheduler(num_steps, _FLOW_SHIFT)
|
||||
engine = Cosmos3DenoiseEngine(
|
||||
transformer=dit,
|
||||
scheduler=fv_sched,
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
spec = Cosmos3VisionSpec(shape=vision_shape, condition_frame_indexes=[])
|
||||
fv_latent = engine.denoise(
|
||||
flat_latent=flat_latent.clone(),
|
||||
timesteps=fv_sched.timesteps,
|
||||
guidance=guidance,
|
||||
specs=[spec],
|
||||
cond_token_ids=cond_ids,
|
||||
uncond_token_ids=uncond_ids,
|
||||
)
|
||||
|
||||
assert fv_latent.shape == fw_latent.shape
|
||||
max_abs = (fw_latent - fv_latent).abs().max().item()
|
||||
print(f"\n[full_denoise {num_steps} steps] final latent max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_latent, fw_latent, atol=1e-4, rtol=1e-3)
|
||||
@@ -0,0 +1,256 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 DiT vs official ``Cosmos3VFMNetwork``.
|
||||
|
||||
Builds a tiny official-framework ``Cosmos3VFMNetwork`` AND a tiny FastVideo
|
||||
``Cosmos3VFMTransformer`` from the SAME tiny config, copies the framework
|
||||
weights into the FastVideo DiT via an explicit framework->fastvideo name map,
|
||||
runs BOTH forwards on identical deterministic inputs (CPU / float32), and
|
||||
asserts ``torch.allclose`` on the vision prediction output (``preds_vision``)
|
||||
and the per-token ``last_hidden_state``.
|
||||
|
||||
The official model is the parity ORACLE. It runs on CPU / float32 via the SDPA
|
||||
monkey-patch in ``test_cosmos3_reference_forward`` (flash2/flash3/natten are
|
||||
CUDA-only). The FastVideo DiT runs natively on CPU with plain SDPA.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_dit_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
# Reuse the reference harness's tiny-model builder + SDPA monkey-patch.
|
||||
from .test_cosmos3_reference_forward import ( # noqa: E402
|
||||
_apply_sdpa_patches,
|
||||
_build_tiny_cosmos3,
|
||||
_build_tiny_packed_seq,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
|
||||
_apply_sdpa_patches()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tiny config shared by both models (must match _build_tiny_cosmos3).
|
||||
# ---------------------------------------------------------------------------
|
||||
def _build_tiny_fastvideo_dit() -> "Cosmos3VFMTransformer": # noqa: F821
|
||||
from fastvideo.configs.models.dits.cosmos3 import (
|
||||
Cosmos3ArchConfig,
|
||||
Cosmos3VideoConfig,
|
||||
)
|
||||
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
|
||||
|
||||
arch = Cosmos3ArchConfig(
|
||||
hidden_size=16,
|
||||
num_hidden_layers=1,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=2,
|
||||
head_dim=8,
|
||||
intermediate_size=32,
|
||||
vocab_size=64,
|
||||
rms_norm_eps=1e-6,
|
||||
attention_bias=False,
|
||||
latent_patch_size=2,
|
||||
latent_channel=16,
|
||||
rope_theta=5_000_000.0,
|
||||
mrope_section=[24, 20, 20],
|
||||
position_embedding_type="3d_rope",
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=4,
|
||||
enable_fps_modulation=False,
|
||||
# Dormant heads present in the checkpoint surface (constructed for
|
||||
# strict-load parity; not exercised by this video-path forward).
|
||||
action_gen=True,
|
||||
action_dim=64,
|
||||
max_action_dim=64,
|
||||
num_embodiment_domains=32,
|
||||
sound_gen=True,
|
||||
sound_dim=64,
|
||||
)
|
||||
cfg = Cosmos3VideoConfig(arch_config=arch)
|
||||
model = Cosmos3VFMTransformer(cfg, hf_config={})
|
||||
return model.to(torch.float32).eval()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Framework -> FastVideo weight name map.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _framework_to_fastvideo_state_dict(vfm, num_layers: int) -> dict[str, torch.Tensor]:
|
||||
"""Translate framework param names into the FastVideo DiT param names.
|
||||
|
||||
Framework (Cosmos3VFMNetwork):
|
||||
language_model.model.{embed_tokens,norm,norm_moe_gen}
|
||||
language_model.lm_head
|
||||
language_model.model.layers.{i}.self_attn.{q,k,v,o}_proj(+ _moe_gen)
|
||||
language_model.model.layers.{i}.self_attn.{q,k}_norm(+ _moe_gen)
|
||||
language_model.model.layers.{i}.{mlp,mlp_moe_gen}.{gate,up,down}_proj
|
||||
language_model.model.layers.{i}.{input,post_attention}_layernorm(+ _moe_gen)
|
||||
vae2llm / llm2vae / time_embedder.mlp.{0,2}
|
||||
|
||||
FastVideo (Cosmos3VFMTransformer):
|
||||
embed_tokens / norm / norm_moe_gen / lm_head
|
||||
layers.{i}.self_attn.{to_q,to_k,to_v,to_out} (und)
|
||||
layers.{i}.self_attn.{add_q,add_k,add_v}_proj / to_add_out (gen)
|
||||
layers.{i}.self_attn.{norm_q,norm_k,norm_added_q,norm_added_k}
|
||||
layers.{i}.{mlp,mlp_moe_gen}.{gate,up,down}_proj
|
||||
layers.{i}.{input,post_attention}_layernorm(+ _moe_gen)
|
||||
proj_in / proj_out / time_embedder.linear_{1,2}
|
||||
"""
|
||||
src = dict(vfm.named_parameters())
|
||||
out: dict[str, torch.Tensor] = {}
|
||||
|
||||
def take(name: str) -> torch.Tensor:
|
||||
return src[name].detach().clone()
|
||||
|
||||
# ---- Top-level backbone ----
|
||||
out["embed_tokens.weight"] = take("language_model.model.embed_tokens.weight")
|
||||
out["norm.weight"] = take("language_model.model.norm.weight")
|
||||
out["norm_moe_gen.weight"] = take("language_model.model.norm_moe_gen.weight")
|
||||
out["lm_head.weight"] = take("language_model.lm_head.weight")
|
||||
|
||||
# ---- Vision adapters ----
|
||||
out["proj_in.weight"] = take("vae2llm.weight")
|
||||
out["proj_in.bias"] = take("vae2llm.bias")
|
||||
out["proj_out.weight"] = take("llm2vae.weight")
|
||||
out["proj_out.bias"] = take("llm2vae.bias")
|
||||
|
||||
# ---- Timestep embedder (mlp.0/mlp.2 -> linear_1/linear_2) ----
|
||||
out["time_embedder.linear_1.weight"] = take("time_embedder.mlp.0.weight")
|
||||
out["time_embedder.linear_1.bias"] = take("time_embedder.mlp.0.bias")
|
||||
out["time_embedder.linear_2.weight"] = take("time_embedder.mlp.2.weight")
|
||||
out["time_embedder.linear_2.bias"] = take("time_embedder.mlp.2.bias")
|
||||
|
||||
# ---- Per layer ----
|
||||
und_attn = {"q_proj": "to_q", "k_proj": "to_k", "v_proj": "to_v", "o_proj": "to_out"}
|
||||
gen_attn = {
|
||||
"q_proj_moe_gen": "add_q_proj",
|
||||
"k_proj_moe_gen": "add_k_proj",
|
||||
"v_proj_moe_gen": "add_v_proj",
|
||||
"o_proj_moe_gen": "to_add_out",
|
||||
}
|
||||
und_norm = {"q_norm": "norm_q", "k_norm": "norm_k"}
|
||||
gen_norm = {"q_norm_moe_gen": "norm_added_q", "k_norm_moe_gen": "norm_added_k"}
|
||||
|
||||
for i in range(num_layers):
|
||||
fw = f"language_model.model.layers.{i}"
|
||||
fv = f"layers.{i}"
|
||||
for s, d in und_attn.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for s, d in gen_attn.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for s, d in und_norm.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for s, d in gen_norm.items():
|
||||
out[f"{fv}.self_attn.{d}.weight"] = take(f"{fw}.self_attn.{s}.weight")
|
||||
for mlp in ("mlp", "mlp_moe_gen"):
|
||||
for proj in ("gate_proj", "up_proj", "down_proj"):
|
||||
out[f"{fv}.{mlp}.{proj}.weight"] = take(f"{fw}.{mlp}.{proj}.weight")
|
||||
for ln in ("input_layernorm", "input_layernorm_moe_gen", "post_attention_layernorm",
|
||||
"post_attention_layernorm_moe_gen"):
|
||||
out[f"{fv}.{ln}.weight"] = take(f"{fw}.{ln}.weight")
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def _copy_weights(vfm, dit) -> None:
|
||||
"""Copy framework weights into the FastVideo DiT (shape-checked)."""
|
||||
mapped = _framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers)
|
||||
dst = dict(dit.named_parameters())
|
||||
# Every mapped tensor must land on an existing FastVideo param with a matching shape.
|
||||
for name, tensor in mapped.items():
|
||||
assert name in dst, f"FastVideo DiT missing param for mapped key {name!r}"
|
||||
assert dst[name].shape == tensor.shape, (f"shape mismatch for {name}: "
|
||||
f"dit={tuple(dst[name].shape)} fw={tuple(tensor.shape)}")
|
||||
with torch.no_grad():
|
||||
for name, tensor in mapped.items():
|
||||
dst[name].copy_(tensor.to(dst[name].dtype))
|
||||
|
||||
|
||||
def _fastvideo_inputs_from_packed_seq(ps) -> dict:
|
||||
"""Build the FastVideo DiT forward kwargs from a framework PackedSequence."""
|
||||
v = ps.vision
|
||||
return dict(
|
||||
text_ids=ps.text_ids,
|
||||
text_indexes=ps.text_indexes,
|
||||
position_ids=ps.position_ids,
|
||||
sequence_length=int(ps.sequence_length),
|
||||
split_lens=list(ps.split_lens),
|
||||
attn_modes=list(ps.attn_modes),
|
||||
vision_tokens=list(v.tokens),
|
||||
vision_token_shapes=list(v.token_shapes),
|
||||
vision_sequence_indexes=v.sequence_indexes,
|
||||
vision_timesteps=v.timesteps,
|
||||
vision_mse_loss_indexes=v.mse_loss_indexes,
|
||||
vision_noisy_frame_indexes=list(v.noisy_frame_indexes),
|
||||
fps_vision=None,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
class TestCosmos3DiTParity:
|
||||
|
||||
def _run_both(self, seed_model: int = 42, seed_data: int = 7):
|
||||
vfm = _build_tiny_cosmos3(seed=seed_model)
|
||||
dit = _build_tiny_fastvideo_dit()
|
||||
_copy_weights(vfm, dit)
|
||||
ps = _build_tiny_packed_seq(n_text=4, seed=seed_data)
|
||||
|
||||
with torch.no_grad():
|
||||
fw_out = vfm(packed_seq=ps)
|
||||
fv_out = dit(**_fastvideo_inputs_from_packed_seq(ps))
|
||||
return fw_out, fv_out
|
||||
|
||||
def test_weight_map_is_complete(self):
|
||||
"""The framework->fastvideo map must cover EVERY FastVideo DiT parameter
|
||||
that is exercised by the video path (i.e. all non-dormant params).
|
||||
|
||||
Dormant action/audio heads have no framework counterpart in this tiny
|
||||
vision-only setup, so they are excluded from the copy; everything else
|
||||
must be covered.
|
||||
"""
|
||||
vfm = _build_tiny_cosmos3(seed=42)
|
||||
dit = _build_tiny_fastvideo_dit()
|
||||
mapped = set(_framework_to_fastvideo_state_dict(vfm, num_layers=dit.num_hidden_layers))
|
||||
dit_params = set(n for n, _ in dit.named_parameters())
|
||||
dormant = {
|
||||
n
|
||||
for n in dit_params
|
||||
if n.startswith(("action_", "audio_"))
|
||||
}
|
||||
uncovered = dit_params - mapped - dormant
|
||||
assert not uncovered, f"FastVideo DiT params not covered by weight map: {sorted(uncovered)}"
|
||||
|
||||
def test_preds_vision_parity(self):
|
||||
fw_out, fv_out = self._run_both()
|
||||
fw_pv = fw_out["preds_vision"][0] # [1, C, T, H, W]
|
||||
fv_pv = fv_out["preds_vision"][0]
|
||||
assert fw_pv.shape == fv_pv.shape, f"shape mismatch: fw={fw_pv.shape} fv={fv_pv.shape}"
|
||||
max_abs = (fw_pv - fv_pv).abs().max().item()
|
||||
print(f"\n[preds_vision] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_pv, fw_pv, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_last_hidden_state_parity(self):
|
||||
fw_out, fv_out = self._run_both()
|
||||
fw_lhs = fw_out["last_hidden_state"] # [N, hidden]
|
||||
fv_lhs = fv_out["last_hidden_state"]
|
||||
assert fw_lhs.shape == fv_lhs.shape, f"shape mismatch: fw={fw_lhs.shape} fv={fv_lhs.shape}"
|
||||
max_abs = (fw_lhs - fv_lhs).abs().max().item()
|
||||
print(f"\n[last_hidden_state] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_lhs, fw_lhs, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_parity_holds_across_seeds(self):
|
||||
"""Re-running with a different random init still matches (not a fluke)."""
|
||||
fw_out, fv_out = self._run_both(seed_model=99, seed_data=13)
|
||||
torch.testing.assert_close(fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3)
|
||||
@@ -0,0 +1,371 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 DiT vs ``Cosmos3VFMNetwork`` (mRoPE).
|
||||
|
||||
Companion to ``test_cosmos3_dit_parity.py`` (which covers ``3d_rope``). This
|
||||
module exercises the rotary mode the REAL ``nvidia/Cosmos3-Nano`` checkpoint
|
||||
uses: ``position_embedding_type="unified_3d_mrope"`` with the real-checkpoint
|
||||
settings (``mrope_section=[24,20,20]``, ``mrope_interleaved=True``,
|
||||
``rope_theta=5e6``, ``unified_3d_mrope_reset_spatial_ids=True``,
|
||||
``temporal_modality_margin=15000``).
|
||||
|
||||
Under unified 3D mRoPE there is NO additive latent position embedding
|
||||
(``latent_pos_embed is None``); all positional information rides on the
|
||||
per-token 3D (T, H, W) rotary embedding. The packed-sequence ``position_ids``
|
||||
are therefore shape ``[3, seq_len]``, built exactly like the framework data
|
||||
packer (``cosmos_framework.data.vfm.sequence_packing``):
|
||||
|
||||
* text tokens broadcast one monotone id across all three axes
|
||||
(``get_3d_mrope_ids_text_tokens``),
|
||||
* the temporal offset is bumped by ``temporal_modality_margin`` at the
|
||||
text->vision boundary,
|
||||
* vision tokens lay out a (T, H, W) grid with spatial ids reset per segment
|
||||
(``get_3d_mrope_ids_vae_tokens`` with ``reset_spatial_indices=True``).
|
||||
|
||||
Both models are built tiny from the SAME config, framework weights are copied
|
||||
into the FastVideo DiT (reusing the ``3d_rope`` test's weight map — the
|
||||
transformer key surface is identical across rotary modes), and BOTH forwards
|
||||
run on identical deterministic CPU / float32 inputs. The official model is the
|
||||
parity ORACLE (run on CPU via the SDPA monkey-patch in
|
||||
``test_cosmos3_reference_forward``).
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_dit_parity_mrope.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
# Reuse the reference harness's SDPA monkey-patch and the 3d_rope parity
|
||||
# test's weight-copy + input-builder helpers (key surface is rotary-agnostic).
|
||||
from .test_cosmos3_dit_parity import ( # noqa: E402
|
||||
_copy_weights,
|
||||
_fastvideo_inputs_from_packed_seq,
|
||||
_framework_to_fastvideo_state_dict,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
|
||||
_apply_sdpa_patches()
|
||||
|
||||
# Real-checkpoint unified_3d_mrope settings (tiny model, real rope constants).
|
||||
_ROPE_THETA = 5_000_000.0
|
||||
_MROPE_SECTION = [24, 20, 20]
|
||||
_MROPE_INTERLEAVED = True
|
||||
_RESET_SPATIAL_IDS = True
|
||||
_TEMPORAL_MODALITY_MARGIN = 15_000
|
||||
_LATENT_PATCH_SIZE = 2
|
||||
_LATENT_CHANNEL = 16
|
||||
_TCF = 4 # temporal compression factor
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tiny model builders (framework + FastVideo) with unified_3d_mrope.
|
||||
# ---------------------------------------------------------------------------
|
||||
_SOUND_DIM = 64
|
||||
_SOUND_LATENT_FPS = 25
|
||||
_ACTION_DIM = 64
|
||||
_NUM_EMBODIMENT_DOMAINS = 32
|
||||
|
||||
|
||||
def _build_tiny_cosmos3_mrope(seed: int = 42, num_layers: int = 2, sound_gen: bool = False,
|
||||
action_gen: bool = False):
|
||||
"""Tiny framework ``Cosmos3VFMNetwork`` with ``unified_3d_mrope``.
|
||||
|
||||
``rope_theta`` / ``rope_scaling`` (carrying ``mrope_section`` +
|
||||
``mrope_interleaved``) are threaded through the materialized text config;
|
||||
``position_embedding_type="unified_3d_mrope"`` leaves ``latent_pos_embed``
|
||||
as ``None`` so positions ride solely on the 3D rotary embedding.
|
||||
|
||||
``sound_gen=True`` additionally builds the sound MoT heads (``sound2llm`` /
|
||||
``llm2sound`` / ``sound_modality_embed``) for the t2vs parity test.
|
||||
"""
|
||||
from cosmos_framework.model.vfm.mot.cosmos3_vfm_network import (
|
||||
Cosmos3VFMNetwork,
|
||||
Cosmos3VFMNetworkConfig,
|
||||
)
|
||||
from cosmos_framework.model.vfm.mot.unified_mot import (
|
||||
Qwen3MoTConfig,
|
||||
Qwen3VLTextForCausalLM,
|
||||
)
|
||||
from cosmos_framework.model.vfm.vlm.qwen3_vl.configuration_qwen3_vl import Qwen3VLConfig
|
||||
|
||||
tiny_text_dict = dict(
|
||||
model_type="qwen3_vl_text",
|
||||
vocab_size=64,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=num_layers,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=2,
|
||||
head_dim=8,
|
||||
rms_norm_eps=1e-6,
|
||||
attention_bias=False,
|
||||
attention_dropout=0.0,
|
||||
rope_theta=_ROPE_THETA,
|
||||
rope_scaling={
|
||||
"rope_type": "default",
|
||||
"mrope_section": _MROPE_SECTION,
|
||||
"mrope_interleaved": _MROPE_INTERLEAVED,
|
||||
},
|
||||
max_position_embeddings=262144,
|
||||
)
|
||||
mot_cfg = Qwen3MoTConfig(
|
||||
config_dict=tiny_text_dict,
|
||||
qk_norm_for_text=True,
|
||||
qk_norm_for_diffusion=True,
|
||||
include_visual=False,
|
||||
)
|
||||
tiny_vlm_cfg = Qwen3VLConfig(text_config=tiny_text_dict)
|
||||
sound_kwargs = dict(
|
||||
sound_gen=True,
|
||||
sound_dim=_SOUND_DIM,
|
||||
temporal_compression_factor_sound=1,
|
||||
sound_latent_fps=_SOUND_LATENT_FPS,
|
||||
) if sound_gen else {}
|
||||
action_kwargs = dict(
|
||||
action_gen=True,
|
||||
action_dim=_ACTION_DIM,
|
||||
num_embodiment_domains=_NUM_EMBODIMENT_DOMAINS,
|
||||
) if action_gen else {}
|
||||
vfm_cfg = Cosmos3VFMNetworkConfig(
|
||||
vision_gen=True,
|
||||
vlm_config=tiny_vlm_cfg,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
latent_downsample_factor=8,
|
||||
latent_channel_size=_LATENT_CHANNEL,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
max_latent_h=16,
|
||||
max_latent_w=16,
|
||||
max_latent_t=8,
|
||||
temporal_compression_factor_vision=_TCF,
|
||||
**sound_kwargs,
|
||||
**action_kwargs,
|
||||
)
|
||||
torch.manual_seed(seed)
|
||||
lm = Qwen3VLTextForCausalLM(config=mot_cfg)
|
||||
vfm = Cosmos3VFMNetwork(language_model=lm, config=vfm_cfg)
|
||||
# inv_freq is a non-persistent buffer; init it on CPU (mirrors from_pretrained).
|
||||
vfm.language_model.model.rotary_emb.init_weights(buffer_device=None)
|
||||
vfm.eval()
|
||||
return vfm
|
||||
|
||||
|
||||
def _build_tiny_fastvideo_dit_mrope(num_layers: int = 2):
|
||||
from fastvideo.configs.models.dits.cosmos3 import (
|
||||
Cosmos3ArchConfig,
|
||||
Cosmos3VideoConfig,
|
||||
)
|
||||
from fastvideo.models.dits.cosmos3 import Cosmos3VFMTransformer
|
||||
|
||||
arch = Cosmos3ArchConfig(
|
||||
hidden_size=16,
|
||||
num_hidden_layers=num_layers,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=2,
|
||||
head_dim=8,
|
||||
intermediate_size=32,
|
||||
vocab_size=64,
|
||||
rms_norm_eps=1e-6,
|
||||
attention_bias=False,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
latent_channel=_LATENT_CHANNEL,
|
||||
rope_theta=_ROPE_THETA,
|
||||
mrope_section=_MROPE_SECTION,
|
||||
mrope_interleaved=_MROPE_INTERLEAVED,
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
enable_fps_modulation=False,
|
||||
# Dormant heads present in the checkpoint surface (constructed for
|
||||
# strict-load parity; not exercised by this video-path forward).
|
||||
action_gen=True,
|
||||
action_dim=64,
|
||||
max_action_dim=64,
|
||||
num_embodiment_domains=32,
|
||||
sound_gen=True,
|
||||
sound_dim=64,
|
||||
)
|
||||
cfg = Cosmos3VideoConfig(arch_config=arch)
|
||||
model = Cosmos3VFMTransformer(cfg, hf_config={})
|
||||
return model.to(torch.float32).eval()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# [3, seq_len] mRoPE position-id builder (mirrors the framework data packer).
|
||||
# ---------------------------------------------------------------------------
|
||||
def _build_mrope_position_ids(n_text: int, grid_t: int, patch_h: int, patch_w: int) -> torch.Tensor:
|
||||
"""Build ``[3, seq_len]`` (T, H, W) mRoPE ids for one text+vision sample.
|
||||
|
||||
Reproduces ``pack_input_sequence`` for a single causal-text + full-vision
|
||||
sample: monotone text ids on all axes, ``+temporal_modality_margin`` at the
|
||||
text->vision boundary, then a reset-spatial (T, H, W) vision grid.
|
||||
"""
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
get_3d_mrope_ids_text_tokens,
|
||||
get_3d_mrope_ids_vae_tokens,
|
||||
)
|
||||
|
||||
offset: int | float = 0
|
||||
text_ids, offset = get_3d_mrope_ids_text_tokens(num_tokens=n_text, temporal_offset=offset)
|
||||
# End of text modality: add the boundary margin before vision.
|
||||
offset += _TEMPORAL_MODALITY_MARGIN
|
||||
vision_ids, offset = get_3d_mrope_ids_vae_tokens(
|
||||
grid_t=grid_t,
|
||||
grid_h=patch_h,
|
||||
grid_w=patch_w,
|
||||
temporal_offset=offset,
|
||||
reset_spatial_indices=_RESET_SPATIAL_IDS,
|
||||
fps=None, # integer positions (enable_fps_modulation=False)
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
return torch.cat([text_ids, vision_ids], dim=1) # [3, seq_len]
|
||||
|
||||
|
||||
def _build_tiny_packed_seq_mrope(
|
||||
*,
|
||||
n_text: int = 6,
|
||||
grid_t: int = 2,
|
||||
latent_h: int = 4,
|
||||
latent_w: int = 4,
|
||||
seed: int = 7,
|
||||
):
|
||||
"""Minimal PackedSequence with ``[3, seq]`` mRoPE position ids.
|
||||
|
||||
Vision latent ``[C, grid_t, latent_h, latent_w]`` patchifies (patch=2) to a
|
||||
``(grid_t, latent_h/2, latent_w/2)`` token grid; all frames are noisy.
|
||||
"""
|
||||
from cosmos_framework.data.vfm.sequence_packing import ModalityData, PackedSequence
|
||||
|
||||
patch_h = latent_h // _LATENT_PATCH_SIZE
|
||||
patch_w = latent_w // _LATENT_PATCH_SIZE
|
||||
n_vision = grid_t * patch_h * patch_w
|
||||
total_len = n_text + n_vision
|
||||
|
||||
torch.manual_seed(seed)
|
||||
vision_tensor = torch.randn(_LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
text_ids = torch.randint(0, 64, (n_text,))
|
||||
position_ids = _build_mrope_position_ids(n_text, grid_t, patch_h, patch_w) # [3, total_len]
|
||||
|
||||
noisy_frame_indexes = torch.arange(grid_t, dtype=torch.long) # all frames noisy
|
||||
vision_mod = ModalityData(
|
||||
sequence_indexes=torch.arange(n_text, total_len, dtype=torch.long),
|
||||
timesteps=torch.full((n_vision,), 500.0),
|
||||
mse_loss_indexes=torch.arange(n_text, total_len, dtype=torch.long),
|
||||
token_shapes=[(grid_t, patch_h, patch_w)],
|
||||
tokens=[vision_tensor],
|
||||
condition_mask=[torch.zeros(grid_t, dtype=torch.long)], # 0 = noisy
|
||||
noisy_frame_indexes=[noisy_frame_indexes],
|
||||
)
|
||||
packed_seq = PackedSequence(
|
||||
sample_lens=[total_len],
|
||||
split_lens=[n_text, n_vision],
|
||||
attn_modes=["causal", "full"],
|
||||
is_image_batch=(grid_t == 1),
|
||||
sequence_length=total_len,
|
||||
text_ids=text_ids,
|
||||
text_indexes=torch.arange(n_text, dtype=torch.long),
|
||||
position_ids=position_ids,
|
||||
vision=vision_mod,
|
||||
)
|
||||
return packed_seq
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
# (grid_t, latent_h, latent_w): a single image, a small video, and a taller
|
||||
# video, to exercise the spatial mRoPE overwrite + gen<->gen full attention.
|
||||
_GRIDS = [
|
||||
pytest.param(1, 8, 8, id="image_1x4x4"),
|
||||
pytest.param(2, 4, 4, id="video_2x2x2"),
|
||||
pytest.param(3, 8, 4, id="video_3x4x2"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3DiTParityMRoPE:
|
||||
|
||||
def _run_both(
|
||||
self,
|
||||
*,
|
||||
grid_t: int,
|
||||
latent_h: int,
|
||||
latent_w: int,
|
||||
seed_model: int = 42,
|
||||
seed_data: int = 7,
|
||||
num_layers: int = 2,
|
||||
):
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=seed_model, num_layers=num_layers)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights(vfm, dit)
|
||||
ps = _build_tiny_packed_seq_mrope(
|
||||
n_text=6, grid_t=grid_t, latent_h=latent_h, latent_w=latent_w, seed=seed_data
|
||||
)
|
||||
with torch.no_grad():
|
||||
fw_out = vfm(packed_seq=ps)
|
||||
fv_out = dit(**_fastvideo_inputs_from_packed_seq(ps))
|
||||
return fw_out, fv_out
|
||||
|
||||
def test_position_ids_are_3xN_mrope(self):
|
||||
"""The packed mRoPE ids are ``[3, seq_len]`` with the text->vision margin."""
|
||||
ps = _build_tiny_packed_seq_mrope(n_text=6, grid_t=2, latent_h=4, latent_w=4)
|
||||
pos = ps.position_ids
|
||||
assert pos.ndim == 2 and pos.shape[0] == 3, f"expected [3, N], got {tuple(pos.shape)}"
|
||||
assert pos.shape[1] == int(ps.sequence_length)
|
||||
# Text axis is monotone 0..5 on all 3 rows; vision temporal jumps by the margin.
|
||||
assert pos[0, :6].tolist() == [0, 1, 2, 3, 4, 5]
|
||||
assert pos[1, :6].tolist() == [0, 1, 2, 3, 4, 5]
|
||||
assert pos[2, :6].tolist() == [0, 1, 2, 3, 4, 5]
|
||||
# First vision token temporal id == last_text_id (5) + margin + 1.
|
||||
assert pos[0, 6].item() == 5 + _TEMPORAL_MODALITY_MARGIN + 1
|
||||
# Reset spatial: first vision token H/W ids are 0.
|
||||
assert pos[1, 6].item() == 0 and pos[2, 6].item() == 0
|
||||
|
||||
def test_no_additive_latent_pos_embed(self):
|
||||
"""unified_3d_mrope must NOT build an additive latent position embedding."""
|
||||
dit = _build_tiny_fastvideo_dit_mrope()
|
||||
assert dit.position_embedding_type == "unified_3d_mrope"
|
||||
assert dit.latent_pos_embed is None
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w"), _GRIDS)
|
||||
def test_preds_vision_parity(self, grid_t, latent_h, latent_w):
|
||||
fw_out, fv_out = self._run_both(grid_t=grid_t, latent_h=latent_h, latent_w=latent_w)
|
||||
fw_pv = fw_out["preds_vision"][0] # [1, C, T, H, W]
|
||||
fv_pv = fv_out["preds_vision"][0]
|
||||
assert fw_pv.shape == fv_pv.shape, f"shape mismatch: fw={fw_pv.shape} fv={fv_pv.shape}"
|
||||
max_abs = (fw_pv - fv_pv).abs().max().item()
|
||||
print(f"\n[preds_vision mrope {grid_t}x{latent_h}x{latent_w}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_pv, fw_pv, atol=1e-4, rtol=1e-3)
|
||||
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w"), _GRIDS)
|
||||
def test_last_hidden_state_parity(self, grid_t, latent_h, latent_w):
|
||||
fw_out, fv_out = self._run_both(grid_t=grid_t, latent_h=latent_h, latent_w=latent_w)
|
||||
fw_lhs = fw_out["last_hidden_state"] # [N, hidden]
|
||||
fv_lhs = fv_out["last_hidden_state"]
|
||||
assert fw_lhs.shape == fv_lhs.shape, f"shape mismatch: fw={fw_lhs.shape} fv={fv_lhs.shape}"
|
||||
max_abs = (fw_lhs - fv_lhs).abs().max().item()
|
||||
print(f"\n[last_hidden_state mrope {grid_t}x{latent_h}x{latent_w}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(fv_lhs, fw_lhs, atol=1e-4, rtol=1e-3)
|
||||
|
||||
def test_parity_holds_across_seeds(self):
|
||||
"""A different random init still matches bit-for-bit (not a fluke)."""
|
||||
fw_out, fv_out = self._run_both(
|
||||
grid_t=2, latent_h=4, latent_w=4, seed_model=99, seed_data=13
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
fv_out["preds_vision"][0], fw_out["preds_vision"][0], atol=1e-4, rtol=1e-3
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
fv_out["last_hidden_state"], fw_out["last_hidden_state"], atol=1e-4, rtol=1e-3
|
||||
)
|
||||
@@ -0,0 +1,80 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 UniPC flow_shift vs the framework.
|
||||
|
||||
The framework selects the UniPC ``shift`` purely from the named resolution
|
||||
bucket the (H, W) belongs to, via ``OmniSampleArgs._RESOLUTION_SHIFT_DEFAULTS``
|
||||
(keyed by the VLM model size — Cosmos3-Nano uses the 8B backbone — and the
|
||||
resolution string), NOT from the task (T2V/I2V/T2I share a shift at a given
|
||||
resolution). FastVideo gets raw pixel ``height``/``width`` and must map back to
|
||||
the same shift.
|
||||
|
||||
This pins ``Cosmos3DenoisingStage._flow_shift_for_resolution`` against the
|
||||
framework's own tables: for every (resolution, aspect) entry in
|
||||
``VIDEO_RES_SIZE_INFO`` whose resolution has an 8B shift default, the FastVideo
|
||||
shift for that exact pixel size must equal the framework default.
|
||||
|
||||
The framework tables are the parity ORACLE.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_flow_shift_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
# The official framework provides the parity oracle for the resolution->pixel
|
||||
# tables. (``cosmos_framework.inference.args`` — which holds the shift constant —
|
||||
# can't be imported here: it transitively requires ``multistorageclient``. The
|
||||
# small shift table is mirrored verbatim below with its source location.)
|
||||
_utils = pytest.importorskip(
|
||||
"cosmos_framework.data.vfm.utils",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import ( # noqa: E402
|
||||
Cosmos3DenoisingStage,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Cosmos3-Nano's VLM backbone is Qwen3-VL-8B (checkpoint config.json).
|
||||
_MODEL_SIZE = "8B"
|
||||
# Verbatim from cosmos_framework.inference.args.OmniSampleArgs
|
||||
# ._RESOLUTION_SHIFT_DEFAULTS (args.py:770), restricted to the 8B rows.
|
||||
_SHIFT_DEFAULTS = {
|
||||
("8B", "256"): 3.0,
|
||||
("8B", "480"): 5.0,
|
||||
("8B", "720"): 10.0,
|
||||
("8B", "768"): 10.0,
|
||||
("32B", "256"): 5.0,
|
||||
("32B", "480"): 5.0,
|
||||
("32B", "720"): 5.0,
|
||||
("32B", "768"): 5.0,
|
||||
}
|
||||
_VIDEO_RES = _utils.VIDEO_RES_SIZE_INFO
|
||||
_IMAGE_RES = _utils.IMAGE_RES_SIZE_INFO
|
||||
|
||||
|
||||
def _cases():
|
||||
seen = set()
|
||||
for resolution, by_aspect in {**_VIDEO_RES, **_IMAGE_RES}.items():
|
||||
key = (_MODEL_SIZE, resolution)
|
||||
if key not in _SHIFT_DEFAULTS:
|
||||
continue
|
||||
expected = _SHIFT_DEFAULTS[key]
|
||||
for aspect, (a, b) in by_aspect.items():
|
||||
cid = f"{resolution}_{aspect.replace(',', '-')}_{a}x{b}"
|
||||
if cid in seen:
|
||||
continue
|
||||
seen.add(cid)
|
||||
yield pytest.param(a, b, expected, id=cid)
|
||||
|
||||
|
||||
class TestCosmos3FlowShiftParity:
|
||||
|
||||
@pytest.mark.parametrize(("dim_a", "dim_b", "expected_shift"), list(_cases()))
|
||||
def test_flow_shift_matches_framework(self, dim_a, dim_b, expected_shift):
|
||||
got = Cosmos3DenoisingStage._flow_shift_for_resolution(dim_a, dim_b)
|
||||
assert got == expected_shift, (
|
||||
f"shift for {dim_a}x{dim_b}: got {got}, framework default {expected_shift}")
|
||||
@@ -0,0 +1,91 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo I2V conditioning pixel video vs the framework.
|
||||
|
||||
The Cosmos3 I2V path conditions on a *static repeat* of the input image. The
|
||||
framework (``cosmos_framework.inference.vision``):
|
||||
|
||||
* ``load_conditioning_image``: aspect-preserving resize + center crop + uint8
|
||||
quantization, then ``/127.5 - 1`` -> ``[3, 1, h, w]`` in [-1, 1];
|
||||
* ``build_conditioned_video_batch``: frame 0 = the image, and every remaining
|
||||
frame **repeats the last conditioning frame** (a static video) -> the clip
|
||||
is then VAE-encoded and only the latent condition frame(s) are kept clean.
|
||||
|
||||
Because the VAE is temporal, zero-filling the non-condition frames (the earlier
|
||||
FastVideo behavior) changes the encoded condition latent, so the repeat-fill is
|
||||
correctness-critical. This pins FastVideo's
|
||||
``Cosmos3DenoisingStage._image_to_video_tensor`` against the framework's
|
||||
image preprocessing + repeat-fill.
|
||||
|
||||
CPU / float32. The framework is the parity ORACLE.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_i2v_conditioning_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
vision = pytest.importorskip(
|
||||
"cosmos_framework.inference.vision",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
from fastvideo.pipelines.stages.cosmos3_stages import ( # noqa: E402
|
||||
Cosmos3DenoisingStage,
|
||||
)
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
|
||||
def _make_image(path, h_in: int, w_in: int, seed: int = 0) -> None:
|
||||
rng = np.random.default_rng(seed)
|
||||
arr = rng.integers(0, 256, size=(h_in, w_in, 3), dtype=np.uint8)
|
||||
Image.fromarray(arr, "RGB").save(path)
|
||||
|
||||
|
||||
# (input H, input W, target H, target W, num_frames)
|
||||
_CASES = [
|
||||
pytest.param(120, 200, 256, 256, 9, id="square_from_landscape"),
|
||||
pytest.param(200, 120, 704, 1280, 13, id="wide_from_portrait"),
|
||||
pytest.param(256, 256, 256, 256, 5, id="same_size"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3I2VConditioningParity:
|
||||
|
||||
@pytest.mark.parametrize(("h_in", "w_in", "h", "w", "num_frames"), _CASES)
|
||||
def test_conditioning_video_matches_framework(self, tmp_path, h_in, w_in, h, w, num_frames):
|
||||
img_path = tmp_path / "cond.png"
|
||||
_make_image(img_path, h_in, w_in)
|
||||
|
||||
# ---- Framework oracle ----
|
||||
# load_conditioning_image -> [3, 1, h, w] in [-1, 1].
|
||||
cond = vision.load_conditioning_image(img_path, target_h=h, target_w=w).float()
|
||||
# Mirror build_conditioned_video_batch (vision.py lines 117-123) in fp32/CPU:
|
||||
# frame 0 = image; remaining frames repeat the last conditioning frame.
|
||||
t_cond = cond.shape[1]
|
||||
expected = torch.zeros(1, 3, num_frames, h, w, dtype=torch.float32)
|
||||
t_fill = min(t_cond, num_frames)
|
||||
expected[0, :, :t_fill] = cond[:, :t_fill]
|
||||
if t_fill < num_frames:
|
||||
expected[0, :, t_fill:] = expected[0, :, t_fill - 1:t_fill].expand(-1, num_frames - t_fill, -1, -1)
|
||||
|
||||
# ---- FastVideo: same PIL image through the stage helper ----
|
||||
pil = Image.open(img_path).convert("RGB")
|
||||
got = Cosmos3DenoisingStage._image_to_video_tensor(
|
||||
pil, num_frames, h, w, torch.device("cpu"), torch.float32)
|
||||
|
||||
assert got.shape == expected.shape, f"shape: got={got.shape} expected={expected.shape}"
|
||||
max_abs = (got - expected).abs().max().item()
|
||||
print(f"\n[i2v_cond {h}x{w} nf={num_frames}] max abs diff = {max_abs:.3e}")
|
||||
torch.testing.assert_close(got, expected)
|
||||
|
||||
# Static repeat (not zero-fill): every frame equals frame 0, and the
|
||||
# frames past frame 0 are non-zero.
|
||||
assert torch.equal(got[0, :, 0], got[0, :, -1]), "non-condition frames must repeat the image"
|
||||
assert got[0, :, 1:].abs().sum() > 0, "non-condition frames must not be zero-filled"
|
||||
@@ -0,0 +1,77 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cosmos3 unified 3D mRoPE position-ID parity (Tier A scaffold).
|
||||
|
||||
Reference: ``vllm_omni/diffusion/models/cosmos3/transformer_cosmos3.py``
|
||||
lines 113-177 (``compute_mrope_position_ids_text`` /
|
||||
``compute_mrope_position_ids_vision``). The reference test asserting
|
||||
these invariants lives at
|
||||
``tests/diffusion/models/cosmos3/test_cosmos3_transformer.py:32-57``.
|
||||
|
||||
The three invariants under test:
|
||||
|
||||
1. Text tokens broadcast the same monotonically-increasing positions
|
||||
across all three (t, h, w) axes. With ``num_tokens=3`` and
|
||||
``temporal_offset=5`` the result is ``[[5,6,7], [5,6,7], [5,6,7]]``
|
||||
and the next-offset is ``8``.
|
||||
|
||||
2. Vision tokens (no FPS modulation) flatten a ``(grid_t, grid_h, grid_w)``
|
||||
position grid in t-major order. With ``(2, 2, 3)`` and offset ``10``
|
||||
the resulting shape is ``(3, 12)`` and the temporal row begins
|
||||
``[10]*6 + [11]*6``; next-offset is ``12``.
|
||||
|
||||
3. FPS-modulated vision tokens scale the temporal axis by
|
||||
``base_fps / temporal_compression_factor / (fps / tcf)``. With
|
||||
``fps=12``, ``base_fps=24``, ``tcf=4``, ``grid_t=2`` the first row is
|
||||
``[10.0, 12.0]``.
|
||||
|
||||
The FastVideo side currently does NOT exist; the test is wrapped in
|
||||
``try/except ImportError`` and skips. Phase 2b replaces the skip with
|
||||
the real import + assertion path.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
|
||||
def test_compute_mrope_position_ids_text_and_vision() -> None:
|
||||
"""Asserts the 3 invariants of unified 3D mRoPE position-ID generation.
|
||||
|
||||
Once FastVideo's ``fastvideo.models.dits.cosmos3`` exports
|
||||
``compute_mrope_position_ids_text`` and
|
||||
``compute_mrope_position_ids_vision``, this test verifies they produce
|
||||
output tensors identical to the vllm-omni reference at
|
||||
transformer_cosmos3.py:113-177.
|
||||
"""
|
||||
try:
|
||||
from fastvideo.models.dits.cosmos3 import ( # type: ignore
|
||||
compute_mrope_position_ids_text,
|
||||
compute_mrope_position_ids_vision,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("FastVideo Cosmos3 not yet implemented (Phase 2b)")
|
||||
|
||||
text_ids, text_offset = compute_mrope_position_ids_text(num_tokens=3, temporal_offset=5)
|
||||
assert text_ids.tolist() == [[5, 6, 7], [5, 6, 7], [5, 6, 7]]
|
||||
assert text_offset == 8
|
||||
|
||||
vision_ids, vision_offset = compute_mrope_position_ids_vision(
|
||||
2, 2, 3, temporal_offset=10, fps=None
|
||||
)
|
||||
assert tuple(vision_ids.shape) == (3, 12)
|
||||
assert vision_ids[0].tolist() == [10] * 6 + [11] * 6
|
||||
assert vision_offset == 12
|
||||
|
||||
modulated_ids, modulated_offset = compute_mrope_position_ids_vision(
|
||||
2,
|
||||
1,
|
||||
1,
|
||||
temporal_offset=10,
|
||||
fps=12.0,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=4,
|
||||
)
|
||||
torch.testing.assert_close(modulated_ids[0], torch.tensor([10.0, 12.0]))
|
||||
assert modulated_offset == 13
|
||||
@@ -0,0 +1,312 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 sequence packing vs the OFFICIAL framework.
|
||||
|
||||
FastVideo's native packer
|
||||
(``fastvideo.pipelines.basic.cosmos3.sequence_packing.pack_cosmos3_video_sequence``)
|
||||
builds the packed-sequence inputs the ``Cosmos3VFMTransformer`` consumes. This
|
||||
test asserts, for the SAME logical inputs (prompt token ids, vision latents,
|
||||
condition-frame indices, diffusion timestep, fps), that FastVideo's packing
|
||||
matches the official ``cosmos_framework.data.vfm.sequence_packing.pack_input_sequence``
|
||||
oracle field-by-field:
|
||||
|
||||
* ``position_ids`` (exact, ``[3, seq]``),
|
||||
* ``text_ids`` / ``text_indexes``,
|
||||
* ``split_lens`` / ``attn_modes`` / ``sample_lens`` / ``sequence_length``,
|
||||
* vision ``sequence_indexes`` / ``token_shapes`` / ``timesteps`` /
|
||||
``mse_loss_indexes`` / ``noisy_frame_indexes`` / ``condition_mask``.
|
||||
|
||||
Coverage spans T2V (no condition frames), I2V (condition frame 0), and T2I
|
||||
(single conditioned frame), across multiple grids, plus a multi-sample batch.
|
||||
|
||||
Then BOTH the framework-packed and FastVideo-packed inputs are fed through the
|
||||
SAME tiny FastVideo DiT (framework weights copied in as in the existing DiT
|
||||
parity tests). Asserting bit-identical DiT output confirms FastVideo's own
|
||||
packing drives the DiT to the same result as the framework's packing.
|
||||
|
||||
The official framework is the parity ORACLE; it runs on CPU / float32 via the
|
||||
SDPA monkey-patch from ``test_cosmos3_reference_forward``.
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_packing_parity.py -q
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# The official framework provides the parity oracle.
|
||||
cosmos_framework = pytest.importorskip(
|
||||
"cosmos_framework",
|
||||
reason="cosmos_framework not installed; run in fv-cosmos3 env.",
|
||||
)
|
||||
|
||||
# Reuse the DiT parity helpers (weight copy + framework->DiT kwarg builder) and
|
||||
# the mRoPE tiny-model builders (real-checkpoint rope constants).
|
||||
from .test_cosmos3_dit_parity import ( # noqa: E402
|
||||
_copy_weights,
|
||||
_fastvideo_inputs_from_packed_seq,
|
||||
)
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_LATENT_CHANNEL,
|
||||
_LATENT_PATCH_SIZE,
|
||||
_MROPE_SECTION,
|
||||
_RESET_SPATIAL_IDS,
|
||||
_ROPE_THETA,
|
||||
_TCF,
|
||||
_TEMPORAL_MODALITY_MARGIN,
|
||||
_build_tiny_cosmos3_mrope,
|
||||
_build_tiny_fastvideo_dit_mrope,
|
||||
)
|
||||
from .test_cosmos3_reference_forward import _apply_sdpa_patches # noqa: E402
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
# Ensure the CPU/float32 SDPA patches are installed (idempotent).
|
||||
_apply_sdpa_patches()
|
||||
|
||||
# Tiny special-token ids (kept < tiny vocab_size=64). The video path appends
|
||||
# eos + start_of_generation after the prompt tokens.
|
||||
_SPECIAL_TOKENS = {
|
||||
"start_of_generation": 60,
|
||||
"end_of_generation": 61,
|
||||
"eos_token_id": 62,
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Builders for the two packers from the SAME logical sample inputs.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _make_vision(grid_t: int, latent_h: int, latent_w: int, seed: int) -> torch.Tensor:
|
||||
"""Deterministic VAE latent ``[1, C, T, H, W]``."""
|
||||
torch.manual_seed(seed)
|
||||
return torch.randn(1, _LATENT_CHANNEL, grid_t, latent_h, latent_w)
|
||||
|
||||
|
||||
def _framework_pack(
|
||||
*,
|
||||
text_ids_per_sample: list[list[int]],
|
||||
visions: list[torch.Tensor],
|
||||
cond_frames_per_sample: list[list[int]],
|
||||
timesteps: list[float],
|
||||
is_image_batch: bool,
|
||||
):
|
||||
from cosmos_framework.data.vfm.sequence_packing import (
|
||||
GenerationDataClean,
|
||||
SequencePlan,
|
||||
pack_input_sequence,
|
||||
)
|
||||
|
||||
gen_data_clean = GenerationDataClean(
|
||||
batch_size=len(visions),
|
||||
is_image_batch=is_image_batch,
|
||||
x0_tokens_vision=list(visions),
|
||||
fps_vision=None,
|
||||
num_vision_items_per_sample=[1] * len(visions),
|
||||
)
|
||||
plans = [
|
||||
SequencePlan(has_text=True, has_vision=True, condition_frame_indexes_vision=list(cf))
|
||||
for cf in cond_frames_per_sample
|
||||
]
|
||||
return pack_input_sequence(
|
||||
sequence_plans=plans,
|
||||
input_text_indexes=[list(t) for t in text_ids_per_sample],
|
||||
gen_data_clean=gen_data_clean,
|
||||
input_timesteps=torch.tensor(timesteps, dtype=torch.float32),
|
||||
special_tokens=_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
position_embedding_type="unified_3d_mrope",
|
||||
unified_3d_mrope_reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
unified_3d_mrope_temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
def _fastvideo_pack(
|
||||
*,
|
||||
text_ids_per_sample: list[list[int]],
|
||||
visions: list[torch.Tensor],
|
||||
cond_frames_per_sample: list[list[int]],
|
||||
timesteps: list[float],
|
||||
):
|
||||
from fastvideo.pipelines.basic.cosmos3.sequence_packing import (
|
||||
Cosmos3SampleInputs,
|
||||
Cosmos3VisionItem,
|
||||
pack_cosmos3_video_sequence,
|
||||
)
|
||||
|
||||
samples = [
|
||||
Cosmos3SampleInputs(
|
||||
text_ids=list(t),
|
||||
vision=Cosmos3VisionItem(latent=v, condition_frame_indexes=list(cf)),
|
||||
timestep=float(ts),
|
||||
)
|
||||
for t, v, cf, ts in zip(text_ids_per_sample, visions, cond_frames_per_sample, timesteps)
|
||||
]
|
||||
return pack_cosmos3_video_sequence(
|
||||
samples,
|
||||
_SPECIAL_TOKENS,
|
||||
latent_patch_size=_LATENT_PATCH_SIZE,
|
||||
include_end_of_generation_token=False,
|
||||
temporal_modality_margin=_TEMPORAL_MODALITY_MARGIN,
|
||||
reset_spatial_ids=_RESET_SPATIAL_IDS,
|
||||
enable_fps_modulation=False,
|
||||
base_fps=24.0,
|
||||
temporal_compression_factor=_TCF,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Field-by-field comparison.
|
||||
# ---------------------------------------------------------------------------
|
||||
def _assert_packs_match(fw, fv) -> None:
|
||||
"""Assert the framework PackedSequence and FastVideo pack agree field-by-field."""
|
||||
# Structure.
|
||||
assert fv.split_lens == list(fw.split_lens), f"split_lens: fv={fv.split_lens} fw={list(fw.split_lens)}"
|
||||
assert fv.attn_modes == list(fw.attn_modes), f"attn_modes: fv={fv.attn_modes} fw={list(fw.attn_modes)}"
|
||||
assert fv.sample_lens == list(fw.sample_lens), f"sample_lens: fv={fv.sample_lens} fw={list(fw.sample_lens)}"
|
||||
assert int(fv.sequence_length) == int(fw.sequence_length)
|
||||
|
||||
# Text.
|
||||
torch.testing.assert_close(fv.text_ids, fw.text_ids.to(torch.long), rtol=0, atol=0)
|
||||
torch.testing.assert_close(fv.text_indexes, fw.text_indexes.to(torch.long), rtol=0, atol=0)
|
||||
|
||||
# position_ids: exact, [3, seq], same dtype.
|
||||
assert fv.position_ids.shape == fw.position_ids.shape, (
|
||||
f"position_ids shape: fv={tuple(fv.position_ids.shape)} fw={tuple(fw.position_ids.shape)}")
|
||||
assert fv.position_ids.dtype == fw.position_ids.dtype, (
|
||||
f"position_ids dtype: fv={fv.position_ids.dtype} fw={fw.position_ids.dtype}")
|
||||
torch.testing.assert_close(fv.position_ids, fw.position_ids, rtol=0, atol=0)
|
||||
|
||||
# Vision.
|
||||
fwv = fw.vision
|
||||
torch.testing.assert_close(fv.vision_sequence_indexes, fwv.sequence_indexes.to(torch.long), rtol=0, atol=0)
|
||||
assert fv.vision_token_shapes == [tuple(s) for s in fwv.token_shapes], (
|
||||
f"token_shapes: fv={fv.vision_token_shapes} fw={[tuple(s) for s in fwv.token_shapes]}")
|
||||
torch.testing.assert_close(fv.vision_timesteps.to(torch.float32), fwv.timesteps.to(torch.float32))
|
||||
torch.testing.assert_close(fv.vision_mse_loss_indexes, fwv.mse_loss_indexes.to(torch.long), rtol=0, atol=0)
|
||||
assert len(fv.vision_noisy_frame_indexes) == len(fwv.noisy_frame_indexes)
|
||||
for a, b in zip(fv.vision_noisy_frame_indexes, fwv.noisy_frame_indexes):
|
||||
torch.testing.assert_close(a.to(torch.long), b.to(torch.long), rtol=0, atol=0)
|
||||
assert len(fv.vision_condition_mask) == len(fwv.condition_mask)
|
||||
for a, b in zip(fv.vision_condition_mask, fwv.condition_mask):
|
||||
torch.testing.assert_close(a.flatten().to(torch.float32), b.flatten().to(torch.float32))
|
||||
|
||||
|
||||
# (grid_t, latent_h, latent_w, n_text, cond_frames, id) — single-sample cases.
|
||||
_CASES = [
|
||||
pytest.param(1, 8, 8, 4, [], id="t2i_1x4x4"),
|
||||
pytest.param(1, 4, 4, 5, [0], id="t2i_cond_1x2x2"),
|
||||
pytest.param(2, 4, 4, 4, [], id="t2v_2x2x2"),
|
||||
pytest.param(3, 8, 4, 6, [], id="t2v_3x4x2"),
|
||||
pytest.param(2, 4, 4, 5, [0], id="i2v_2x2x2"),
|
||||
pytest.param(3, 4, 4, 4, [0], id="i2v_3x2x2"),
|
||||
]
|
||||
|
||||
|
||||
class TestCosmos3PackingParity:
|
||||
|
||||
# -- Field-by-field packing parity -------------------------------------
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
|
||||
def test_packing_fields_match_framework(self, grid_t, latent_h, latent_w, n_text, cond):
|
||||
torch.manual_seed(0)
|
||||
text_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
vision = _make_vision(grid_t, latent_h, latent_w, seed=123)
|
||||
timestep = 500.0
|
||||
|
||||
fw = _framework_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
is_image_batch=(grid_t == 1),
|
||||
)
|
||||
fv = _fastvideo_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
)
|
||||
_assert_packs_match(fw, fv)
|
||||
|
||||
def test_packing_fields_match_framework_multi_sample(self):
|
||||
"""A batch of two samples (T2V + I2V) packs identically to the framework."""
|
||||
torch.manual_seed(1)
|
||||
t0 = torch.randint(0, 60, (3,)).tolist()
|
||||
t1 = torch.randint(0, 60, (5,)).tolist()
|
||||
v0 = _make_vision(2, 4, 4, seed=11)
|
||||
v1 = _make_vision(2, 4, 4, seed=22)
|
||||
kwargs = dict(
|
||||
text_ids_per_sample=[t0, t1],
|
||||
visions=[v0, v1],
|
||||
cond_frames_per_sample=[[], [0]],
|
||||
timesteps=[500.0, 250.0],
|
||||
)
|
||||
fw = _framework_pack(is_image_batch=False, **kwargs)
|
||||
fv = _fastvideo_pack(**kwargs)
|
||||
_assert_packs_match(fw, fv)
|
||||
|
||||
# -- End-to-end: FastVideo packing drives the DiT identically ----------
|
||||
@pytest.mark.parametrize(("grid_t", "latent_h", "latent_w", "n_text", "cond"), _CASES)
|
||||
def test_fastvideo_packing_drives_dit_like_framework(self, grid_t, latent_h, latent_w, n_text, cond):
|
||||
"""Feed BOTH the framework-packed and FastVideo-packed inputs through the
|
||||
SAME FastVideo DiT (framework weights copied in); assert identical output.
|
||||
"""
|
||||
num_layers = 2
|
||||
torch.manual_seed(0)
|
||||
text_ids = torch.randint(0, 60, (n_text,)).tolist()
|
||||
vision = _make_vision(grid_t, latent_h, latent_w, seed=123)
|
||||
timestep = 500.0
|
||||
|
||||
fw_pack = _framework_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
is_image_batch=(grid_t == 1),
|
||||
)
|
||||
fv_pack = _fastvideo_pack(
|
||||
text_ids_per_sample=[text_ids],
|
||||
visions=[vision],
|
||||
cond_frames_per_sample=[cond],
|
||||
timesteps=[timestep],
|
||||
)
|
||||
# Guard: the two packs must agree before we trust the DiT comparison.
|
||||
_assert_packs_match(fw_pack, fv_pack)
|
||||
|
||||
# One DiT instance, framework weights copied in (parity oracle weights).
|
||||
vfm = _build_tiny_cosmos3_mrope(seed=42, num_layers=num_layers)
|
||||
dit = _build_tiny_fastvideo_dit_mrope(num_layers=num_layers)
|
||||
_copy_weights(vfm, dit)
|
||||
|
||||
with torch.no_grad():
|
||||
out_fw = dit(**_fastvideo_inputs_from_packed_seq(fw_pack))
|
||||
out_fv = dit(**fv_pack.to_dit_kwargs())
|
||||
|
||||
# last_hidden_state must be bit-identical.
|
||||
lhs_fw = out_fw["last_hidden_state"]
|
||||
lhs_fv = out_fv["last_hidden_state"]
|
||||
assert lhs_fw.shape == lhs_fv.shape
|
||||
max_abs_lhs = (lhs_fw - lhs_fv).abs().max().item()
|
||||
print(f"\n[packing->dit {grid_t}x{latent_h}x{latent_w} cond={cond}] "
|
||||
f"last_hidden_state max abs diff = {max_abs_lhs:.3e}")
|
||||
torch.testing.assert_close(lhs_fv, lhs_fw, rtol=0, atol=0)
|
||||
|
||||
# preds_vision must be bit-identical when there are noisy frames to
|
||||
# predict. (A fully-conditioned clip has no noisy patches, so the DiT
|
||||
# emits no "preds_vision" — both packs agree the mse-loss set is empty,
|
||||
# already asserted by the field-parity guard above.)
|
||||
has_preds = fv_pack.vision_mse_loss_indexes.numel() > 0
|
||||
assert ("preds_vision" in out_fw) == has_preds
|
||||
assert ("preds_vision" in out_fv) == has_preds
|
||||
if has_preds:
|
||||
pv_fw = out_fw["preds_vision"][0]
|
||||
pv_fv = out_fv["preds_vision"][0]
|
||||
assert pv_fw.shape == pv_fv.shape
|
||||
max_abs_pv = (pv_fw - pv_fv).abs().max().item()
|
||||
print(f"[packing->dit {grid_t}x{latent_h}x{latent_w} cond={cond}] "
|
||||
f"preds_vision max abs diff = {max_abs_pv:.3e}")
|
||||
torch.testing.assert_close(pv_fv, pv_fw, rtol=0, atol=0)
|
||||
@@ -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,159 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: FastVideo Cosmos3 text reasoning vs the framework.
|
||||
|
||||
The omni model's reasoning (VLM text generation) uses ONLY the und (causal)
|
||||
pathway weights (no ``_moe_gen``) + ``embed_tokens`` / ``norm`` / ``lm_head``;
|
||||
the generation pathway and all multimodal embedders are bypassed. FastVideo's
|
||||
native ``Cosmos3VFMTransformer`` already contains exactly those (the und branch
|
||||
of the dual-pathway forward + ``lm_head``), so a text-only forward + ``lm_head``
|
||||
reproduces the framework reasoner.
|
||||
|
||||
This pins:
|
||||
* **prefill logits** — native text-only forward + ``lm_head`` vs the framework
|
||||
``language_model.model.reasoner_forward`` + ``lm_head`` (per-position); and
|
||||
* **greedy generation** — ``cosmos3_generate_reasoner_text`` vs the framework
|
||||
``generate_reasoner_text(do_sample=False)`` (token-for-token).
|
||||
|
||||
Framework is the parity ORACLE (CPU/float32 via SDPA monkey-patch). Text-only
|
||||
(image-conditioned reasoning additionally needs the Qwen3-VL ``vision_encoder``,
|
||||
tracked separately).
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_reasoning_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 fastvideo.pipelines.basic.cosmos3.cosmos3_pipeline import ( # noqa: E402
|
||||
cosmos3_generate_reasoner_text,
|
||||
)
|
||||
|
||||
from .test_cosmos3_dit_parity import _copy_weights # noqa: E402
|
||||
from .test_cosmos3_dit_parity_mrope import ( # noqa: E402
|
||||
_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()
|
||||
|
||||
|
||||
def _framework_prefill_logits(vfm, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
"""Framework reasoner text-only prefill logits ``[1, T, vocab]`` (the oracle).
|
||||
|
||||
Mirrors ``_impl_generate_reasoner_text`` prefill: ``model.reasoner_forward``
|
||||
then ``lm_head`` (here over ALL positions, not just the last)."""
|
||||
from cosmos_framework.model.vfm.mot.unified_mot import ReasonerKVCache
|
||||
|
||||
causal_lm = vfm.language_model
|
||||
model = causal_lm.model
|
||||
cache = ReasonerKVCache.empty(num_layers=len(model.layers))
|
||||
hidden = model.reasoner_forward(input_ids.unsqueeze(0), cache=cache) # [1,T,hidden]
|
||||
return causal_lm.lm_head(hidden) # [1,T,vocab]
|
||||
|
||||
|
||||
def _native_prefill_logits(dit, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
"""Native text-only forward + ``lm_head`` -> ``[T, vocab]``."""
|
||||
n = input_ids.numel()
|
||||
pos = torch.arange(n).unsqueeze(0).expand(3, -1).contiguous()
|
||||
out = dit(
|
||||
text_ids=input_ids.to(torch.long),
|
||||
text_indexes=torch.arange(n),
|
||||
position_ids=pos,
|
||||
sequence_length=n,
|
||||
split_lens=[n],
|
||||
attn_modes=["causal"],
|
||||
vision_tokens=[],
|
||||
vision_token_shapes=[],
|
||||
vision_sequence_indexes=torch.empty(0, dtype=torch.long),
|
||||
vision_timesteps=torch.empty(0),
|
||||
vision_mse_loss_indexes=torch.empty(0, dtype=torch.long),
|
||||
vision_noisy_frame_indexes=[],
|
||||
)
|
||||
return dit.lm_head(out["last_hidden_state"]) # [T, vocab]
|
||||
|
||||
|
||||
def _diffs(a, b):
|
||||
d = (a - b).abs()
|
||||
return d.max().item(), d.mean().item()
|
||||
|
||||
|
||||
class TestCosmos3ReasoningParity:
|
||||
|
||||
def _build(self, seed_model=42, num_layers=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)
|
||||
return vfm, dit
|
||||
|
||||
@pytest.mark.parametrize("n_text", [4, 8, 12])
|
||||
def test_reasoner_prefill_logits_match_framework(self, n_text):
|
||||
vfm, dit = self._build()
|
||||
torch.manual_seed(n_text)
|
||||
input_ids = torch.randint(0, 60, (n_text,))
|
||||
with torch.no_grad():
|
||||
fw = _framework_prefill_logits(vfm, input_ids)[0] # [T, vocab]
|
||||
fv = _native_prefill_logits(dit, input_ids) # [T, vocab]
|
||||
assert fw.shape == fv.shape, f"shape fw={fw.shape} fv={fv.shape}"
|
||||
mx, mn = _diffs(fv, fw)
|
||||
print(f"\n[reasoner_prefill n={n_text}] logits max abs diff = {mx:.3e} mean abs diff = {mn:.3e}")
|
||||
torch.testing.assert_close(fv, fw, atol=1e-4, rtol=1e-3)
|
||||
# The greedy decisions (argmax per position) must agree exactly.
|
||||
assert torch.equal(fv.argmax(-1), fw.argmax(-1))
|
||||
|
||||
@pytest.mark.parametrize("seed", [0, 1, 2, 3])
|
||||
def test_greedy_reasoning_matches_framework(self, seed):
|
||||
vfm, dit = self._build()
|
||||
torch.manual_seed(seed)
|
||||
input_ids = torch.randint(0, 60, (6,))
|
||||
fw = vfm.generate_reasoner_text(
|
||||
input_ids.unsqueeze(0), max_new_tokens=8, do_sample=False, return_only_new_tokens=True,
|
||||
)[0].tolist()
|
||||
fv = cosmos3_generate_reasoner_text(dit, input_ids.tolist(), max_new_tokens=8)
|
||||
print(f"\n[greedy_reason seed={seed}] framework={fw} native={fv}")
|
||||
assert fv == fw, f"token mismatch: native={fv} framework={fw}"
|
||||
|
||||
def test_deepstack_reasoner_forward_matches_framework(self):
|
||||
"""The deepstack reasoner backbone (image-conditioned reasoning's one new
|
||||
native piece) matches the framework ``reasoner_forward`` given identical
|
||||
prefill inputs (inputs_embeds + positions + per-layer deepstack embeds +
|
||||
visual mask)."""
|
||||
from cosmos_framework.model.vfm.mot.unified_mot import ReasonerKVCache
|
||||
|
||||
vfm, dit = self._build()
|
||||
hidden = dit.hidden_size
|
||||
n_layers = dit.num_hidden_layers
|
||||
seq = 10
|
||||
torch.manual_seed(5)
|
||||
inputs_embeds = torch.randn(seq, hidden)
|
||||
position_ids = torch.arange(seq).unsqueeze(0).expand(3, -1).contiguous() # [3, seq]
|
||||
visual_pos_mask = torch.zeros(seq, dtype=torch.bool)
|
||||
visual_pos_mask[2:6] = True # 4 "image" tokens
|
||||
k = int(visual_pos_mask.sum())
|
||||
deepstack = [torch.randn(k, hidden) for _ in range(n_layers)] # one per layer
|
||||
|
||||
model = vfm.language_model.model
|
||||
cache = ReasonerKVCache.empty(num_layers=len(model.layers))
|
||||
with torch.no_grad():
|
||||
hid_fw = model.reasoner_forward(
|
||||
input_ids=None,
|
||||
inputs_embeds=inputs_embeds.unsqueeze(0),
|
||||
position_ids=position_ids.unsqueeze(1), # [3, B=1, seq]
|
||||
visual_pos_masks=visual_pos_mask.unsqueeze(0),
|
||||
deepstack_visual_embeds=deepstack,
|
||||
cache=cache,
|
||||
)[0] # [seq, hidden]
|
||||
hid_fv = dit.reason_forward(inputs_embeds, position_ids, deepstack, visual_pos_mask) # [seq, hidden]
|
||||
assert hid_fw.shape == hid_fv.shape, f"shape fw={hid_fw.shape} fv={hid_fv.shape}"
|
||||
mx, mn = _diffs(hid_fv, hid_fw)
|
||||
print(f"\n[deepstack_reasoner] hidden max abs diff = {mx:.3e} mean abs diff = {mn:.3e}")
|
||||
torch.testing.assert_close(hid_fv, hid_fw, atol=1e-4, rtol=1e-3)
|
||||
@@ -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.
|
||||
@@ -0,0 +1,96 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Numerical-parity test: Cosmos3 vision_encoder (transformers) vs the framework.
|
||||
|
||||
Image-conditioned reasoning needs the Qwen3-VL ``vision_encoder``. The Cosmos3
|
||||
checkpoint's ``vision_encoder`` is a standard ``transformers`` ``Qwen3VLVisionModel``
|
||||
(``architectures: ["Qwen3VLVisionModel"]``); the framework ships its own copy
|
||||
(``cosmos_framework.model.vfm.vlm.qwen3_vl.qwen3_vl.Qwen3VLVisionModel``). Like the
|
||||
Qwen2 tokenizer, FastVideo reuses the ``transformers`` model (no diffusers); this
|
||||
pins it bit-for-bit against the framework's implementation.
|
||||
|
||||
Both built tiny from the SAME vision config; ``transformers`` weights copied into
|
||||
the framework model; identical ``(hidden_states, grid_thw)`` forward. The
|
||||
framework is the parity ORACLE.
|
||||
|
||||
(Separately verified on the REAL 1.15 GB checkpoint: both strict-load and produce
|
||||
identical ``[N, out_hidden]`` embeds, max abs diff = 0.0.)
|
||||
|
||||
Run:
|
||||
cd <worktree> && <fv-cosmos3 python> -m pytest \
|
||||
tests/local_tests/cosmos3/test_cosmos3_vision_encoder_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.",
|
||||
)
|
||||
transformers = pytest.importorskip("transformers")
|
||||
|
||||
pytestmark = [pytest.mark.local]
|
||||
|
||||
_TINY_VISION_CFG = dict(
|
||||
depth=2,
|
||||
hidden_size=32,
|
||||
num_heads=2,
|
||||
intermediate_size=64,
|
||||
patch_size=16,
|
||||
temporal_patch_size=2,
|
||||
spatial_merge_size=2,
|
||||
in_channels=3,
|
||||
out_hidden_size=64,
|
||||
num_position_embeddings=64,
|
||||
deepstack_visual_indexes=[1],
|
||||
hidden_act="gelu_pytorch_tanh",
|
||||
initializer_range=0.02,
|
||||
)
|
||||
|
||||
|
||||
def _diffs(a, b):
|
||||
d = (a - b).abs()
|
||||
return d.max().item(), d.mean().item()
|
||||
|
||||
|
||||
def _to_tensor(out):
|
||||
return out[0] if isinstance(out, (tuple, list)) else out
|
||||
|
||||
|
||||
class TestCosmos3VisionEncoderParity:
|
||||
|
||||
@pytest.mark.parametrize(("grid_h", "grid_w"), [(2, 2), (4, 4), (4, 6)])
|
||||
def test_vision_encoder_matches_framework(self, grid_h, grid_w):
|
||||
from cosmos_framework.model.vfm.vlm.qwen3_vl.configuration_qwen3_vl import (
|
||||
Qwen3VLVisionConfig as FwVisionConfig,
|
||||
)
|
||||
from cosmos_framework.model.vfm.vlm.qwen3_vl.qwen3_vl import (
|
||||
Qwen3VLVisionModel as FwVisionModel,
|
||||
)
|
||||
from transformers import Qwen3VLVisionModel as TfVisionModel
|
||||
from transformers.models.qwen3_vl.configuration_qwen3_vl import (
|
||||
Qwen3VLVisionConfig as TfVisionConfig,
|
||||
)
|
||||
|
||||
torch.manual_seed(0)
|
||||
tf_model = TfVisionModel(TfVisionConfig(**_TINY_VISION_CFG)).eval()
|
||||
fw_model = FwVisionModel(FwVisionConfig(**_TINY_VISION_CFG)).eval()
|
||||
# transformers is the unit under test; copy its weights into the framework
|
||||
# oracle (identical key surface — both are the same Qwen3-VL ViT).
|
||||
missing, unexpected = fw_model.load_state_dict(tf_model.state_dict(), strict=False)
|
||||
assert not missing and not unexpected, f"key mismatch: missing={missing[:3]} unexpected={unexpected[:3]}"
|
||||
|
||||
in_dim = _TINY_VISION_CFG["in_channels"] * _TINY_VISION_CFG["temporal_patch_size"] * (
|
||||
_TINY_VISION_CFG["patch_size"] ** 2)
|
||||
seq = grid_h * grid_w
|
||||
hidden_states = torch.randn(seq, in_dim)
|
||||
grid_thw = torch.tensor([[1, grid_h, grid_w]], dtype=torch.long)
|
||||
|
||||
with torch.no_grad():
|
||||
tf_out = _to_tensor(tf_model(hidden_states, grid_thw))
|
||||
fw_out = _to_tensor(fw_model(hidden_states, grid_thw))
|
||||
assert tf_out.shape == fw_out.shape, f"shape tf={tf_out.shape} fw={fw_out.shape}"
|
||||
mx, mn = _diffs(tf_out, fw_out)
|
||||
print(f"\n[vision_encoder grid={grid_h}x{grid_w}] embeds max abs diff = {mx:.3e} mean abs diff = {mn:.3e}")
|
||||
torch.testing.assert_close(tf_out, fw_out, atol=1e-5, rtol=1e-4)
|
||||
Reference in New Issue
Block a user