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