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
+23
View File
@@ -0,0 +1,23 @@
"""Composable mixing layers: attn and ffn both map [B,T,D] -> [B,T,D].
Depth mixing (AttnRes) is not a layer_specs kind. CausalLM reads
``config.attnres`` (off | full | block) and wraps DecoderBlock sublayers.
"""
from .block import DecoderBlock, build_attn, build_ffn
from .kda_attn import KDAAttention
from .latent_moe import LatentMoE
from .mla import GatedMLA
from .rmsnorm import RMSNorm
from .swiglu import SwiGLUMLP
__all__ = [
"DecoderBlock",
"GatedMLA",
"KDAAttention",
"LatentMoE",
"RMSNorm",
"SwiGLUMLP",
"build_attn",
"build_ffn",
]
+519
View File
@@ -0,0 +1,519 @@
"""
Attention Residual in one file
Reference:
Kimi Team, Guangyu Chen, Yu Zhang, Jianlin Su, Weixin Xu, Siyuan Pan,
Yaoyu Wang, Yucheng Wang, Guanduo Chen, et al.
"Attention Residuals." arXiv:2603.15031, 2026.
https://arxiv.org/abs/2603.15031
This module is a compact PyTorch reference implementation of:
- Full AttnRes
- Block AttnRes
- two-phase inter/intra-block computation from the paper
CausalLM wires Full/Block stacks when ``config.attnres`` is ``full`` or
``block``. Standard residual (``x += attn; x += ffn``) is ``attnres="off"``.
"""
import torch
import torch.nn.functional as F
from einops import rearrange
from torch import Tensor, nn
ATTNRES_MODES = ("off", "full", "block")
def exists(x):
return x is not None
def validate_attnres(mode: str, block_size: int | None) -> None:
if mode not in ATTNRES_MODES:
raise ValueError(f"attnres must be one of {ATTNRES_MODES}, got {mode!r}")
if block_size is not None and block_size < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {block_size}")
def atomic_block_size(num_hidden_layers: int, attnres_block_size: int | None) -> int:
"""DecoderBlocks per AttnRes block, converted to attn|ffn atomic layers.
``None`` targets about 8 blocks: ``max(1, ceil(L / 8))`` DecoderBlocks.
"""
layers_per_block = (
attnres_block_size
if attnres_block_size is not None
else max(1, (num_hidden_layers + 7) // 8)
)
if layers_per_block < 1:
raise ValueError(f"attnres_block_size must be >= 1, got {layers_per_block}")
return layers_per_block * 2
class BorrowedSubLayer(nn.Module):
"""``fn(norm(x))`` without registering ``norm``/``fn`` (owned by DecoderBlock)."""
def __init__(self, norm: nn.Module, fn: nn.Module):
super().__init__()
self._borrowed = (norm, fn)
def forward(self, x: Tensor) -> Tensor:
norm, fn = self._borrowed
return fn(norm(x))
def rms(x: Tensor, eps: float):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + eps)
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: Tensor) -> Tensor:
return rms(x, self.eps) * self.weight
class DepthResidual(nn.Module):
"""
h_l = sum_i softmax_i(w_l^T RMSNorm(v_i))*v_i
Keep query and RMSNorm gain separate
Since q^T (gamma * RMS(v)) == (q * gamma)^T RMS(v),
we can fold gamma into q for scoring.
"""
def __init__(self, dim: int, eps: float = 1e-8, zero_init: bool = True):
super().__init__()
self.query = nn.Parameter(torch.zeros(dim))
self.norm = RMSNorm(dim, eps=eps)
if not zero_init:
nn.init.normal_(self.query, std=0.02)
def effective_query(self) -> Tensor:
return (self.query * self.norm.weight).float()
def logits(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources) # [n, b, t, d]
q = self.effective_query() # [d]
k = rms(sources.float(), self.norm.eps) # [n, b, t, d]
return torch.einsum("d, n b t d -> n b t", q, k)
def forward(self, sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
sources = stack_layers(sources)
weights = self.logits(sources).softmax(dim=0)
out = torch.einsum("n b t, n b t d -> b t d", weights, sources.float())
return out.to(sources.dtype)
class DepthResidualList(nn.Module):
def __init__(self, dim: int, depth: int, eps: float, zero_init: bool = True):
super().__init__()
# for L layers (depth), create depth residual modules
self.layers = nn.ModuleList(
[DepthResidual(dim, eps=eps, zero_init=zero_init) for _ in range(depth)]
)
def __getitem__(self, idx: int) -> DepthResidual:
return self.layers[idx]
def __iter__(self):
return iter(self.layers)
def __len__(self):
return len(self.layers)
# attnres stacks
class FullAttnResStack(nn.Module):
"""
Full AttnRes over atomic layers
eg: f_1,...,f_L
Each entry in `layers` should already be a full atomic layer fxn
x -> f_l(x)
"""
def __init__(
self,
dim: int,
layers,
*,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(dim, depth, eps, zero_init_queries)
self.final_residual = (
DepthResidual(dim, eps, zero_init_queries) if is_final_aggregate else None
)
def forward_naive(self, x: Tensor) -> Tensor:
sources = [x]
for layer, residual in zip(self.layers, self.residuals):
h = residual(sources)
out = layer(h)
sources.append(out)
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward_two_phase(self, x: Tensor, schedule_block_size: int) -> Tensor:
assert schedule_block_size > 0
sources = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + schedule_block_size, depth)
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(sources)
inter_stats = attn_with_stats(queries, inter_sources, self.eps)
local_outputs = [] # outputs of intra-block
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter_stats.select(local_idx)
if len(local_outputs) > 0:
intra_sources = stack_layers(local_outputs)
intra = attn_with_stats(
queries[local_idx : local_idx + 1], intra_sources, self.eps
).select(0)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
local_outputs.append(out)
sources.append(out)
start = end
return (
self.final_residual(sources) if exists(self.final_residual) else sources[-1]
)
def forward(self, x: Tensor, schedule_block_size: int | None = None) -> Tensor:
if schedule_block_size is None:
return self.forward_naive(x)
return self.forward_two_phase(x, schedule_block_size)
class BlockAttnResStack(nn.Module):
"""
Block AttnRes over atomic layers
`block_size` is in atomic layers, not Transformer blocks.
Eg: block_size=8 -> 4 transformer blocks when layers alternate attn/MLP
The default forward path is the two-phase algorithm from the paper:
phase 1: batch inter-block attn from all queries in the block
phase 2: merge the evolving intra-block partial sum with online softmax
"""
def __init__(
self,
dim: int,
layers,
*,
block_size: int,
eps: float = 1e-8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
):
super().__init__()
self.layers = nn.ModuleList(list(layers))
assert len(self.layers) > 0
assert block_size > 0
self.block_size = block_size
self.eps = eps
depth = len(self.layers)
self.residuals = DepthResidualList(
dim, depth, eps=eps, zero_init=zero_init_queries
)
self.final_residual = (
DepthResidual(dim, eps=eps, zero_init=zero_init_queries)
if is_final_aggregate
else None
)
def forward_naive(self, x: Tensor) -> Tensor:
blocks = [x] # b_0=embedding/input representation
partial = None
for layer_idx, (layer, residual) in enumerate(
zip(self.layers, self.residuals), start=1
):
sources = blocks if partial is None else blocks + [partial]
h = residual(sources)
out = layer(h)
partial = out if partial is None else (partial + out)
if (layer_idx % self.block_size == 0) or (layer_idx == len(self.layers)):
blocks.append(partial)
partial = None
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
def _run_block_two_phase(
self, blocks: list[Tensor], start: int, end: int
) -> Tensor:
queries = torch.stack(
[self.residuals[i].effective_query() for i in range(start, end)], dim=0
)
inter_sources = stack_layers(blocks)
inter = attn_with_stats(queries, inter_sources, self.eps)
partial = None
for local_idx, layer_idx in enumerate(range(start, end)):
stats = inter.select(local_idx)
if partial is not None:
intra = single_source_stats(queries[local_idx], partial, self.eps)
stats = merge_attn_stats(stats, intra)
h = stats.normalized()
out = self.layers[layer_idx](h)
partial = out if partial is None else (partial + out)
return partial
def forward(self, x: Tensor) -> Tensor:
blocks = [x]
depth = len(self.layers)
start = 0
while start < depth:
end = min(start + self.block_size, depth)
blocks.append(self._run_block_two_phase(blocks, start, end))
start = end
return (
self.final_residual(blocks) if exists(self.final_residual) else blocks[-1]
)
# helpers
def stack_layers(sources: Tensor | list[Tensor] | tuple[Tensor, ...]) -> Tensor:
if isinstance(sources, Tensor):
assert sources.ndim == 4, f"expected [n, b, t, d] got {tuple(sources.shape)}"
return sources
assert len(sources) > 0, "needs at least one source"
return torch.stack(tuple(sources), dim=0)
class SingleAttnStats:
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [b,t,d]
self.max = max # [b,t]
self.denom = denom # [b,t]
def normalized(self) -> Tensor:
return self.numer / self.denom[..., None]
class AttnStats:
# store the numerator => e^{s_{j}-m} * v_j where m is the max score so far
# store the max m = max(s_j)
# store the denominator sum_j e^{s_{j}-m}
def __init__(self, numer: Tensor, denom: Tensor, max: Tensor):
self.numer = numer # [q,b,t,d]
self.max = max # [q,b,t]
self.denom = denom # [q,b,t]
def select(self, idx: int) -> "SingleAttnStats":
return SingleAttnStats(self.numer[idx], self.denom[idx], self.max[idx])
def attn_with_stats(queries: Tensor, sources: Tensor, eps: float = 1e-8) -> AttnStats:
"""
queries: [q, d]
sources: [n, b, t, d]
Returns the following for online softmax:
numer = sum_i exp(logit_i - m)*v_i
m = max_i logit_i
denom = sum_i exp(logit_i - m)
"""
normed = rms(sources, eps)
logits = torch.einsum("q d, n b t d -> q n b t", queries, normed)
m = logits.amax(dim=1)
weights = torch.exp(logits - m[:, None])
numer = torch.einsum("q n b t, n b t d -> q b t d", weights, sources)
denom = weights.sum(dim=1)
return AttnStats(numer, denom, m)
def single_source_stats(
query: Tensor, source: Tensor, eps: float = 1e-8
) -> SingleAttnStats:
score = torch.einsum("d, b t d -> b t", query, rms(source, eps))
denom = torch.ones_like(score)
return SingleAttnStats(source, denom, score)
def merge_attn_stats(a: SingleAttnStats, b: SingleAttnStats) -> SingleAttnStats:
m = torch.maximum(a.max, b.max)
wa = torch.exp(a.max - m)
wb = torch.exp(b.max - m)
numer = wa[..., None] * a.numer + wb[..., None] * b.numer
denom = wa * a.denom + wb * b.denom
return SingleAttnStats(numer, denom, m)
# transformer
class PreNorm(nn.Module):
def __init__(self, dim: int, fn: nn.Module, eps: float = 1e-8):
super().__init__()
self.norm = RMSNorm(dim, eps=eps)
self.fn = fn
def forward(self, x: Tensor) -> Tensor:
return self.fn(self.norm(x))
class CausalAttention(nn.Module):
def __init__(
self, dim: int, heads: int = 8, dim_head: int = 64, dropout: float = 0.0
):
super().__init__()
inner_dim = heads * dim_head
self.heads = heads
self.dim_head = dim_head
self.dropout = dropout
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
def forward(self, x: Tensor) -> Tensor:
q, k, v = self.to_qkv(x).chunk(3, dim=-1)
def split_heads(y: Tensor) -> Tensor:
return rearrange(y, "b t (h d) -> b h t d", h=self.heads)
q, k, v = map(split_heads, (q, k, v))
out = F.scaled_dot_product_attention(
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
)
out = rearrange(out, "b h t d -> b t (h d)")
return self.to_out(out)
class SwiGLU(nn.Module):
def __init__(self, dim: int, mult: int = 4, dropout: float = 0.0):
# dropout not needed unless training on a smaller training data
super().__init__()
inner_dim = dim * mult
self.to_hidden = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
self.dropout = nn.Dropout(dropout)
def forward(self, x: Tensor) -> Tensor:
gate, value = self.to_hidden(x).chunk(2, dim=-1)
x = F.silu(gate) * value
x = self.dropout(x)
return self.to_out(x)
class AttnResTransformer(nn.Module):
"""
Small GPT-style reference model using AttnRes
Using plain PyTorch: tok/pos embedding, alternating
causal attn, SwiGLU MLP layers, final norm, output head.
"""
def __init__(
self,
*,
num_tokens: int,
dim: int,
depth: int,
max_seq_len: int,
heads: int = 8,
dim_head: int = 64,
ff_mult: int = 4,
attn_dropout: float = 0.0,
ff_dropout: float = 0.0,
attnres: str = "block", # full or block
block_size: int = 8,
zero_init_queries: bool = True,
is_final_aggregate: bool = True,
eps: float = 1e-8,
):
super().__init__()
assert attnres in {"full", "block"}
self.max_seq_len = max_seq_len
self.attnres = attnres
self.token_emb = nn.Embedding(num_tokens, dim)
self.pos_emb = nn.Embedding(max_seq_len, dim)
atomic_layers = []
for _ in range(depth):
atomic_layers.append(
PreNorm(dim, CausalAttention(dim, heads, dim_head, attn_dropout), eps)
)
atomic_layers.append(PreNorm(dim, SwiGLU(dim, ff_mult, ff_dropout), eps))
if attnres == "full":
self.backbone = FullAttnResStack(
dim,
atomic_layers,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
else:
self.backbone = BlockAttnResStack(
dim,
atomic_layers,
block_size=block_size,
eps=eps,
zero_init_queries=zero_init_queries,
is_final_aggregate=is_final_aggregate,
)
self.final_norm = RMSNorm(dim, eps)
self.to_logits = nn.Linear(dim, num_tokens, bias=False)
def forward(self, ids: Tensor, schedule_block_size: int | None = None) -> Tensor:
b, t = ids.shape
assert t <= self.max_seq_len
pos = torch.arange(t, device=ids.device)
x = self.token_emb(ids) + self.pos_emb(pos)[None, :, :]
if self.attnres == "full":
x = self.backbone(x, schedule_block_size=schedule_block_size)
else:
x = self.backbone(x)
x = self.final_norm(x)
return self.to_logits(x)
__all__ = [
"ATTNRES_MODES",
"RMSNorm",
"DepthResidual",
"DepthResidualList",
"FullAttnResStack",
"BlockAttnResStack",
"BorrowedSubLayer",
"PreNorm",
"CausalAttention",
"SwiGLU",
"AttnResTransformer",
"atomic_block_size",
"validate_attnres",
]
+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))
+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))
+118
View File
@@ -0,0 +1,118 @@
"""Stable LatentMoE (K3): shared 全宽 + routed 半宽专家 + SiTU-GLU + Top-k.
对照 learning/kimi-k3-notes §Stable LatentMoE:
z = W_down(x) [B, T, ℓ] ℓ = d/2 latent 接口宽
u = Σ_{i∈Top-k(x)} p_i E_i^rt(z) [B, T, ℓ] routed 专家只在 ℓ 上算
y = Σ_j E_j^sh(x) + W_up RMSNorm(u) [B, T, d] shared 全宽
SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2)
||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出.
E: R^in → R^in (内部中间维 d_ff).
Router: Top-k logits 基于全宽 x (笔记 Topk(x)); 归一化权重取 softmax(topk).
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
class SiTU(nn.Module):
"""SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维."""
def __init__(self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0):
super().__init__()
self.beta1, self.beta2 = beta1, beta2
self.w_g = nn.Linear(dim_in, dim_ff, bias=False)
self.w_u = nn.Linear(dim_in, dim_ff, bias=False)
self.w_o = nn.Linear(dim_ff, dim_in, bias=False)
def forward(self, x: torch.Tensor):
wg = self.w_g(x)
g = self.beta1 * torch.tanh(wg / self.beta1) * torch.sigmoid(wg)
u = self.beta2 * torch.tanh(self.w_u(x) / self.beta2)
return self.w_o(g * u)
class LatentMoE(nn.Module):
def __init__(
self,
hidden_size: int,
latent_size: int,
n_routed: int,
top_k: int,
n_shared: int,
d_ff: int,
beta1: float = 4.0,
beta2: float = 25.0,
):
super().__init__()
self.latent_size = latent_size
self.n_routed = n_routed
self.top_k = top_k
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
self.router = nn.Linear(hidden_size, n_routed, bias=False) # Top-k logits
self.shared = nn.ModuleList(
[SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)]
)
self.experts = nn.ModuleList(
[SiTU(latent_size, d_ff, beta1, beta2) for _ in range(n_routed)]
)
self.norm = RMSNorm(latent_size)
self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑
self.last_route_ids: torch.Tensor | None = None
@classmethod
def from_config(cls, config) -> LatentMoE:
return cls(
config.hidden_size,
config.moe_latent_size,
config.n_routed,
config.top_k,
config.n_shared,
config.moe_d_ff,
config.situ_beta1,
config.situ_beta2,
)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
z = self.down(x) # [B, T, ℓ]
logits = self.router(x) # [B, T, n_routed]
topk = torch.topk(logits, self.top_k, dim=-1)
ids = topk.indices # [B, T, k]
self.last_route_ids = ids.detach()
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
# 向量化 routed: 预计算全部专家输出, 按 token 的 Top-k id 取
all_out = torch.stack([e(z) for e in self.experts]) # [R, B, T, ℓ]
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, self.n_routed, self.latent_size)
u = torch.zeros(B, T, self.latent_size, device=x.device, dtype=x.dtype)
for i in range(self.top_k):
idx = ids[:, :, i].reshape(B * T) # [B*T]
sel = all_out[torch.arange(B * T, device=x.device), idx] # [B*T, ℓ]
u += probs[:, :, i : i + 1] * sel.reshape(B, T, self.latent_size)
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
return shared_out + self.up(self.norm(u))
def moe_route_frac(model: nn.Module) -> torch.Tensor | None:
"""Mean expert occupancy over LatentMoE layers from the last forward."""
hists: list[torch.Tensor] = []
n_routed: int | None = None
for module in model.modules():
if not isinstance(module, LatentMoE) or module.last_route_ids is None:
continue
n_routed = module.n_routed
ids = module.last_route_ids.reshape(-1)
hists.append(torch.bincount(ids, minlength=n_routed).float())
if not hists or n_routed is None:
return None
stacked = torch.stack(hists).sum(0)
return stacked / stacked.sum().clamp_min(1.0)
+97
View File
@@ -0,0 +1,97 @@
"""Gated MLA (K3): NoPE, latent KV compression, matrix absorption, full-rank output gate.
K3 相对 DeepSeek MLA 的三个改动 (对照 learning/kimi-k3-notes):
1. NoPE — 不显式 RoPE; 位置感交给夹层 KDA 的 decay/gate。
2. 矩阵吸收 — 训练/推理都不解压 K/V: q 吸收 W_UK 后直接与 latent c 内积,
输出先在 latent 加权再乘 W_UV 还原 (v2 吸收版)。
3. Full-rank 输出门 — y = W_o[ σ(W_g x) ⊙ õ ]。
形状 (小规模 toy, d 为 hidden):
c = RMSNorm(kv_down(x)) [B, T, r] latent
q = q_up(RMSNorm(q_down(x))) [B, T, H, d_q] d_q = d_nope (NoPE)
W_UK = kv_up[.., :H*d_q].view(H,d_q,r) W_UV = kv_up[.., H*d_q:].view(H,d_v,r)
score = (q @ W_UK^T) @ c^T [B, H, T, T] causal
õ = (softmax(score) @ c) @ W_UV^T [B, T, H, d_v]
y = o_proj( σ(W_g x) ⊙ õ_head ) [B, T, d]
"""
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
class GatedMLA(nn.Module):
def __init__(
self,
hidden_size: int,
num_heads: int,
kv_lora_rank: int,
q_lora_rank: int,
qk_nope_head_dim: int,
v_head_dim: int,
):
super().__init__()
self.hidden_size = hidden_size
self.num_heads = num_heads
self.qk_nope_head_dim = qk_nope_head_dim
self.v_head_dim = v_head_dim
# Q 低秩路径 (NoPE, 只有 nope 段)
self.q_down = nn.Linear(hidden_size, q_lora_rank, bias=False)
self.q_norm = RMSNorm(q_lora_rank)
self.q_up = nn.Linear(q_lora_rank, num_heads * qk_nope_head_dim, bias=False)
# KV latent 压缩 + 解压 (W_UK | W_UV 拼接在同一矩阵里)
self.kv_down = nn.Linear(hidden_size, kv_lora_rank, bias=False)
self.kv_norm = RMSNorm(kv_lora_rank)
self.kv_up = nn.Linear(
kv_lora_rank, num_heads * (qk_nope_head_dim + v_head_dim), bias=False
)
# Full-rank 输出门: σ(W_g x) 与 õ (H*d_v 维) 逐元素相乘
self.gate = nn.Linear(hidden_size, num_heads * v_head_dim, bias=False)
self.o_proj = nn.Linear(num_heads * v_head_dim, hidden_size, bias=False)
@classmethod
def from_config(cls, config) -> GatedMLA:
return cls(
config.hidden_size,
config.num_heads,
config.kv_lora_rank,
config.q_lora_rank,
config.qk_nope_head_dim,
config.v_head_dim,
)
def forward(self, x: torch.Tensor):
B, T, _ = x.shape
H, r = self.num_heads, self.kv_up.in_features
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
w = self.kv_up.weight # [H*(d_q+d_v), r]
w_uk = w[: H * self.qk_nope_head_dim].view(H, self.qk_nope_head_dim, r)
w_uv = w[H * self.qk_nope_head_dim :].view(H, self.v_head_dim, r)
# 吸收 W_UK 进 query: score = (q @ W_UK^T) @ c^T
q_absorb = torch.einsum("bthd,hdj->bthj", q, w_uk) # [B, T, H, r]
scores = torch.einsum("bthj,bsj->bhts", q_absorb, c) # [B, H, T, T]
mask = torch.triu(
torch.ones(T, T, dtype=torch.bool, device=x.device), diagonal=1
)
scores = scores.masked_fill(mask, float("-inf"))
attn = F.softmax(scores, dim=-1) # [B, H, T, T]
# 先在 latent 加权, 再乘 W_UV^T 还原 v —— 永不解压
latent_out = torch.einsum("bhts,bsj->bhtj", attn, c) # [B, H, T, r]
o_heads = torch.einsum("bhtj,hvj->bhtv", latent_out, w_uv) # [B, H, T, d_v]
o_heads = o_heads.transpose(1, 2).reshape(B, T, H * self.v_head_dim)
gate = torch.sigmoid(self.gate(x)) # [B, T, H*d_v]
return self.o_proj(gate * o_heads) # [B, T, d]
+17
View File
@@ -0,0 +1,17 @@
"""RMSNorm used by attention, FFN, and the final LM stem."""
from __future__ import annotations
import torch
from torch import nn
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(dim))
self.eps = eps
def forward(self, x: torch.Tensor):
dtype = x.dtype
x = x.float()
return (x * torch.rsqrt(x.square().mean(-1, keepdim=True) + self.eps)).to(dtype) * self.weight
+20
View File
@@ -0,0 +1,20 @@
"""SwiGLU FFN: x [B,T,D] -> y [B,T,D]."""
from __future__ import annotations
import torch.nn.functional as F
from torch import nn
class SwiGLUMLP(nn.Module):
def __init__(self, hidden_size: int, intermediate_size: int):
super().__init__()
self.w1 = nn.Linear(hidden_size, intermediate_size, bias=False)
self.w3 = nn.Linear(hidden_size, intermediate_size, bias=False)
self.w2 = nn.Linear(intermediate_size, hidden_size, bias=False)
@classmethod
def from_config(cls, config) -> SwiGLUMLP:
return cls(config.hidden_size, config.intermediate_size)
def forward(self, x):
return self.w2(F.silu(self.w1(x)) * self.w3(x))