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.
This commit is contained in:
@@ -0,0 +1,145 @@
|
||||
"""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"
|
||||
Reference in New Issue
Block a user