File size: 12,566 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 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 | """Tiny dense decoder with V4.1-style shared global KV + per-layer sliding window, and Engram.
Attention (simplified CSA2, compression ratio 1, no indexer):
* Layers are grouped (``kv_group`` layers per group). The first layer of a group is in
"Full" mode: it projects the global K/V for the whole group. The other layers are in
"Reuse" mode: they only compute their own queries and reuse the group's global K/V.
* Every layer also projects its own local K/V, visible only inside a ``swa_window`` window.
* Each query does ONE softmax over the union of global (causal) and local (windowed) keys,
as in V4.1. Implemented with FlexAttention over the concatenated [global | local] keys.
So the decode KV cache is: one global K/V per group (grows with context) + one local K/V
ring buffer of ``swa_window`` entries per layer (constant size).
Engram (Cheng et al. 2026, as used in V4.1): hashed n-gram embeddings over a compressed
token id space, fused into the residual stream with a context-aware sigmoid gate. The short
causal conv is omitted, as in V4.1.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
import torch
import torch.nn as nn
import torch.nn.functional as F
try:
from torch.nn.attention.flex_attention import create_block_mask, flex_attention
except ImportError: # pragma: no cover
flex_attention = None
@dataclass
class ModelConfig:
vocab_size: int = 32768
d_model: int = 640
n_layers: int = 16
n_heads: int = 10
head_dim: int = 64
n_kv_heads: int = 2 # heads of both the shared global KV and the per-layer local KV
kv_group: int = 4 # layers per shared-global-KV group (1 Full + kv_group-1 Reuse)
swa_window: int = 128
ffn_mult: float = 4.0 # ReLU^2 MLP hidden = ffn_mult * d_model
rope_base: float = 10000.0
max_seq_len: int = 8192
logit_softcap: float = 15.0
# Engram
engram_layers: tuple[int, ...] = () # e.g. (1,) ; empty disables Engram
engram_orders: tuple[int, ...] = (2, 3)
engram_heads: int = 8
engram_head_dim: int = 64
engram_rows_per_head: int = 262144 # rounded up to distinct primes per (order, head)
def layer_mode(self, i: int) -> str:
return "full" if i % self.kv_group == 0 else "reuse"
@property
def ffn_hidden(self) -> int:
return int(self.ffn_mult * self.d_model) // 64 * 64
def rms_norm(x: torch.Tensor) -> torch.Tensor:
return F.rms_norm(x, (x.size(-1),))
class Rotary(nn.Module):
def __init__(self, dim: int, max_len: int, base: float):
super().__init__()
inv = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim))
t = torch.arange(max_len, dtype=torch.float32)
f = torch.outer(t, inv)
self.register_buffer("cos", f.cos(), persistent=False)
self.register_buffer("sin", f.sin(), persistent=False)
def forward(self, x: torch.Tensor, pos: torch.Tensor | None = None) -> torch.Tensor:
# x: (B, H, T, D). pos: optional (T,) absolute positions (decode).
T = x.size(-2)
cos = self.cos[:T] if pos is None else self.cos[pos]
sin = self.sin[:T] if pos is None else self.sin[pos]
cos, sin = cos.to(x.dtype), sin.to(x.dtype)
x1, x2 = x.chunk(2, dim=-1)
return torch.cat((x1 * cos - x2 * sin, x1 * sin + x2 * cos), dim=-1)
def next_prime(n: int) -> int:
def is_prime(k: int) -> bool:
if k < 2:
return False
if k % 2 == 0:
return k == 2
r = int(k**0.5)
return all(k % d for d in range(3, r + 1, 2))
while not is_prime(n):
n += 1
return n
class Engram(nn.Module):
"""Hashed n-gram memory with context-aware gating (one module = one insertion layer)."""
def __init__(self, cfg: ModelConfig, seed: int):
super().__init__()
self.cfg = cfg
g = torch.Generator().manual_seed(seed)
primes, p = [], cfg.engram_rows_per_head
for _ in cfg.engram_orders:
for _ in range(cfg.engram_heads):
p = next_prime(p + 1)
primes.append(p)
self.register_buffer("primes", torch.tensor(primes, dtype=torch.int64), persistent=False)
offs = torch.tensor([0] + primes[:-1], dtype=torch.int64).cumsum(0)
self.register_buffer("offsets", offs, persistent=False)
max_n = max(cfg.engram_orders)
# odd multipliers per (order, head, position-in-ngram)
mult = torch.randint(1, 2**30, (len(cfg.engram_orders), cfg.engram_heads, max_n), generator=g) * 2 + 1
self.register_buffer("mult", mult.to(torch.int64), persistent=False)
self.table = nn.Embedding(sum(primes), cfg.engram_head_dim)
mem_dim = len(cfg.engram_orders) * cfg.engram_heads * cfg.engram_head_dim
self.w_k = nn.Linear(mem_dim, cfg.d_model, bias=False)
self.w_v = nn.Linear(mem_dim, cfg.d_model, bias=False)
nn.init.normal_(self.table.weight, std=0.02)
nn.init.zeros_(self.w_v.weight)
@torch.no_grad()
def addresses(self, cids: torch.Tensor) -> torch.Tensor:
"""cids: (B, T) compressed token ids -> (B, T, n_orders*heads) global row indices."""
B, T = cids.shape
max_n = max(self.cfg.engram_orders)
pad = cids.new_full((B, max_n - 1), 0)
ext = torch.cat([pad, cids], dim=1).long()
# shifted[k] = token at t-k
shifted = torch.stack([ext[:, max_n - 1 - k: max_n - 1 - k + T] for k in range(max_n)], dim=-1) # B,T,max_n
out = []
for oi, n in enumerate(self.cfg.engram_orders):
m = self.mult[oi, :, :n] # heads, n
mix = shifted[..., None, 0] * m[:, 0] # B,T,heads
for k in range(1, n):
mix = mix ^ (shifted[..., None, k] * m[:, k])
out.append(mix)
mix = torch.cat(out, dim=-1) # B,T,n_orders*heads
return torch.remainder(mix, self.primes) + self.offsets
def forward(self, h: torch.Tensor, addr: torch.Tensor) -> torch.Tensor:
mem = self.table(addr).flatten(-2) # B,T,mem_dim
mem = mem.to(h.dtype)
k = self.w_k(mem)
v = self.w_v(mem)
g = (rms_norm(h) * rms_norm(k)).sum(-1, keepdim=True) / math.sqrt(h.size(-1))
g = torch.sigmoid(g.abs().clamp_min(1e-6).sqrt() * g.sign())
return g * v
class Attention(nn.Module):
def __init__(self, cfg: ModelConfig, mode: str):
super().__init__()
self.cfg, self.mode = cfg, mode
H, Hk, D = cfg.n_heads, cfg.n_kv_heads, cfg.head_dim
self.w_q = nn.Linear(cfg.d_model, H * D, bias=False)
self.w_q._muon_heads = H # head-wise Muon (V4.1)
self.w_kv_local = nn.Linear(cfg.d_model, 2 * Hk * D, bias=False)
if mode == "full":
self.w_kv_global = nn.Linear(cfg.d_model, 2 * Hk * D, bias=False)
self.w_o = nn.Linear(H * D, cfg.d_model, bias=False)
nn.init.zeros_(self.w_o.weight)
def _kv(self, lin: nn.Linear, x: torch.Tensor, rope: Rotary, pos=None):
B, T, _ = x.shape
Hk, D = self.cfg.n_kv_heads, self.cfg.head_dim
k, v = lin(x).view(B, T, 2, Hk, D).permute(2, 0, 3, 1, 4)
return rope(rms_norm(k), pos), v
def forward(self, x, rope: Rotary, shared: dict, block_mask):
B, T, _ = x.shape
H, D = self.cfg.n_heads, self.cfg.head_dim
q = self.w_q(x).view(B, T, H, D).transpose(1, 2)
q = rope(rms_norm(q))
if self.mode == "full":
shared["k"], shared["v"] = self._kv(self.w_kv_global, x, rope)
kl, vl = self._kv(self.w_kv_local, x, rope)
k = torch.cat([shared["k"], kl], dim=2)
v = torch.cat([shared["v"], vl], dim=2)
dt = v.dtype # rms_norm autocasts to fp32; flex backward needs one dtype
y = flex_attention(q.to(dt), k.to(dt), v, block_mask=block_mask, enable_gqa=True)
return self.w_o(y.transpose(1, 2).reshape(B, T, H * D))
class MLP(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.up = nn.Linear(cfg.d_model, cfg.ffn_hidden, bias=False)
self.down = nn.Linear(cfg.ffn_hidden, cfg.d_model, bias=False)
nn.init.zeros_(self.down.weight)
def forward(self, x):
return self.down(F.relu(self.up(x)).square())
class Block(nn.Module):
def __init__(self, cfg: ModelConfig, i: int):
super().__init__()
self.attn = Attention(cfg, cfg.layer_mode(i))
self.mlp = MLP(cfg)
self.engram = Engram(cfg, seed=1000 + i) if i in cfg.engram_layers else None
def forward(self, x, rope, shared, block_mask, addr):
if self.engram is not None:
x = x + self.engram(x, addr)
x = x + self.attn(rms_norm(x), rope, shared, block_mask)
return x + self.mlp(rms_norm(x))
def make_block_mask(doc: torch.Tensor, window: int):
"""doc: (B, T) int document ids. Keys are [global(T) | local(T)]."""
B, T = doc.shape
def mask_mod(b, h, q, kv):
is_local = kv >= T
j = torch.where(is_local, kv - T, kv)
same = doc[b, q] == doc[b, j]
causal = q >= j
in_win = (q - j) < window
return same & causal & (~is_local | in_win)
return create_block_mask(mask_mod, B, None, T, 2 * T, device=doc.device)
class TinyAgentLM(nn.Module):
def __init__(self, cfg: ModelConfig):
super().__init__()
self.cfg = cfg
self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
self.blocks = nn.ModuleList([Block(cfg, i) for i in range(cfg.n_layers)])
self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
self.rope = Rotary(cfg.head_dim, cfg.max_seq_len, cfg.rope_base)
# token id -> compressed id for Engram (identity until a tokenizer map is loaded)
self.register_buffer("cid_map", torch.arange(cfg.vocab_size), persistent=True)
nn.init.normal_(self.embed.weight, std=1.0)
nn.init.zeros_(self.head.weight)
def engram_modules(self):
return [b.engram for b in self.blocks if b.engram is not None]
def forward(self, idx: torch.Tensor, doc: torch.Tensor, block_mask=None, targets=None):
"""idx, doc: (B, T). Returns softcapped float32 logits (B, T, V), or the mean CE loss
over targets != -1 when targets is given (lets torch.compile fuse head + softcap + CE)."""
if block_mask is None:
block_mask = make_block_mask(doc, self.cfg.swa_window)
addr = None
ems = self.engram_modules()
if ems:
cids = self.cid_map[idx]
x = self.embed(idx)
shared: dict = {}
for b in self.blocks:
a = b.engram.addresses(cids) if b.engram is not None else None
x = b(x, self.rope, shared, block_mask, a)
logits = self.head(rms_norm(x)).float()
c = self.cfg.logit_softcap
logits = c * torch.tanh(logits / c)
if targets is None:
return logits
return F.cross_entropy(logits.view(-1, logits.size(-1)), targets.reshape(-1), ignore_index=-1)
def param_counts(self) -> dict:
emb = self.embed.weight.numel() + self.head.weight.numel()
eng = sum(p.numel() for m in self.engram_modules() for p in [m.table.weight])
total = sum(p.numel() for p in self.parameters())
return {"total": total, "embedding": emb, "engram_tables": eng, "backbone": total - emb - eng}
def kv_bytes_per_token(self, context: int, bytes_per_value: float = 2.0) -> float:
"""Decode KV cache per token of context, amortized (global grows, local is a fixed ring)."""
c = self.cfg
groups = math.ceil(c.n_layers / c.kv_group)
glob = groups * 2 * c.n_kv_heads * c.head_dim
local = c.n_layers * 2 * c.n_kv_heads * c.head_dim * min(c.swa_window, context) / max(context, 1)
return (glob + local) * bytes_per_value
def dense_reference_attention_mask(doc: torch.Tensor, window: int) -> torch.Tensor:
"""Boolean (B, 1, T, 2T) mask equal to make_block_mask's mask_mod, for tests."""
B, T = doc.shape
q = torch.arange(T, device=doc.device)[:, None]
kv = torch.arange(2 * T, device=doc.device)[None, :]
is_local = kv >= T
j = torch.where(is_local, kv - T, kv) # (1, 2T)
same = doc[:, :, None] == doc[:, j[0]][:, None, :] # (B, T, 2T)
m = same & (q >= j) & (~is_local | ((q - j) < window))
return m[:, None]
|