Skip to content

dreamx_world

Classes

fastvideo.models.dits.dreamx_world.DreamXPropeSelfAttention

DreamXPropeSelfAttention(dim: int, attn_dim: int, num_heads: int, qk_norm: str | bool = True, eps: float = 1e-06, quant_config: QuantizationConfig | None = None, prefix: str = '')

Bases: Module

DreamX-World parallel PRoPE camera self-attention branch.

Source code in fastvideo/models/dits/dreamx_world.py
def __init__(self,
             dim: int,
             attn_dim: int,
             num_heads: int,
             qk_norm: str | bool = True,
             eps: float = 1e-6,
             quant_config: QuantizationConfig | None = None,
             prefix: str = ""):
    super().__init__()
    assert attn_dim % num_heads == 0
    self.attn_dim = attn_dim
    self.num_heads = num_heads
    self.head_dim = attn_dim // num_heads
    self.qk_norm = qk_norm

    self.q_proj = ReplicatedLinear(dim,
                                   attn_dim,
                                   quant_config=quant_config,
                                   prefix=f"{prefix}.q_proj")
    self.k_proj = ReplicatedLinear(dim,
                                   attn_dim,
                                   quant_config=quant_config,
                                   prefix=f"{prefix}.k_proj")
    self.v_proj = ReplicatedLinear(dim,
                                   attn_dim,
                                   quant_config=quant_config,
                                   prefix=f"{prefix}.v_proj")
    self.out_proj = ReplicatedLinear(attn_dim,
                                     dim,
                                     quant_config=quant_config,
                                     prefix=f"{prefix}.out_proj")

    if qk_norm == "rms_norm":
        self.norm_q = RMSNorm(self.head_dim, eps=eps)
        self.norm_k = RMSNorm(self.head_dim, eps=eps)
    elif qk_norm in (True, "rms_norm_across_heads"):
        self.norm_q = RMSNorm(attn_dim, eps=eps)
        self.norm_k = RMSNorm(attn_dim, eps=eps)
    elif qk_norm is False:
        self.norm_q = nn.Identity()
        self.norm_k = nn.Identity()
    else:
        raise ValueError(f"Unsupported qk_norm for DreamX PRoPE: {qk_norm}")

    nn.init.zeros_(self.out_proj.weight)
    if self.out_proj.bias is not None:
        nn.init.zeros_(self.out_proj.bias)

    self.attn = LocalAttention(
        num_heads=num_heads,
        head_size=self.head_dim,
        dropout_rate=0,
        softmax_scale=None,
        causal=False,
        supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
                                      AttentionBackendEnum.TORCH_SDPA))

Functions: