Files
K3/kda/layers/latent_moe.py
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

119 lines
4.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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)