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)
|