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