"""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"