File size: 9,159 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
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
"""Incremental decoding with the small KV cache, for batched multi-turn rollouts.

Cache layout (what makes it small):
  * global K/V: one per layer GROUP (written by the group's Full layer), grows with context;
  * local K/V: one ring buffer of `swa_window` slots per layer, constant size.
Each row has its own length, so episodes can be at different points: `append` takes a
right-padded chunk of new tokens per row (prefill, tool results, or one decoded token).

No host syncs inside a step: padded tokens are written to a scratch slot (last index) instead of
being filtered out with boolean masks, lengths are mirrored on the CPU, and attention is computed
directly against the cache with grouped queries (no K/V concatenation or head repetition).
"""
from __future__ import annotations

import math

import torch

from tiny_agent.model import TinyAgentLM, rms_norm

CHUNK = 64  # max new tokens per row processed at once (bounds the score tensor during prefill)


class KVCache:
    def __init__(self, model: TinyAgentLM, B: int, max_len: int, device, dtype=torch.bfloat16):
        c = model.cfg
        self.cfg, self.B, self.max_len, self.device = c, B, max_len, device
        G = math.ceil(c.n_layers / c.kv_group)
        Hk, D, W = c.n_kv_heads, c.head_dim, c.swa_window
        self.gk = torch.zeros(G, B, Hk, max_len + 1, D, device=device, dtype=dtype)   # +1 scratch slot
        self.gv = torch.zeros_like(self.gk)
        self.lk = torch.zeros(c.n_layers, B, Hk, W + 1, D, device=device, dtype=dtype)  # +1 scratch slot
        self.lv = torch.zeros_like(self.lk)
        self.lpos = torch.full((B, W + 1), -1, device=device, dtype=torch.long)
        self.lens = torch.zeros(B, device=device, dtype=torch.long)
        self.lens_cpu = torch.zeros(B, dtype=torch.long)
        self.tokens = torch.zeros(B, max_len + 1, device=device, dtype=torch.long)    # history (Engram)

    def nbytes(self) -> int:
        return sum(t.numel() * t.element_size() for t in (self.gk, self.gv, self.lk, self.lv))

    def reset(self, rows) -> None:
        """Empty these rows so a new sequence can start there (old K/V is masked by length)."""
        rows = torch.as_tensor(rows, dtype=torch.long)
        self.lens_cpu[rows] = 0
        rows = rows.to(self.device)
        self.lens[rows] = 0
        self.lpos[rows] = -1


def _rope_pos(rope, x, pos):
    """x: (B, H, C, D); pos: (B, C) per-row absolute positions."""
    cos = rope.cos[pos][:, None].to(x.dtype)
    sin = rope.sin[pos][:, None].to(x.dtype)
    x1, x2 = x.chunk(2, dim=-1)
    return torch.cat((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1)


@torch.no_grad()
def append(model: TinyAgentLM, cache: KVCache, idx: torch.Tensor, n_new, rows=None) -> torch.Tensor:
    """Feed a right-padded chunk idx (B, C); row b has n_new[b] real tokens (0 = untouched).
    n_new may be a CPU tensor/list (preferred: avoids a device sync). Returns logits (B, V) at
    each row's last real token (unchanged garbage for rows with n_new == 0).
    rows: optional cache rows that idx refers to (a sub-batch, e.g. only rows that need prefill);
    by default idx covers all cache rows in order."""
    n_cpu = torch.as_tensor(n_new, dtype=torch.long).cpu()
    B, C = idx.shape
    if rows is not None:
        rows = torch.as_tensor(rows, dtype=torch.long).cpu()
    lens_cpu = cache.lens_cpu if rows is None else cache.lens_cpu[rows]
    if int((lens_cpu + n_cpu).max()) > cache.max_len:
        raise RuntimeError("KV cache full")
    out = None
    for s in range(0, C, CHUNK):
        n_sub = (n_cpu - s).clamp(0, CHUNK)
        if int(n_sub.max()) == 0:
            break
        width = int(n_sub.max())
        logits = _append_chunk(model, cache, idx[:, s:s + width], n_sub, rows)
        if out is None:
            out = logits
        else:
            out = torch.where((n_sub > 0).to(out.device)[:, None], logits, out)
    return out


def _append_chunk(model, cache: KVCache, idx, n_cpu, rows=None):
    c = model.cfg
    B, C = idx.shape
    dev = idx.device
    H, Hk, D, W = c.n_heads, c.n_kv_heads, c.head_dim, c.swa_window
    rep = H // Hk
    n = n_cpu.to(dev, non_blocking=True)
    # full batch: views of the cache; sub-batch: gathered copies of just those rows
    if rows is None:
        rdev = torch.arange(B, device=dev)
        sel = lambda t: t                                                       # noqa: E731
        lens, lens_cpu = cache.lens, cache.lens_cpu
    else:
        rdev = rows.to(dev, non_blocking=True)
        sel = lambda t: t[rdev]                                                 # noqa: E731
        lens, lens_cpu = cache.lens[rdev], cache.lens_cpu[rows]
    ar = torch.arange(C, device=dev)
    valid = ar[None, :] < n[:, None]                                            # (B, C)
    pos = lens[:, None] + ar[None, :]                                           # (B, C)
    new_lens = lens + n
    new_lens_cpu = lens_cpu + n_cpu
    Lg = int(new_lens_cpu.max())
    bidx = rdev[:, None].expand(B, C)
    gpos = torch.where(valid, pos, torch.full_like(pos, cache.max_len))        # scratch for padding
    keep = valid & (pos >= (new_lens[:, None] - W))                            # last W real tokens
    slot = torch.where(keep, pos % W, torch.full_like(pos, W))                 # scratch slot W

    # Engram addresses need the previous n-1 tokens of each row
    addrs = {}
    if model.engram_modules():
        max_n = max(c.engram_orders)
        back = torch.arange(-(max_n - 1), 0, device=dev)
        prev_pos = lens[:, None] + back[None, :]
        prev = torch.where(prev_pos >= 0, sel(cache.tokens).gather(1, prev_pos.clamp_min(0)), torch.zeros_like(prev_pos))
        cids = model.cid_map[torch.cat([prev, idx], dim=1)]
        for i, b in enumerate(model.blocks):
            if b.engram is not None:
                addrs[i] = b.engram.addresses(cids)[:, max_n - 1:]
    cache.tokens[bidx, gpos] = idx

    kpos = torch.arange(Lg, device=dev)
    gmask = (kpos[None, None, :] <= pos[:, :, None]) & (kpos[None, None, :] < new_lens[:, None, None])  # B,C,Lg
    lpos_all = torch.cat([sel(cache.lpos), torch.where(valid, pos, torch.full_like(pos, -1))], dim=1)  # B,W+1+C
    lmask = (lpos_all[:, None, :] >= 0) & (lpos_all[:, None, :] <= pos[:, :, None]) & \
            (pos[:, :, None] - lpos_all[:, None, :] < W)
    mask = torch.cat([gmask, lmask], dim=2)[:, None, None]                      # B,1,1,C,L
    neg = torch.finfo(torch.float32).min / 2
    scale = 1.0 / math.sqrt(D)

    x = model.embed(idx)
    rope = model.rope
    group = -1
    for li, blk in enumerate(model.blocks):
        if blk.engram is not None:
            x = x + blk.engram(x, addrs[li])
        a = blk.attn
        h = rms_norm(x)
        q = _rope_pos(rope, rms_norm(a.w_q(h).view(B, C, H, D).transpose(1, 2)), pos)
        if a.mode == "full":
            group += 1
            k, v = a.w_kv_global(h).view(B, C, 2, Hk, D).permute(2, 0, 3, 1, 4)
            k = _rope_pos(rope, rms_norm(k), pos)
            cache.gk[group][bidx, :, gpos] = k.transpose(1, 2).to(cache.gk.dtype)
            cache.gv[group][bidx, :, gpos] = v.transpose(1, 2).to(cache.gv.dtype)
        kl, vl = a.w_kv_local(h).view(B, C, 2, Hk, D).permute(2, 0, 3, 1, 4)
        kl = _rope_pos(rope, rms_norm(kl), pos).to(cache.lk.dtype)
        vl = vl.to(cache.lv.dtype)
        # grouped-query attention straight against the cache: queries (B, Hk, rep*C, D)
        q2 = q.to(cache.gk.dtype).reshape(B, Hk, rep * C, D)
        if rows is None:
            Kg, Vg = cache.gk[group][:, :, :Lg], cache.gv[group][:, :, :Lg]
        else:
            Kg, Vg = cache.gk[group][rdev, :, :Lg], cache.gv[group][rdev, :, :Lg]
        Kl = torch.cat([sel(cache.lk[li]), kl], dim=2)                          # B,Hk,W+1+C,D (small)
        Vl = torch.cat([sel(cache.lv[li]), vl], dim=2)
        s = torch.cat([q2 @ Kg.transpose(-1, -2), q2 @ Kl.transpose(-1, -2)], dim=-1).float() * scale
        s = s.view(B, Hk, rep, C, -1).masked_fill(~mask, neg)
        p = torch.softmax(s, dim=-1).to(Vg.dtype).view(B, Hk, rep * C, -1)
        y = p[..., :Lg] @ Vg + p[..., Lg:] @ Vl                                 # B,Hk,rep*C,D
        y = y.view(B, H, C, D).transpose(1, 2).reshape(B, C, H * D)
        x = x + a.w_o(y.to(x.dtype))
        x = x + blk.mlp(rms_norm(x))
        cache.lk[li][bidx, :, slot] = kl.transpose(1, 2)
        cache.lv[li][bidx, :, slot] = vl.transpose(1, 2)
    cache.lpos[bidx, slot] = torch.where(keep, pos, torch.full_like(pos, -1))
    if rows is None:
        cache.lens, cache.lens_cpu = new_lens, new_lens_cpu
    else:
        cache.lens[rdev] = new_lens
        cache.lens_cpu[rows] = new_lens_cpu
    last = (n - 1).clamp_min(0)
    xl = x[torch.arange(B, device=dev), last]
    logits = model.head(rms_norm(xl)).float()
    sc = c.logit_softcap
    return sc * torch.tanh(logits / sc)


def sample(logits: torch.Tensor, temperature: float = 1.0, generator=None) -> torch.Tensor:
    if temperature <= 0:
        return logits.argmax(-1)
    p = torch.softmax(logits / temperature, dim=-1)
    return torch.multinomial(p, 1, generator=generator).squeeze(-1)