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:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+99
View File
@@ -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))