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