CLI default attnres=off was overwriting block checkpoints on resume so later loads hit Unexpected key(s). Only apply flags the user passed. Track next_chunk so a budget-exit save does not skip the untrained yield. 0.5b now uses 01-ai/Yi-6B (64k); refuse resume when the ckpt tokenizer does not match.
306 lines
12 KiB
Python
306 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.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"
|