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