File size: 3,652 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 | """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
|