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