Download tinychess/model.py from cazyundee/training: direct link, hf CLI and curl.
- Browser
- Download file 25.7 kB
-
https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/model.py
- Command line
-
hf download hf://spaces/cazyundee/training/tinychess/model.py
-
curl -L -o model.py https://huggingface.co/spaces/cazyundee/training/resolve/main/tinychess/model.py
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"] | |
| 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 | |