tiny-agent-112m / code /tests /test_generate.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
3.65 kB
"""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