LatentMoE: K3 sigmoid routing and Switch aux/z-loss

Route with σ(W_r x), Top-k(s+b), then L1-normalize over the selected set.
Add Switch/GShard aux and router z-loss into train_k3 and train_sft.
Wiki parquet URLs honor HF_ENDPOINT for mirrored downloads.
This commit is contained in:
dela
2026-08-25 19:50:07 +08:00
parent d1da0816f2
commit 7a12f61de1
8 changed files with 381 additions and 82 deletions
+83 -11
View File
@@ -4,7 +4,7 @@ import torch.nn.functional as F
import pytest
from kda.layers.kda_attn import KDAAttention
from kda.layers.latent_moe import LatentMoE
from kda.layers.latent_moe import LatentMoE, moe_router_losses
from kda.layers.mla import GatedMLA
from kda.models.causal_lm import CausalLM
from kda.models.k3_config import K3Config
@@ -89,12 +89,9 @@ def test_preset_0_5b_schedule():
def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor:
"""Old dense path: run every routed expert, then gather top-k."""
"""Dense path: run every routed expert, then gather K3 sigmoid-norm top-k."""
z = moe.down(x)
logits = moe.router(x)
topk = torch.topk(logits, moe.top_k, dim=-1)
ids = topk.indices
probs = F.softmax(topk.values, dim=-1)
ids, probs = moe._route(moe.router(x))
all_out = torch.stack([expert(z) for expert in moe.experts])
B, T, _ = x.shape
all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, moe.n_routed, moe.latent_size)
@@ -113,20 +110,22 @@ def test_moe_router_activates_topk_only():
x = torch.randn(2, 6, 32)
with torch.no_grad():
y = moe(x)
logits = moe.router(x)
topk = torch.topk(logits, moe.top_k, dim=-1)
scores = torch.sigmoid(moe.router(x))
ids, probs = moe._route(moe.router(x))
z = moe.down(x)
# 手算: 只有 top-k 专家输出被加权, 再经 shared + up(norm(u))
expected_u = torch.zeros(2, 6, moe.latent_size)
all_out = torch.stack([e(z) for e in moe.experts]) # [R,B,T,ℓ]
probs = F.softmax(topk.values, dim=-1)
for i in range(moe.top_k):
idx = topk.indices[:, :, i]
idx = ids[:, :, i]
for b in range(2):
for t in range(6):
expected_u[b, t] += probs[b, t, i] * all_out[idx[b, t], b, t]
shared = torch.stack([e(x) for e in moe.shared]).sum(0)
expected_y = shared + moe.up(moe.norm(expected_u))
selected = scores.gather(-1, ids)
torch.testing.assert_close(
probs, selected / selected.sum(-1, keepdim=True).clamp_min(1e-9)
)
torch.testing.assert_close(y, expected_y, atol=1e-5, rtol=1e-5)
assert moe.last_route_ids is not None
assert moe.last_route_ids.shape[-1] == moe.top_k
@@ -188,6 +187,79 @@ def test_moe_unselected_experts_have_zero_grad():
assert torch.equal(param.grad, torch.zeros_like(param.grad))
def test_moe_aux_loss_penalizes_collapse():
torch.manual_seed(0)
moe = LatentMoE(
hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24,
aux_loss_coef=1.0, z_loss_coef=1.0,
)
moe(torch.randn(4, 16, 32))
spread = float(moe.last_aux_loss.detach())
with torch.no_grad():
moe.router.weight.zero_()
moe.router.weight[0] = 1.0
moe.router.weight[1] = 0.5
moe(torch.ones(4, 16, 32))
collapsed = float(moe.last_aux_loss.detach())
assert collapsed > spread
assert collapsed > 1.2
assert float(moe.last_z_loss.detach()) > 0
def test_moe_aux_loss_updates_router_only():
torch.manual_seed(1)
moe = LatentMoE(
hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24,
aux_loss_coef=1.0, z_loss_coef=1.0,
)
moe(torch.randn(2, 8, 32))
aux, z_loss = moe_router_losses(moe)
(aux + z_loss).backward()
assert moe.router.weight.grad is not None
assert moe.router.weight.grad.abs().sum() > 0
for expert in moe.experts:
assert expert.w_g.weight.grad is None
assert expert.w_u.weight.grad is None
assert expert.w_o.weight.grad is None
def test_moe_router_losses_sums_layers():
cfg = K3Config(
hidden_size=32, num_hidden_layers=2, num_heads=4, head_dim=8,
chunk_size=4, vocab_size=64, moe_latent_size=16, moe_d_ff=16,
n_routed=4, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8,
moe_aux_loss_coef=1.0, moe_z_loss_coef=1.0,
)
model = CausalLM(cfg)
model(torch.randint(0, cfg.vocab_size, (2, 8)))
aux, z_loss = moe_router_losses(model)
layers = [module for module in model.modules() if isinstance(module, LatentMoE)]
assert len(layers) == 2
torch.testing.assert_close(aux, layers[0].last_aux_loss + layers[1].last_aux_loss)
torch.testing.assert_close(z_loss, layers[0].last_z_loss + layers[1].last_z_loss)
def test_moe_aux_backward_with_checkpoint():
cfg = K3Config(
hidden_size=32, num_hidden_layers=2, num_heads=4, head_dim=8,
chunk_size=4, vocab_size=64, moe_latent_size=16, moe_d_ff=16,
n_routed=4, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8,
gradient_checkpointing=True, moe_aux_loss_coef=1.0, moe_z_loss_coef=1.0,
)
model = CausalLM(cfg)
model.train()
tokens = torch.randint(0, cfg.vocab_size, (2, 8))
task = model(tokens, labels=tokens)
aux, z_loss = moe_router_losses(model)
(task + aux + z_loss).backward()
router_grad = sum(
param.grad.abs().sum().item()
for name, param in model.named_parameters()
if "router" in name and param.grad is not None
)
assert router_grad > 0
def test_k3_causal_future_does_not_change_past_logits():
torch.manual_seed(51)
cfg = K3Config(hidden_size=64, num_hidden_layers=4, num_heads=4, head_dim=8,