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