"""KDA attention: project q/k/v/g/beta, run chunk_kda, project back to D.""" from __future__ import annotations import torch from torch import nn from ..ops.api import chunk_kda class KDAAttention(nn.Module): """Mixing module: x [B,T,D] -> y [B,T,D].""" def __init__( self, hidden_size: int, num_heads: int, num_value_heads: int, head_dim: int, *, chunk_size: int = 16, initializer_range: float = 0.02, use_gate_in_kernel: bool = True, use_qk_l2norm_in_kernel: bool = True, use_beta_sigmoid_in_kernel: bool = True, lower_bound: float | None = -5.0, kda_backend: str = "reference", ): super().__init__() if num_value_heads % num_heads: raise ValueError("num_value_heads must be divisible by num_heads") self.hidden_size = hidden_size self.num_heads = num_heads self.num_value_heads = num_value_heads self.head_dim = head_dim self.chunk_size = chunk_size self.initializer_range = initializer_range self.use_gate_in_kernel = use_gate_in_kernel self.use_qk_l2norm_in_kernel = use_qk_l2norm_in_kernel self.use_beta_sigmoid_in_kernel = use_beta_sigmoid_in_kernel self.lower_bound = lower_bound self.kda_backend = kda_backend H, HV, K, V = num_heads, num_value_heads, head_dim, head_dim self.q_proj = nn.Linear(hidden_size, H * K, bias=False) self.k_proj = nn.Linear(hidden_size, H * K, bias=False) self.v_proj = nn.Linear(hidden_size, HV * V, bias=False) self.g_proj = nn.Linear(hidden_size, HV * K, bias=False) self.beta_proj = nn.Linear(hidden_size, HV, bias=False) self.o_proj = nn.Linear(HV * V, hidden_size, bias=False) self.A_log = nn.Parameter(torch.zeros(HV)) # With safe_gate=-5, bias=-4 starts at g≈-0.09 (about 91% state retention). self.dt_bias = nn.Parameter(torch.full((HV, K), -4.0)) self.apply(self._init_weights) @classmethod def from_config(cls, config) -> KDAAttention: return cls( hidden_size=config.hidden_size, num_heads=config.num_heads, num_value_heads=getattr(config, "num_value_heads", config.num_heads), head_dim=config.head_dim, chunk_size=config.chunk_size, initializer_range=config.initializer_range, use_gate_in_kernel=config.use_gate_in_kernel, use_qk_l2norm_in_kernel=config.use_qk_l2norm_in_kernel, use_beta_sigmoid_in_kernel=config.use_beta_sigmoid_in_kernel, lower_bound=config.lower_bound, kda_backend=config.kda_backend, ) def _init_weights(self, module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, std=self.initializer_range) def forward(self, x: torch.Tensor): B, T, _ = x.shape H, HV, K, V = self.num_heads, self.num_value_heads, self.head_dim, self.head_dim q = self.q_proj(x).view(B, T, H, K) k = self.k_proj(x).view(B, T, H, K) v = self.v_proj(x).view(B, T, HV, V) g_raw = self.g_proj(x).view(B, T, HV, K) beta_raw = self.beta_proj(x).view(B, T, HV) o, _ = chunk_kda( q, k, v, g_raw, beta_raw, A_log=self.A_log, dt_bias=self.dt_bias, use_qk_l2norm_in_kernel=self.use_qk_l2norm_in_kernel, use_gate_in_kernel=self.use_gate_in_kernel, use_beta_sigmoid_in_kernel=self.use_beta_sigmoid_in_kernel, safe_gate=self.lower_bound is not None, lower_bound=self.lower_bound, chunk_size=self.chunk_size, backend=self.kda_backend, ) return self.o_proj(o.reshape(B, T, HV * V))