Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
@@ -0,0 +1,99 @@
|
||||
"""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))
|
||||
Reference in New Issue
Block a user