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,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"
|
||||
Reference in New Issue
Block a user