Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
100 lines
3.8 KiB
Python
100 lines
3.8 KiB
Python
"""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))
|