Download code/tests/test_model.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 4.38 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tests/test_model.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tests/test_model.py
-
curl -L -o test_model.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tests/test_model.py
4.38 kB
| """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] | |