import torch import torch.nn.functional as F from kda.models.causal_lm import CausalLM from kda.models.config import KDAConfig from kda.models.k3_config import K3Config def _tiny(): torch.manual_seed(4) 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, kda_backend="reference", ) return CausalLM(cfg), cfg def test_ignore_index_skips_masked_positions(): model, _ = _tiny() tokens = torch.tensor([[1, 2, 3, 4]]) labels = tokens.clone() labels[:, 1:3] = -100 with torch.no_grad(): logits = model(tokens) actual = model(tokens, labels=labels) expected = F.cross_entropy( logits[:, :-1].reshape(-1, logits.size(-1)), labels[:, 1:].reshape(-1), ignore_index=-100, ) torch.testing.assert_close(actual, expected) def test_gradient_checkpointing_matches_eager_grad(): torch.manual_seed(8) tokens = torch.randint(0, 32, (2, 8)) m1, cfg = _tiny() m2 = CausalLM(cfg) m2.load_state_dict(m1.state_dict()) m2.gradient_checkpointing = True m1.train() m2.train() l1 = m1(tokens, labels=tokens) l2 = m2(tokens, labels=tokens) torch.testing.assert_close(l1, l2, atol=1e-5, rtol=1e-5) l1.backward() l2.backward() for p1, p2 in zip(m1.parameters(), m2.parameters()): if p1.grad is None: assert p2.grad is None continue torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4) def test_0_5b_preset_enables_checkpointing(): assert K3Config.preset("0.5b").gradient_checkpointing is True assert K3Config.preset("toy").gradient_checkpointing is False