Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c28b0c9d89 | ||
|
|
feed03b456 |
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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()}"
|
||||
Reference in New Issue
Block a user