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