training / tinychess /model.py
cazyundee's picture
tinychess: full research artifact - self-play RL, ablations, compute scaling, faithfulness, precision, plasticity
227f251 verified
Raw History Blame Contribute Delete
25.7 kB
"""
TinyChess: a ~100K-parameter recurrent chess reasoning substrate.
Core ideas (see docs/RESEARCH_LOG.md):
* 64 persistent per-square latent states, never collapsed to one vector.
* ONE shared recurrent core applied n times: parameters are reused, depth is free.
* Learned memory slots participate in the same attention as the squares.
* A router mixes candidate operations per step.
* ACT-style learned halting gives adaptive depth.
* Compositional move embeddings scored against the *legal* candidate set only.
Everything is instrumented: `forward(..., trace=True)` returns h0..hn, memory,
router weights, halting probabilities and per-step move latents.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field, asdict
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from .encoding import (N_PIECE_TOKENS, N_SQ_EXTRA, N_GLOBAL, MOVE_FIELD_ORDER,
MOVE_FIELD_SIZES)
from .quant import QuantPolicy, maybe_quant
OPS = ["attn", "local", "mlp", "mem"]
@dataclass
class TinyChessConfig:
d_model: int = 64
n_heads: int = 4
d_ff: int = 128
d_attn_enc: int = 48
d_ff_enc: int = 64
n_mem: int = 8
d_move: int = 32
max_steps: int = 12
min_steps: int = 1
halt_threshold: float = 0.95
ponder_cost: float = 1e-2
n_refine: int = 2
use_router: bool = True
router_temp: float = 1.0
router_topk: int = 0 # 0 = dense mixture
use_memory: bool = True
use_local: bool = True
use_attn: bool = True
use_halting: bool = True
use_refine: bool = True
compositional_moves: bool = True
n_value_bins: int = 1 # 1 => scalar tanh value
dropout: float = 0.0
# structural plasticity
plastic: bool = False
# thought channel
thought: bool = False
# arch family for ablations: 'recurrent' | 'mlp' | 'transformer'
family: str = "recurrent"
n_layers: int = 2 # only for 'transformer'/'mlp' baselines
def to_dict(self):
return asdict(self)
# ---------------------------------------------------------------------------
# helpers
# ---------------------------------------------------------------------------
class RMSNorm(nn.Module):
"""Cheaper than LayerNorm (no bias), keeps the parameter budget for compute."""
def __init__(self, d):
super().__init__()
self.g = nn.Parameter(torch.ones(d))
def forward(self, x):
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + 1e-6) * self.g
def _neighbour_index():
"""[64, 8] index of the 8 king-neighbours of each square; self if off-board."""
idx = torch.zeros(64, 8, dtype=torch.long)
valid = torch.zeros(64, 8)
dirs = [(1, 0), (-1, 0), (0, 1), (0, -1), (1, 1), (1, -1), (-1, 1), (-1, -1)]
for s in range(64):
r, f = divmod(s, 8)
for k, (dr, df) in enumerate(dirs):
rr, ff = r + dr, f + df
if 0 <= rr < 8 and 0 <= ff < 8:
idx[s, k] = rr * 8 + ff
valid[s, k] = 1.0
else:
idx[s, k] = s
return idx, valid
def _ray_index():
"""[64, 4, 7] sliding-ray neighbours (rank, file, diag, anti-diag), self-padded."""
idx = torch.zeros(64, 4, 7, dtype=torch.long)
valid = torch.zeros(64, 4, 7)
axes = [((0, 1), (0, -1)), ((1, 0), (-1, 0)), ((1, 1), (-1, -1)), ((1, -1), (-1, 1))]
for s in range(64):
r, f = divmod(s, 8)
for a, (d1, d2) in enumerate(axes):
slot = 0
for (dr, df) in (d1, d2):
for step in range(1, 8):
rr, ff = r + dr * step, f + df * step
if not (0 <= rr < 8 and 0 <= ff < 8) or slot >= 7:
break
idx[s, a, slot] = rr * 8 + ff
valid[s, a, slot] = 1.0
slot += 1
for j in range(slot, 7):
idx[s, a, j] = s
return idx, valid
# ---------------------------------------------------------------------------
# Board encoder (bidirectional / jointly contextualised)
# ---------------------------------------------------------------------------
class BoardEncoder(nn.Module):
"""Produces 64 persistent per-square latent states.
'Bidirectional' = every square attends to every other square of the CURRENT
position. No future information is used anywhere.
"""
def __init__(self, cfg: TinyChessConfig):
super().__init__()
d = cfg.d_model
self.cfg = cfg
self.piece = nn.Embedding(N_PIECE_TOKENS, d)
# factorised square identity: file x rank instead of a 64xd table
self.file_emb = nn.Embedding(8, d)
self.rank_emb = nn.Embedding(8, d)
self.extra = nn.Linear(N_SQ_EXTRA, d, bias=False)
self.glob = nn.Linear(N_GLOBAL, d)
self.norm_in = RMSNorm(d)
da = cfg.d_attn_enc
self.qkv = nn.Linear(d, 3 * da, bias=False)
self.proj = nn.Linear(da, d, bias=False)
self.norm_a = RMSNorm(d)
self.ff = nn.Sequential(nn.Linear(d, cfg.d_ff_enc), nn.GELU(),
nn.Linear(cfg.d_ff_enc, d, bias=False))
self.norm_f = RMSNorm(d)
self.register_buffer("sq_ids", torch.arange(64), persistent=False)
self.register_buffer("file_ids", torch.arange(64) % 8, persistent=False)
self.register_buffer("rank_ids", torch.arange(64) // 8, persistent=False)
def forward(self, squares, extras, glob):
# squares [B,64] long, extras [B,64,E] float, glob [B,G] float
x = self.piece(squares) + (self.file_emb(self.file_ids) + self.rank_emb(self.rank_ids))[None]
x = x + self.extra(extras)
x = x + self.glob(glob)[:, None, :]
x = self.norm_in(x)
B = x.shape[0]
h = self.cfg.n_heads
hd = self.cfg.d_attn_enc // h
q, k, v = self.qkv(x).chunk(3, -1)
q = q.view(B, 64, h, hd).transpose(1, 2)
k = k.view(B, 64, h, hd).transpose(1, 2)
v = v.view(B, 64, h, hd).transpose(1, 2)
a = F.scaled_dot_product_attention(q, k, v).transpose(1, 2).reshape(B, 64, -1)
a = self.proj(a)
x = self.norm_a(x + a)
x = self.norm_f(x + self.ff(x))
return x
# ---------------------------------------------------------------------------
# The shared recurrent core
# ---------------------------------------------------------------------------
class RecurrentCore(nn.Module):
"""Applied n times with the SAME parameters.
Operations available at every step (mixed by the router):
attn - global attention over 64 squares + memory slots
local - structured chess mixing (king-neighbours + sliding rays)
mlp - pointwise nonlinearity
mem - explicit memory read + gated write
"""
def __init__(self, cfg: TinyChessConfig):
super().__init__()
d, h = cfg.d_model, cfg.n_heads
self.cfg = cfg
self.d, self.h = d, h
self.hd = d // h
# -- global attention (squares + memory as tokens) --
if cfg.use_attn:
self.qkv = nn.Linear(d, 3 * d, bias=False)
self.proj = nn.Linear(d, d, bias=False)
self.tok_type = nn.Parameter(torch.zeros(2, d)) # square vs memory tag
# -- local structured mixing --
if cfg.use_local:
nb, nbv = _neighbour_index()
ry, ryv = _ray_index()
self.register_buffer("nb_idx", nb, persistent=False)
self.register_buffer("nb_valid", nbv, persistent=False)
self.register_buffer("ray_idx", ry, persistent=False)
self.register_buffer("ray_valid", ryv, persistent=False)
# per-direction gates (cheap) + one shared mixing matrix
self.nb_gate = nn.Parameter(torch.zeros(8, d))
self.ray_gate = nn.Parameter(torch.zeros(4, d))
self.ray_decay = nn.Parameter(torch.zeros(4, 7))
self.local_proj = nn.Linear(2 * d, d, bias=False)
# flattened gather indices: [64*8] and [64*28]
self._p_nb_cache = None
# -- pointwise MLP (plastic: hidden units can be masked) --
self.ff1 = nn.Linear(d, cfg.d_ff)
self.ff2 = nn.Linear(cfg.d_ff, d, bias=False)
self.register_buffer("ff_mask", torch.ones(cfg.d_ff), persistent=True)
# -- memory --
if cfg.use_memory:
self.mem0 = nn.Parameter(torch.randn(cfg.n_mem, d) * 0.02)
self.mem_r = nn.Linear(d, d, bias=False) # query from squares
self.mem_w1 = nn.Linear(2 * d, 32, bias=False) # factorised write
self.mem_w2 = nn.Linear(32, 2 * d, bias=False)
self.register_buffer("mem_mask", torch.ones(cfg.n_mem), persistent=True)
# -- router --
self.ops = [o for o in OPS
if (o != "attn" or cfg.use_attn)
and (o != "local" or cfg.use_local)
and (o != "mem" or cfg.use_memory)]
if cfg.use_router:
self.router = nn.Linear(2 * d, len(self.ops))
# -- halting --
if cfg.use_halting:
self.halt = nn.Linear(2 * d, 1)
self.norm1 = RMSNorm(d)
self.norm2 = RMSNorm(d)
self.step_emb = nn.Parameter(torch.zeros(cfg.max_steps + 1, d))
# ---- individual operations -------------------------------------------------
def _attn(self, x, mem):
B, N, d = x.shape
if mem is not None:
toks = torch.cat([x + self.tok_type[0], mem + self.tok_type[1]], 1)
else:
toks = x + self.tok_type[0]
T = toks.shape[1]
q, k, v = self.qkv(toks).chunk(3, -1)
q = q.view(B, T, self.h, self.hd).transpose(1, 2)
k = k.view(B, T, self.h, self.hd).transpose(1, 2)
v = v.view(B, T, self.h, self.hd).transpose(1, 2)
o = F.scaled_dot_product_attention(q, k, v)
o = o.transpose(1, 2).reshape(B, T, d)
o = self.proj(o)
return o[:, :N], (o[:, N:] if mem is not None else None)
def _local(self, x):
"""Structured chess mixing: king-neighbours + sliding rays.
Key identity: each direction d contributes gate[d] * (P_d @ x), where
P_d is a fixed 64x64 permutation-like matrix and gate[d] is a per-channel
vector. Summing over directions is therefore
sum_d (P_d @ x) * gate_d
For the rays the 7 distance slots collapse into P_a (a = 4 axes) because
the decay weight depends only on (square, axis, slot), not on channels.
We precompute P_nb [8,64,64] and P_ray [4,64,64] ONCE per forward from
the current gates, then use two batched matmuls. This avoids the
[B,64,4,7,d] intermediate entirely.
"""
B, N, d = x.shape
# --- king neighbours ---
# P_nb[k] is a 0/1 matrix selecting neighbour k of each square
nb = torch.einsum("knm,bmd->bknd", self._P_nb(), x) # [B,8,64,d]
nb = torch.einsum("bknd,kd->bnd", nb, torch.tanh(self.nb_gate))
# --- sliding rays (decay folded into the matrix) ---
ry = torch.einsum("anm,bmd->band", self._P_ray(), x) # [B,4,64,d]
ry = torch.einsum("band,ad->bnd", ry, torch.tanh(self.ray_gate))
return self.local_proj(torch.cat([nb, ry], -1))
def _P_nb(self):
"""[8,64,64] neighbour selection matrices (cached; no grad path)."""
if getattr(self, "_p_nb_cache", None) is None:
P = torch.zeros(8, 64, 64)
for k in range(8):
P[k, torch.arange(64), self.nb_idx[:, k]] = self.nb_valid[:, k]
self._p_nb_cache = P
return self._p_nb_cache
def _P_ray(self):
"""[4,64,64] ray matrices with the learned decay folded in.
Depends on self.ray_decay, so it is rebuilt every call (cheap: 4x64x7
scatter) and keeps the gradient to ray_decay.
"""
w = self.ray_valid * torch.sigmoid(self.ray_decay).unsqueeze(0) # [64,4,7]
P = torch.zeros(4, 64, 64, device=w.device, dtype=w.dtype)
rows = torch.arange(64, device=w.device).view(64, 1).expand(64, 7)
for a in range(4):
P[a] = P[a].index_put((rows.reshape(-1), self.ray_idx[:, a, :].reshape(-1)),
w[:, a, :].reshape(-1), accumulate=True)
return P
def _mlp(self, x):
h = F.gelu(self.ff1(x)) * self.ff_mask
return self.ff2(h)
def _mem_read(self, x, mem):
q = self.mem_r(x) # [B,64,d]
att = torch.einsum("bnd,bmd->bnm", q, mem) / math.sqrt(self.d)
att = att + torch.log(self.mem_mask.clamp_min(1e-9))[None, None]
w = att.softmax(-1)
return torch.einsum("bnm,bmd->bnd", w, mem), w
def _mem_write(self, x, mem):
summ = x.mean(1, keepdim=True).expand(-1, mem.shape[1], -1)
gz = self.mem_w2(F.gelu(self.mem_w1(torch.cat([mem, summ], -1))))
gate, cand = gz.chunk(2, -1)
gate = torch.sigmoid(gate)
new = mem * (1 - gate) + torch.tanh(cand) * gate
return new * self.mem_mask[None, :, None], gate
# ---- one recurrent step ----------------------------------------------------
def forward(self, x, mem, step: int, qp: Optional[QuantPolicy] = None,
collect: Optional[dict] = None):
d = self.d
xn = self.norm1(x) + self.step_emb[min(step, self.cfg.max_steps)]
summary = torch.cat([xn.mean(1), xn.amax(1)], -1) # [B,2d]
# router decides which operations matter this step
if self.cfg.use_router:
logits = self.router(summary) / self.cfg.router_temp
if self.cfg.router_topk and self.cfg.router_topk < len(self.ops):
k = self.cfg.router_topk
thresh = logits.topk(k, -1).values[:, -1:]
logits = logits.masked_fill(logits < thresh, float("-inf"))
w = logits.softmax(-1)
else:
w = x.new_full((x.shape[0], len(self.ops)), 1.0 / len(self.ops))
delta = torch.zeros_like(x)
mem_out = mem
mem_attn = None
for i, op in enumerate(self.ops):
wi = w[:, i][:, None, None]
if op == "attn":
o, mdelta = self._attn(xn, mem)
o = maybe_quant(o, qp, "attn")
delta = delta + wi * o
if mdelta is not None and mem is not None:
mem_out = mem_out + w[:, i][:, None, None] * mdelta
elif op == "local":
delta = delta + wi * maybe_quant(self._local(xn), qp, "local")
elif op == "mlp":
delta = delta + wi * maybe_quant(self._mlp(self.norm2(x)), qp, "mlp")
elif op == "mem" and mem is not None:
r, mem_attn = self._mem_read(xn, mem_out)
delta = delta + wi * maybe_quant(r, qp, "mem")
mem_out, _ = self._mem_write(xn, mem_out)
x_new = x + delta
p_halt = None
if self.cfg.use_halting:
s2 = torch.cat([x_new.mean(1), x_new.amax(1)], -1)
p_halt = torch.sigmoid(self.halt(s2)).squeeze(-1)
if collect is not None:
collect.setdefault("router", []).append(w.detach())
if p_halt is not None:
collect.setdefault("halt", []).append(p_halt.detach())
if mem_attn is not None:
collect.setdefault("mem_attn", []).append(mem_attn.detach())
return x_new, mem_out, p_halt
# ---------------------------------------------------------------------------
# Compositional move embeddings
# ---------------------------------------------------------------------------
class MoveEmbedder(nn.Module):
"""Moves are built from structural components, not a flat 4096-way vocabulary."""
def __init__(self, cfg: TinyChessConfig):
super().__init__()
self.cfg = cfg
dm = cfg.d_move
if cfg.compositional_moves:
self.tables = nn.ModuleList([
nn.Embedding(MOVE_FIELD_SIZES[f], dm) for f in MOVE_FIELD_ORDER])
self.mix = nn.Linear(dm, dm, bias=False)
else:
# ablation: flat from*to vocabulary (deliberately bigger, for comparison)
self.flat = nn.Embedding(64 * 64, dm)
self.norm = RMSNorm(dm)
def forward(self, fields):
# fields [B,L,7] long
if self.cfg.compositional_moves:
e = 0
for i, t in enumerate(self.tables):
e = e + t(fields[..., i])
e = e + self.mix(F.gelu(e))
else:
e = self.flat(fields[..., 0] * 64 + fields[..., 1])
return self.norm(e)
# ---------------------------------------------------------------------------
# Full model
# ---------------------------------------------------------------------------
class TinyChess(nn.Module):
def __init__(self, cfg: TinyChessConfig):
super().__init__()
self.cfg = cfg
d = cfg.d_model
self.encoder = BoardEncoder(cfg)
if cfg.family == "recurrent":
self.core = RecurrentCore(cfg)
elif cfg.family == "transformer":
self.blocks = nn.ModuleList([RecurrentCore(cfg) for _ in range(cfg.n_layers)])
elif cfg.family == "mlp":
self.blocks = nn.ModuleList([
nn.Sequential(RMSNorm(d), nn.Linear(d, cfg.d_ff), nn.GELU(),
nn.Linear(cfg.d_ff, d)) for _ in range(cfg.n_layers)])
else:
raise ValueError(cfg.family)
self.move_emb = MoveEmbedder(cfg)
self.readout = nn.Linear(d, cfg.d_move, bias=False) # h -> z_move
self.sq_to_move = nn.Linear(d, cfg.d_move, bias=False) # per-square -> move space
if cfg.use_refine:
self.refine = nn.GRUCell(cfg.d_move, cfg.d_move)
self.move_scale = nn.Parameter(torch.tensor(1.0))
self.value = nn.Sequential(nn.Linear(2 * d, 24), nn.GELU(), nn.Linear(24, cfg.n_value_bins))
self.norm_out = RMSNorm(d)
# ---- parameter accounting -------------------------------------------------
def param_report(self) -> dict:
groups = {}
for name, p in self.named_parameters():
top = name.split(".")[0]
groups[top] = groups.get(top, 0) + p.numel()
total = sum(groups.values())
active = self.active_params()
return {"groups": groups, "total": total, "active": active}
def active_params(self) -> int:
"""Parameters that are actually live given plasticity masks."""
total = sum(p.numel() for p in self.parameters())
core = getattr(self, "core", None)
if core is None:
return total
dead_ff = int((core.ff_mask == 0).sum())
# each dead hidden unit removes: ff1 row (d+1) and ff2 column (d)
total -= dead_ff * (self.cfg.d_model * 2 + 1)
if self.cfg.use_memory:
dead_m = int((core.mem_mask == 0).sum())
total -= dead_m * self.cfg.d_model
return total
# ---- forward ---------------------------------------------------------------
def forward(self, squares, extras, glob, cand_fields=None, cand_mask=None,
steps: Optional[int] = None, adaptive: bool = False,
trace: bool = False, qp: Optional[QuantPolicy] = None):
"""
squares [B,64] long; extras [B,64,E]; glob [B,G]
cand_fields [B,L,7] long; cand_mask [B,L] bool (legal candidates only)
Returns dict with logits, value, trajectory info.
"""
B = squares.shape[0]
x = self.encoder(squares, extras, glob)
collect: dict = {} if trace else None
traj = [x.detach()] if trace else None
cfg = self.cfg
mem = None
if cfg.family == "recurrent":
if cfg.use_memory:
mem = self.core.mem0[None].expand(B, -1, -1).contiguous()
n = steps if steps is not None else cfg.max_steps
n = max(cfg.min_steps, min(n, cfg.max_steps))
if adaptive and cfg.use_halting:
x, mem, info = self._act_loop(x, mem, n, qp, collect, traj)
else:
halts = []
for t in range(n):
x, mem, ph = self.core(x, mem, t, qp, collect)
if trace:
traj.append(x.detach())
if ph is not None:
halts.append(ph)
info = {"n_steps": torch.full((B,), float(n), device=x.device),
"ponder": torch.zeros(B, device=x.device),
"halt_probs": torch.stack(halts, 1) if halts else None}
else:
for blk in self.blocks:
if cfg.family == "transformer":
x, mem, _ = blk(x, None, 0, qp, collect)
else:
x = x + blk(x)
if trace:
traj.append(x.detach())
info = {"n_steps": torch.full((B,), float(len(self.blocks)), device=x.device),
"ponder": torch.zeros(B, device=x.device), "halt_probs": None}
x = self.norm_out(x)
pooled = torch.cat([x.mean(1), x.amax(1)], -1) # [B,2d]
value = self.value(pooled)
out = {"h": x, "pooled": pooled, "value": value, "memory": mem, **info}
if trace:
out["trajectory"] = traj
out["collect"] = collect
if cand_fields is None:
return out
z = self.readout(pooled[:, :self.cfg.d_model]) # [B,dm]
cand = self.move_emb(cand_fields) # [B,L,dm]
# squares contribute directly: a move's from/to squares index the board state
sq_m = self.sq_to_move(x) # [B,64,dm]
idx_f = cand_fields[..., 0].clamp(0, 63)
idx_t = cand_fields[..., 1].clamp(0, 63)
cand = cand + torch.gather(sq_m, 1, idx_f[..., None].expand(-1, -1, sq_m.shape[-1]))
cand = cand + torch.gather(sq_m, 1, idx_t[..., None].expand(-1, -1, sq_m.shape[-1]))
zs = [z]
if cfg.use_refine and cfg.n_refine > 0:
ctx = z
for _ in range(cfg.n_refine):
z = self.refine(ctx, z)
zs.append(z)
logits = torch.einsum("bd,bld->bl", z, cand) * self.move_scale / math.sqrt(cfg.d_move)
if cand_mask is not None:
logits = logits.masked_fill(~cand_mask, float("-inf"))
out["logits"] = logits
out["z_move"] = z
if trace:
out["z_steps"] = [zz.detach() for zz in zs]
out["refine_logits"] = [
(torch.einsum("bd,bld->bl", zz, cand) * self.move_scale / math.sqrt(cfg.d_move)
).masked_fill(~cand_mask, float("-inf")).detach() if cand_mask is not None else None
for zz in zs]
return out
# ---- ACT / adaptive depth ---------------------------------------------------
def _act_loop(self, x, mem, n_max, qp, collect, traj):
"""Graves-style ACT.
Each batch element accumulates halting mass until it exceeds
`halt_threshold`; the step at which that happens gets the *remainder*
weight. The returned state is the halting-weighted mean of the visited
states, so gradients flow into the halting unit. `ponder` is the
expected number of steps and is what the compute penalty charges for.
"""
B = x.shape[0]
dev = x.device
still = torch.ones(B, device=dev) # 1 while the element is running
acc = torch.zeros(B, device=dev) # accumulated halting mass
x_out = torch.zeros_like(x)
mem_out = torch.zeros_like(mem) if mem is not None else None
n_steps = torch.zeros(B, device=dev) # discrete steps actually taken
ponder = torch.zeros(B, device=dev) # differentiable remainder term
thr = self.cfg.halt_threshold
halts = []
for t in range(n_max):
x, mem, ph = self.core(x, mem, t, qp, collect)
if traj is not None:
traj.append(x.detach())
if ph is None:
ph = torch.full((B,), 1.0 / n_max, device=dev)
halts.append(ph)
forced_min = t < (self.cfg.min_steps - 1)
is_last = (t == n_max - 1)
new_acc = acc + ph
# an element finishes this step if it crosses the threshold, or time is up
finish = ((new_acc > thr) | is_last) & (~torch.tensor(forced_min, device=dev))
finish = finish & (still > 0)
w = torch.where(finish, (1.0 - acc).clamp(min=0.0), ph) * still
x_out = x_out + w[:, None, None] * x
if mem_out is not None:
mem_out = mem_out + w[:, None, None] * mem
n_steps = n_steps + still
ponder = ponder + still * (1.0 - acc).clamp(min=0.0)
acc = torch.where(finish, acc, new_acc)
still = torch.where(finish, torch.zeros_like(still), still)
if float(still.max()) == 0.0:
break
info = {"n_steps": n_steps, "ponder": ponder,
"halt_probs": torch.stack(halts, 1)}
return x_out, (mem_out if mem_out is not None else mem), info
def build_model(cfg: TinyChessConfig) -> TinyChess:
m = TinyChess(cfg)
for p in m.parameters():
if p.dim() > 1 and p.requires_grad:
pass
return m