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