minimax-h3 / diffusers /models /transformers /transformer_cosmos3.py
multimodalart's picture
multimodalart HF Staff
Sync the split MiniMax-H3 Spaces
186aa49 verified
Raw
History Blame Contribute Delete
42.6 kB
# Copyright 2025 The NVIDIA Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import math
from dataclasses import dataclass
import torch
import torch.nn as nn
from ...configuration_utils import ConfigMixin, register_to_config
from ...loaders import PeftAdapterMixin
from ...utils import BaseOutput
from ..attention import AttentionMixin, AttentionModuleMixin
from ..attention_dispatch import dispatch_attention_fn
from ..embeddings import TimestepEmbedding, Timesteps
from ..modeling_utils import ModelMixin
from ..normalization import RMSNorm
@dataclass
class Cosmos3OmniTransformerOutput(BaseOutput):
"""Output of [`Cosmos3OmniTransformer`].
Args:
sample (`list[torch.Tensor]`):
Per-item vision velocity predictions.
sound (`list[torch.Tensor]`, *optional*):
Per-item sound velocity predictions when sound generation is enabled.
action (`list[torch.Tensor]`, *optional*):
Per-item action velocity predictions when action generation is enabled.
"""
sample: list[torch.Tensor]
sound: list[torch.Tensor] | None = None
action: list[torch.Tensor] | None = None
class Cosmos3AttnProcessor:
"""Dual-pathway attention processor for Cosmos3.
Projects, normalizes, applies rotary position embeddings, then runs separate causal (understanding) and full
(generation) attention pathways. The generation pathway cross-attends to both und and gen keys/values.
"""
_attention_backend = None
_parallel_config = None
def __call__(
self,
attn: "Cosmos3PackedMoTAttention",
und_seq: torch.Tensor,
gen_seq: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
# Per-pathway projections
q_und = attn.to_q(und_seq).view(-1, attn.num_attention_heads, attn.head_dim)
k_und = attn.to_k(und_seq).view(-1, attn.num_key_value_heads, attn.head_dim)
v_und = attn.to_v(und_seq).view(-1, attn.num_key_value_heads, attn.head_dim)
q_gen = attn.add_q_proj(gen_seq).view(-1, attn.num_attention_heads, attn.head_dim)
k_gen = attn.add_k_proj(gen_seq).view(-1, attn.num_key_value_heads, attn.head_dim)
v_gen = attn.add_v_proj(gen_seq).view(-1, attn.num_key_value_heads, attn.head_dim)
q_und = attn.norm_q(q_und)
k_und = attn.norm_k(k_und)
k_und_for_gen = attn.k_norm_und_for_gen(k_und) if attn.k_norm_und_for_gen is not None else k_und
q_gen = attn.norm_added_q(q_gen)
k_gen = attn.norm_added_k(k_gen)
# Apply rotary position embeddings per pathway
cos_und, sin_und, cos_gen, sin_gen = rotary_emb
cos_und = cos_und.unsqueeze(1)
sin_und = sin_und.unsqueeze(1)
q_und = q_und * cos_und + _rotate_half(q_und) * sin_und
k_und = k_und * cos_und + _rotate_half(k_und) * sin_und
k_und_for_gen = k_und_for_gen * cos_und + _rotate_half(k_und_for_gen) * sin_und
cos_gen = cos_gen.unsqueeze(1)
sin_gen = sin_gen.unsqueeze(1)
q_gen = q_gen * cos_gen + _rotate_half(q_gen) * sin_gen
k_gen = k_gen * cos_gen + _rotate_half(k_gen) * sin_gen
# Causal pathway (understanding): und tokens self-attend with causal masking.
causal_out = dispatch_attention_fn(
q_und.unsqueeze(0),
k_und.unsqueeze(0),
v_und.unsqueeze(0),
is_causal=True,
enable_gqa=True,
backend=self._attention_backend,
parallel_config=self._parallel_config,
)
causal_out = causal_out.squeeze(0).flatten(-2, -1)
# Full pathway (generation): gen tokens cross-attend to all (und + gen) keys/values.
all_k = torch.cat([k_und_for_gen, k_gen], dim=0)
all_v = torch.cat([v_und, v_gen], dim=0)
full_out = dispatch_attention_fn(
q_gen.unsqueeze(0),
all_k.unsqueeze(0),
all_v.unsqueeze(0),
is_causal=False,
enable_gqa=True,
backend=self._attention_backend,
parallel_config=self._parallel_config,
)
full_out = full_out.squeeze(0).flatten(-2, -1)
# Per-pathway output projection
und_out = attn.to_out(causal_out)
gen_out = attn.to_add_out(full_out)
return und_out, gen_out
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
half = x.shape[-1] // 2
return torch.cat((-x[..., half:], x[..., :half]), dim=-1)
class Cosmos3VLTextRotaryEmbedding(nn.Module):
def __init__(self, head_dim: int, rope_theta: float, rope_axes_dim: tuple[int, int, int]):
super().__init__()
inv_freq = 1.0 / (rope_theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.rope_axes_dim = rope_axes_dim
def apply_interleaved_mrope(self, freqs, rope_axes_dim):
"""Reorganize chunked [TTT...HHH...WWW] frequency layout into interleaved
[THTHWHTHW...TT], preserving frequency continuity across the 3 grids."""
freqs_t = freqs[0]
for dim, offset in enumerate((1, 2), start=1): # H, W
length = rope_axes_dim[dim] * 3
idx = slice(offset, length, 3)
freqs_t[..., idx] = freqs[dim, ..., idx]
return freqs_t
def forward(self, position_ids, device, dtype):
if position_ids.ndim == 2:
position_ids = position_ids[None, ...].expand(3, position_ids.shape[0], -1) # [3,B,N]
inv_freq_expanded = (
self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1).to(device)
) # [3,B,head_dim//2,1]
position_ids_expanded = position_ids[:, :, None, :].float() # [3,B,1,N]
# Disable autocast so the position-id matmul runs in float32: under an ambient autocast it would run in
# bfloat16, which cannot represent consecutive integers past 256, collapsing positions onto the same
# frequency and degrading the rotary embedding.
with torch.autocast(device_type=position_ids.device.type, enabled=False):
freqs = inv_freq_expanded @ position_ids_expanded
freqs = freqs.transpose(2, 3) # [3,B,N,head_dim//2]
freqs = self.apply_interleaved_mrope(freqs, self.rope_axes_dim) # [B,N,head_dim//2]
emb = torch.cat((freqs, freqs), dim=-1) # [B,N,head_dim]
return emb.cos().to(dtype=dtype), emb.sin().to(dtype=dtype) # each: [B,N,head_dim]
class Cosmos3NemotronRMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
input_dtype = hidden_states.dtype
hidden_states = hidden_states.float()
variance = hidden_states.pow(2).mean(-1, keepdim=True)
hidden_states = hidden_states * torch.rsqrt(variance + self.eps)
return (self.weight.float() * hidden_states).to(input_dtype)
class Cosmos3VLTextMLP(nn.Module):
def __init__(self, hidden_size: int, intermediate_size: int, hidden_act: str = "silu"):
super().__init__()
if hidden_act not in ("relu2", "silu"):
raise ValueError(f"Cosmos3 only supports `hidden_act` values 'relu2' and 'silu', got {hidden_act!r}.")
self.hidden_act = hidden_act
if hidden_act == "silu":
self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False)
self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
self.act_fn = nn.SiLU() if hidden_act == "silu" else None
def forward(self, x):
if self.hidden_act == "relu2":
return self.down_proj(torch.relu(self.up_proj(x)).square())
return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))
class DomainAwareLinear(nn.Module):
"""Linear projection with one weight/bias pair per embodiment domain."""
def __init__(self, input_size: int, output_size: int, num_domains: int) -> None:
super().__init__()
self.input_size = input_size
self.output_size = output_size
self.num_domains = num_domains
self.fc = nn.Embedding(self.num_domains, self.output_size * self.input_size)
self.bias = nn.Embedding(self.num_domains, self.output_size)
def forward(self, x: torch.Tensor, domain_id: torch.Tensor) -> torch.Tensor:
if domain_id.ndim == 0:
domain_id = domain_id.unsqueeze(0)
domain_id = domain_id.to(device=x.device, dtype=torch.long).reshape(-1)
if x.shape[0] != domain_id.shape[0]:
raise ValueError(
"Cosmos3 action domain_id batch size must match action tokens: "
f"tokens={x.shape[0]}, domain_id={domain_id.shape[0]}."
)
if torch.any((domain_id < 0) | (domain_id >= self.num_domains)):
raise ValueError(f"Cosmos3 action domain_id must be in [0, {self.num_domains}), got {domain_id.tolist()}.")
weight = self.fc(domain_id).view(domain_id.shape[0], self.input_size, self.output_size)
bias = self.bias(domain_id).view(domain_id.shape[0], self.output_size)
if x.ndim == 2:
return torch.bmm(x.unsqueeze(1), weight).squeeze(1) + bias
if x.ndim == 3:
return torch.bmm(x, weight) + bias.unsqueeze(1)
raise ValueError(f"Cosmos3 DomainAwareLinear expected rank-2 or rank-3 input, got {tuple(x.shape)}.")
class Cosmos3PackedMoTAttention(nn.Module, AttentionModuleMixin):
"""Dual-pathway packed attention with separate projections for the understanding and generation token streams."""
_default_processor_cls = Cosmos3AttnProcessor
_available_processors = [Cosmos3AttnProcessor]
_supports_qkv_fusion = False
def __init__(
self,
hidden_size: int,
head_dim: int,
num_attention_heads: int,
num_key_value_heads: int,
attention_bias: bool,
rms_norm_eps: float,
qk_norm_for_text: bool = True,
use_und_k_norm_for_gen: bool = False,
norm_type: str = "rms_norm",
processor=None,
):
super().__init__()
self.hidden_size = hidden_size
self.head_dim = head_dim
self.num_attention_heads = num_attention_heads
self.num_key_value_heads = num_key_value_heads
self.num_key_value_groups = num_attention_heads // num_key_value_heads
# Understanding pathway. norm_q / norm_k are applied per-head (only on
# head_dim), so no reshape is needed after them.
self.to_q = nn.Linear(hidden_size, num_attention_heads * head_dim, bias=attention_bias)
self.to_k = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias)
self.to_v = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias)
self.to_out = nn.Linear(num_attention_heads * head_dim, hidden_size, bias=attention_bias)
if not qk_norm_for_text:
self.norm_q = nn.Identity()
self.norm_k = nn.Identity()
elif norm_type == "nemotron_rms_norm":
self.norm_q = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps)
self.norm_k = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps)
else:
self.norm_q = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.norm_k = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False)
if use_und_k_norm_for_gen and not qk_norm_for_text:
if norm_type == "nemotron_rms_norm":
self.k_norm_und_for_gen = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps)
else:
self.k_norm_und_for_gen = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False)
else:
self.k_norm_und_for_gen = None
# Generation pathway
self.add_q_proj = nn.Linear(hidden_size, num_attention_heads * head_dim, bias=attention_bias)
self.add_k_proj = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias)
self.add_v_proj = nn.Linear(hidden_size, num_key_value_heads * head_dim, bias=attention_bias)
self.to_add_out = nn.Linear(num_attention_heads * head_dim, hidden_size, bias=attention_bias)
if norm_type == "nemotron_rms_norm":
self.norm_added_q = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps)
self.norm_added_k = Cosmos3NemotronRMSNorm(head_dim, eps=rms_norm_eps)
else:
self.norm_added_q = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.norm_added_k = RMSNorm(head_dim, eps=rms_norm_eps, elementwise_affine=True, bias=False)
if processor is None:
processor = self._default_processor_cls()
self.set_processor(processor)
def forward(
self,
und_seq: torch.Tensor,
gen_seq: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
return self.processor(self, und_seq, gen_seq, rotary_emb)
class Cosmos3VLTextMoTDecoderLayer(nn.Module):
"""Cosmos3 text MoT decoder layer for the Qwen3 and Nemotron dense backbones."""
def __init__(
self,
hidden_size: int,
head_dim: int,
num_attention_heads: int,
num_key_value_heads: int,
intermediate_size: int,
attention_bias: bool,
rms_norm_eps: float,
hidden_act: str = "silu",
qk_norm_for_text: bool = True,
use_und_k_norm_for_gen: bool = False,
):
super().__init__()
self.hidden_size = hidden_size
norm_type = "nemotron_rms_norm" if hidden_act == "relu2" else "rms_norm"
self.self_attn = Cosmos3PackedMoTAttention(
hidden_size=hidden_size,
head_dim=head_dim,
num_attention_heads=num_attention_heads,
num_key_value_heads=num_key_value_heads,
attention_bias=attention_bias,
rms_norm_eps=rms_norm_eps,
qk_norm_for_text=qk_norm_for_text,
use_und_k_norm_for_gen=use_und_k_norm_for_gen,
norm_type=norm_type,
)
self.mlp = Cosmos3VLTextMLP(
hidden_size=hidden_size, intermediate_size=intermediate_size, hidden_act=hidden_act
)
self.mlp_moe_gen = Cosmos3VLTextMLP(
hidden_size=hidden_size, intermediate_size=intermediate_size, hidden_act=hidden_act
)
if norm_type == "nemotron_rms_norm":
self.input_layernorm = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps)
self.input_layernorm_moe_gen = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps)
self.post_attention_layernorm = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps)
self.post_attention_layernorm_moe_gen = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps)
else:
self.input_layernorm = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.input_layernorm_moe_gen = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.post_attention_layernorm = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.post_attention_layernorm_moe_gen = RMSNorm(
hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False
)
def forward(
self,
und_seq: torch.Tensor,
gen_seq: torch.Tensor,
rotary_emb: tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
und_norm = self.input_layernorm(und_seq)
gen_norm = self.input_layernorm_moe_gen(gen_seq)
und_attn_out, gen_attn_out = self.self_attn(und_norm, gen_norm, rotary_emb)
residual_und = und_seq + und_attn_out
residual_gen = gen_seq + gen_attn_out
mlp_out_und = self.mlp(self.post_attention_layernorm(residual_und))
mlp_out_gen = self.mlp_moe_gen(self.post_attention_layernorm_moe_gen(residual_gen))
return residual_und + mlp_out_und, residual_gen + mlp_out_gen
class Cosmos3OmniTransformer(ModelMixin, ConfigMixin, PeftAdapterMixin, AttentionMixin):
_supports_gradient_checkpointing = True
_no_split_modules = ["Cosmos3VLTextMoTDecoderLayer"]
_repeated_blocks = ["Cosmos3VLTextMoTDecoderLayer"]
_skip_layerwise_casting_patterns = ["embed_tokens", "time_embedder", "norm"]
_keep_in_fp32_modules = ["time_embedder"]
# Optional context-parallelism seams. They default to ``None`` (no-op) so the
# model itself carries no CP logic. `forward` applies `_cp_shard_fn` to the
# per-pathway hidden states + rotary embeddings before the decoder layers, and
# `_cp_gather_fn` to the per-pathway outputs after the final norm. An external
# helper (see `examples/cosmos3/cosmos_parallel.py`) sets these to
# shard/gather across a device mesh and installs a context-parallel attention
# processor — the packed dual-pathway + GQA + ragged-length structure cannot be
# expressed as diffusers' declarative `_cp_plan`, so CP lives outside the model.
_cp_shard_fn = None
_cp_gather_fn = None
# `dtype` is injected into init_dict by ModelMixin.from_pretrained (configuration_utils.py:289),
# so __init__ must accept it. Excluding it here keeps save_pretrained from writing it into
# config.json — the value is a load-time runtime hint, not part of the model architecture.
ignore_for_config = ["dtype"]
@register_to_config
def __init__(
self,
attention_bias: bool = False,
attention_dropout: float = 0.0,
dtype: str = "bfloat16", # required by the loader (see `ignore_for_config` above); not read here
head_dim: int = 128,
hidden_size: int = 4096,
intermediate_size: int = 12288,
base_fps: int = 24,
enable_fps_modulation: bool = True,
latent_channel: int = 48,
unified_3d_mrope_reset_spatial_ids: bool = True,
unified_3d_mrope_temporal_modality_margin: int = 15000,
latent_patch_size: int = 2,
num_attention_heads: int = 32,
num_hidden_layers: int = 36,
num_key_value_heads: int = 8,
patch_latent_dim: int = 192,
rms_norm_eps: float = 1e-6,
rope_scaling: dict | None = None,
rope_theta: float = 5000000.0,
action_dim: int | None = None,
action_gen: bool = False,
num_embodiment_domains: int = 32,
sound_dim: int | None = None,
sound_gen: bool = False,
sound_latent_fps: float = 25.0,
timestep_scale: float = 0.001,
vocab_size: int = 151936,
hidden_act: str = "silu",
qk_norm_for_text: bool = True,
use_und_k_norm_for_gen: bool = False,
rope_axes_dim: tuple[int, int, int] | list[int] | None = None,
):
super().__init__()
if rope_axes_dim is None:
rope_axes_dim = (
rope_scaling.get("mrope_section", [24, 20, 20]) if rope_scaling is not None else [24, 20, 20]
)
self.register_to_config(rope_axes_dim=rope_axes_dim)
# Text-model layers live directly on the transformer (flat layout). The published
# checkpoint must be re-keyed with the leading `model.` prefix stripped — see
# scripts/build_flat_layout_repo.py for the rewrite.
self.embed_tokens = nn.Embedding(vocab_size, hidden_size)
self.layers = nn.ModuleList(
[
Cosmos3VLTextMoTDecoderLayer(
hidden_size=hidden_size,
head_dim=head_dim,
num_attention_heads=num_attention_heads,
num_key_value_heads=num_key_value_heads,
intermediate_size=intermediate_size,
attention_bias=attention_bias,
rms_norm_eps=rms_norm_eps,
hidden_act=hidden_act,
qk_norm_for_text=qk_norm_for_text,
use_und_k_norm_for_gen=use_und_k_norm_for_gen,
)
for _ in range(num_hidden_layers)
]
)
if hidden_act == "relu2":
self.norm = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps)
self.norm_moe_gen = Cosmos3NemotronRMSNorm(hidden_size, eps=rms_norm_eps)
else:
self.norm = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.norm_moe_gen = RMSNorm(hidden_size, eps=rms_norm_eps, elementwise_affine=True, bias=False)
self.rotary_emb = Cosmos3VLTextRotaryEmbedding(
head_dim=head_dim, rope_theta=rope_theta, rope_axes_dim=rope_axes_dim
)
# Modality projection heads + timestep embedding.
self.vocab_size = vocab_size
self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False)
self.proj_in = nn.Linear(patch_latent_dim, hidden_size, bias=True)
self.proj_out = nn.Linear(hidden_size, patch_latent_dim, bias=True)
self.time_proj = Timesteps(num_channels=256, flip_sin_to_cos=True, downscale_freq_shift=0)
self.time_embedder = TimestepEmbedding(in_channels=256, time_embed_dim=hidden_size)
self.action_gen = action_gen
self.action_dim = action_dim
self.num_embodiment_domains = num_embodiment_domains
if action_gen:
if self.action_dim is None:
raise ValueError("`action_dim` must be provided when `action_gen=True`.")
self.action_proj_in = DomainAwareLinear(self.action_dim, hidden_size, self.num_embodiment_domains)
self.action_proj_out = DomainAwareLinear(hidden_size, self.action_dim, self.num_embodiment_domains)
self.action_modality_embed = nn.Parameter(torch.zeros(hidden_size))
if sound_gen:
if sound_dim is None:
raise ValueError("`sound_dim` must be provided when `sound_gen=True`.")
self.audio_proj_in = nn.Linear(sound_dim, hidden_size, bias=True)
self.audio_proj_out = nn.Linear(hidden_size, sound_dim, bias=True)
self.audio_modality_embed = nn.Parameter(torch.zeros(hidden_size))
self.gradient_checkpointing = False
# -------------------------------------------------------------------------
# Pure-tensor packing/unpacking helpers (no layer state).
# -------------------------------------------------------------------------
def _apply_timestep_embeds_to_noisy_tokens(
self,
packed_tokens: torch.Tensor,
packed_timestep_embeds: torch.Tensor,
noisy_frame_indexes: list[torch.Tensor],
token_shapes: list[tuple[int, ...]],
) -> torch.Tensor:
start_noisy_index = 0
flattened_noisy_frame_indexes: list[torch.Tensor] = []
for noisy_indexes_i, token_shape_i in zip(noisy_frame_indexes, token_shapes):
spatial_numel_i = math.prod(token_shape_i[1:])
spatial_indexes_i = torch.arange(spatial_numel_i, device=packed_tokens.device)
# Broadcast [N, 1] + [spatial_numel_i] → [N, spatial_numel_i]
frame_offsets = (noisy_indexes_i * spatial_numel_i).unsqueeze(-1) + spatial_indexes_i + start_noisy_index
flattened_noisy_frame_indexes.append(frame_offsets.flatten())
start_noisy_index += token_shape_i[0] * spatial_numel_i
flattened = torch.cat(flattened_noisy_frame_indexes, dim=0).unsqueeze(-1).expand(-1, packed_tokens.shape[1])
return packed_tokens.scatter_add(dim=0, index=flattened, src=packed_timestep_embeds)
def _patchify_and_pack_latents(
self,
tokens_vision: list[torch.Tensor],
) -> tuple[torch.Tensor, list[tuple[int, int, int]]]:
p = self.config.latent_patch_size
latent_channel = self.config.latent_channel
packed_latent: list[torch.Tensor] = []
original_latent_shapes: list[tuple[int, int, int]] = []
for latent in tokens_vision:
latent = latent.squeeze(0) # [C, T, H, W]
_, t_actual, h_actual, w_actual = latent.shape
original_latent_shapes.append((t_actual, h_actual, w_actual))
h_padded = ((h_actual + p - 1) // p) * p
w_padded = ((w_actual + p - 1) // p) * p
if h_padded != h_actual or w_padded != w_actual:
padded = torch.zeros(
(latent_channel, t_actual, h_padded, w_padded),
device=latent.device,
dtype=latent.dtype,
)
padded[:, :, :h_actual, :w_actual] = latent
latent = padded
h_patches = h_padded // p
w_patches = w_padded // p
latent = latent.reshape(latent_channel, t_actual, h_patches, p, w_patches, p)
latent = torch.einsum("cthpwq->thwpqc", latent).reshape(-1, p * p * latent_channel)
packed_latent.append(latent)
return torch.cat(packed_latent, dim=0), original_latent_shapes
def _unpatchify_and_unpack_latents(
self,
packed_mse_preds: torch.Tensor,
token_shapes_vision: list[tuple[int, int, int]],
noisy_frame_indexes_vision: list[torch.Tensor],
original_latent_shapes: list[tuple[int, int, int]],
) -> list[torch.Tensor]:
p = self.config.latent_patch_size
latent_channel = self.config.latent_channel
unpatchified_latents: list[torch.Tensor] = []
start_idx = 0
for token_shape, noisy_frame_indexes, original_shape in zip(
token_shapes_vision, noisy_frame_indexes_vision, original_latent_shapes
):
t_c = token_shape[0]
_, h_orig, w_orig = original_shape
h_padded = ((h_orig + p - 1) // p) * p
w_padded = ((w_orig + p - 1) // p) * p
h_patches = h_padded // p
w_patches = w_padded // p
t_n = len(noisy_frame_indexes)
output_tensor = torch.zeros(
(latent_channel, t_c, h_orig, w_orig),
device=packed_mse_preds.device,
dtype=packed_mse_preds.dtype,
)
num_patches = t_n * h_patches * w_patches
if num_patches > 0:
end_idx = start_idx + num_patches
latent_patches = packed_mse_preds[start_idx:end_idx]
latent_patches = latent_patches.reshape(t_n, h_patches, w_patches, p, p, latent_channel)
latent = torch.einsum("thwpqc->cthpwq", latent_patches)
latent = latent.reshape(latent_channel, t_n, h_patches * p, w_patches * p)
latent = latent[:, :, :h_orig, :w_orig]
output_tensor[:, noisy_frame_indexes] = latent
start_idx = end_idx
unpatchified_latents.append(output_tensor.unsqueeze(0))
return unpatchified_latents
def _pack_sound_latents(
self,
tokens_sound: list[torch.Tensor],
token_shapes_sound: list[tuple[int, int, int]],
) -> torch.Tensor:
"""List of ``[C, T]`` tensors → packed ``[total_T, C]`` tensor."""
return torch.cat(
[sound[:, : shape[0]].permute(1, 0) for sound, shape in zip(tokens_sound, token_shapes_sound)],
dim=0,
)
def _unpack_sound_latents(
self,
packed_preds: torch.Tensor,
token_shapes_sound: list[tuple[int, int, int]],
noisy_frame_indexes_sound: list[torch.Tensor],
) -> list[torch.Tensor]:
"""Packed ``[total_noisy_T, C]`` predictions → list of ``[C, T]`` tensors (zeros at conditioned positions)."""
sound_dim = self.config.sound_dim
unpacked: list[torch.Tensor] = []
start_idx = 0
for shape, noisy_idxs in zip(token_shapes_sound, noisy_frame_indexes_sound):
T = shape[0]
output = torch.zeros((sound_dim, T), device=packed_preds.device, dtype=packed_preds.dtype)
t_n = len(noisy_idxs)
if t_n > 0:
output[:, noisy_idxs] = packed_preds[start_idx : start_idx + t_n].T
start_idx += t_n
unpacked.append(output)
return unpacked
def _pack_action_latents(
self,
tokens_action: list[torch.Tensor],
token_shapes_action: list[tuple[int, int, int]],
domain_ids_action: list[torch.Tensor],
) -> tuple[torch.Tensor, torch.Tensor]:
"""List of ``[T, D]`` tensors → packed ``[total_T, D]`` plus per-token domain ids."""
packed: list[torch.Tensor] = []
domain_ids: list[torch.Tensor] = []
for action, shape, domain_id in zip(tokens_action, token_shapes_action, domain_ids_action):
token_count = shape[0]
packed.append(action[:token_count])
domain_ids.append(domain_id.reshape(1).expand(token_count))
return torch.cat(packed, dim=0), torch.cat(domain_ids, dim=0)
def _unpack_action_latents(
self,
packed_preds: torch.Tensor,
token_shapes_action: list[tuple[int, int, int]],
noisy_frame_indexes_action: list[torch.Tensor],
) -> list[torch.Tensor]:
"""Packed ``[total_noisy_T, D]`` predictions → list of ``[T, D]`` tensors."""
unpacked: list[torch.Tensor] = []
start_idx = 0
for shape, noisy_idxs in zip(token_shapes_action, noisy_frame_indexes_action):
T = shape[0]
output = torch.zeros((T, self.action_dim), device=packed_preds.device, dtype=packed_preds.dtype)
t_n = len(noisy_idxs)
if t_n > 0:
output[noisy_idxs] = packed_preds[start_idx : start_idx + t_n]
start_idx += t_n
unpacked.append(output)
return unpacked
# -------------------------------------------------------------------------
# forward: full per-step pass — encode text/vision/sound/action → run layers →
# decode vision/sound/action. Pipeline calls this once per CFG pass.
# -------------------------------------------------------------------------
def forward(
self,
input_ids: torch.Tensor,
text_indexes: torch.Tensor,
position_ids: torch.Tensor,
und_len: int,
sequence_length: int,
vision_tokens: list[torch.Tensor],
vision_token_shapes: list[tuple[int, int, int]],
vision_sequence_indexes: torch.Tensor,
vision_mse_loss_indexes: torch.Tensor,
vision_timesteps: torch.Tensor,
vision_noisy_frame_indexes: list[torch.Tensor],
sound_tokens: list[torch.Tensor] | None = None,
sound_token_shapes: list[tuple[int, int, int]] | None = None,
sound_sequence_indexes: torch.Tensor | None = None,
sound_mse_loss_indexes: torch.Tensor | None = None,
sound_timesteps: torch.Tensor | None = None,
sound_noisy_frame_indexes: list[torch.Tensor] | None = None,
action_tokens: list[torch.Tensor] | None = None,
action_token_shapes: list[tuple[int, int, int]] | None = None,
action_sequence_indexes: torch.Tensor | None = None,
action_mse_loss_indexes: torch.Tensor | None = None,
action_timesteps: torch.Tensor | None = None,
action_noisy_frame_indexes: list[torch.Tensor] | None = None,
action_domain_ids: list[torch.Tensor] | None = None,
return_dict: bool = True,
) -> (
Cosmos3OmniTransformerOutput | tuple[list[torch.Tensor], list[torch.Tensor] | None, list[torch.Tensor] | None]
):
"""Run a full denoising-step forward pass.
Args:
input_ids: Text token IDs placed at ``text_indexes`` in the joint sequence.
text_indexes: Indices of text tokens in the joint sequence.
position_ids: ``[3, sequence_length]`` mRoPE position IDs for the full joint sequence.
und_len: Length of the causal text (understanding) prefix; generation tokens follow.
sequence_length: Total length of the joint packed sequence.
vision_tokens: Per-item vision latent tensors before patchify.
vision_token_shapes: Patch grid shapes ``(T, H, W)`` per vision item.
vision_sequence_indexes: Indices of vision tokens in the joint sequence.
vision_mse_loss_indexes: Indices used to read vision predictions after the backbone.
vision_timesteps: Per-patch diffusion timesteps for vision tokens.
vision_noisy_frame_indexes: Noisy frame indices per vision item.
sound_tokens: Optional sound latent tensors before packing.
sound_token_shapes: Optional patch grid shapes for sound items.
sound_sequence_indexes: Optional indices of sound tokens in the joint sequence.
sound_mse_loss_indexes: Optional indices used to read sound predictions.
sound_timesteps: Optional per-token diffusion timesteps for sound.
sound_noisy_frame_indexes: Optional noisy frame indices per sound item.
action_tokens: Optional action latent tensors before packing.
action_token_shapes: Optional patch grid shapes ``(T, H, W)`` per action item.
action_sequence_indexes: Optional indices of action tokens in the joint sequence.
action_mse_loss_indexes: Optional indices used to read action predictions after the backbone.
action_timesteps: Optional per-token diffusion timesteps for action tokens.
action_noisy_frame_indexes: Optional noisy frame indices per action item.
action_domain_ids: Optional per-item domain IDs selecting the action head weights.
return_dict: Whether to return a [`Cosmos3OmniTransformerOutput`] instead of a tuple.
Returns:
A [`Cosmos3OmniTransformerOutput`] or a tuple of per-modality prediction lists. Optional modalities return
``None`` when their inputs are omitted.
"""
has_sound = sound_tokens is not None and sound_sequence_indexes is not None
has_action = action_tokens is not None and action_sequence_indexes is not None
# Embed text tokens into the joint hidden_states buffer at their sequence positions.
packed_text_embedding = self.embed_tokens(input_ids)
target_dtype = packed_text_embedding.dtype
hidden_states = packed_text_embedding.new_zeros(size=(sequence_length, self.config.hidden_size))
hidden_states[text_indexes] = packed_text_embedding
# Patchify + project vision latents, then add timestep embeddings to noisy frames.
packed_tokens_vision, original_latent_shapes = self._patchify_and_pack_latents(vision_tokens)
packed_tokens_vision = self.proj_in(packed_tokens_vision)
timesteps_vision = vision_timesteps * self.config.timestep_scale
time_embedder_dtype = next(self.time_embedder.parameters()).dtype
packed_timestep_embeds_vision = self.time_embedder(self.time_proj(timesteps_vision).to(time_embedder_dtype))
packed_timestep_embeds_vision = packed_timestep_embeds_vision.to(target_dtype)
packed_tokens_vision = self._apply_timestep_embeds_to_noisy_tokens(
packed_tokens=packed_tokens_vision,
packed_timestep_embeds=packed_timestep_embeds_vision,
noisy_frame_indexes=vision_noisy_frame_indexes,
token_shapes=vision_token_shapes,
)
hidden_states[vision_sequence_indexes] = packed_tokens_vision
# Pack + project sound latents (when present); all sound frames are noisy.
if has_sound:
packed_tokens_sound = self._pack_sound_latents(sound_tokens, sound_token_shapes).to(target_dtype)
packed_tokens_sound = self.audio_proj_in(packed_tokens_sound) + self.audio_modality_embed
timesteps_sound = sound_timesteps * self.config.timestep_scale
packed_timestep_embeds_sound = self.time_embedder(self.time_proj(timesteps_sound).to(time_embedder_dtype))
packed_timestep_embeds_sound = packed_timestep_embeds_sound.to(target_dtype)
packed_tokens_sound = self._apply_timestep_embeds_to_noisy_tokens(
packed_tokens=packed_tokens_sound,
packed_timestep_embeds=packed_timestep_embeds_sound,
noisy_frame_indexes=sound_noisy_frame_indexes,
token_shapes=sound_token_shapes,
)
hidden_states[sound_sequence_indexes] = packed_tokens_sound
# Pack + project action latents (when present). Domain ids select the action head weights.
if has_action:
packed_tokens_action, per_token_domain_ids = self._pack_action_latents(
action_tokens, action_token_shapes, action_domain_ids
)
packed_tokens_action = packed_tokens_action.to(target_dtype)
per_token_domain_ids = per_token_domain_ids.to(device=packed_tokens_action.device)
packed_tokens_action = self.action_proj_in(packed_tokens_action, per_token_domain_ids)
packed_tokens_action = packed_tokens_action + self.action_modality_embed
if action_mse_loss_indexes.numel() > 0:
timesteps_action = action_timesteps * self.config.timestep_scale
packed_timestep_embeds_action = self.time_embedder(
self.time_proj(timesteps_action).to(time_embedder_dtype)
)
packed_timestep_embeds_action = packed_timestep_embeds_action.to(target_dtype)
packed_tokens_action = self._apply_timestep_embeds_to_noisy_tokens(
packed_tokens=packed_tokens_action,
packed_timestep_embeds=packed_timestep_embeds_action,
noisy_frame_indexes=action_noisy_frame_indexes,
token_shapes=action_token_shapes,
)
hidden_states[action_sequence_indexes] = packed_tokens_action
# Compute rotary embeddings once for the joint sequence, then slice into und/gen halves.
_meta_tensor = torch.tensor([], dtype=hidden_states.dtype, device=hidden_states.device)
cos, sin = self.rotary_emb(
position_ids=position_ids.unsqueeze(0) if position_ids.ndim == 1 else position_ids.unsqueeze(1),
device=hidden_states.device,
dtype=hidden_states.dtype,
)
# cos, sin: [1, N, head_dim] (1-D pos_ids) or [3, 1, N, head_dim] (mrope pos_ids)
cos = cos.squeeze(0)
sin = sin.squeeze(0)
und_seq = hidden_states[:und_len]
gen_seq = hidden_states[und_len:]
rotary_emb = (cos[:und_len], sin[:und_len], cos[und_len:], sin[und_len:])
# Optional context-parallelism shard seam (no-op unless set by an external
# helper, e.g. `examples/cosmos3/cosmos_parallel.py`). When set, it
# shards each pathway's sequence and rotary embeddings across a device mesh, so
# the decoder layers below run on local sequence shards.
if self._cp_shard_fn is not None:
und_seq, gen_seq, rotary_emb = self._cp_shard_fn(und_seq, gen_seq, rotary_emb)
for decoder_layer in self.layers:
if torch.is_grad_enabled() and self.gradient_checkpointing:
und_seq, gen_seq = self._gradient_checkpointing_func(
decoder_layer.__call__, und_seq, gen_seq, rotary_emb
)
else:
und_seq, gen_seq = decoder_layer(und_seq, gen_seq, rotary_emb)
und_out = self.norm(und_seq)
gen_out = self.norm_moe_gen(gen_seq)
# Optional context-parallelism gather seam: re-gather the full per-pathway
# sequence on every rank (and drop the padding) before the global-index decode
# below, since the downstream indexes address positions in the unpadded joint
# sequence. No-op unless `_cp_shard_fn`'s counterpart is set.
if self._cp_gather_fn is not None:
und_out, gen_out = self._cp_gather_fn(und_out, gen_out)
last_hidden_state = torch.cat([und_out, gen_out], dim=0)
# Decode vision predictions from the joint hidden state.
preds_vision_packed = self.proj_out(last_hidden_state[vision_mse_loss_indexes])
preds_vision = self._unpatchify_and_unpack_latents(
preds_vision_packed,
token_shapes_vision=vision_token_shapes,
noisy_frame_indexes_vision=vision_noisy_frame_indexes,
original_latent_shapes=original_latent_shapes,
)
preds_sound: list[torch.Tensor] | None = None
if has_sound:
preds_sound_packed = self.audio_proj_out(last_hidden_state[sound_mse_loss_indexes])
preds_sound = self._unpack_sound_latents(preds_sound_packed, sound_token_shapes, sound_noisy_frame_indexes)
preds_action: list[torch.Tensor] | None = None
if has_action:
per_noisy_domain_ids = [
domain_id.reshape(1).expand(len(noisy_idxs))
for domain_id, noisy_idxs in zip(action_domain_ids, action_noisy_frame_indexes)
]
per_noisy_domain_ids = torch.cat(per_noisy_domain_ids, dim=0).to(device=last_hidden_state.device)
preds_action_packed = self.action_proj_out(
last_hidden_state[action_mse_loss_indexes], per_noisy_domain_ids
)
preds_action = self._unpack_action_latents(
preds_action_packed, action_token_shapes, action_noisy_frame_indexes
)
if not return_dict:
return preds_vision, preds_sound, preds_action
return Cosmos3OmniTransformerOutput(sample=preds_vision, sound=preds_sound, action=preds_action)