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
+51
View File
@@ -0,0 +1,51 @@
"""Decoder block: x += attn(norm(x)); x += ffn(norm(x)).
attn/ffn are any modules with forward: [B,T,D] -> [B,T,D].
"""
from __future__ import annotations
from torch import nn
from .kda_attn import KDAAttention
from .latent_moe import LatentMoE
from .mla import GatedMLA
from .rmsnorm import RMSNorm
from .swiglu import SwiGLUMLP
def build_attn(config, kind: str) -> nn.Module:
if kind == "kda":
return KDAAttention.from_config(config)
if kind == "mla":
return GatedMLA.from_config(config)
raise ValueError(f"unknown attn kind: {kind}")
def build_ffn(config, kind: str) -> nn.Module:
if kind == "swiglu":
return SwiGLUMLP.from_config(config)
if kind == "moe":
return LatentMoE.from_config(config)
raise ValueError(f"unknown ffn kind: {kind}")
class DecoderBlock(nn.Module):
def __init__(self, hidden_size: int, norm_eps: float, attn: nn.Module, ffn: nn.Module):
super().__init__()
self.attn_norm = RMSNorm(hidden_size, norm_eps)
self.attn = attn
self.ffn_norm = RMSNorm(hidden_size, norm_eps)
self.ffn = ffn
@classmethod
def from_spec(cls, config, attn_kind: str, ffn_kind: str) -> DecoderBlock:
return cls(
config.hidden_size,
config.norm_eps,
build_attn(config, attn_kind),
build_ffn(config, ffn_kind),
)
def forward(self, x):
x = x + self.attn(self.attn_norm(x))
return x + self.ffn(self.ffn_norm(x))