"""Causal-language-model behavior independent of toy memorization.""" import torch import torch.nn.functional as F from kda.layers.kda_attn import KDAAttention from kda.layers.swiglu import SwiGLUMLP from kda.models.causal_lm import CausalLM from kda.models.config import KDAConfig def _model(): torch.manual_seed(51) config = KDAConfig( hidden_size=16, num_hidden_layers=1, num_heads=2, num_value_heads=2, head_dim=4, chunk_size=4, vocab_size=32, intermediate_size=32, kda_backend="reference", ) return CausalLM(config).eval() def test_kda_schedule_and_unified_stem(): config = KDAConfig(num_hidden_layers=2) assert config.layer_specs() == [("kda", "swiglu"), ("kda", "swiglu")] model = CausalLM(config) assert isinstance(model.blocks[0].attn, KDAAttention) assert isinstance(model.blocks[0].ffn, SwiGLUMLP) def test_future_token_does_not_change_past_logits(): model = _model() first = torch.tensor([[1, 2, 3, 4]]) second = torch.tensor([[1, 2, 3, 9]]) with torch.no_grad(): first_logits = model(first) second_logits = model(second) torch.testing.assert_close(first_logits[:, :3], second_logits[:, :3]) def test_attention_reads_operator_flags_from_config(): torch.manual_seed(52) config = KDAConfig( hidden_size=16, num_hidden_layers=1, num_heads=2, num_value_heads=2, head_dim=4, chunk_size=4, vocab_size=32, intermediate_size=32, use_gate_in_kernel=False, use_qk_l2norm_in_kernel=False, use_beta_sigmoid_in_kernel=False, lower_bound=None, kda_backend="reference", ) model = CausalLM(config).eval() with torch.no_grad(): logits = model(torch.tensor([[1, 2, 3, 4]])) assert logits.shape == (1, 4, 32) def test_loss_is_shifted_next_token_cross_entropy(): model = _model() tokens = torch.tensor([[1, 2, 3, 4]]) with torch.no_grad(): logits = model(tokens) actual = model(tokens, labels=tokens) expected = F.cross_entropy( logits[:, :-1].reshape(-1, logits.size(-1)), tokens[:, 1:].reshape(-1), ) torch.testing.assert_close(actual, expected)