Download code/tiny_agent/generate.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 9.16 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/generate.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tiny_agent/generate.py
-
curl -L -o generate.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/generate.py
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) | |
| 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) | |