Files
K3/tests/integration/test_k3_arch.py
T
dela 8442f92c58 Keep LatentMoE routed bmm in activation dtype under bf16 autocast
Python float scales and fp32 expert weights promoted SiTU outputs to
fp32, so index_add mixed BFloat16 dest with Float source and crashed
the 0.5b run. Cast packed weights and gate scalars to z.dtype.
2026-08-25 20:22:51 +08:00

305 lines
12 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.
"""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.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"