Skip to content

vllm_omni.diffusion.models.anima.anima_text_conditioner

ANIMA_TEXT_CONDITIONER_CONFIG module-attribute

ANIMA_TEXT_CONDITIONER_CONFIG = {
    "source_dim": 1024,
    "target_dim": 1024,
    "model_dim": 1024,
    "num_layers": 6,
    "num_attention_heads": 16,
    "mlp_ratio": 4.0,
    "target_vocab_size": 32128,
    "use_self_attention": True,
    "use_layer_norm": False,
    "min_sequence_length": 512,
}

AnimaRotaryEmbedding

Bases: Module

forward

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

AnimaTextConditioner

Bases: Module

blocks instance-attribute

blocks = nn.ModuleList(
    [
        AnimaTextConditionerBlock(
            source_dim=source_dim,
            model_dim=model_dim,
            num_attention_heads=num_attention_heads,
            mlp_ratio=mlp_ratio,
            use_self_attention=use_self_attention,
            use_layer_norm=use_layer_norm,
            prefix=f"blocks.{i}",
        )
        for i in range(num_layers)
    ]
)

config instance-attribute

config = SimpleNamespace(
    source_dim=source_dim,
    target_dim=target_dim,
    model_dim=model_dim,
    num_layers=num_layers,
    num_attention_heads=num_attention_heads,
    mlp_ratio=mlp_ratio,
    target_vocab_size=target_vocab_size,
    use_self_attention=use_self_attention,
    use_layer_norm=use_layer_norm,
    min_sequence_length=min_sequence_length,
    extra_config=kwargs,
)

dtype property

dtype: dtype

embed instance-attribute

embed = nn.Embedding(target_vocab_size, target_dim)

gradient_checkpointing instance-attribute

gradient_checkpointing = False

in_proj instance-attribute

in_proj = (
    nn.Linear(target_dim, model_dim)
    if model_dim != target_dim
    else nn.Identity()
)

norm instance-attribute

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

out_proj instance-attribute

out_proj = nn.Linear(model_dim, target_dim)

rotary_emb instance-attribute

rotary_emb = AnimaRotaryEmbedding(
    model_dim // num_attention_heads
)

forward

forward(
    source_hidden_states: Tensor,
    target_input_ids: Tensor,
    target_attention_mask: Tensor | None = None,
    source_attention_mask: Tensor | None = None,
) -> Tensor

load_weights

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

AnimaTextConditionerAttention

Bases: Module

attention_head_dim instance-attribute

attention_head_dim = attention_head_dim

attn instance-attribute

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

k_norm instance-attribute

k_norm = nn.RMSNorm(attention_head_dim, eps=1e-06)

k_proj instance-attribute

k_proj = nn.Linear(context_dim, inner_dim, bias=False)

num_attention_heads instance-attribute

num_attention_heads = num_attention_heads

o_proj instance-attribute

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

q_norm instance-attribute

q_norm = nn.RMSNorm(attention_head_dim, eps=1e-06)

q_proj instance-attribute

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

v_proj instance-attribute

v_proj = nn.Linear(context_dim, inner_dim, bias=False)

forward

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

AnimaTextConditionerBlock

Bases: Module

cross_attn instance-attribute

cross_attn = AnimaTextConditionerAttention(
    query_dim=model_dim,
    context_dim=source_dim,
    num_attention_heads=num_attention_heads,
    attention_head_dim=model_dim // num_attention_heads,
    prefix=f"{prefix}.cross_attn"
    if prefix
    else "cross_attn",
)

mlp instance-attribute

mlp = nn.Sequential(
    nn.Linear(model_dim, int(model_dim * mlp_ratio)),
    nn.GELU(),
    nn.Linear(int(model_dim * mlp_ratio), model_dim),
)

norm_cross_attn instance-attribute

norm_cross_attn = norm_cls(model_dim, **norm_kwargs)

norm_mlp instance-attribute

norm_mlp = norm_cls(model_dim, **norm_kwargs)

norm_self_attn instance-attribute

norm_self_attn = norm_cls(model_dim, **norm_kwargs)

self_attn instance-attribute

self_attn = AnimaTextConditionerAttention(
    query_dim=model_dim,
    context_dim=model_dim,
    num_attention_heads=num_attention_heads,
    attention_head_dim=model_dim // num_attention_heads,
    prefix=f"{prefix}.self_attn" if prefix else "self_attn",
)

use_self_attention instance-attribute

use_self_attention = use_self_attention

forward

forward(
    hidden_states: Tensor,
    encoder_hidden_states: Tensor,
    target_attention_mask: Tensor | None = None,
    source_attention_mask: Tensor | None = None,
    position_embeddings: tuple[Tensor, Tensor]
    | None = None,
    source_position_embeddings: tuple[Tensor, Tensor]
    | None = None,
) -> Tensor

is_conditioner_block_module

is_conditioner_block_module(
    name: str, module: Module
) -> bool