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,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))
|
||||
Reference in New Issue
Block a user