Skip to content

vllm_omni.diffusion.models.anima.anima_transformer

ANIMA_TRANSFORMER_CONFIG module-attribute

ANIMA_TRANSFORMER_CONFIG = {
    "in_channels": 16,
    "out_channels": 16,
    "num_attention_heads": 16,
    "attention_head_dim": 128,
    "num_layers": 28,
    "mlp_ratio": 4.0,
    "text_embed_dim": 1024,
    "adaln_lora_dim": 256,
    "max_size": (128, 240, 240),
    "patch_size": (1, 2, 2),
    "rope_scale": (1.0, 4.0, 4.0),
    "concat_padding_mask": True,
    "extra_pos_embed_type": None,
}

AnimaTransformer3DModel

Bases: Module

config instance-attribute

config = SimpleNamespace(
    in_channels=in_channels,
    out_channels=out_channels,
    num_attention_heads=num_attention_heads,
    attention_head_dim=attention_head_dim,
    num_layers=num_layers,
    mlp_ratio=mlp_ratio,
    text_embed_dim=text_embed_dim,
    adaln_lora_dim=adaln_lora_dim,
    max_size=max_size,
    patch_size=patch_size,
    rope_scale=rope_scale,
    concat_padding_mask=concat_padding_mask,
    extra_pos_embed_type=extra_pos_embed_type,
    use_crossattn_projection=use_crossattn_projection,
    crossattn_proj_in_channels=crossattn_proj_in_channels,
    encoder_hidden_states_channels=encoder_hidden_states_channels,
    controlnet_block_every_n=controlnet_block_every_n,
    img_context_dim_in=img_context_dim_in,
    img_context_num_tokens=img_context_num_tokens,
    img_context_dim_out=img_context_dim_out,
    extra_config=kwargs,
)

crossattn_proj instance-attribute

crossattn_proj: Module | None = None

dtype property

dtype: dtype

gradient_checkpointing instance-attribute

gradient_checkpointing = False

img_context_proj instance-attribute

img_context_proj: Module | None = None

learnable_pos_embed instance-attribute

learnable_pos_embed: (
    CosmosLearnablePositionalEmbed | None
) = None

norm_out instance-attribute

norm_out = CosmosAdaLayerNorm(hidden_size, adaln_lora_dim)

patch_embed instance-attribute

patch_embed = CosmosPatchEmbed(
    patch_embed_in_channels,
    hidden_size,
    patch_size,
    bias=False,
)

proj_out instance-attribute

proj_out = nn.Linear(
    hidden_size,
    patch_size[0]
    * patch_size[1]
    * patch_size[2]
    * out_channels,
    bias=False,
)

rope instance-attribute

rope = CosmosRotaryPosEmbed(
    hidden_size=attention_head_dim,
    max_size=max_size,
    patch_size=patch_size,
    rope_scale=rope_scale,
)

time_embed instance-attribute

time_embed = CosmosEmbedding(hidden_size, hidden_size)

transformer_blocks instance-attribute

transformer_blocks = nn.ModuleList(
    [
        CosmosTransformerBlock(
            num_attention_heads=num_attention_heads,
            attention_head_dim=attention_head_dim,
            cross_attention_dim=text_embed_dim,
            mlp_ratio=mlp_ratio,
            adaln_lora_dim=adaln_lora_dim,
            out_bias=False,
            img_context=img_context_dim_in is not None
            and img_context_dim_in > 0,
            prefix=f"transformer_blocks.{i}",
        )
        for i in range(num_layers)
    ]
)

forward

forward(
    hidden_states: Tensor,
    timestep: Tensor,
    encoder_hidden_states: Tensor
    | tuple[Tensor, Tensor | None],
    block_controlnet_hidden_states: list[Tensor]
    | None = None,
    attention_mask: Tensor
    | tuple[Tensor | None, Tensor | None]
    | None = None,
    fps: float | None = None,
    condition_mask: Tensor | None = None,
    padding_mask: Tensor | None = None,
    return_dict: bool = True,
) -> Transformer2DModelOutput | tuple[Tensor]

load_weights

load_weights(
    weights: Iterable[tuple[str, Tensor]],
) -> set[str]

CosmosAdaLayerNorm

Bases: Module

activation instance-attribute

activation = nn.SiLU()

embedding_dim instance-attribute

embedding_dim = in_features

linear_1 instance-attribute

linear_1 = nn.Linear(
    in_features, hidden_features, bias=False
)

linear_2 instance-attribute

linear_2 = nn.Linear(
    hidden_features, 2 * in_features, bias=False
)

norm instance-attribute

norm = nn.LayerNorm(
    in_features, elementwise_affine=False, eps=1e-06
)

forward

forward(
    hidden_states: Tensor,
    embedded_timestep: Tensor,
    temb: Tensor | None = None,
) -> Tensor

CosmosAdaLayerNormZero

Bases: Module

activation instance-attribute

activation = nn.SiLU()

linear_1 instance-attribute

linear_1 = (
    nn.Identity()
    if hidden_features is None
    else nn.Linear(
        in_features, hidden_features, bias=False
    )
)

linear_2 instance-attribute

linear_2 = nn.Linear(
    hidden_features or in_features,
    3 * in_features,
    bias=False,
)

norm instance-attribute

norm = nn.LayerNorm(
    in_features, elementwise_affine=False, eps=1e-06
)

forward

forward(
    hidden_states: Tensor,
    embedded_timestep: Tensor,
    temb: Tensor | None = None,
) -> tuple[Tensor, Tensor]

CosmosAttention

Bases: Module

attn instance-attribute

attn = Attention(
    num_heads=heads,
    head_size=dim_head,
    softmax_scale=1.0 / dim_head**0.5,
    causal=False,
    num_kv_heads=heads,
    prefix=prefix,
)

dim_head instance-attribute

dim_head = dim_head

heads instance-attribute

heads = heads

img_attn instance-attribute

img_attn = Attention(
    num_heads=heads,
    head_size=dim_head,
    softmax_scale=1.0 / dim_head**0.5,
    causal=False,
    num_kv_heads=heads,
    prefix=f"{prefix}.img_attn" if prefix else "img_attn",
)

img_context instance-attribute

img_context = img_context

k_img instance-attribute

k_img = nn.Linear(query_dim, inner_dim, bias=False)

k_img_norm instance-attribute

k_img_norm = nn.RMSNorm(dim_head, eps=1e-06)

norm_k instance-attribute

norm_k = nn.RMSNorm(dim_head, eps=1e-06)

norm_q instance-attribute

norm_q = nn.RMSNorm(dim_head, eps=1e-06)

out_dim instance-attribute

out_dim = inner_dim

q_img instance-attribute

q_img = nn.Linear(query_dim, inner_dim, bias=False)

q_img_norm instance-attribute

q_img_norm = nn.RMSNorm(dim_head, eps=1e-06)

query_dim instance-attribute

query_dim = query_dim

to_k instance-attribute

to_k = nn.Linear(
    cross_attention_dim, inner_dim, bias=False
)

to_out instance-attribute

to_out = nn.ModuleList(
    [
        nn.Linear(inner_dim, query_dim, bias=out_bias),
        nn.Identity(),
    ]
)

to_q instance-attribute

to_q = nn.Linear(query_dim, inner_dim, bias=False)

to_v instance-attribute

to_v = nn.Linear(
    cross_attention_dim, inner_dim, bias=False
)

v_img instance-attribute

v_img = nn.Linear(query_dim, inner_dim, bias=False)

forward

forward(
    hidden_states: Tensor,
    encoder_hidden_states: Tensor
    | tuple[Tensor, Tensor | None]
    | None = None,
    attention_mask: Tensor
    | tuple[Tensor | None, Tensor | None]
    | None = None,
    image_rotary_emb: tuple[Tensor, Tensor] | None = None,
) -> Tensor

CosmosEmbedding

Bases: Module

norm instance-attribute

norm = nn.RMSNorm(embedding_dim, eps=1e-06)

t_embedder instance-attribute

t_embedder = CosmosTimestepEmbedding(
    embedding_dim, condition_dim
)

time_proj instance-attribute

time_proj = Timesteps(
    embedding_dim,
    flip_sin_to_cos=True,
    downscale_freq_shift=0.0,
)

forward

forward(
    hidden_states: Tensor, timestep: Tensor
) -> tuple[Tensor, Tensor]

CosmosFeedForward

Bases: Module

net instance-attribute

net = nn.ModuleList(
    [
        CosmosGELU(dim, inner_dim, bias=bias),
        nn.Identity(),
        nn.Linear(inner_dim, dim, bias=bias),
    ]
)

forward

forward(hidden_states: Tensor) -> Tensor

CosmosGELU

Bases: Module

proj instance-attribute

proj = nn.Linear(dim, inner_dim, bias=bias)

forward

forward(hidden_states: Tensor) -> Tensor

CosmosLearnablePositionalEmbed

Bases: Module

eps instance-attribute

eps = eps

max_size instance-attribute

max_size = [
    size // patch
    for size, patch in zip(max_size, patch_size)
]

patch_size instance-attribute

patch_size = patch_size

pos_emb_h instance-attribute

pos_emb_h = nn.Parameter(
    torch.zeros(self.max_size[1], hidden_size)
)

pos_emb_t instance-attribute

pos_emb_t = nn.Parameter(
    torch.zeros(self.max_size[0], hidden_size)
)

pos_emb_w instance-attribute

pos_emb_w = nn.Parameter(
    torch.zeros(self.max_size[2], hidden_size)
)

forward

forward(hidden_states: Tensor) -> Tensor

CosmosPatchEmbed

Bases: Module

patch_size instance-attribute

patch_size = patch_size

proj instance-attribute

proj = nn.Linear(
    in_channels
    * patch_size[0]
    * patch_size[1]
    * patch_size[2],
    out_channels,
    bias=bias,
)

forward

forward(hidden_states: Tensor) -> Tensor

CosmosRotaryPosEmbed

Bases: Module

base_fps instance-attribute

base_fps = base_fps

dim_h instance-attribute

dim_h = hidden_size // 6 * 2

dim_t instance-attribute

dim_t = hidden_size - self.dim_h - self.dim_w

dim_w instance-attribute

dim_w = hidden_size // 6 * 2

h_ntk_factor instance-attribute

h_ntk_factor = rope_scale[1] ** (
    self.dim_h / (self.dim_h - 2)
)

max_size instance-attribute

max_size = [
    size // patch
    for size, patch in zip(max_size, patch_size)
]

patch_size instance-attribute

patch_size = patch_size

t_ntk_factor instance-attribute

t_ntk_factor = rope_scale[0] ** (
    self.dim_t / (self.dim_t - 2)
)

w_ntk_factor instance-attribute

w_ntk_factor = rope_scale[2] ** (
    self.dim_w / (self.dim_w - 2)
)

forward

forward(
    hidden_states: Tensor, fps: float | None = None
) -> tuple[Tensor, Tensor]

CosmosTimestepEmbedding

Bases: Module

activation instance-attribute

activation = nn.SiLU()

linear_1 instance-attribute

linear_1 = nn.Linear(in_features, out_features, bias=False)

linear_2 instance-attribute

linear_2 = nn.Linear(
    out_features, 3 * out_features, bias=False
)

forward

forward(timesteps: Tensor) -> Tensor

CosmosTransformerBlock

Bases: Module

after_proj instance-attribute

after_proj = (
    nn.Linear(hidden_size, hidden_size)
    if after_proj
    else None
)

attn1 instance-attribute

attn1 = CosmosAttention(
    query_dim=hidden_size,
    cross_attention_dim=None,
    heads=num_attention_heads,
    dim_head=attention_head_dim,
    out_bias=out_bias,
    prefix=f"{prefix}.attn1" if prefix else "attn1",
)

attn2 instance-attribute

attn2 = CosmosAttention(
    query_dim=hidden_size,
    cross_attention_dim=cross_attention_dim,
    heads=num_attention_heads,
    dim_head=attention_head_dim,
    out_bias=out_bias,
    img_context=img_context,
    prefix=f"{prefix}.attn2" if prefix else "attn2",
)

before_proj instance-attribute

before_proj = (
    nn.Linear(hidden_size, hidden_size)
    if before_proj
    else None
)

ff instance-attribute

ff = CosmosFeedForward(
    hidden_size, mult=mlp_ratio, bias=out_bias
)

img_context instance-attribute

img_context = img_context

norm1 instance-attribute

norm1 = CosmosAdaLayerNormZero(
    in_features=hidden_size, hidden_features=adaln_lora_dim
)

norm2 instance-attribute

norm2 = CosmosAdaLayerNormZero(
    in_features=hidden_size, hidden_features=adaln_lora_dim
)

norm3 instance-attribute

norm3 = CosmosAdaLayerNormZero(
    in_features=hidden_size, hidden_features=adaln_lora_dim
)

forward

forward(
    hidden_states: Tensor,
    encoder_hidden_states: Tensor
    | tuple[Tensor, Tensor | None],
    embedded_timestep: Tensor,
    temb: Tensor | None = None,
    image_rotary_emb: tuple[Tensor, Tensor] | None = None,
    extra_pos_emb: Tensor | None = None,
    attention_mask: Tensor
    | tuple[Tensor | None, Tensor | None]
    | None = None,
    controlnet_residual: Tensor | None = None,
    latents: Tensor | None = None,
    block_idx: int | None = None,
) -> Tensor | tuple[Tensor, Tensor]