Files
K3/tests/integration/test_k3_arch.py
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

146 lines
6.0 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
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 test_moe_router_activates_topk_only():
from kda.layers.latent_moe import LatentMoE
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)
logits = moe.router(x)
topk = torch.topk(logits, moe.top_k, dim=-1)
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]
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))
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
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"