Files
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

177 lines
5.7 KiB
Python
Raw Permalink 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.
"""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"