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))