"""The incremental decoder must reproduce the training forward pass exactly (up to float error), with rows at different lengths and chunks longer than the sliding window.""" import torch from tiny_agent.generate import KVCache, append from tiny_agent.model import TinyAgentLM from tests.test_model import DEV, small_cfg def test_decode_matches_forward(): cfg = small_cfg(swa_window=8, kv_group=2, n_layers=4) torch.manual_seed(0) m = TinyAgentLM(cfg).to(DEV).float().eval() with torch.no_grad(): for p in m.parameters(): p.normal_(0, 0.05) B, L = 3, 60 seqs = torch.randint(0, cfg.vocab_size, (B, L), device=DEV) # reference: full forward per row with torch.no_grad(): ref = m(seqs, torch.zeros(B, L, dtype=torch.long, device=DEV)) # (B, L, V) cache = KVCache(m, B, 128, DEV, dtype=torch.float32) # schedule of chunk sizes per row (different per row; includes chunks > window and zeros) plans = [[20, 1, 1, 13, 1, 24], [5, 30, 0, 1, 1, 23], [1, 1, 1, 40, 16, 1]] done = [0] * B for t in range(len(plans[0])): n = torch.tensor([plans[b][t] for b in range(B)], device=DEV) C = int(n.max()) chunk = torch.zeros(B, C, dtype=torch.long, device=DEV) for b in range(B): chunk[b, : n[b]] = seqs[b, done[b]: done[b] + n[b]] logits = append(m, cache, chunk, n) for b in range(B): if n[b] > 0: done[b] += int(n[b]) exp = ref[b, done[b] - 1] err = (logits[b] - exp).abs().max().item() assert err < 2e-3, (t, b, err) assert done == [L, L, L] assert torch.equal(cache.tokens[:, :L], seqs) def test_cache_is_small(): cfg = small_cfg(swa_window=128, kv_group=4, n_layers=16, n_kv_heads=2, head_dim=64) m = TinyAgentLM(cfg) per_tok = m.kv_bytes_per_token(8192) full_mha = cfg.n_layers * 2 * cfg.n_heads * cfg.head_dim * 2 # plain per-layer MHA cache, bf16 assert per_tok < full_mha / 4 def test_subbatch_rows_and_slot_reuse(): """Feeding only some rows (rows=...) and reusing a reset row for a new sequence must give the same logits as a full forward pass (catches stale local-window positions or Engram history).""" cfg = small_cfg(swa_window=8, kv_group=2, n_layers=4) torch.manual_seed(1) m = TinyAgentLM(cfg).to(DEV).float().eval() with torch.no_grad(): for p in m.parameters(): p.normal_(0, 0.05) B, L = 4, 40 seqs = torch.randint(0, cfg.vocab_size, (B + 1, L), device=DEV) with torch.no_grad(): ref = m(seqs, torch.zeros(B + 1, L, dtype=torch.long, device=DEV)) cache = KVCache(m, B, 128, DEV, dtype=torch.float32) done = [0] * B which = list(range(B)) # which sequence each row is decoding def feed(rows, ns): C = max(ns) chunk = torch.zeros(len(rows), C, dtype=torch.long, device=DEV) for j, (r, n) in enumerate(zip(rows, ns)): chunk[j, :n] = seqs[which[r], done[r]: done[r] + n] logits = append(m, cache, chunk, torch.tensor(ns), rows=rows) for j, (r, n) in enumerate(zip(rows, ns)): if n: done[r] += n err = (logits[j] - ref[which[r], done[r] - 1]).abs().max().item() assert err < 2e-3, (r, done[r], err) feed([0, 2], [12, 5]) feed([1, 2, 3], [20, 1, 9]) feed([3, 0], [1, 1]) # row 2 gets a new sequence part-way through: reset and refill it cache.reset([2]) which[2], done[2] = B, 0 feed([2], [17]) feed([0, 1, 2, 3], [3, 1, 1, 15]) feed([2], [22]) assert done[2] == L