Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
177 lines
5.7 KiB
Python
177 lines
5.7 KiB
Python
"""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"
|