"""K3 架构复现测试: MLA 吸收等价, LatentMoE 路由, hybrid pattern, 因果性, overfit.""" import torch import torch.nn.functional as F import pytest from kda.layers.kda_attn import KDAAttention 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 def _mla(d=64, H=4, r=16, q_r=32, d_q=16, d_v=16): torch.manual_seed(7) return GatedMLA(d, H, r, q_r, d_q, d_v) def _naive_mla(x, module: GatedMLA): """解压版参考: 标准 attention (吸收版数学上应与它逐位一致).""" B, T, _ = x.shape H, r = module.num_heads, module.kv_up.in_features c = module.kv_norm(module.kv_down(x)) # [B,T,r] q = module.q_up(module.q_norm(module.q_down(x))).view(B, T, H, module.qk_nope_head_dim) w = module.kv_up.weight w_uk = w[: H * module.qk_nope_head_dim].view(H, module.qk_nope_head_dim, r) w_uv = w[H * module.qk_nope_head_dim :].view(H, module.v_head_dim, r) k = torch.einsum("btj,hvj->bthv", c, w_uk) # 解压 K v = torch.einsum("btj,hvj->bthv", c, w_uv) # 解压 V scores = torch.einsum("bthv,bshv->bhts", q, k) # [B,H,T,T] mask = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1) scores = scores.masked_fill(mask, float("-inf")) attn = F.softmax(scores, dim=-1) o = torch.einsum("bhts,bshv->bthv", attn, v) # [B,T,H,d_v] o = o.reshape(B, T, H * module.v_head_dim) gate = torch.sigmoid(module.gate(x)) return module.o_proj(gate * o) def test_mla_absorption_matches_unrolled(): m = _mla().eval() x = torch.randn(3, 12, 64) with torch.no_grad(): absorbed = m(x) unrolled = _naive_mla(x, m) torch.testing.assert_close(absorbed, unrolled, atol=1e-5, rtol=1e-5) def test_mla_absorption_matches_unrolled_grad(): """吸收版与解压版的梯度也应一致 (fwd+bwd 双重验证).""" m1, m2 = _mla(), _mla() m2.load_state_dict(m1.state_dict()) x = torch.randn(2, 8, 64) l1 = m1(x).square().mean() l2 = _naive_mla(x, m2).square().mean() l1.backward() l2.backward() for (n1, p1), (n2, p2) in zip(m1.named_parameters(), m2.named_parameters()): torch.testing.assert_close(p1.grad, p2.grad, atol=1e-5, rtol=1e-5) def test_hybrid_layer_pattern(): cfg = K3Config(num_hidden_layers=4) assert cfg.layer_types() == ["kda", "kda", "kda", "mla"] cfg8 = K3Config(num_hidden_layers=8) assert cfg8.layer_types() == ["kda", "kda", "kda", "mla"] * 2 # 末层强制 MLA: L=5 → 层 3 MLA + 层 4 (末层) MLA cfg5 = K3Config(num_hidden_layers=5) assert cfg5.layer_types() == ["kda", "kda", "kda", "mla", "mla"] assert cfg.layer_specs() == [("kda", "moe"), ("kda", "moe"), ("kda", "moe"), ("mla", "moe")] model = CausalLM(K3Config(num_hidden_layers=4, hidden_size=32, moe_d_ff=16, moe_latent_size=16)) assert isinstance(model.blocks[0].attn, KDAAttention) assert isinstance(model.blocks[3].attn, GatedMLA) assert isinstance(model.blocks[0].ffn, LatentMoE) def test_preset_0_5b_schedule(): cfg = K3Config.preset("0.5b") assert cfg.hidden_size == 768 assert cfg.num_heads * cfg.head_dim == cfg.hidden_size assert cfg.num_hidden_layers == 24 assert cfg.vocab_size == 64000 assert cfg.tie_word_embeddings assert cfg.chunk_size == 64 assert cfg.gradient_checkpointing is True assert cfg.moe_latent_size == cfg.hidden_size // 2 types = cfg.layer_types() assert types.count("mla") == 6 assert types[-1] == "mla" assert cfg.layer_specs()[3] == ("mla", "moe") def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor: """Dense path: run every routed expert, then gather K3 sigmoid-norm top-k.""" z = moe.down(x) 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) u = z.new_zeros(B, T, moe.latent_size) for i in range(moe.top_k): idx = ids[:, :, i].reshape(B * T) sel = all_out[torch.arange(B * T, device=x.device), idx] u = u + probs[:, :, i : i + 1] * sel.reshape(B, T, moe.latent_size) shared = torch.stack([expert(x) for expert in moe.shared]).sum(0) return shared + moe.up(moe.norm(u)) def test_moe_router_activates_topk_only(): torch.manual_seed(3) moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) x = torch.randn(2, 6, 32) with torch.no_grad(): y = moe(x) scores = torch.sigmoid(moe.router(x)) ids, probs = moe._route(moe.router(x)) z = moe.down(x) expected_u = torch.zeros(2, 6, moe.latent_size) all_out = torch.stack([e(z) for e in moe.experts]) # [R,B,T,ℓ] for i in range(moe.top_k): 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 counts = torch.bincount(moe.last_route_ids.reshape(-1), minlength=moe.n_routed) assert moe.last_capacity == int(counts.max()) def test_moe_sparse_matches_dense_fwd_bwd(): torch.manual_seed(3) moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) x = torch.randn(2, 6, 32) y_sparse = moe(x) y_dense = _dense_moe_forward(moe, x) torch.testing.assert_close(y_sparse, y_dense, atol=1e-5, rtol=1e-5) moe.zero_grad(set_to_none=True) xs = x.detach().requires_grad_(True) moe(xs).square().mean().backward() grads_s = { name: param.grad.detach().clone() for name, param in moe.named_parameters() if param.grad is not None } dx_s = xs.grad.detach().clone() moe.zero_grad(set_to_none=True) xd = x.detach().requires_grad_(True) _dense_moe_forward(moe, xd).square().mean().backward() grads_d = { name: param.grad.detach().clone() for name, param in moe.named_parameters() if param.grad is not None } dx_d = xd.grad.detach().clone() torch.testing.assert_close(dx_s, dx_d, atol=1e-5, rtol=1e-5) assert grads_s.keys() == grads_d.keys() for name in grads_s: torch.testing.assert_close(grads_s[name], grads_d[name], atol=1e-5, rtol=1e-5) def test_moe_unselected_experts_have_zero_grad(): torch.manual_seed(0) moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) with torch.no_grad(): moe.router.weight.zero_() moe.router.weight[0] = 1.0 moe.router.weight[1] = 0.5 x = torch.ones(2, 4, 32) moe(x).square().mean().backward() selected = set(moe.last_route_ids.reshape(-1).tolist()) assert selected == {0, 1} for idx, expert in enumerate(moe.experts): for param in (expert.w_g.weight, expert.w_u.weight, expert.w_o.weight): assert param.grad is not None if idx in selected: assert param.grad.abs().sum() > 0 else: assert torch.equal(param.grad, torch.zeros_like(param.grad)) def test_routed_u_bf16_index_add_matches_z_dtype(): """Python float * bf16 promotes to fp32; index_add must still land in z.dtype.""" torch.manual_seed(0) moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) z = torch.randn(2, 6, 16, dtype=torch.bfloat16) logits = torch.randn(2, 6, 8, dtype=torch.bfloat16) ids, probs = moe._route(logits) u = moe._routed_u(z, ids, probs) assert u.dtype == torch.bfloat16 u.float().square().mean().backward() assert any(p.grad is not None and p.grad.abs().sum() > 0 for p in moe.experts[0].parameters()) 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, chunk_size=4, vocab_size=64, moe_latent_size=32, moe_d_ff=24, n_routed=8, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8) m = CausalLM(cfg).eval() with torch.no_grad(): a = m(torch.tensor([[1, 2, 3, 4]])) b = m(torch.tensor([[1, 2, 3, 9]])) torch.testing.assert_close(a[:, :3], b[:, :3], atol=1e-6, rtol=0) def test_k3_small_model_overfits_single_batch(): """K3 混合架构单 batch overfit 冒烟: loss < 0.5 (收敛即架构可训).""" torch.manual_seed(30) cfg = K3Config(hidden_size=64, num_hidden_layers=2, num_heads=4, head_dim=8, chunk_size=4, vocab_size=64, moe_latent_size=32, moe_d_ff=24, n_routed=8, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8) m = CausalLM(cfg) x = torch.randint(0, cfg.vocab_size, (2, 16)) optim = torch.optim.AdamW(m.parameters(), lr=3e-3) final = None for step in range(200): optim.zero_grad() loss = m(x, labels=x) loss.backward() optim.step() final = loss.item() assert final < 0.5, f"final loss {final:.4f} >= 0.5"