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:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+176
View File
@@ -0,0 +1,176 @@
"""AttnRes depth mixer: switch, no double-register, causality, two-phase match."""
from dataclasses import asdict
import pytest
import torch
from kda.layers.attn_res import BlockAttnResStack, FullAttnResStack, atomic_block_size
from kda.layers.kda_attn import KDAAttention
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
from kda.models.k3_config import K3Config
from kda.training.toy import load_ckpt, save_ckpt
def _tiny_k3(**kwargs):
defaults = dict(
hidden_size=32,
num_hidden_layers=4,
num_heads=4,
head_dim=8,
chunk_size=4,
vocab_size=64,
moe_latent_size=16,
moe_d_ff=16,
n_routed=4,
top_k=2,
n_shared=1,
kv_lora_rank=8,
q_lora_rank=16,
qk_nope_head_dim=8,
v_head_dim=8,
)
defaults.update(kwargs)
return K3Config(**defaults)
def test_default_attnres_is_off():
cfg = K3Config(num_hidden_layers=2)
assert cfg.attnres == "off"
model = CausalLM(cfg)
assert model.mixer is None
assert isinstance(model.blocks[0].attn, KDAAttention)
def test_invalid_attnres_rejected():
with pytest.raises(ValueError, match="attnres"):
K3Config(attnres="yes")
with pytest.raises(ValueError, match="attnres_block_size"):
K3Config(attnres="block", attnres_block_size=0)
@pytest.mark.parametrize("mode, stack_cls", [("block", BlockAttnResStack), ("full", FullAttnResStack)])
def test_mixer_kind_and_atomic_count(mode, stack_cls):
cfg = _tiny_k3(attnres=mode, attnres_block_size=2)
model = CausalLM(cfg)
assert model.attnres == mode
assert isinstance(model.mixer, stack_cls)
assert len(model.mixer.layers) == 2 * cfg.num_hidden_layers
if mode == "block":
assert model.mixer.block_size == 4 # 2 DecoderBlocks × attn|ffn
def test_auto_block_size_targets_eight_blocks():
assert atomic_block_size(24, None) == 6 # 3 DecoderBlocks × 2
assert atomic_block_size(4, None) == 2
assert atomic_block_size(93, None) == 24 # 12 DecoderBlocks × 2, K3 S=12
def test_no_duplicate_parameter_ids():
model = CausalLM(_tiny_k3(attnres="block"))
ids = [id(p) for p in model.parameters()]
assert len(ids) == len(set(ids))
names = [n for n, _ in model.named_parameters()]
assert len(names) == len(set(names))
residual_names = [n for n in names if "residuals" in n or "final_residual" in n]
assert residual_names
block_names = [n for n in names if n.startswith("blocks.")]
mixer_weight_names = [
n for n in names if n.startswith("mixer.layers.") and "query" not in n and "norm" not in n
]
assert block_names
assert mixer_weight_names == []
def test_off_and_block_differ_at_same_seed():
torch.manual_seed(0)
off = CausalLM(_tiny_k3(attnres="off"))
torch.manual_seed(0)
on = CausalLM(_tiny_k3(attnres="block"))
x = torch.tensor([[1, 2, 3, 4]])
with torch.no_grad():
assert not torch.allclose(off(x), on(x))
def test_block_two_phase_matches_naive():
torch.manual_seed(4)
model = CausalLM(_tiny_k3(attnres="block", attnres_block_size=2)).eval()
x = torch.randint(0, 64, (2, 8))
with torch.no_grad():
emb = model.embedding(x)
naive = model.mixer.forward_naive(emb)
two_phase = model.mixer(emb)
torch.testing.assert_close(naive, two_phase, atol=1e-5, rtol=1e-5)
@pytest.mark.parametrize("mode", ["block", "full"])
def test_attnres_is_still_causal(mode):
torch.manual_seed(51)
model = CausalLM(_tiny_k3(attnres=mode, num_hidden_layers=2)).eval()
with torch.no_grad():
a = model(torch.tensor([[1, 2, 3, 4]]))
b = model(torch.tensor([[1, 2, 3, 9]]))
torch.testing.assert_close(a[:, :3], b[:, :3], atol=1e-5, rtol=1e-5)
def test_kda_config_block_runs():
cfg = KDAConfig(
hidden_size=16,
num_hidden_layers=2,
num_heads=2,
num_value_heads=2,
head_dim=4,
chunk_size=4,
vocab_size=32,
intermediate_size=32,
attnres="block",
attnres_block_size=1,
)
model = CausalLM(cfg)
logits = model(torch.tensor([[1, 2, 3, 4]]))
assert logits.shape == (1, 4, 32)
def test_attnres_ckpt_roundtrip(tmp_path):
cfg = _tiny_k3(attnres="block", attnres_block_size=2)
model = CausalLM(cfg)
path = str(tmp_path / "attnres.pt")
save_ckpt(model, cfg, path)
loaded, loaded_cfg = load_ckpt(path)
assert loaded_cfg.attnres == "block"
assert loaded_cfg.attnres_block_size == 2
torch.manual_seed(1)
x = torch.randint(0, cfg.vocab_size, (2, 8))
with torch.no_grad():
torch.testing.assert_close(model(x), loaded(x), atol=1e-5, rtol=1e-5)
def test_old_ckpt_without_attnres_stays_off(tmp_path):
cfg = _tiny_k3()
payload = asdict(cfg)
payload.pop("attnres")
payload.pop("attnres_block_size")
payload.pop("attnres_zero_init_queries")
payload.pop("attnres_final_aggregate")
model = CausalLM(cfg)
path = str(tmp_path / "legacy.pt")
torch.save({"model_state": model.state_dict(), "config": payload}, path)
_, loaded_cfg = load_ckpt(path)
assert loaded_cfg.attnres == "off"
assert loaded_cfg.attnres_block_size is None
def test_attnres_block_overfits_single_batch():
torch.manual_seed(30)
cfg = _tiny_k3(attnres="block", num_hidden_layers=2, attnres_block_size=1)
model = CausalLM(cfg)
x = torch.randint(0, cfg.vocab_size, (2, 8))
optim = torch.optim.AdamW(model.parameters(), lr=3e-3)
final = None
for _ in range(200):
optim.zero_grad()
loss = model(x, labels=x)
loss.backward()
optim.step()
final = loss.item()
assert final < 0.5, f"final loss {final:.4f} >= 0.5"