Compare commits

...
Author SHA1 Message Date
SolitaryThinker c28b0c9d89 run lint 2025-07-06 01:16:42 -07:00
SolitaryThinker feed03b456 cosmos2 dit 2025-07-06 01:12:05 -07:00
8 changed files with 1229 additions and 2 deletions
+6 -1
View File
@@ -1,5 +1,10 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.v1.configs.models.dits.cosmos import CosmosConfig
from fastvideo.v1.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.v1.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.v1.configs.models.dits.wanvideo import WanVideoConfig
__all__ = ["HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig"]
__all__ = [
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig", "CosmosConfig"
]
+104
View File
@@ -0,0 +1,104 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_transformer_blocks(n: str, m) -> bool:
return "transformer_blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class CosmosArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_transformer_blocks])
_param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embed\.(.*)$": r"patch_embed.\1",
r"^time_embed\.time_proj\.(.*)$": r"time_embed.time_proj.\1",
r"^time_embed\.t_embedder\.(.*)$": r"time_embed.t_embedder.\1",
r"^time_embed\.norm\.(.*)$": r"time_embed.norm.\1",
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"transformer_blocks.\1.attn1.norm_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"transformer_blocks.\1.attn1.norm_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.norm_q\.(.*)$":
r"transformer_blocks.\1.attn2.norm_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.norm_k\.(.*)$":
r"transformer_blocks.\1.attn2.norm_k.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0\.proj\.(.*)$":
r"transformer_blocks.\1.ff.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2\.(.*)$":
r"transformer_blocks.\1.ff.fc_out.\2",
r"^norm_out\.(.*)$": r"norm_out.\1",
r"^proj_out\.(.*)$": r"proj_out.\1",
})
_lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^transformer_blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"transformer_blocks.\1.attn1.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"transformer_blocks.\1.attn1.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"transformer_blocks.\1.attn1.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn1\.to_out\.(.*)$":
r"transformer_blocks.\1.attn1.to_out.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_q\.(.*)$":
r"transformer_blocks.\1.attn2.to_q.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_k\.(.*)$":
r"transformer_blocks.\1.attn2.to_k.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_v\.(.*)$":
r"transformer_blocks.\1.attn2.to_v.\2",
r"^transformer_blocks\.(\d+)\.attn2\.to_out\.(.*)$":
r"transformer_blocks.\1.attn2.to_out.\2",
r"^transformer_blocks\.(\d+)\.ff\.(.*)$":
r"transformer_blocks.\1.ff.\2",
})
# Cosmos-specific config parameters based on transformer_cosmos.py
in_channels: int = 16
out_channels: int = 16
num_attention_heads: int = 16
attention_head_dim: int = 128
num_layers: int = 28
mlp_ratio: float = 4.0
text_embed_dim: int = 1024
adaln_lora_dim: int = 256
max_size: tuple[int, int, int] = (128, 240, 240)
patch_size: tuple[int, int, int] = (1, 2, 2)
rope_scale: tuple[float, float, float] = (1.0, 4.0, 4.0)
concat_padding_mask: bool = True
extra_pos_embed_type: str | None = None
qk_norm: str = "rms_norm"
eps: float = 1e-6
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.in_channels
@dataclass
class CosmosConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=CosmosArchConfig)
prefix: str = "Cosmos"
+33
View File
@@ -0,0 +1,33 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import CacheParams
@dataclass
class CosmosTeaCacheParams(CacheParams):
cache_type: str = "teacache"
teacache_thresh: float = 0.0
use_ret_steps: bool = True
ret_steps_coeffs: list[float] = field(default_factory=list)
non_ret_steps_coeffs: list[float] = field(default_factory=list)
@property
def coefficients(self) -> list[float]:
if self.use_ret_steps:
return self.ret_steps_coeffs
else:
return self.non_ret_steps_coeffs
@property
def ret_steps(self) -> int:
if self.use_ret_steps:
return 5 * 2
else:
return 1 * 2
def get_cutoff_steps(self, num_inference_steps: int) -> int:
if self.use_ret_steps:
return num_inference_steps * 2
else:
return num_inference_steps * 2 - 2
+54
View File
@@ -44,6 +44,60 @@ def _rotate_gptj(x: torch.Tensor) -> torch.Tensor:
return x.flatten(-2)
def apply_rotary_emb(
x: torch.Tensor,
freqs_cis: torch.Tensor | tuple[torch.Tensor],
use_real: bool = True,
use_real_unbind_dim: int = -1,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply rotary embeddings to input tensors using the given frequency tensor. This function applies rotary embeddings
to the given query or key 'x' tensors using the provided frequency tensor 'freqs_cis'. The input tensors are
reshaped as complex numbers, and the frequency tensor is reshaped for broadcasting compatibility. The resulting
tensors contain rotary embeddings and are returned as real tensors.
Args:
x (`torch.Tensor`):
Query or key tensor to apply rotary embeddings. [B, H, S, D] xk (torch.Tensor): Key tensor to apply
freqs_cis (`Tuple[torch.Tensor]`): Precomputed frequency tensor for complex exponentials. ([S, D], [S, D],)
Returns:
Tuple[torch.Tensor, torch.Tensor]: Tuple of modified query tensor and key tensor with rotary embeddings.
"""
if use_real:
cos, sin = freqs_cis # [S, D]
cos = cos[None, None]
sin = sin[None, None]
cos, sin = cos.to(x.device), sin.to(x.device)
if use_real_unbind_dim == -1:
# Used for flux, cogvideox, hunyuan-dit
x_real, x_imag = x.reshape(*x.shape[:-1], -1,
2).unbind(-1) # [B, S, H, D//2]
x_rotated = torch.stack([-x_imag, x_real], dim=-1).flatten(3)
elif use_real_unbind_dim == -2:
# Used for Stable Audio, OmniGen, CogView4 and Cosmos
x_real, x_imag = x.reshape(*x.shape[:-1], 2,
-1).unbind(-2) # [B, S, H, D//2]
x_rotated = torch.cat([-x_imag, x_real], dim=-1)
else:
raise ValueError(
f"`use_real_unbind_dim={use_real_unbind_dim}` but should be -1 or -2."
)
out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
return out
else:
# used for lumina
x_rotated = torch.view_as_complex(x.float().reshape(
*x.shape[:-1], -1, 2))
freqs_cis = freqs_cis.unsqueeze(2)
x_out = torch.view_as_real(x_rotated * freqs_cis).flatten(3)
return x_out.type_as(x)
def _apply_rotary_emb(
x: torch.Tensor,
cos: torch.Tensor,
+77
View File
@@ -173,3 +173,80 @@ def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
def get_timestep_embedding(
timesteps: torch.Tensor,
embedding_dim: int,
flip_sin_to_cos: bool = False,
downscale_freq_shift: float = 1,
scale: float = 1,
max_period: int = 10000,
) -> torch.Tensor:
"""
This matches the implementation in Denoising Diffusion Probabilistic Models: Create sinusoidal timestep embeddings.
Args
timesteps (torch.Tensor):
a 1-D Tensor of N indices, one per batch element. These may be fractional.
embedding_dim (int):
the dimension of the output.
flip_sin_to_cos (bool):
Whether the embedding order should be `cos, sin` (if True) or `sin, cos` (if False)
downscale_freq_shift (float):
Controls the delta between frequencies between dimensions
scale (float):
Scaling factor applied to the embeddings.
max_period (int):
Controls the maximum frequency of the embeddings
Returns
torch.Tensor: an [N x dim] Tensor of positional embeddings.
"""
assert len(timesteps.shape) == 1, "Timesteps should be a 1d-array"
half_dim = embedding_dim // 2
exponent = -math.log(max_period) * torch.arange(
start=0, end=half_dim, dtype=torch.float32, device=timesteps.device)
exponent = exponent / (half_dim - downscale_freq_shift)
emb = torch.exp(exponent)
emb = timesteps[:, None].float() * emb[None, :]
# scale embeddings
emb = scale * emb
# concat sine and cosine embeddings
emb = torch.cat([torch.sin(emb), torch.cos(emb)], dim=-1)
# flip sine and cosine embeddings
if flip_sin_to_cos:
emb = torch.cat([emb[:, half_dim:], emb[:, :half_dim]], dim=-1)
# zero pad
if embedding_dim % 2 == 1:
emb = torch.nn.functional.pad(emb, (0, 1, 0, 0))
return emb
class Timesteps(nn.Module):
def __init__(self,
num_channels: int,
flip_sin_to_cos: bool,
downscale_freq_shift: float,
scale: int = 1):
super().__init__()
self.num_channels = num_channels
self.flip_sin_to_cos = flip_sin_to_cos
self.downscale_freq_shift = downscale_freq_shift
self.scale = scale
def forward(self, timesteps: torch.Tensor) -> torch.Tensor:
t_emb = get_timestep_embedding(
timesteps,
self.num_channels,
flip_sin_to_cos=self.flip_sin_to_cos,
downscale_freq_shift=self.downscale_freq_shift,
scale=self.scale,
)
return t_emb
+731
View File
@@ -0,0 +1,731 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import numpy as np
import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits.cosmos import CosmosConfig
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.mlp import MLP
from fastvideo.v1.layers.rotary_embedding import apply_rotary_emb
from fastvideo.v1.layers.visual_embedding import Timesteps
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import AttentionBackendEnum
class CosmosPatchEmbed(nn.Module):
def __init__(self,
in_channels: int,
out_channels: int,
patch_size: tuple[int, int, int],
bias: bool = True) -> None:
super().__init__()
self.patch_size = patch_size
self.proj = nn.Linear(in_channels * patch_size[0] * patch_size[1] *
patch_size[2],
out_channels,
bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
hidden_states = hidden_states.reshape(batch_size, num_channels,
num_frames // p_t, p_t,
height // p_h, p_h, width // p_w,
p_w)
hidden_states = hidden_states.permute(0, 2, 4, 6, 1, 3, 5,
7).flatten(4, 7)
hidden_states = self.proj(hidden_states)
return hidden_states
class CosmosTimestepEmbedding(nn.Module):
def __init__(self, in_features: int, out_features: int) -> None:
super().__init__()
self.linear_1 = nn.Linear(in_features, out_features, bias=False)
self.activation = nn.SiLU()
self.linear_2 = nn.Linear(out_features, 3 * out_features, bias=False)
def forward(self, timesteps: torch.Tensor) -> torch.Tensor:
emb = self.linear_1(timesteps)
emb = self.activation(emb)
emb = self.linear_2(emb)
return emb
class CosmosEmbedding(nn.Module):
def __init__(self, embedding_dim: int, condition_dim: int) -> None:
super().__init__()
self.time_proj = Timesteps(embedding_dim,
flip_sin_to_cos=True,
downscale_freq_shift=0.0)
self.t_embedder = CosmosTimestepEmbedding(embedding_dim, condition_dim)
self.norm = RMSNorm(embedding_dim, eps=1e-6)
def forward(self, hidden_states: torch.Tensor,
timestep: torch.LongTensor) -> torch.Tensor:
timesteps_proj = self.time_proj(timestep).type_as(hidden_states)
temb = self.t_embedder(timesteps_proj)
embedded_timestep = self.norm(timesteps_proj)
return temb, embedded_timestep
class CosmosAdaLayerNorm(nn.Module):
def __init__(self, in_features: int, hidden_features: int) -> None:
super().__init__()
self.embedding_dim = in_features
self.activation = nn.SiLU()
self.norm = nn.LayerNorm(in_features,
elementwise_affine=False,
eps=1e-6)
self.linear_1 = nn.Linear(in_features, hidden_features, bias=False)
self.linear_2 = nn.Linear(hidden_features, 2 * in_features, bias=False)
def forward(self,
hidden_states: torch.Tensor,
embedded_timestep: torch.Tensor,
temb: torch.Tensor | None = None) -> torch.Tensor:
embedded_timestep = self.activation(embedded_timestep)
embedded_timestep = self.linear_1(embedded_timestep)
embedded_timestep = self.linear_2(embedded_timestep)
if temb is not None:
embedded_timestep = embedded_timestep + temb[..., :2 *
self.embedding_dim]
shift, scale = embedded_timestep.chunk(2, dim=-1)
hidden_states = self.norm(hidden_states)
if embedded_timestep.ndim == 2:
shift, scale = (x.unsqueeze(1) for x in (shift, scale))
hidden_states = hidden_states * (1 + scale) + shift
return hidden_states
class CosmosAdaLayerNormZero(nn.Module):
def __init__(self,
in_features: int,
hidden_features: int | None = None) -> None:
super().__init__()
self.norm = nn.LayerNorm(in_features,
elementwise_affine=False,
eps=1e-6)
self.activation = nn.SiLU()
if hidden_features is None:
self.linear_1 = nn.Identity()
else:
self.linear_1 = nn.Linear(in_features, hidden_features, bias=False)
self.linear_2 = nn.Linear(
hidden_features if hidden_features is not None else in_features,
3 * in_features,
bias=False)
def forward(
self,
hidden_states: torch.Tensor,
embedded_timestep: torch.Tensor,
temb: torch.Tensor | None = None,
) -> torch.Tensor:
embedded_timestep = self.activation(embedded_timestep)
embedded_timestep = self.linear_1(embedded_timestep)
embedded_timestep = self.linear_2(embedded_timestep)
if temb is not None:
embedded_timestep = embedded_timestep + temb
shift, scale, gate = embedded_timestep.chunk(3, dim=-1)
hidden_states = self.norm(hidden_states)
if embedded_timestep.ndim == 2:
shift, scale, gate = (x.unsqueeze(1) for x in (shift, scale, gate))
hidden_states = hidden_states * (1 + scale) + shift
return hidden_states, gate
class CosmosSelfAttention(nn.Module):
def __init__(self,
dim: int,
num_heads: int,
qk_norm=True,
eps=1e-6,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "") -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.qk_norm = qk_norm
self.eps = eps
# layers
self.to_q = ReplicatedLinear(dim, dim, bias=False)
self.to_k = ReplicatedLinear(dim, dim, bias=False)
self.to_v = ReplicatedLinear(dim, dim, bias=False)
self.to_out = ReplicatedLinear(dim, dim, bias=False)
self.norm_q = RMSNorm(self.head_dim,
eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(self.head_dim,
eps=eps) if qk_norm else nn.Identity()
# Attention mechanism
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=prefix)
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
image_rotary_emb: torch.Tensor | None = None) -> torch.Tensor:
if encoder_hidden_states is None:
encoder_hidden_states = hidden_states
# Get QKV
query, _ = self.to_q(hidden_states)
key, _ = self.to_k(encoder_hidden_states)
value, _ = self.to_v(encoder_hidden_states)
# Reshape for multi-head attention
query = query.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
key = key.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
value = value.unflatten(2, (self.num_heads, -1)).transpose(1, 2)
# Apply normalization
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
# Apply RoPE if provided
if image_rotary_emb is not None:
query = apply_rotary_emb(query,
image_rotary_emb,
use_real=True,
use_real_unbind_dim=-2)
key = apply_rotary_emb(key,
image_rotary_emb,
use_real=True,
use_real_unbind_dim=-2)
# Attention computation
attn_output, _ = self.attn(query, key, value)
# attn_output = attn_output.flatten(2)
attn_output = attn_output.transpose(1, 2).flatten(2, 3).type_as(query)
# Output projection
attn_output, _ = self.to_out(attn_output)
return attn_output
class CosmosCrossAttention(nn.Module):
def __init__(self,
dim: int,
cross_attention_dim: int,
num_heads: int,
qk_norm=True,
eps=1e-6,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "") -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.cross_attention_dim = cross_attention_dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.qk_norm = qk_norm
self.eps = eps
# layers
self.to_q = ReplicatedLinear(dim, dim, bias=False)
self.to_k = ReplicatedLinear(cross_attention_dim, dim, bias=False)
self.to_v = ReplicatedLinear(cross_attention_dim, dim, bias=False)
self.to_out = ReplicatedLinear(dim, dim, bias=False)
self.norm_q = RMSNorm(self.head_dim,
eps=eps) if qk_norm else nn.Identity()
self.norm_k = RMSNorm(self.head_dim,
eps=eps) if qk_norm else nn.Identity()
# Attention mechanism
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends)
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
attention_mask: torch.Tensor | None = None) -> torch.Tensor:
# Get QKV
query, _ = self.to_q(hidden_states)
key, _ = self.to_k(encoder_hidden_states)
value, _ = self.to_v(encoder_hidden_states)
# Reshape for multi-head attention
query = query.unflatten(2, (self.num_heads, -1))
key = key.unflatten(2, (self.num_heads, -1))
value = value.unflatten(2, (self.num_heads, -1))
# Apply normalization
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
# Attention computation
attn_output = self.attn(query, key, value)
attn_output = attn_output.flatten(2, 3).type_as(query)
# Output projection
attn_output, _ = self.to_out(attn_output)
return attn_output
class CosmosTransformerBlock(nn.Module):
def __init__(
self,
num_attention_heads: int,
attention_head_dim: int,
cross_attention_dim: int,
mlp_ratio: float = 4.0,
adaln_lora_dim: int = 256,
qk_norm: str = "rms_norm",
out_bias: bool = False,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
) -> None:
super().__init__()
hidden_size = num_attention_heads * attention_head_dim
self.norm1 = CosmosAdaLayerNormZero(in_features=hidden_size,
hidden_features=adaln_lora_dim)
self.attn1 = CosmosSelfAttention(
dim=hidden_size,
num_heads=num_attention_heads,
qk_norm=(qk_norm == "rms_norm"),
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.norm2 = CosmosAdaLayerNormZero(in_features=hidden_size,
hidden_features=adaln_lora_dim)
self.attn2 = CosmosCrossAttention(
dim=hidden_size,
cross_attention_dim=cross_attention_dim,
num_heads=num_attention_heads,
qk_norm=(qk_norm == "rms_norm"),
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn2")
self.norm3 = CosmosAdaLayerNormZero(in_features=hidden_size,
hidden_features=adaln_lora_dim)
self.ff = MLP(hidden_size,
int(hidden_size * mlp_ratio),
act_type="gelu",
bias=False)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
embedded_timestep: torch.Tensor,
temb: torch.Tensor | None = None,
image_rotary_emb: torch.Tensor | None = None,
extra_pos_emb: torch.Tensor | None = None,
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
if extra_pos_emb is not None:
hidden_states = hidden_states + extra_pos_emb
# 1. Self Attention
norm_hidden_states, gate = self.norm1(hidden_states, embedded_timestep,
temb)
attn_output = self.attn1(norm_hidden_states,
image_rotary_emb=image_rotary_emb)
hidden_states = hidden_states + gate * attn_output
# 2. Cross Attention
norm_hidden_states, gate = self.norm2(hidden_states, embedded_timestep,
temb)
attn_output = self.attn2(norm_hidden_states,
encoder_hidden_states=encoder_hidden_states,
attention_mask=attention_mask)
hidden_states = hidden_states + gate * attn_output
# 3. Feed Forward
norm_hidden_states, gate = self.norm3(hidden_states, embedded_timestep,
temb)
ff_output = self.ff(norm_hidden_states)
hidden_states = hidden_states + gate * ff_output
return hidden_states
class CosmosRotaryPosEmbed(nn.Module):
def __init__(
self,
hidden_size: int,
max_size: tuple[int, int, int] = (128, 240, 240),
patch_size: tuple[int, int, int] = (1, 2, 2),
base_fps: int = 24,
rope_scale: tuple[float, float, float] = (2.0, 1.0, 1.0),
) -> None:
super().__init__()
self.max_size = [
size // patch
for size, patch in zip(max_size, patch_size, strict=False)
]
self.patch_size = patch_size
self.base_fps = base_fps
self.dim_h = hidden_size // 6 * 2
self.dim_w = hidden_size // 6 * 2
self.dim_t = hidden_size - self.dim_h - self.dim_w
self.h_ntk_factor = rope_scale[1]**(self.dim_h / (self.dim_h - 2))
self.w_ntk_factor = rope_scale[2]**(self.dim_w / (self.dim_w - 2))
self.t_ntk_factor = rope_scale[0]**(self.dim_t / (self.dim_t - 2))
def forward(self,
hidden_states: torch.Tensor,
fps: int | None = None) -> tuple[torch.Tensor, torch.Tensor]:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
pe_size = [
num_frames // self.patch_size[0], height // self.patch_size[1],
width // self.patch_size[2]
]
device = hidden_states.device
h_theta = 10000.0 * self.h_ntk_factor
w_theta = 10000.0 * self.w_ntk_factor
t_theta = 10000.0 * self.t_ntk_factor
seq = torch.arange(max(self.max_size),
device=device,
dtype=torch.float32)
dim_h_range = (
torch.arange(0, self.dim_h, 2, device=device,
dtype=torch.float32)[:(self.dim_h // 2)] / self.dim_h)
dim_w_range = (
torch.arange(0, self.dim_w, 2, device=device,
dtype=torch.float32)[:(self.dim_w // 2)] / self.dim_w)
dim_t_range = (
torch.arange(0, self.dim_t, 2, device=device,
dtype=torch.float32)[:(self.dim_t // 2)] / self.dim_t)
h_spatial_freqs = 1.0 / (h_theta**dim_h_range)
w_spatial_freqs = 1.0 / (w_theta**dim_w_range)
temporal_freqs = 1.0 / (t_theta**dim_t_range)
emb_h = torch.outer(seq[:pe_size[1]],
h_spatial_freqs)[None, :, None, :].repeat(
pe_size[0], 1, pe_size[2], 1)
emb_w = torch.outer(seq[:pe_size[2]],
w_spatial_freqs)[None, None, :, :].repeat(
pe_size[0], pe_size[1], 1, 1)
# Apply sequence scaling in temporal dimension
if fps is None:
# Images
emb_t = torch.outer(seq[:pe_size[0]], temporal_freqs)
else:
# Videos
emb_t = torch.outer(seq[:pe_size[0]] / fps * self.base_fps,
temporal_freqs)
emb_t = emb_t[:, None, None, :].repeat(1, pe_size[1], pe_size[2], 1)
freqs = torch.cat([emb_t, emb_h, emb_w] * 2, dim=-1).flatten(0,
2).float()
cos = torch.cos(freqs)
sin = torch.sin(freqs)
return cos, sin
class CosmosLearnablePositionalEmbed(nn.Module):
def __init__(
self,
hidden_size: int,
max_size: tuple[int, int, int],
patch_size: tuple[int, int, int],
eps: float = 1e-6,
) -> None:
super().__init__()
self.max_size = [
size // patch
for size, patch in zip(max_size, patch_size, strict=False)
]
self.patch_size = patch_size
self.eps = eps
self.pos_emb_t = nn.Parameter(torch.zeros(self.max_size[0],
hidden_size))
self.pos_emb_h = nn.Parameter(torch.zeros(self.max_size[1],
hidden_size))
self.pos_emb_w = nn.Parameter(torch.zeros(self.max_size[2],
hidden_size))
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, num_channels, num_frames, height, width = hidden_states.shape
pe_size = [
num_frames // self.patch_size[0], height // self.patch_size[1],
width // self.patch_size[2]
]
emb_t = self.pos_emb_t[:pe_size[0]][None, :, None, None, :].repeat(
batch_size, 1, pe_size[1], pe_size[2], 1)
emb_h = self.pos_emb_h[:pe_size[1]][None, None, :, None, :].repeat(
batch_size, pe_size[0], 1, pe_size[2], 1)
emb_w = self.pos_emb_w[:pe_size[2]][None, None, None, :, :].repeat(
batch_size, pe_size[0], pe_size[1], 1, 1)
emb = emb_t + emb_h + emb_w
emb = emb.flatten(1, 3)
norm = torch.linalg.vector_norm(emb,
dim=-1,
keepdim=True,
dtype=torch.float32)
norm = torch.add(self.eps,
norm,
alpha=np.sqrt(norm.numel() / emb.numel()))
return (emb / norm).type_as(hidden_states)
class CosmosTransformer3DModel(BaseDiT):
_fsdp_shard_conditions = CosmosConfig()._fsdp_shard_conditions
_compile_conditions = CosmosConfig()._compile_conditions
_supported_attention_backends = CosmosConfig()._supported_attention_backends
_param_names_mapping = CosmosConfig()._param_names_mapping
_lora_param_names_mapping = CosmosConfig()._lora_param_names_mapping
def __init__(self, config: CosmosConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.max_size = config.max_size
self.rope_scale = config.rope_scale
self.concat_padding_mask = config.concat_padding_mask
self.extra_pos_embed_type = config.extra_pos_embed_type
# 1. Patch Embedding
patch_embed_in_channels = config.in_channels + 1 if config.concat_padding_mask else config.in_channels
self.patch_embed = CosmosPatchEmbed(patch_embed_in_channels,
inner_dim,
config.patch_size,
bias=False)
# 2. Positional Embedding
self.rope = CosmosRotaryPosEmbed(hidden_size=config.attention_head_dim,
max_size=config.max_size,
patch_size=config.patch_size,
rope_scale=config.rope_scale)
self.learnable_pos_embed = None
if config.extra_pos_embed_type == "learnable":
self.learnable_pos_embed = CosmosLearnablePositionalEmbed(
hidden_size=inner_dim,
max_size=config.max_size,
patch_size=config.patch_size,
)
# 3. Time Embedding
self.time_embed = CosmosEmbedding(inner_dim, inner_dim)
# 4. Transformer Blocks
self.transformer_blocks = nn.ModuleList([
CosmosTransformerBlock(
num_attention_heads=config.num_attention_heads,
attention_head_dim=config.attention_head_dim,
cross_attention_dim=config.text_embed_dim,
mlp_ratio=config.mlp_ratio,
adaln_lora_dim=config.adaln_lora_dim,
qk_norm=config.qk_norm,
out_bias=False,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.transformer_blocks.{i}",
) for i in range(config.num_layers)
])
# 5. Output norm & projection
self.norm_out = CosmosAdaLayerNorm(inner_dim, config.adaln_lora_dim)
self.proj_out = nn.Linear(inner_dim,
config.out_channels *
math.prod(config.patch_size),
bias=False)
self.gradient_checkpointing = False
# For TeaCache
self.previous_e0_even = None
self.previous_e0_odd = None
self.previous_residual_even = None
self.previous_residual_odd = None
self.is_even = True
self.should_calc_even = True
self.should_calc_odd = True
self.accumulated_rel_l1_distance_even = 0
self.accumulated_rel_l1_distance_odd = 0
self.cnt = 0
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
attention_mask: torch.Tensor | None = None,
fps: int | None = None,
condition_mask: torch.Tensor | None = None,
padding_mask: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
forward_batch = get_forward_context().forward_batch
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
batch_size, num_channels, num_frames, height, width = hidden_states.shape
# 1. Concatenate padding mask if needed & prepare attention mask
if condition_mask is not None:
hidden_states = torch.cat([hidden_states, condition_mask], dim=1)
if self.concat_padding_mask and padding_mask is not None:
from torchvision import transforms
padding_mask = transforms.functional.resize(
padding_mask,
list(hidden_states.shape[-2:]),
interpolation=transforms.InterpolationMode.NEAREST)
hidden_states = torch.cat([
hidden_states,
padding_mask.unsqueeze(2).repeat(batch_size, 1, num_frames, 1,
1)
],
dim=1)
# # Resize padding mask to match hidden states spatial dimensions
# padding_mask_resized = F.interpolate(
# padding_mask.float().unsqueeze(1),
# size=(height, width),
# mode='nearest'
# ).squeeze(1)
# hidden_states = torch.cat(
# [hidden_states, padding_mask_resized.unsqueeze(1).unsqueeze(2).repeat(1, 1, num_frames, 1, 1)], dim=1
# )
if attention_mask is not None:
attention_mask = attention_mask.unsqueeze(1).unsqueeze(
1) # [B, 1, 1, S]
# 2. Generate positional embeddings
image_rotary_emb = self.rope(hidden_states, fps=fps)
extra_pos_emb = self.learnable_pos_embed(
hidden_states) if self.extra_pos_embed_type == "learnable" else None
# 3. Patchify input
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
hidden_states = self.patch_embed(hidden_states)
hidden_states = hidden_states.flatten(
1, 3) # [B, T, H, W, C] -> [B, THW, C] codespell:ignore
# 4. Timestep embeddings
if timestep.ndim == 1:
temb, embedded_timestep = self.time_embed(hidden_states, timestep)
elif timestep.ndim == 5:
assert timestep.shape == (batch_size, 1, num_frames, 1, 1), (
f"Expected timestep to have shape [B, 1, T, 1, 1], but got {timestep.shape}"
)
timestep = timestep.flatten()
temb, embedded_timestep = self.time_embed(hidden_states, timestep)
# We can do this because num_frames == post_patch_num_frames, as p_t is 1
temb, embedded_timestep = (
x.view(batch_size, post_patch_num_frames, 1, 1,
-1).expand(-1, -1, post_patch_height, post_patch_width,
-1).flatten(1, 3)
for x in (temb, embedded_timestep)
) # [BT, C] -> [B, T, 1, 1, C] -> [B, T, H, W, C] -> [B, THW, C] codespell:ignore
else:
raise ValueError(f"Unsupported timestep shape: {timestep.shape}")
# 6. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.transformer_blocks:
hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
embedded_timestep,
temb,
image_rotary_emb,
extra_pos_emb,
attention_mask,
)
else:
for block in self.transformer_blocks:
hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
embedded_timestep=embedded_timestep,
temb=temb,
image_rotary_emb=image_rotary_emb,
extra_pos_emb=extra_pos_emb,
attention_mask=attention_mask,
)
# 7. Output norm & projection & unpatchify
hidden_states = self.norm_out(hidden_states, embedded_timestep, temb)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.unflatten(2, (p_h, p_w, p_t, -1))
hidden_states = hidden_states.unflatten(
1, (post_patch_num_frames, post_patch_height, post_patch_width))
# NOTE: The permutation order here is not the inverse operation of what happens when patching as usually expected.
# It might be a source of confusion to the reader, but this is correct
hidden_states = hidden_states.permute(0, 7, 1, 6, 2, 4, 3, 5)
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return hidden_states
+2 -1
View File
@@ -23,7 +23,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel")
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel")
}
_IMAGE_TO_VIDEO_DIT_MODELS = {
@@ -0,0 +1,222 @@
# SPDX-License-Identifier: Apache-2.0
import os
import numpy as np
import pytest
import torch
from diffusers import CosmosTransformer3DModel
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TransformerLoader
from fastvideo.v1.utils import maybe_download_model
from fastvideo.v1.configs.models.dits import CosmosConfig
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29504"
BASE_MODEL_PATH = "nvidia/Cosmos-Predict2-2B-Text2Image"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_cosmos2_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
pipeline_config=PipelineConfig(dit_config=CosmosConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
model1 = CosmosTransformer3DModel.from_pretrained(
TRANSFORMER_PATH, device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
seq_len = 30
# Video latents [B, C, T, H, W] - Cosmos2 specific dimensions
hidden_states = torch.randn(batch_size,
16,
1, # Single frame for image generation
32, # Height patches
32, # Width patches
device=device,
dtype=precision)
# Text embeddings [B, L, D] - Cosmos2 uses T5 embeddings with 1024 dim
encoder_hidden_states = torch.randn(batch_size,
seq_len,
1024, # T5 embedding dimension
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
# padding mask
padding_mask = hidden_states.new_zeros(1, 1, 32, 32, device=device, dtype=precision)
# print(padding_mask.shape)
forward_batch = ForwardBatch(
data_type="dummy",
)
with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
padding_mask=padding_mask,
return_dict=False,
)[0]
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
padding_mask=padding_mask)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
@pytest.mark.usefixtures("distributed_setup")
def test_cosmos2_transformer_video2world():
"""Test Cosmos2 Video2World variant"""
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
# Use Video2World model path
base_model_path = "nvidia/Cosmos-Predict2-2B-Video2World"
model_path = maybe_download_model(base_model_path,
local_dir=os.path.join(
'data', base_model_path))
transformer_path = os.path.join(model_path, "transformer")
args = FastVideoArgs(model_path=transformer_path,
use_cpu_offload=False,
pipeline_config=PipelineConfig(dit_config=CosmosConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(transformer_path, "", args).to(device, dtype=precision)
model1 = CosmosTransformer3DModel.from_pretrained(
transformer_path, device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
seq_len = 30
# Video latents [B, C, T, H, W] - Video2World has additional condition channel
hidden_states = torch.randn(batch_size,
17, # 16 + 1 for condition channel
8, # Multiple frames for video
32, # Height patches
32, # Width patches
device=device,
dtype=precision)
# Text embeddings [B, L, D]
encoder_hidden_states = torch.randn(batch_size,
seq_len,
1024, # T5 embedding dimension
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(
data_type="dummy",
)
with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
return_dict=False,
)[0]
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-1, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"