"""Correctness tests for the model. Run: source env.sh && $TA_PY -m pytest -q tests/""" import math import torch import torch.nn.functional as F from tiny_agent.model import ModelConfig, TinyAgentLM, dense_reference_attention_mask, make_block_mask from tiny_agent.optim import build_optimizers DEV = "xpu" if torch.xpu.is_available() else "cpu" def small_cfg(**kw): base = dict(vocab_size=512, d_model=128, n_layers=4, n_heads=4, head_dim=32, n_kv_heads=2, kv_group=2, swa_window=8, max_seq_len=256, engram_layers=(1,), engram_rows_per_head=1000, engram_heads=2, engram_head_dim=16) base.update(kw) return ModelConfig(**base) def random_model(cfg): torch.manual_seed(0) m = TinyAgentLM(cfg).to(DEV) # zero-inits make outputs trivially constant; randomize everything for the tests with torch.no_grad(): for p in m.parameters(): p.normal_(0, 0.05) return m def docs(B, T, cuts=(37, 80)): d = torch.zeros(B, T, dtype=torch.long) for c in cuts: d[:, c:] += 1 return d.to(DEV) def test_flex_matches_dense_reference(): torch.manual_seed(0) B, H, Hk, T, D, W = 2, 4, 2, 128, 32, 8 q = torch.randn(B, H, T, D, device=DEV) k = torch.randn(B, Hk, 2 * T, D, device=DEV) v = torch.randn(B, Hk, 2 * T, D, device=DEV) doc = docs(B, T) from torch.nn.attention.flex_attention import flex_attention y = flex_attention(q, k, v, block_mask=make_block_mask(doc, W), enable_gqa=True) m = dense_reference_attention_mask(doc, W) ref = F.scaled_dot_product_attention(q, k.repeat_interleave(H // Hk, 1), v.repeat_interleave(H // Hk, 1), attn_mask=m) assert torch.allclose(y, ref, atol=2e-4, rtol=1e-3), (y - ref).abs().max() def test_mask_semantics(): doc = docs(1, 64, cuts=(20,)) m = dense_reference_attention_mask(doc, 8)[0, 0] T = 64 assert m[30, 25] and not m[30, 31] # global causal assert not m[30, 10] # other document assert m[30, T + 25] and not m[30, T + 22] # local window 8: 23..30 visible assert not m[30, T + 31] # local causal def test_causal_and_doc_isolation(): cfg = small_cfg() m = random_model(cfg).eval() B, T = 2, 128 idx = torch.randint(0, cfg.vocab_size, (B, T), device=DEV) doc = docs(B, T) with torch.no_grad(): a = m(idx, doc) idx2 = idx.clone() idx2[:, 60] = (idx2[:, 60] + 1) % cfg.vocab_size b = m(idx2, doc) assert torch.allclose(a[:, :60], b[:, :60], atol=1e-4) # nothing before t=60 changes assert not torch.allclose(a[:, 60:80], b[:, 60:80]) # same doc after t changes assert torch.allclose(a[:, 80:], b[:, 80:], atol=1e-4) # later document is isolated def test_engram_addresses_deterministic_and_in_range(): cfg = small_cfg() m = random_model(cfg) e = m.engram_modules()[0] ids = torch.randint(0, cfg.vocab_size, (2, 50), device=DEV) a1, a2 = e.addresses(ids), e.addresses(ids) assert torch.equal(a1, a2) assert a1.min() >= 0 and a1.max() < e.table.num_embeddings # bigram head address at t depends only on tokens t-1, t ids2 = ids.clone() ids2[:, 10] += 1 b = e.addresses(ids2) n_heads = cfg.engram_heads assert torch.equal(a1[:, :10], b[:, :10]) and torch.equal(a1[:, 12:, :n_heads], b[:, 12:, :n_heads]) def test_kv_reuse_groups(): cfg = small_cfg(n_layers=8, kv_group=4) m = TinyAgentLM(cfg) modes = [b.attn.mode for b in m.blocks] assert modes == ["full", "reuse", "reuse", "reuse"] * 2 assert not hasattr(m.blocks[1].attn, "w_kv_global") def test_overfit_one_batch(): cfg = small_cfg() torch.manual_seed(0) m = TinyAgentLM(cfg).to(DEV) opts = build_optimizers(m, lr=1e-2) idx = torch.randint(0, cfg.vocab_size, (4, 65), device=DEV) doc = torch.zeros(4, 64, dtype=torch.long, device=DEV) losses = [] for _ in range(60): logits = m(idx[:, :-1], doc) loss = F.cross_entropy(logits.reshape(-1, cfg.vocab_size), idx[:, 1:].reshape(-1)) loss.backward() for o in opts: o.step() o.zero_grad(set_to_none=True) losses.append(loss.item()) assert losses[0] > math.log(cfg.vocab_size) - 0.5 assert losses[-1] < 0.5 * losses[0], losses[::10]