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