Upload folder using huggingface_hub
Browse files- chess_io.py +1 -18
- engine_uci.py +4 -18
- model.py +58 -126
- search.py +35 -102
- training/data_pipeline.py +21 -43
- training/match.py +12 -19
- training/pretrain.py +106 -93
- training/run_selfplay_parallel.py +24 -23
- training/selfplay_finetune.py +29 -59
chess_io.py
CHANGED
|
@@ -1,15 +1,4 @@
|
|
| 1 |
-
"""Move encoding shared by data pipeline, model, search, and the UCI engine.
|
| 2 |
-
|
| 3 |
-
Move representation: every move is (from_square, to_square, promo_idx).
|
| 4 |
-
promo_idx in {0: none, 1: q, 2: r, 3: b, 4: n}. The from/to squares alone give a
|
| 5 |
-
4096-way index (from_sq * 64 + to_sq) that covers every geometrically possible
|
| 6 |
-
square pair -- no separate move-vocabulary table needed, and it trivially
|
| 7 |
-
covers every legal move in any position since legality is enforced by masking,
|
| 8 |
-
not by the encoding.
|
| 9 |
-
"""
|
| 10 |
-
|
| 11 |
from __future__ import annotations
|
| 12 |
-
|
| 13 |
import chess
|
| 14 |
|
| 15 |
PROMO_PIECES = [None, chess.QUEEN, chess.ROOK, chess.BISHOP, chess.KNIGHT]
|
|
@@ -19,10 +8,9 @@ NUM_PROMO = len(PROMO_PIECES)
|
|
| 19 |
|
| 20 |
|
| 21 |
def move_to_ids(move: chess.Move) -> tuple[int, int]:
|
| 22 |
-
"""Return (from_to_id in [0,4096), promo_id in [0,5))."""
|
| 23 |
from_to_id = move.from_square * 64 + move.to_square
|
| 24 |
promo_id = PROMO_INDEX[move.promotion]
|
| 25 |
-
return from_to_id, promo_id
|
| 26 |
|
| 27 |
|
| 28 |
def ids_to_move(from_to_id: int, promo_id: int) -> chess.Move:
|
|
@@ -35,10 +23,5 @@ def legal_move_ids(board: chess.Board) -> list[tuple[int, int]]:
|
|
| 35 |
|
| 36 |
|
| 37 |
def find_legal_move(board: chess.Board, from_to_id: int, promo_id: int) -> chess.Move | None:
|
| 38 |
-
"""Match a predicted (from_to_id, promo_id) against the board's actual legal moves.
|
| 39 |
-
|
| 40 |
-
Needed because promo_id alone doesn't disambiguate underpromotion capture direction
|
| 41 |
-
edge cases are already fully specified by from/to, so this is just a legality check.
|
| 42 |
-
"""
|
| 43 |
candidate = ids_to_move(from_to_id, promo_id)
|
| 44 |
return candidate if candidate in board.legal_moves else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import chess
|
| 3 |
|
| 4 |
PROMO_PIECES = [None, chess.QUEEN, chess.ROOK, chess.BISHOP, chess.KNIGHT]
|
|
|
|
| 8 |
|
| 9 |
|
| 10 |
def move_to_ids(move: chess.Move) -> tuple[int, int]:
|
|
|
|
| 11 |
from_to_id = move.from_square * 64 + move.to_square
|
| 12 |
promo_id = PROMO_INDEX[move.promotion]
|
| 13 |
+
return (from_to_id, promo_id)
|
| 14 |
|
| 15 |
|
| 16 |
def ids_to_move(from_to_id: int, promo_id: int) -> chess.Move:
|
|
|
|
| 23 |
|
| 24 |
|
| 25 |
def find_legal_move(board: chess.Board, from_to_id: int, promo_id: int) -> chess.Move | None:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
candidate = ids_to_move(from_to_id, promo_id)
|
| 27 |
return candidate if candidate in board.legal_moves else None
|
engine_uci.py
CHANGED
|
@@ -1,26 +1,16 @@
|
|
| 1 |
-
"""UCI protocol loop for the ChessModel entry.
|
| 2 |
-
|
| 3 |
-
Safety-net design: this runs and plays LEGAL moves even with no trained checkpoint
|
| 4 |
-
present (falls back to a policy-free shallow search / random legal move). The
|
| 5 |
-
trained ChessMamba net is loaded opportunistically -- if the checkpoint file isn't
|
| 6 |
-
there yet, the engine still works, just weaker. This guarantees there is always
|
| 7 |
-
something submittable regardless of how training goes under the deadline.
|
| 8 |
-
"""
|
| 9 |
-
|
| 10 |
from __future__ import annotations
|
| 11 |
-
|
| 12 |
import os
|
| 13 |
import random
|
| 14 |
import sys
|
| 15 |
-
|
| 16 |
import chess
|
| 17 |
-
|
| 18 |
from search import Searcher
|
| 19 |
|
| 20 |
ENGINE_NAME = "ChessMamba"
|
| 21 |
ENGINE_AUTHOR = "CatGirlHtfi"
|
| 22 |
-
DEFAULT_CKPT = os.environ.get(
|
| 23 |
-
|
|
|
|
|
|
|
| 24 |
|
| 25 |
|
| 26 |
def send(msg: str) -> None:
|
|
@@ -31,14 +21,12 @@ def send(msg: str) -> None:
|
|
| 31 |
def main() -> None:
|
| 32 |
board = chess.Board()
|
| 33 |
searcher = Searcher(checkpoint_path=DEFAULT_CKPT, policy_only=POLICY_ONLY)
|
| 34 |
-
|
| 35 |
for line in sys.stdin:
|
| 36 |
line = line.strip()
|
| 37 |
if not line:
|
| 38 |
continue
|
| 39 |
tokens = line.split()
|
| 40 |
cmd = tokens[0]
|
| 41 |
-
|
| 42 |
if cmd == "uci":
|
| 43 |
send(f"id name {ENGINE_NAME}")
|
| 44 |
send(f"id author {ENGINE_AUTHOR}")
|
|
@@ -64,7 +52,6 @@ def main() -> None:
|
|
| 64 |
break
|
| 65 |
elif cmd == "stop":
|
| 66 |
pass
|
| 67 |
-
# unknown commands are ignored, per UCI spec
|
| 68 |
|
| 69 |
|
| 70 |
def _handle_position(board: chess.Board, tokens: list[str]) -> None:
|
|
@@ -79,7 +66,6 @@ def _handle_position(board: chess.Board, tokens: list[str]) -> None:
|
|
| 79 |
idx = moves_idx
|
| 80 |
else:
|
| 81 |
return
|
| 82 |
-
|
| 83 |
if idx < len(tokens) and tokens[idx] == "moves":
|
| 84 |
for uci_move in tokens[idx + 1 :]:
|
| 85 |
board.push_uci(uci_move)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import os
|
| 3 |
import random
|
| 4 |
import sys
|
|
|
|
| 5 |
import chess
|
|
|
|
| 6 |
from search import Searcher
|
| 7 |
|
| 8 |
ENGINE_NAME = "ChessMamba"
|
| 9 |
ENGINE_AUTHOR = "CatGirlHtfi"
|
| 10 |
+
DEFAULT_CKPT = os.environ.get(
|
| 11 |
+
"CHESSMAMBA_CKPT", os.path.join(os.path.dirname(__file__), "ckpt", "model.pt")
|
| 12 |
+
)
|
| 13 |
+
POLICY_ONLY = os.environ.get("CHESSMAMBA_POLICY_ONLY", "") == "1"
|
| 14 |
|
| 15 |
|
| 16 |
def send(msg: str) -> None:
|
|
|
|
| 21 |
def main() -> None:
|
| 22 |
board = chess.Board()
|
| 23 |
searcher = Searcher(checkpoint_path=DEFAULT_CKPT, policy_only=POLICY_ONLY)
|
|
|
|
| 24 |
for line in sys.stdin:
|
| 25 |
line = line.strip()
|
| 26 |
if not line:
|
| 27 |
continue
|
| 28 |
tokens = line.split()
|
| 29 |
cmd = tokens[0]
|
|
|
|
| 30 |
if cmd == "uci":
|
| 31 |
send(f"id name {ENGINE_NAME}")
|
| 32 |
send(f"id author {ENGINE_AUTHOR}")
|
|
|
|
| 52 |
break
|
| 53 |
elif cmd == "stop":
|
| 54 |
pass
|
|
|
|
| 55 |
|
| 56 |
|
| 57 |
def _handle_position(board: chess.Board, tokens: list[str]) -> None:
|
|
|
|
| 66 |
idx = moves_idx
|
| 67 |
else:
|
| 68 |
return
|
|
|
|
| 69 |
if idx < len(tokens) and tokens[idx] == "moves":
|
| 70 |
for uci_move in tokens[idx + 1 :]:
|
| 71 |
board.push_uci(uci_move)
|
model.py
CHANGED
|
@@ -1,17 +1,5 @@
|
|
| 1 |
-
"""ChessMamba: a from-scratch selective-state-space (S6) model over move-history
|
| 2 |
-
sequences, with a policy head (masked to legal moves) and a value head.
|
| 3 |
-
|
| 4 |
-
Unlike NNUE (static sparse-feature eval + alpha-beta) or AlphaZero/Lc0 (CNN/ResNet
|
| 5 |
-
over board planes + MCTS), this treats a chess game as a causal sequence of moves
|
| 6 |
-
and scans it with a Mamba-style selective SSM, matching how "Chess Llama" and
|
| 7 |
-
DeepMind's search-free transformer frame the problem but swapping attention for a
|
| 8 |
-
recurrent scan. See PLAN.md for the full rationale.
|
| 9 |
-
"""
|
| 10 |
-
|
| 11 |
from __future__ import annotations
|
| 12 |
-
|
| 13 |
import math
|
| 14 |
-
|
| 15 |
import torch
|
| 16 |
import torch.nn as nn
|
| 17 |
import torch.nn.functional as F
|
|
@@ -19,11 +7,12 @@ from torch.utils.checkpoint import checkpoint
|
|
| 19 |
|
| 20 |
NUM_FROM_TO = 4096
|
| 21 |
NUM_PROMO = 5
|
| 22 |
-
MAX_PLIES = 96
|
| 23 |
|
| 24 |
|
| 25 |
class RMSNorm(nn.Module):
|
| 26 |
-
|
|
|
|
| 27 |
super().__init__()
|
| 28 |
self.eps = eps
|
| 29 |
self.weight = nn.Parameter(torch.ones(dim))
|
|
@@ -34,22 +23,11 @@ class RMSNorm(nn.Module):
|
|
| 34 |
|
| 35 |
|
| 36 |
def parallel_scan(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
|
| 37 |
-
"""Hillis-Steele (doubling) associative scan for the affine recurrence
|
| 38 |
-
h_t = a_t * h_{t-1} + b_t, h_{-1} = 0, computed along dim=1 (time).
|
| 39 |
-
|
| 40 |
-
a, b: (B, L, ...). Returns h: same shape, h[:, t] = the recurrence's value at t.
|
| 41 |
-
|
| 42 |
-
This is the same mathematical trick real Mamba's CUDA parallel-scan kernel uses
|
| 43 |
-
to make a sequential recurrence run in O(log L) dependent steps instead of O(L):
|
| 44 |
-
the affine maps x -> a*x + b compose associatively, (a1,b1) then (a2,b2) gives
|
| 45 |
-
(a2*a1, a2*b1+b2), so prefix-composing all the way to t gives h_t directly (since
|
| 46 |
-
composing the identity initial state h_{-1}=0 just picks out the b component).
|
| 47 |
-
"""
|
| 48 |
L = a.shape[1]
|
| 49 |
d = 1
|
| 50 |
while d < L:
|
| 51 |
-
a_prev, b_prev = a[:, :-d], b[:, :-d]
|
| 52 |
-
a_cur, b_cur = a[:, d:], b[:, d:]
|
| 53 |
new_a = a_cur * a_prev
|
| 54 |
new_b = a_cur * b_prev + b_cur
|
| 55 |
a = torch.cat([a[:, :d], new_a], dim=1)
|
|
@@ -59,17 +37,6 @@ def parallel_scan(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
|
|
| 59 |
|
| 60 |
|
| 61 |
class S6Block(nn.Module):
|
| 62 |
-
"""Selective-scan SSM block (the core of Mamba), parallel-scan implementation.
|
| 63 |
-
|
| 64 |
-
The recurrence h_t = A_bar_t * h_{t-1} + Bx_t is sequential by nature, but it's
|
| 65 |
-
also an affine recurrence, so it can be computed via an O(log L) doubling scan
|
| 66 |
-
(see `parallel_scan`) instead of an O(L) Python loop. Measured impact: the
|
| 67 |
-
naive per-timestep loop (depth=10, ~96 timesteps, plus the recompute pass from
|
| 68 |
-
gradient checkpointing) bottlenecked training at ~11s/step on an A100 -- almost
|
| 69 |
-
entirely fixed per-launch overhead on a chain of ~1920 tiny sequential kernel
|
| 70 |
-
launches per step that batch size can't parallelize away, since batching only
|
| 71 |
-
parallelizes *within* a timestep, not across the dependency chain.
|
| 72 |
-
"""
|
| 73 |
|
| 74 |
def __init__(self, dim: int, state_dim: int = 16, expand: int = 2):
|
| 75 |
super().__init__()
|
|
@@ -77,79 +44,56 @@ class S6Block(nn.Module):
|
|
| 77 |
self.dim = dim
|
| 78 |
self.inner_dim = inner_dim
|
| 79 |
self.state_dim = state_dim
|
| 80 |
-
|
| 81 |
-
self.
|
| 82 |
-
# input-dependent SSM parameters, projected from x
|
| 83 |
-
self.x_proj = nn.Linear(inner_dim, state_dim * 2 + inner_dim, bias=False) # -> (B, C, delta_raw)
|
| 84 |
self.dt_bias = nn.Parameter(torch.zeros(inner_dim))
|
| 85 |
-
|
| 86 |
-
# A: per-channel state matrix, log-parameterized for stability (stays negative after -exp)
|
| 87 |
A = torch.arange(1, state_dim + 1, dtype=torch.float32).unsqueeze(0).repeat(inner_dim, 1)
|
| 88 |
self.A_log = nn.Parameter(torch.log(A))
|
| 89 |
self.D = nn.Parameter(torch.ones(inner_dim))
|
| 90 |
-
|
| 91 |
self.out_proj = nn.Linear(inner_dim, dim, bias=False)
|
| 92 |
|
| 93 |
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 94 |
-
# x: (B, L, dim)
|
| 95 |
B_, L, _ = x.shape
|
| 96 |
xz = self.in_proj(x)
|
| 97 |
-
x_in, gate = xz.chunk(2, dim=-1)
|
| 98 |
x_in = F.silu(x_in)
|
| 99 |
-
|
| 100 |
-
x_dbl = self.x_proj(x_in) # (B, L, state_dim*2 + inner_dim)
|
| 101 |
Bmat, Cmat, delta_raw = torch.split(
|
| 102 |
x_dbl, [self.state_dim, self.state_dim, self.inner_dim], dim=-1
|
| 103 |
)
|
| 104 |
-
delta = F.softplus(delta_raw + self.dt_bias)
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
|
| 110 |
-
Bx = (delta * x_in).unsqueeze(-1) * Bmat.unsqueeze(2) # (B, L, inner_dim, state_dim)
|
| 111 |
-
|
| 112 |
-
h = parallel_scan(A_bar, Bx) # (B, L, inner_dim, state_dim), h[:, t] = hidden state after step t
|
| 113 |
-
|
| 114 |
-
y = (h * Cmat.unsqueeze(2)).sum(-1) + self.D * x_in # (B, L, inner_dim)
|
| 115 |
y = y * F.silu(gate)
|
| 116 |
return self.out_proj(y)
|
| 117 |
|
| 118 |
-
def step(
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
O(1) per move instead of re-scanning the whole move history from scratch.
|
| 122 |
-
This is the exact per-step formula `parallel_scan` is mathematically
|
| 123 |
-
equivalent to (verified against it in tests), just applied once instead of
|
| 124 |
-
via the doubling trick, since there's nothing to parallelize over a single step.
|
| 125 |
-
|
| 126 |
-
x_t: (B, dim). h_prev: (B, inner_dim, state_dim) or None (treated as zero,
|
| 127 |
-
i.e. the initial state before any moves).
|
| 128 |
-
"""
|
| 129 |
xz = self.in_proj(x_t)
|
| 130 |
-
x_in, gate = xz.chunk(2, dim=-1)
|
| 131 |
x_in = F.silu(x_in)
|
| 132 |
-
|
| 133 |
x_dbl = self.x_proj(x_in)
|
| 134 |
Bmat, Cmat, delta_raw = torch.split(
|
| 135 |
x_dbl, [self.state_dim, self.state_dim, self.inner_dim], dim=-1
|
| 136 |
)
|
| 137 |
-
delta = F.softplus(delta_raw + self.dt_bias)
|
| 138 |
-
|
| 139 |
-
|
| 140 |
-
|
| 141 |
-
Bx = (delta * x_in).unsqueeze(-1) * Bmat.unsqueeze(1) # (B, inner_dim, state_dim)
|
| 142 |
-
|
| 143 |
if h_prev is None:
|
| 144 |
h_prev = x_t.new_zeros(x_t.shape[0], self.inner_dim, self.state_dim)
|
| 145 |
h_new = A_bar * h_prev + Bx
|
| 146 |
-
|
| 147 |
-
y = (h_new * Cmat.unsqueeze(1)).sum(-1) + self.D * x_in # (B, inner_dim)
|
| 148 |
y = y * F.silu(gate)
|
| 149 |
-
return self.out_proj(y), h_new
|
| 150 |
|
| 151 |
|
| 152 |
class MambaBlock(nn.Module):
|
|
|
|
| 153 |
def __init__(self, dim: int, state_dim: int = 16, expand: int = 2):
|
| 154 |
super().__init__()
|
| 155 |
self.norm = RMSNorm(dim)
|
|
@@ -158,12 +102,15 @@ class MambaBlock(nn.Module):
|
|
| 158 |
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 159 |
return x + self.ssm(self.norm(x))
|
| 160 |
|
| 161 |
-
def step(
|
|
|
|
|
|
|
| 162 |
y, h_new = self.ssm.step(self.norm(x_t), h_prev)
|
| 163 |
-
return x_t + y, h_new
|
| 164 |
|
| 165 |
|
| 166 |
class ChessMamba(nn.Module):
|
|
|
|
| 167 |
def __init__(
|
| 168 |
self,
|
| 169 |
dim: int = 256,
|
|
@@ -176,32 +123,25 @@ class ChessMamba(nn.Module):
|
|
| 176 |
super().__init__()
|
| 177 |
self.dim = dim
|
| 178 |
self.max_plies = max_plies
|
| 179 |
-
# gradient checkpointing recomputes each block's forward during backward instead of
|
| 180 |
-
# retaining all depth * seq_len intermediate scan tensors at once -- without this,
|
| 181 |
-
# backprop through a ~96-step sequential scan stacked depth=10 deep OOMs even an 80GB
|
| 182 |
-
# GPU (verified: batch 1024 at dim=384/depth=10 exhausted 80GB with checkpoint off).
|
| 183 |
self.use_checkpoint = use_checkpoint
|
| 184 |
-
|
| 185 |
self.from_embed = nn.Embedding(64, dim)
|
| 186 |
self.to_embed = nn.Embedding(64, dim)
|
| 187 |
self.promo_embed = nn.Embedding(NUM_PROMO, dim)
|
| 188 |
-
self.pos_embed = nn.Embedding(max_plies + 1, dim)
|
| 189 |
self.start_token = nn.Parameter(torch.zeros(1, 1, dim))
|
| 190 |
-
|
| 191 |
self.blocks = nn.ModuleList([MambaBlock(dim, state_dim, expand) for _ in range(depth)])
|
| 192 |
self.norm_f = RMSNorm(dim)
|
| 193 |
-
|
| 194 |
self.policy_head = nn.Linear(dim, NUM_FROM_TO)
|
| 195 |
self.promo_head = nn.Linear(dim, NUM_PROMO)
|
| 196 |
self.value_head = nn.Linear(dim, 1)
|
| 197 |
-
|
| 198 |
nn.init.normal_(self.from_embed.weight, std=0.02)
|
| 199 |
nn.init.normal_(self.to_embed.weight, std=0.02)
|
| 200 |
nn.init.normal_(self.promo_embed.weight, std=0.02)
|
| 201 |
nn.init.normal_(self.pos_embed.weight, std=0.02)
|
| 202 |
|
| 203 |
-
def embed_moves(
|
| 204 |
-
|
|
|
|
| 205 |
B_ = from_ids.shape[0]
|
| 206 |
start = self.start_token.expand(B_, 1, -1)
|
| 207 |
if from_ids.shape[1] == 0:
|
|
@@ -212,17 +152,13 @@ class ChessMamba(nn.Module):
|
|
| 212 |
positions = torch.arange(tok.shape[1], device=tok.device).unsqueeze(0)
|
| 213 |
return tok + self.pos_embed(positions)
|
| 214 |
|
| 215 |
-
def forward(
|
| 216 |
-
|
| 217 |
-
|
| 218 |
-
|
| 219 |
-
|
| 220 |
-
|
| 221 |
-
|
| 222 |
-
position (index == lengths[i], since index 0 is the <start> token). No attention
|
| 223 |
-
mask needed, unlike a transformer. If None, the last position is used (for live
|
| 224 |
-
incremental play where the input is exactly the current position, no padding).
|
| 225 |
-
"""
|
| 226 |
x = self.embed_moves(from_ids, to_ids, promo_ids)
|
| 227 |
for block in self.blocks:
|
| 228 |
if self.use_checkpoint and self.training:
|
|
@@ -231,21 +167,14 @@ class ChessMamba(nn.Module):
|
|
| 231 |
x = block(x)
|
| 232 |
x = self.norm_f(x)
|
| 233 |
if lengths is None:
|
| 234 |
-
pooled = x[:, -1]
|
| 235 |
else:
|
| 236 |
idx = lengths.view(-1, 1, 1).expand(-1, 1, x.shape[-1])
|
| 237 |
pooled = x.gather(1, idx).squeeze(1)
|
| 238 |
policy_logits = self.policy_head(pooled)
|
| 239 |
promo_logits = self.promo_head(pooled)
|
| 240 |
value = torch.tanh(self.value_head(pooled))
|
| 241 |
-
return policy_logits, promo_logits, value
|
| 242 |
-
|
| 243 |
-
# -- incremental (single-move) inference path, for live/search play --
|
| 244 |
-
# Search visits many positions per move (up to top_k^depth); recomputing the full
|
| 245 |
-
# windowed forward pass at every node was the actual search-speed bottleneck. These
|
| 246 |
-
# methods carry per-block hidden state across moves so each new node costs O(1)
|
| 247 |
-
# (one token through `depth` layers) instead of O(window length). State is
|
| 248 |
-
# (next_position_index, [per-block hidden state or None]).
|
| 249 |
|
| 250 |
@torch.no_grad()
|
| 251 |
def init_incremental(self, device: torch.device | str = "cpu"):
|
|
@@ -257,14 +186,14 @@ class ChessMamba(nn.Module):
|
|
| 257 |
block_states.append(s)
|
| 258 |
x = self.norm_f(x)
|
| 259 |
outputs = self._heads(x)
|
| 260 |
-
return (1, block_states), outputs
|
| 261 |
|
| 262 |
@torch.no_grad()
|
| 263 |
def step_move(self, from_sq: int, to_sq: int, promo_id: int, state):
|
| 264 |
pos_idx, block_states = state
|
| 265 |
device = self.from_embed.weight.device
|
| 266 |
pos_idx_clamped = min(pos_idx, self.max_plies)
|
| 267 |
-
idx = lambda v: torch.tensor([v], device=device)
|
| 268 |
x = (
|
| 269 |
self.from_embed(idx(from_sq))
|
| 270 |
+ self.to_embed(idx(to_sq))
|
|
@@ -277,33 +206,36 @@ class ChessMamba(nn.Module):
|
|
| 277 |
new_block_states.append(ns)
|
| 278 |
x = self.norm_f(x)
|
| 279 |
outputs = self._heads(x)
|
| 280 |
-
return (pos_idx + 1, new_block_states), outputs
|
| 281 |
|
| 282 |
@torch.no_grad()
|
| 283 |
-
def build_incremental_state(
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
|
|
|
|
|
|
|
|
|
| 287 |
state, outputs = self.init_incremental(device)
|
| 288 |
for f, t, p in zip(from_list, to_list, promo_list):
|
| 289 |
state, outputs = self.step_move(f, t, p, state)
|
| 290 |
-
return state, outputs
|
| 291 |
|
| 292 |
def _heads(self, x: torch.Tensor):
|
| 293 |
policy_logits = self.policy_head(x)
|
| 294 |
promo_logits = self.promo_head(x)
|
| 295 |
value = torch.tanh(self.value_head(x))
|
| 296 |
-
return policy_logits, promo_logits, value
|
| 297 |
|
| 298 |
|
| 299 |
def count_params(model: nn.Module) -> int:
|
| 300 |
-
return sum(p.numel() for p in model.parameters())
|
| 301 |
|
| 302 |
|
| 303 |
if __name__ == "__main__":
|
| 304 |
m = ChessMamba(dim=256, depth=8)
|
| 305 |
print(f"params: {count_params(m):,}")
|
| 306 |
-
B_, L = 4, 20
|
| 307 |
from_ids = torch.randint(0, 64, (B_, L))
|
| 308 |
to_ids = torch.randint(0, 64, (B_, L))
|
| 309 |
promo_ids = torch.zeros(B_, L, dtype=torch.long)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import math
|
|
|
|
| 3 |
import torch
|
| 4 |
import torch.nn as nn
|
| 5 |
import torch.nn.functional as F
|
|
|
|
| 7 |
|
| 8 |
NUM_FROM_TO = 4096
|
| 9 |
NUM_PROMO = 5
|
| 10 |
+
MAX_PLIES = 96
|
| 11 |
|
| 12 |
|
| 13 |
class RMSNorm(nn.Module):
|
| 14 |
+
|
| 15 |
+
def __init__(self, dim: int, eps: float = 1e-05):
|
| 16 |
super().__init__()
|
| 17 |
self.eps = eps
|
| 18 |
self.weight = nn.Parameter(torch.ones(dim))
|
|
|
|
| 23 |
|
| 24 |
|
| 25 |
def parallel_scan(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
L = a.shape[1]
|
| 27 |
d = 1
|
| 28 |
while d < L:
|
| 29 |
+
a_prev, b_prev = (a[:, :-d], b[:, :-d])
|
| 30 |
+
a_cur, b_cur = (a[:, d:], b[:, d:])
|
| 31 |
new_a = a_cur * a_prev
|
| 32 |
new_b = a_cur * b_prev + b_cur
|
| 33 |
a = torch.cat([a[:, :d], new_a], dim=1)
|
|
|
|
| 37 |
|
| 38 |
|
| 39 |
class S6Block(nn.Module):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
|
| 41 |
def __init__(self, dim: int, state_dim: int = 16, expand: int = 2):
|
| 42 |
super().__init__()
|
|
|
|
| 44 |
self.dim = dim
|
| 45 |
self.inner_dim = inner_dim
|
| 46 |
self.state_dim = state_dim
|
| 47 |
+
self.in_proj = nn.Linear(dim, inner_dim * 2, bias=False)
|
| 48 |
+
self.x_proj = nn.Linear(inner_dim, state_dim * 2 + inner_dim, bias=False)
|
|
|
|
|
|
|
| 49 |
self.dt_bias = nn.Parameter(torch.zeros(inner_dim))
|
|
|
|
|
|
|
| 50 |
A = torch.arange(1, state_dim + 1, dtype=torch.float32).unsqueeze(0).repeat(inner_dim, 1)
|
| 51 |
self.A_log = nn.Parameter(torch.log(A))
|
| 52 |
self.D = nn.Parameter(torch.ones(inner_dim))
|
|
|
|
| 53 |
self.out_proj = nn.Linear(inner_dim, dim, bias=False)
|
| 54 |
|
| 55 |
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
|
|
| 56 |
B_, L, _ = x.shape
|
| 57 |
xz = self.in_proj(x)
|
| 58 |
+
x_in, gate = xz.chunk(2, dim=-1)
|
| 59 |
x_in = F.silu(x_in)
|
| 60 |
+
x_dbl = self.x_proj(x_in)
|
|
|
|
| 61 |
Bmat, Cmat, delta_raw = torch.split(
|
| 62 |
x_dbl, [self.state_dim, self.state_dim, self.inner_dim], dim=-1
|
| 63 |
)
|
| 64 |
+
delta = F.softplus(delta_raw + self.dt_bias)
|
| 65 |
+
A = -torch.exp(self.A_log)
|
| 66 |
+
A_bar = torch.exp(delta.unsqueeze(-1) * A.view(1, 1, self.inner_dim, self.state_dim))
|
| 67 |
+
Bx = (delta * x_in).unsqueeze(-1) * Bmat.unsqueeze(2)
|
| 68 |
+
h = parallel_scan(A_bar, Bx)
|
| 69 |
+
y = (h * Cmat.unsqueeze(2)).sum(-1) + self.D * x_in
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 70 |
y = y * F.silu(gate)
|
| 71 |
return self.out_proj(y)
|
| 72 |
|
| 73 |
+
def step(
|
| 74 |
+
self, x_t: torch.Tensor, h_prev: torch.Tensor | None
|
| 75 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
xz = self.in_proj(x_t)
|
| 77 |
+
x_in, gate = xz.chunk(2, dim=-1)
|
| 78 |
x_in = F.silu(x_in)
|
|
|
|
| 79 |
x_dbl = self.x_proj(x_in)
|
| 80 |
Bmat, Cmat, delta_raw = torch.split(
|
| 81 |
x_dbl, [self.state_dim, self.state_dim, self.inner_dim], dim=-1
|
| 82 |
)
|
| 83 |
+
delta = F.softplus(delta_raw + self.dt_bias)
|
| 84 |
+
A = -torch.exp(self.A_log)
|
| 85 |
+
A_bar = torch.exp(delta.unsqueeze(-1) * A.unsqueeze(0))
|
| 86 |
+
Bx = (delta * x_in).unsqueeze(-1) * Bmat.unsqueeze(1)
|
|
|
|
|
|
|
| 87 |
if h_prev is None:
|
| 88 |
h_prev = x_t.new_zeros(x_t.shape[0], self.inner_dim, self.state_dim)
|
| 89 |
h_new = A_bar * h_prev + Bx
|
| 90 |
+
y = (h_new * Cmat.unsqueeze(1)).sum(-1) + self.D * x_in
|
|
|
|
| 91 |
y = y * F.silu(gate)
|
| 92 |
+
return (self.out_proj(y), h_new)
|
| 93 |
|
| 94 |
|
| 95 |
class MambaBlock(nn.Module):
|
| 96 |
+
|
| 97 |
def __init__(self, dim: int, state_dim: int = 16, expand: int = 2):
|
| 98 |
super().__init__()
|
| 99 |
self.norm = RMSNorm(dim)
|
|
|
|
| 102 |
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 103 |
return x + self.ssm(self.norm(x))
|
| 104 |
|
| 105 |
+
def step(
|
| 106 |
+
self, x_t: torch.Tensor, h_prev: torch.Tensor | None
|
| 107 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 108 |
y, h_new = self.ssm.step(self.norm(x_t), h_prev)
|
| 109 |
+
return (x_t + y, h_new)
|
| 110 |
|
| 111 |
|
| 112 |
class ChessMamba(nn.Module):
|
| 113 |
+
|
| 114 |
def __init__(
|
| 115 |
self,
|
| 116 |
dim: int = 256,
|
|
|
|
| 123 |
super().__init__()
|
| 124 |
self.dim = dim
|
| 125 |
self.max_plies = max_plies
|
|
|
|
|
|
|
|
|
|
|
|
|
| 126 |
self.use_checkpoint = use_checkpoint
|
|
|
|
| 127 |
self.from_embed = nn.Embedding(64, dim)
|
| 128 |
self.to_embed = nn.Embedding(64, dim)
|
| 129 |
self.promo_embed = nn.Embedding(NUM_PROMO, dim)
|
| 130 |
+
self.pos_embed = nn.Embedding(max_plies + 1, dim)
|
| 131 |
self.start_token = nn.Parameter(torch.zeros(1, 1, dim))
|
|
|
|
| 132 |
self.blocks = nn.ModuleList([MambaBlock(dim, state_dim, expand) for _ in range(depth)])
|
| 133 |
self.norm_f = RMSNorm(dim)
|
|
|
|
| 134 |
self.policy_head = nn.Linear(dim, NUM_FROM_TO)
|
| 135 |
self.promo_head = nn.Linear(dim, NUM_PROMO)
|
| 136 |
self.value_head = nn.Linear(dim, 1)
|
|
|
|
| 137 |
nn.init.normal_(self.from_embed.weight, std=0.02)
|
| 138 |
nn.init.normal_(self.to_embed.weight, std=0.02)
|
| 139 |
nn.init.normal_(self.promo_embed.weight, std=0.02)
|
| 140 |
nn.init.normal_(self.pos_embed.weight, std=0.02)
|
| 141 |
|
| 142 |
+
def embed_moves(
|
| 143 |
+
self, from_ids: torch.Tensor, to_ids: torch.Tensor, promo_ids: torch.Tensor
|
| 144 |
+
) -> torch.Tensor:
|
| 145 |
B_ = from_ids.shape[0]
|
| 146 |
start = self.start_token.expand(B_, 1, -1)
|
| 147 |
if from_ids.shape[1] == 0:
|
|
|
|
| 152 |
positions = torch.arange(tok.shape[1], device=tok.device).unsqueeze(0)
|
| 153 |
return tok + self.pos_embed(positions)
|
| 154 |
|
| 155 |
+
def forward(
|
| 156 |
+
self,
|
| 157 |
+
from_ids: torch.Tensor,
|
| 158 |
+
to_ids: torch.Tensor,
|
| 159 |
+
promo_ids: torch.Tensor,
|
| 160 |
+
lengths: torch.Tensor | None = None,
|
| 161 |
+
):
|
|
|
|
|
|
|
|
|
|
|
|
|
| 162 |
x = self.embed_moves(from_ids, to_ids, promo_ids)
|
| 163 |
for block in self.blocks:
|
| 164 |
if self.use_checkpoint and self.training:
|
|
|
|
| 167 |
x = block(x)
|
| 168 |
x = self.norm_f(x)
|
| 169 |
if lengths is None:
|
| 170 |
+
pooled = x[:, -1]
|
| 171 |
else:
|
| 172 |
idx = lengths.view(-1, 1, 1).expand(-1, 1, x.shape[-1])
|
| 173 |
pooled = x.gather(1, idx).squeeze(1)
|
| 174 |
policy_logits = self.policy_head(pooled)
|
| 175 |
promo_logits = self.promo_head(pooled)
|
| 176 |
value = torch.tanh(self.value_head(pooled))
|
| 177 |
+
return (policy_logits, promo_logits, value)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 178 |
|
| 179 |
@torch.no_grad()
|
| 180 |
def init_incremental(self, device: torch.device | str = "cpu"):
|
|
|
|
| 186 |
block_states.append(s)
|
| 187 |
x = self.norm_f(x)
|
| 188 |
outputs = self._heads(x)
|
| 189 |
+
return ((1, block_states), outputs)
|
| 190 |
|
| 191 |
@torch.no_grad()
|
| 192 |
def step_move(self, from_sq: int, to_sq: int, promo_id: int, state):
|
| 193 |
pos_idx, block_states = state
|
| 194 |
device = self.from_embed.weight.device
|
| 195 |
pos_idx_clamped = min(pos_idx, self.max_plies)
|
| 196 |
+
idx = lambda v: torch.tensor([v], device=device)
|
| 197 |
x = (
|
| 198 |
self.from_embed(idx(from_sq))
|
| 199 |
+ self.to_embed(idx(to_sq))
|
|
|
|
| 206 |
new_block_states.append(ns)
|
| 207 |
x = self.norm_f(x)
|
| 208 |
outputs = self._heads(x)
|
| 209 |
+
return ((pos_idx + 1, new_block_states), outputs)
|
| 210 |
|
| 211 |
@torch.no_grad()
|
| 212 |
+
def build_incremental_state(
|
| 213 |
+
self,
|
| 214 |
+
from_list: list[int],
|
| 215 |
+
to_list: list[int],
|
| 216 |
+
promo_list: list[int],
|
| 217 |
+
device: torch.device | str = "cpu",
|
| 218 |
+
):
|
| 219 |
state, outputs = self.init_incremental(device)
|
| 220 |
for f, t, p in zip(from_list, to_list, promo_list):
|
| 221 |
state, outputs = self.step_move(f, t, p, state)
|
| 222 |
+
return (state, outputs)
|
| 223 |
|
| 224 |
def _heads(self, x: torch.Tensor):
|
| 225 |
policy_logits = self.policy_head(x)
|
| 226 |
promo_logits = self.promo_head(x)
|
| 227 |
value = torch.tanh(self.value_head(x))
|
| 228 |
+
return (policy_logits, promo_logits, value)
|
| 229 |
|
| 230 |
|
| 231 |
def count_params(model: nn.Module) -> int:
|
| 232 |
+
return sum((p.numel() for p in model.parameters()))
|
| 233 |
|
| 234 |
|
| 235 |
if __name__ == "__main__":
|
| 236 |
m = ChessMamba(dim=256, depth=8)
|
| 237 |
print(f"params: {count_params(m):,}")
|
| 238 |
+
B_, L = (4, 20)
|
| 239 |
from_ids = torch.randint(0, 64, (B_, L))
|
| 240 |
to_ids = torch.randint(0, 64, (B_, L))
|
| 241 |
promo_ids = torch.zeros(B_, L, dtype=torch.long)
|
search.py
CHANGED
|
@@ -1,70 +1,28 @@
|
|
| 1 |
-
"""Shallow policy-guided negamax search on top of ChessMamba.
|
| 2 |
-
|
| 3 |
-
Why search at all, given a trained policy/value net: a small, data-limited model
|
| 4 |
-
playing pure policy-argmax would blunder tactically (DeepMind's search-free
|
| 5 |
-
transformer needed 270M params + 10M Stockfish-annotated positions to get away
|
| 6 |
-
with zero search -- we have neither). Plan: expand the top-K policy-ranked moves
|
| 7 |
-
at each node, shallow alpha-beta (small fixed depth) with the value head as leaf
|
| 8 |
-
eval. This protects the Elo/win-loss score, one of the tournament's three scoring
|
| 9 |
-
axes alongside uniqueness.
|
| 10 |
-
|
| 11 |
-
Search visits many nodes per move (up to top_k^depth). Each node used to pay for a
|
| 12 |
-
full windowed forward pass (recomputing the whole move-history scan from scratch,
|
| 13 |
-
`model.py`'s pre-incremental design) -- the actual search-speed bottleneck. Nodes
|
| 14 |
-
here instead carry `model.py`'s incremental (state, outputs) pair: computing a
|
| 15 |
-
child node is one O(1) step through the model, not an O(window length) rescan.
|
| 16 |
-
Falls back to a random legal move if no checkpoint is present, so the engine is
|
| 17 |
-
always submittable regardless of how training goes under the deadline.
|
| 18 |
-
"""
|
| 19 |
-
|
| 20 |
from __future__ import annotations
|
| 21 |
-
|
| 22 |
import os
|
| 23 |
import random
|
| 24 |
import time
|
| 25 |
-
|
| 26 |
import torch
|
| 27 |
|
| 28 |
-
# search does single-position (batch=1) sequential incremental steps -- torch's default
|
| 29 |
-
# intra-op multi-threading mainly benefits large batched tensor ops, not this workload
|
| 30 |
-
# (confirmed by the CPU-beats-CUDA benchmark: latency-bound, not throughput-bound), so
|
| 31 |
-
# there's little to lose here. There's a lot to lose from NOT setting this though: two
|
| 32 |
-
# engine processes on the same test box (e.g. testing our own engine against itself, or
|
| 33 |
-
# any match where both sides run locally) each default to using every core, causing the
|
| 34 |
-
# same thread-oversubscription/contention bug already fixed in selfplay_finetune.py --
|
| 35 |
-
# caught when a diagnostic match silently ran ~15x slower than its movetime budget implied.
|
| 36 |
torch.set_num_threads(1)
|
| 37 |
-
|
| 38 |
import chess
|
| 39 |
-
|
| 40 |
from chess_io import move_to_ids, PROMO_INDEX
|
| 41 |
from model import ChessMamba, MAX_PLIES
|
| 42 |
|
| 43 |
DEFAULT_TOP_K = 10
|
| 44 |
-
DEFAULT_MAX_DEPTH = 6
|
| 45 |
DEFAULT_MOVETIME_S = 3.0
|
| 46 |
-
DEFAULT_MAX_QDEPTH = 6
|
| 47 |
-
DEFAULT_QS_TOP_K = 6
|
| 48 |
-
|
| 49 |
-
# fallback budgeting when the GUI sends a real clock (wtime/btime/winc/binc) instead of
|
| 50 |
-
# a fixed movetime -- simple "assume ~N moves left" heuristic, same idea most small
|
| 51 |
-
# engines use. Only engages when movetime is absent; a fixed movetime always wins.
|
| 52 |
CLOCK_MOVES_DIVISOR = 30
|
| 53 |
MIN_CLOCK_MOVETIME_S = 0.05
|
| 54 |
-
# hard ceiling applied AFTER any endgame extension, regardless of source (movetime,
|
| 55 |
-
# clock heuristic, or default) -- Ender's stated tournament rule is "no more than 20
|
| 56 |
-
# seconds" per move; without this, a movetime near 20000ms plus the 1.6x endgame
|
| 57 |
-
# extension would silently blow past that (e.g. 20s * 1.6 = 32s).
|
| 58 |
HARD_MOVETIME_CEILING_S = 19.0
|
| 59 |
-
|
| 60 |
-
# rough material count, used ONLY to decide how much search time an endgame position
|
| 61 |
-
# deserves -- a resource-allocation heuristic, not a position evaluation (the value head
|
| 62 |
-
# still does all actual evaluating). Standard piece-value scale, nothing unusual.
|
| 63 |
_MATERIAL_VALUE = {chess.PAWN: 1, chess.KNIGHT: 3, chess.BISHOP: 3, chess.ROOK: 5, chess.QUEEN: 9}
|
| 64 |
-
_ENDGAME_MATERIAL_THRESHOLD = 14
|
| 65 |
|
| 66 |
|
| 67 |
class Searcher:
|
|
|
|
| 68 |
def __init__(
|
| 69 |
self,
|
| 70 |
checkpoint_path: str | None = None,
|
|
@@ -75,17 +33,7 @@ class Searcher:
|
|
| 75 |
device: str = "cpu",
|
| 76 |
policy_only: bool = False,
|
| 77 |
):
|
| 78 |
-
# policy_only: skip negamax entirely, just play the top policy-ranked legal move.
|
| 79 |
-
# Used to diagnose whether the search is actually a genuine improvement over raw
|
| 80 |
-
# policy -- expert iteration (self-play fine-tuning) only helps if it is; if search
|
| 81 |
-
# doesn't reliably beat policy-only, self-play is just distilling the network's own
|
| 82 |
-
# blind spots back into itself, not a real correction.
|
| 83 |
self.policy_only = policy_only
|
| 84 |
-
# NOT auto-detecting/defaulting to cuda even when available: search does many
|
| 85 |
-
# single-position (batch=1) incremental steps, and un-benchmarked GPU use for
|
| 86 |
-
# tiny sequential calls like this is as likely to be slower (kernel-launch and
|
| 87 |
-
# host<->device transfer overhead per node) as faster. Only flip the default
|
| 88 |
-
# after actually measuring it on real target-class hardware -- see PLAN.md.
|
| 89 |
self.top_k = top_k
|
| 90 |
self.max_depth = max_depth
|
| 91 |
self.max_qdepth = max_qdepth
|
|
@@ -99,7 +47,7 @@ class Searcher:
|
|
| 99 |
self.model.load_state_dict(ckpt["model"])
|
| 100 |
self.model.eval()
|
| 101 |
self.model.to(self.device)
|
| 102 |
-
self.root_node = None
|
| 103 |
self._root_key: tuple | None = None
|
| 104 |
|
| 105 |
def reset(self) -> None:
|
|
@@ -107,24 +55,21 @@ class Searcher:
|
|
| 107 |
self._root_key = None
|
| 108 |
|
| 109 |
def sync(self, board: chess.Board) -> None:
|
| 110 |
-
"""Rebuild the root node from the board's actual move history, matching what a
|
| 111 |
-
windowed forward() call would see. Cheap (<=MAX_PLIES incremental steps), and
|
| 112 |
-
cached against the move list so redundant calls for the same position are free."""
|
| 113 |
if self.model is None:
|
| 114 |
return
|
| 115 |
window = board.move_stack[-MAX_PLIES:]
|
| 116 |
-
key = tuple(m.uci() for m in window)
|
| 117 |
if key == self._root_key:
|
| 118 |
return
|
| 119 |
self._root_key = key
|
| 120 |
from_list = [m.from_square for m in window]
|
| 121 |
to_list = [m.to_square for m in window]
|
| 122 |
promo_list = [move_to_ids(m)[1] for m in window]
|
| 123 |
-
self.root_node = self.model.build_incremental_state(
|
|
|
|
|
|
|
| 124 |
|
| 125 |
def _expand(self, board: chess.Board, node, move: chess.Move):
|
| 126 |
-
"""Pushes `move` on `board` and returns the child node. Caller must board.pop()
|
| 127 |
-
when done (the state itself needs no explicit undo -- it's just discarded)."""
|
| 128 |
state, _ = node
|
| 129 |
board.push(move)
|
| 130 |
promo_id = PROMO_INDEX[move.promotion]
|
|
@@ -134,30 +79,28 @@ class Searcher:
|
|
| 134 |
@staticmethod
|
| 135 |
def _is_endgame(board: chess.Board) -> bool:
|
| 136 |
total = sum(
|
| 137 |
-
|
| 138 |
-
|
| 139 |
-
|
|
|
|
|
|
|
| 140 |
)
|
| 141 |
return total < _ENDGAME_MATERIAL_THRESHOLD
|
| 142 |
|
| 143 |
def _policy_value(self, board: chess.Board, node):
|
| 144 |
-
"""Return (priors: dict[chess.Move -> float], value: float from side-to-move perspective)."""
|
| 145 |
if board.is_checkmate():
|
| 146 |
-
return {}, -1.0
|
| 147 |
if board.is_stalemate() or board.is_insufficient_material() or board.can_claim_draw():
|
| 148 |
-
return {}, 0.0
|
| 149 |
-
|
| 150 |
legal = list(board.legal_moves)
|
| 151 |
if not legal:
|
| 152 |
-
return {}, 0.0
|
| 153 |
if self.model is None:
|
| 154 |
uniform = 1.0 / len(legal)
|
| 155 |
-
return {m: uniform for m in legal}, 0.0
|
| 156 |
-
|
| 157 |
_, (policy_logits, promo_logits, value) = node
|
| 158 |
policy_logits = policy_logits[0]
|
| 159 |
promo_logits = promo_logits[0]
|
| 160 |
-
|
| 161 |
scores = {}
|
| 162 |
for mv in legal:
|
| 163 |
from_to_id = mv.from_square * 64 + mv.to_square
|
|
@@ -165,24 +108,19 @@ class Searcher:
|
|
| 165 |
if mv.promotion is not None:
|
| 166 |
score += promo_logits[PROMO_INDEX[mv.promotion]].item()
|
| 167 |
scores[mv] = score
|
| 168 |
-
|
| 169 |
max_score = max(scores.values())
|
| 170 |
exp_scores = {m: torch.exp(torch.tensor(s - max_score)).item() for m, s in scores.items()}
|
| 171 |
total = sum(exp_scores.values()) or 1.0
|
| 172 |
priors = {m: s / total for m, s in exp_scores.items()}
|
| 173 |
-
return priors, value.item()
|
| 174 |
|
| 175 |
-
def _quiescence(
|
| 176 |
-
|
| 177 |
-
|
| 178 |
-
classic horizon-effect blunder for any fixed-depth search). Move ordering reuses
|
| 179 |
-
the same policy priors the rest of the search uses -- no separate hand-crafted
|
| 180 |
-
piece-value ordering, consistent with evaluating everything through the net."""
|
| 181 |
if board.is_checkmate():
|
| 182 |
return -1.0
|
| 183 |
if board.is_stalemate() or board.is_insufficient_material() or board.can_claim_draw():
|
| 184 |
return 0.0
|
| 185 |
-
|
| 186 |
priors, stand_pat = self._policy_value(board, node)
|
| 187 |
if time.monotonic() > deadline or qdepth >= self.max_qdepth:
|
| 188 |
return stand_pat
|
|
@@ -190,12 +128,10 @@ class Searcher:
|
|
| 190 |
return beta
|
| 191 |
if stand_pat > alpha:
|
| 192 |
alpha = stand_pat
|
| 193 |
-
|
| 194 |
loud = [mv for mv in priors if board.is_capture(mv) or mv.promotion is not None]
|
| 195 |
if not loud:
|
| 196 |
return alpha
|
| 197 |
loud.sort(key=lambda mv: -priors[mv])
|
| 198 |
-
|
| 199 |
for mv in loud[: self.qs_top_k]:
|
| 200 |
child = self._expand(board, node, mv)
|
| 201 |
score = -self._quiescence(board, child, -beta, -alpha, deadline, qdepth + 1)
|
|
@@ -208,7 +144,9 @@ class Searcher:
|
|
| 208 |
break
|
| 209 |
return alpha
|
| 210 |
|
| 211 |
-
def _negamax(
|
|
|
|
|
|
|
| 212 |
if board.is_checkmate():
|
| 213 |
return -1.0
|
| 214 |
if board.is_stalemate() or board.is_insufficient_material() or board.can_claim_draw():
|
|
@@ -218,12 +156,10 @@ class Searcher:
|
|
| 218 |
return value
|
| 219 |
if depth == 0:
|
| 220 |
return self._quiescence(board, node, alpha, beta, deadline)
|
| 221 |
-
|
| 222 |
priors, _ = self._policy_value(board, node)
|
| 223 |
if not priors:
|
| 224 |
return 0.0
|
| 225 |
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: self.top_k]
|
| 226 |
-
|
| 227 |
best = -float("inf")
|
| 228 |
for mv, _ in ranked:
|
| 229 |
child = self._expand(board, node, mv)
|
|
@@ -241,39 +177,37 @@ class Searcher:
|
|
| 241 |
legal = list(board.legal_moves)
|
| 242 |
if not legal:
|
| 243 |
return None
|
| 244 |
-
|
| 245 |
movetime_s = DEFAULT_MOVETIME_S
|
| 246 |
if "movetime" in go_tokens:
|
| 247 |
movetime_s = int(go_tokens[go_tokens.index("movetime") + 1]) / 1000.0
|
| 248 |
elif "wtime" in go_tokens or "btime" in go_tokens:
|
| 249 |
time_key = "wtime" if board.turn == chess.WHITE else "btime"
|
| 250 |
inc_key = "winc" if board.turn == chess.WHITE else "binc"
|
| 251 |
-
own_time_ms =
|
| 252 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 253 |
movetime_s = own_time_ms / 1000.0 / CLOCK_MOVES_DIVISOR + own_inc_ms / 1000.0
|
| 254 |
movetime_s = max(MIN_CLOCK_MOVETIME_S, movetime_s)
|
| 255 |
if self._is_endgame(board):
|
| 256 |
-
movetime_s *= 1.6
|
| 257 |
movetime_s = min(movetime_s, HARD_MOVETIME_CEILING_S)
|
| 258 |
deadline = time.monotonic() + movetime_s
|
| 259 |
-
|
| 260 |
if self.model is None:
|
| 261 |
return random.choice(legal)
|
| 262 |
-
|
| 263 |
self.sync(board)
|
| 264 |
node = self.root_node
|
| 265 |
-
|
| 266 |
priors, _ = self._policy_value(board, node)
|
| 267 |
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: self.top_k]
|
| 268 |
-
|
| 269 |
best_move = ranked[0][0] if ranked else random.choice(legal)
|
| 270 |
if self.policy_only:
|
| 271 |
return best_move
|
| 272 |
best_score = -float("inf")
|
| 273 |
-
|
| 274 |
depth = 1
|
| 275 |
while depth <= self.max_depth and time.monotonic() < deadline:
|
| 276 |
-
alpha, beta = -float("inf"), float("inf")
|
| 277 |
current_best_move = None
|
| 278 |
current_best_score = -float("inf")
|
| 279 |
for mv, _ in ranked:
|
|
@@ -288,7 +222,6 @@ class Searcher:
|
|
| 288 |
if time.monotonic() > deadline:
|
| 289 |
break
|
| 290 |
if current_best_move is not None and time.monotonic() <= deadline:
|
| 291 |
-
best_move, best_score = current_best_move, current_best_score
|
| 292 |
depth += 1
|
| 293 |
-
|
| 294 |
return best_move
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import os
|
| 3 |
import random
|
| 4 |
import time
|
|
|
|
| 5 |
import torch
|
| 6 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 7 |
torch.set_num_threads(1)
|
|
|
|
| 8 |
import chess
|
|
|
|
| 9 |
from chess_io import move_to_ids, PROMO_INDEX
|
| 10 |
from model import ChessMamba, MAX_PLIES
|
| 11 |
|
| 12 |
DEFAULT_TOP_K = 10
|
| 13 |
+
DEFAULT_MAX_DEPTH = 6
|
| 14 |
DEFAULT_MOVETIME_S = 3.0
|
| 15 |
+
DEFAULT_MAX_QDEPTH = 6
|
| 16 |
+
DEFAULT_QS_TOP_K = 6
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
CLOCK_MOVES_DIVISOR = 30
|
| 18 |
MIN_CLOCK_MOVETIME_S = 0.05
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
HARD_MOVETIME_CEILING_S = 19.0
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
_MATERIAL_VALUE = {chess.PAWN: 1, chess.KNIGHT: 3, chess.BISHOP: 3, chess.ROOK: 5, chess.QUEEN: 9}
|
| 21 |
+
_ENDGAME_MATERIAL_THRESHOLD = 14
|
| 22 |
|
| 23 |
|
| 24 |
class Searcher:
|
| 25 |
+
|
| 26 |
def __init__(
|
| 27 |
self,
|
| 28 |
checkpoint_path: str | None = None,
|
|
|
|
| 33 |
device: str = "cpu",
|
| 34 |
policy_only: bool = False,
|
| 35 |
):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
self.policy_only = policy_only
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 37 |
self.top_k = top_k
|
| 38 |
self.max_depth = max_depth
|
| 39 |
self.max_qdepth = max_qdepth
|
|
|
|
| 47 |
self.model.load_state_dict(ckpt["model"])
|
| 48 |
self.model.eval()
|
| 49 |
self.model.to(self.device)
|
| 50 |
+
self.root_node = None
|
| 51 |
self._root_key: tuple | None = None
|
| 52 |
|
| 53 |
def reset(self) -> None:
|
|
|
|
| 55 |
self._root_key = None
|
| 56 |
|
| 57 |
def sync(self, board: chess.Board) -> None:
|
|
|
|
|
|
|
|
|
|
| 58 |
if self.model is None:
|
| 59 |
return
|
| 60 |
window = board.move_stack[-MAX_PLIES:]
|
| 61 |
+
key = tuple((m.uci() for m in window))
|
| 62 |
if key == self._root_key:
|
| 63 |
return
|
| 64 |
self._root_key = key
|
| 65 |
from_list = [m.from_square for m in window]
|
| 66 |
to_list = [m.to_square for m in window]
|
| 67 |
promo_list = [move_to_ids(m)[1] for m in window]
|
| 68 |
+
self.root_node = self.model.build_incremental_state(
|
| 69 |
+
from_list, to_list, promo_list, device=self.device
|
| 70 |
+
)
|
| 71 |
|
| 72 |
def _expand(self, board: chess.Board, node, move: chess.Move):
|
|
|
|
|
|
|
| 73 |
state, _ = node
|
| 74 |
board.push(move)
|
| 75 |
promo_id = PROMO_INDEX[move.promotion]
|
|
|
|
| 79 |
@staticmethod
|
| 80 |
def _is_endgame(board: chess.Board) -> bool:
|
| 81 |
total = sum(
|
| 82 |
+
(
|
| 83 |
+
_MATERIAL_VALUE.get(p.piece_type, 0)
|
| 84 |
+
for p in board.piece_map().values()
|
| 85 |
+
if p.piece_type != chess.KING
|
| 86 |
+
)
|
| 87 |
)
|
| 88 |
return total < _ENDGAME_MATERIAL_THRESHOLD
|
| 89 |
|
| 90 |
def _policy_value(self, board: chess.Board, node):
|
|
|
|
| 91 |
if board.is_checkmate():
|
| 92 |
+
return ({}, -1.0)
|
| 93 |
if board.is_stalemate() or board.is_insufficient_material() or board.can_claim_draw():
|
| 94 |
+
return ({}, 0.0)
|
|
|
|
| 95 |
legal = list(board.legal_moves)
|
| 96 |
if not legal:
|
| 97 |
+
return ({}, 0.0)
|
| 98 |
if self.model is None:
|
| 99 |
uniform = 1.0 / len(legal)
|
| 100 |
+
return ({m: uniform for m in legal}, 0.0)
|
|
|
|
| 101 |
_, (policy_logits, promo_logits, value) = node
|
| 102 |
policy_logits = policy_logits[0]
|
| 103 |
promo_logits = promo_logits[0]
|
|
|
|
| 104 |
scores = {}
|
| 105 |
for mv in legal:
|
| 106 |
from_to_id = mv.from_square * 64 + mv.to_square
|
|
|
|
| 108 |
if mv.promotion is not None:
|
| 109 |
score += promo_logits[PROMO_INDEX[mv.promotion]].item()
|
| 110 |
scores[mv] = score
|
|
|
|
| 111 |
max_score = max(scores.values())
|
| 112 |
exp_scores = {m: torch.exp(torch.tensor(s - max_score)).item() for m, s in scores.items()}
|
| 113 |
total = sum(exp_scores.values()) or 1.0
|
| 114 |
priors = {m: s / total for m, s in exp_scores.items()}
|
| 115 |
+
return (priors, value.item())
|
| 116 |
|
| 117 |
+
def _quiescence(
|
| 118 |
+
self, board: chess.Board, node, alpha: float, beta: float, deadline: float, qdepth: int = 0
|
| 119 |
+
) -> float:
|
|
|
|
|
|
|
|
|
|
| 120 |
if board.is_checkmate():
|
| 121 |
return -1.0
|
| 122 |
if board.is_stalemate() or board.is_insufficient_material() or board.can_claim_draw():
|
| 123 |
return 0.0
|
|
|
|
| 124 |
priors, stand_pat = self._policy_value(board, node)
|
| 125 |
if time.monotonic() > deadline or qdepth >= self.max_qdepth:
|
| 126 |
return stand_pat
|
|
|
|
| 128 |
return beta
|
| 129 |
if stand_pat > alpha:
|
| 130 |
alpha = stand_pat
|
|
|
|
| 131 |
loud = [mv for mv in priors if board.is_capture(mv) or mv.promotion is not None]
|
| 132 |
if not loud:
|
| 133 |
return alpha
|
| 134 |
loud.sort(key=lambda mv: -priors[mv])
|
|
|
|
| 135 |
for mv in loud[: self.qs_top_k]:
|
| 136 |
child = self._expand(board, node, mv)
|
| 137 |
score = -self._quiescence(board, child, -beta, -alpha, deadline, qdepth + 1)
|
|
|
|
| 144 |
break
|
| 145 |
return alpha
|
| 146 |
|
| 147 |
+
def _negamax(
|
| 148 |
+
self, board: chess.Board, node, depth: int, alpha: float, beta: float, deadline: float
|
| 149 |
+
) -> float:
|
| 150 |
if board.is_checkmate():
|
| 151 |
return -1.0
|
| 152 |
if board.is_stalemate() or board.is_insufficient_material() or board.can_claim_draw():
|
|
|
|
| 156 |
return value
|
| 157 |
if depth == 0:
|
| 158 |
return self._quiescence(board, node, alpha, beta, deadline)
|
|
|
|
| 159 |
priors, _ = self._policy_value(board, node)
|
| 160 |
if not priors:
|
| 161 |
return 0.0
|
| 162 |
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: self.top_k]
|
|
|
|
| 163 |
best = -float("inf")
|
| 164 |
for mv, _ in ranked:
|
| 165 |
child = self._expand(board, node, mv)
|
|
|
|
| 177 |
legal = list(board.legal_moves)
|
| 178 |
if not legal:
|
| 179 |
return None
|
|
|
|
| 180 |
movetime_s = DEFAULT_MOVETIME_S
|
| 181 |
if "movetime" in go_tokens:
|
| 182 |
movetime_s = int(go_tokens[go_tokens.index("movetime") + 1]) / 1000.0
|
| 183 |
elif "wtime" in go_tokens or "btime" in go_tokens:
|
| 184 |
time_key = "wtime" if board.turn == chess.WHITE else "btime"
|
| 185 |
inc_key = "winc" if board.turn == chess.WHITE else "binc"
|
| 186 |
+
own_time_ms = (
|
| 187 |
+
int(go_tokens[go_tokens.index(time_key) + 1]) if time_key in go_tokens else 0
|
| 188 |
+
)
|
| 189 |
+
own_inc_ms = (
|
| 190 |
+
int(go_tokens[go_tokens.index(inc_key) + 1]) if inc_key in go_tokens else 0
|
| 191 |
+
)
|
| 192 |
movetime_s = own_time_ms / 1000.0 / CLOCK_MOVES_DIVISOR + own_inc_ms / 1000.0
|
| 193 |
movetime_s = max(MIN_CLOCK_MOVETIME_S, movetime_s)
|
| 194 |
if self._is_endgame(board):
|
| 195 |
+
movetime_s *= 1.6
|
| 196 |
movetime_s = min(movetime_s, HARD_MOVETIME_CEILING_S)
|
| 197 |
deadline = time.monotonic() + movetime_s
|
|
|
|
| 198 |
if self.model is None:
|
| 199 |
return random.choice(legal)
|
|
|
|
| 200 |
self.sync(board)
|
| 201 |
node = self.root_node
|
|
|
|
| 202 |
priors, _ = self._policy_value(board, node)
|
| 203 |
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: self.top_k]
|
|
|
|
| 204 |
best_move = ranked[0][0] if ranked else random.choice(legal)
|
| 205 |
if self.policy_only:
|
| 206 |
return best_move
|
| 207 |
best_score = -float("inf")
|
|
|
|
| 208 |
depth = 1
|
| 209 |
while depth <= self.max_depth and time.monotonic() < deadline:
|
| 210 |
+
alpha, beta = (-float("inf"), float("inf"))
|
| 211 |
current_best_move = None
|
| 212 |
current_best_score = -float("inf")
|
| 213 |
for mv, _ in ranked:
|
|
|
|
| 222 |
if time.monotonic() > deadline:
|
| 223 |
break
|
| 224 |
if current_best_move is not None and time.monotonic() <= deadline:
|
| 225 |
+
best_move, best_score = (current_best_move, current_best_score)
|
| 226 |
depth += 1
|
|
|
|
| 227 |
return best_move
|
training/data_pipeline.py
CHANGED
|
@@ -1,45 +1,30 @@
|
|
| 1 |
-
"""Streams Lichess monthly rated-game dumps straight off the network (no full
|
| 2 |
-
download-then-decompress -- decompresses and parses on the fly, stops once enough
|
| 3 |
-
games are collected) and packs them into fixed-shape numpy shards for training.
|
| 4 |
-
|
| 5 |
-
Each training example = one ply: (move history up to MAX_PLIES back, the move
|
| 6 |
-
actually played, the eventual game result from the mover's perspective). This
|
| 7 |
-
mirrors exactly what search.py's encode_history() builds at inference time.
|
| 8 |
-
|
| 9 |
-
Run on the box with real network access / disk, not necessarily locally:
|
| 10 |
-
python3 data_pipeline.py --out-dir /root/chess/data --target-games 500000
|
| 11 |
-
"""
|
| 12 |
-
|
| 13 |
from __future__ import annotations
|
| 14 |
-
|
| 15 |
import argparse
|
| 16 |
import io
|
| 17 |
import os
|
| 18 |
import time
|
| 19 |
import urllib.request
|
| 20 |
-
|
| 21 |
import chess
|
| 22 |
import chess.pgn
|
| 23 |
import numpy as np
|
| 24 |
import zstandard
|
| 25 |
-
|
| 26 |
from chess_io import move_to_ids
|
| 27 |
from model import MAX_PLIES
|
| 28 |
|
| 29 |
BASE_URL = "https://database.lichess.org/standard/lichess_db_standard_rated_{}.pgn.zst"
|
| 30 |
MIN_ELO = 1800
|
| 31 |
MIN_PLIES = 10
|
| 32 |
-
SHARD_SIZE =
|
| 33 |
-
SOFT_K = 10
|
| 34 |
|
| 35 |
|
| 36 |
def candidate_months(start_year: int = 2026, start_month: int = 7, count: int = 30):
|
| 37 |
-
y, m = start_year, start_month
|
| 38 |
for _ in range(count):
|
| 39 |
yield f"{y:04d}-{m:02d}"
|
| 40 |
m -= 1
|
| 41 |
if m == 0:
|
| 42 |
-
m, y = 12, y - 1
|
| 43 |
|
| 44 |
|
| 45 |
def find_available_month() -> str:
|
|
@@ -85,7 +70,6 @@ def game_is_ok(game: chess.pgn.Game) -> bool:
|
|
| 85 |
|
| 86 |
|
| 87 |
def result_value(result: str, side_to_move: bool) -> int:
|
| 88 |
-
"""side_to_move: chess.WHITE (True) or chess.BLACK (False). Returns +1/0/-1 from that side's view."""
|
| 89 |
if result == "1/2-1/2":
|
| 90 |
return 0
|
| 91 |
white_won = result == "1-0"
|
|
@@ -99,10 +83,8 @@ def game_to_examples(game: chess.pgn.Game):
|
|
| 99 |
hist_from: list[int] = []
|
| 100 |
hist_to: list[int] = []
|
| 101 |
hist_promo: list[int] = []
|
| 102 |
-
|
| 103 |
if len(list(game.mainline_moves())) < MIN_PLIES:
|
| 104 |
return
|
| 105 |
-
|
| 106 |
for move in game.mainline_moves():
|
| 107 |
side_to_move = board.turn
|
| 108 |
from_to_id, promo_id = move_to_ids(move)
|
|
@@ -111,7 +93,6 @@ def game_to_examples(game: chess.pgn.Game):
|
|
| 111 |
window_promo = hist_promo[-MAX_PLIES:]
|
| 112 |
value = result_value(result, side_to_move)
|
| 113 |
yield (window_from, window_to, window_promo, from_to_id, promo_id, value)
|
| 114 |
-
|
| 115 |
hist_from.append(move.from_square)
|
| 116 |
hist_to.append(move.to_square)
|
| 117 |
hist_promo.append(promo_id)
|
|
@@ -119,6 +100,7 @@ def game_to_examples(game: chess.pgn.Game):
|
|
| 119 |
|
| 120 |
|
| 121 |
class ShardWriter:
|
|
|
|
| 122 |
def __init__(self, out_dir: str, shard_size: int = SHARD_SIZE):
|
| 123 |
os.makedirs(out_dir, exist_ok=True)
|
| 124 |
self.out_dir = out_dir
|
|
@@ -135,20 +117,21 @@ class ShardWriter:
|
|
| 135 |
self.target_from_to = np.zeros((n,), dtype=np.int16)
|
| 136 |
self.target_promo = np.zeros((n,), dtype=np.int8)
|
| 137 |
self.values = np.zeros((n,), dtype=np.int8)
|
| 138 |
-
# distributional policy target (see class docstring on `add`) -- degenerate
|
| 139 |
-
# one-hot by default so ordinary hard-label writes (data_pipeline.py's own
|
| 140 |
-
# game_to_examples) are unaffected and unambiguous to a soft-aware loss.
|
| 141 |
self.soft_move_ids = np.full((n, SOFT_K), -1, dtype=np.int16)
|
| 142 |
self.soft_weights = np.zeros((n, SOFT_K), dtype=np.float32)
|
| 143 |
self.count = 0
|
| 144 |
|
| 145 |
-
def add(
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
target_from_to
|
| 151 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 152 |
i = self.count
|
| 153 |
L = len(window_from)
|
| 154 |
if L:
|
|
@@ -194,18 +177,17 @@ class ShardWriter:
|
|
| 194 |
def main():
|
| 195 |
ap = argparse.ArgumentParser()
|
| 196 |
ap.add_argument("--out-dir", default="/root/chess/data")
|
| 197 |
-
ap.add_argument("--target-games", type=int, default=
|
| 198 |
-
ap.add_argument(
|
|
|
|
|
|
|
| 199 |
args = ap.parse_args()
|
| 200 |
-
|
| 201 |
url = find_available_month()
|
| 202 |
writer = ShardWriter(args.out_dir)
|
| 203 |
-
|
| 204 |
kept_games = 0
|
| 205 |
seen_games = 0
|
| 206 |
total_examples = 0
|
| 207 |
t0 = time.monotonic()
|
| 208 |
-
|
| 209 |
for game in stream_games(url):
|
| 210 |
seen_games += 1
|
| 211 |
if not game_is_ok(game):
|
|
@@ -214,19 +196,15 @@ def main():
|
|
| 214 |
writer.add(*ex)
|
| 215 |
total_examples += 1
|
| 216 |
kept_games += 1
|
| 217 |
-
|
| 218 |
if kept_games % 2000 == 0:
|
| 219 |
dt = time.monotonic() - t0
|
| 220 |
print(
|
| 221 |
-
f"[data_pipeline] kept={kept_games} seen={seen_games} examples={total_examples} "
|
| 222 |
-
f"elapsed={dt:.0f}s rate={kept_games/dt:.1f} games/s"
|
| 223 |
)
|
| 224 |
-
|
| 225 |
if kept_games >= args.target_games:
|
| 226 |
break
|
| 227 |
if args.target_examples and total_examples >= args.target_examples:
|
| 228 |
break
|
| 229 |
-
|
| 230 |
writer.flush()
|
| 231 |
print(f"[data_pipeline] DONE kept={kept_games} seen={seen_games} examples={total_examples}")
|
| 232 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import argparse
|
| 3 |
import io
|
| 4 |
import os
|
| 5 |
import time
|
| 6 |
import urllib.request
|
|
|
|
| 7 |
import chess
|
| 8 |
import chess.pgn
|
| 9 |
import numpy as np
|
| 10 |
import zstandard
|
|
|
|
| 11 |
from chess_io import move_to_ids
|
| 12 |
from model import MAX_PLIES
|
| 13 |
|
| 14 |
BASE_URL = "https://database.lichess.org/standard/lichess_db_standard_rated_{}.pgn.zst"
|
| 15 |
MIN_ELO = 1800
|
| 16 |
MIN_PLIES = 10
|
| 17 |
+
SHARD_SIZE = 200000
|
| 18 |
+
SOFT_K = 10
|
| 19 |
|
| 20 |
|
| 21 |
def candidate_months(start_year: int = 2026, start_month: int = 7, count: int = 30):
|
| 22 |
+
y, m = (start_year, start_month)
|
| 23 |
for _ in range(count):
|
| 24 |
yield f"{y:04d}-{m:02d}"
|
| 25 |
m -= 1
|
| 26 |
if m == 0:
|
| 27 |
+
m, y = (12, y - 1)
|
| 28 |
|
| 29 |
|
| 30 |
def find_available_month() -> str:
|
|
|
|
| 70 |
|
| 71 |
|
| 72 |
def result_value(result: str, side_to_move: bool) -> int:
|
|
|
|
| 73 |
if result == "1/2-1/2":
|
| 74 |
return 0
|
| 75 |
white_won = result == "1-0"
|
|
|
|
| 83 |
hist_from: list[int] = []
|
| 84 |
hist_to: list[int] = []
|
| 85 |
hist_promo: list[int] = []
|
|
|
|
| 86 |
if len(list(game.mainline_moves())) < MIN_PLIES:
|
| 87 |
return
|
|
|
|
| 88 |
for move in game.mainline_moves():
|
| 89 |
side_to_move = board.turn
|
| 90 |
from_to_id, promo_id = move_to_ids(move)
|
|
|
|
| 93 |
window_promo = hist_promo[-MAX_PLIES:]
|
| 94 |
value = result_value(result, side_to_move)
|
| 95 |
yield (window_from, window_to, window_promo, from_to_id, promo_id, value)
|
|
|
|
| 96 |
hist_from.append(move.from_square)
|
| 97 |
hist_to.append(move.to_square)
|
| 98 |
hist_promo.append(promo_id)
|
|
|
|
| 100 |
|
| 101 |
|
| 102 |
class ShardWriter:
|
| 103 |
+
|
| 104 |
def __init__(self, out_dir: str, shard_size: int = SHARD_SIZE):
|
| 105 |
os.makedirs(out_dir, exist_ok=True)
|
| 106 |
self.out_dir = out_dir
|
|
|
|
| 117 |
self.target_from_to = np.zeros((n,), dtype=np.int16)
|
| 118 |
self.target_promo = np.zeros((n,), dtype=np.int8)
|
| 119 |
self.values = np.zeros((n,), dtype=np.int8)
|
|
|
|
|
|
|
|
|
|
| 120 |
self.soft_move_ids = np.full((n, SOFT_K), -1, dtype=np.int16)
|
| 121 |
self.soft_weights = np.zeros((n, SOFT_K), dtype=np.float32)
|
| 122 |
self.count = 0
|
| 123 |
|
| 124 |
+
def add(
|
| 125 |
+
self,
|
| 126 |
+
window_from,
|
| 127 |
+
window_to,
|
| 128 |
+
window_promo,
|
| 129 |
+
target_from_to,
|
| 130 |
+
target_promo,
|
| 131 |
+
value,
|
| 132 |
+
soft_move_ids=None,
|
| 133 |
+
soft_weights=None,
|
| 134 |
+
):
|
| 135 |
i = self.count
|
| 136 |
L = len(window_from)
|
| 137 |
if L:
|
|
|
|
| 177 |
def main():
|
| 178 |
ap = argparse.ArgumentParser()
|
| 179 |
ap.add_argument("--out-dir", default="/root/chess/data")
|
| 180 |
+
ap.add_argument("--target-games", type=int, default=500000)
|
| 181 |
+
ap.add_argument(
|
| 182 |
+
"--target-examples", type=int, default=0, help="0 = no cap, stop on target-games instead"
|
| 183 |
+
)
|
| 184 |
args = ap.parse_args()
|
|
|
|
| 185 |
url = find_available_month()
|
| 186 |
writer = ShardWriter(args.out_dir)
|
|
|
|
| 187 |
kept_games = 0
|
| 188 |
seen_games = 0
|
| 189 |
total_examples = 0
|
| 190 |
t0 = time.monotonic()
|
|
|
|
| 191 |
for game in stream_games(url):
|
| 192 |
seen_games += 1
|
| 193 |
if not game_is_ok(game):
|
|
|
|
| 196 |
writer.add(*ex)
|
| 197 |
total_examples += 1
|
| 198 |
kept_games += 1
|
|
|
|
| 199 |
if kept_games % 2000 == 0:
|
| 200 |
dt = time.monotonic() - t0
|
| 201 |
print(
|
| 202 |
+
f"[data_pipeline] kept={kept_games} seen={seen_games} examples={total_examples} elapsed={dt:.0f}s rate={kept_games / dt:.1f} games/s"
|
|
|
|
| 203 |
)
|
|
|
|
| 204 |
if kept_games >= args.target_games:
|
| 205 |
break
|
| 206 |
if args.target_examples and total_examples >= args.target_examples:
|
| 207 |
break
|
|
|
|
| 208 |
writer.flush()
|
| 209 |
print(f"[data_pipeline] DONE kept={kept_games} seen={seen_games} examples={total_examples}")
|
| 210 |
|
training/match.py
CHANGED
|
@@ -1,24 +1,13 @@
|
|
| 1 |
-
"""Minimal UCI-vs-UCI match driver, for informal calibration -- not a submission
|
| 2 |
-
artifact. Runs two engines as subprocesses speaking the UCI protocol over
|
| 3 |
-
stdin/stdout, alternates moves on a shared board, reports the result. Useful for
|
| 4 |
-
sanity-checking a checkpoint against a known reference (e.g. Sunfish) or against
|
| 5 |
-
itself, since we have no access to a real rating pool before the tournament.
|
| 6 |
-
|
| 7 |
-
Usage:
|
| 8 |
-
python3 match.py --white "python3 engine_uci.py" --black "python3 sunfish_uci.py" --movetime 200
|
| 9 |
-
"""
|
| 10 |
-
|
| 11 |
from __future__ import annotations
|
| 12 |
-
|
| 13 |
import argparse
|
| 14 |
import shlex
|
| 15 |
import subprocess
|
| 16 |
import time
|
| 17 |
-
|
| 18 |
import chess
|
| 19 |
|
| 20 |
|
| 21 |
class UciEngine:
|
|
|
|
| 22 |
def __init__(self, command: str, cwd: str | None = None):
|
| 23 |
self.proc = subprocess.Popen(
|
| 24 |
shlex.split(command),
|
|
@@ -74,12 +63,13 @@ class UciEngine:
|
|
| 74 |
self.proc.kill()
|
| 75 |
|
| 76 |
|
| 77 |
-
def play_game(
|
|
|
|
|
|
|
| 78 |
white = UciEngine(white_cmd)
|
| 79 |
black = UciEngine(black_cmd)
|
| 80 |
white.new_game()
|
| 81 |
black.new_game()
|
| 82 |
-
|
| 83 |
board = chess.Board()
|
| 84 |
moves_uci: list[str] = []
|
| 85 |
result = "*"
|
|
@@ -90,7 +80,9 @@ def play_game(white_cmd: str, black_cmd: str, movetime_ms: int, max_plies: int =
|
|
| 90 |
move = chess.Move.from_uci(mv_uci)
|
| 91 |
if move not in board.legal_moves:
|
| 92 |
result = "0-1" if board.turn == chess.WHITE else "1-0"
|
| 93 |
-
print(
|
|
|
|
|
|
|
| 94 |
break
|
| 95 |
board.push(move)
|
| 96 |
moves_uci.append(mv_uci)
|
|
@@ -98,11 +90,11 @@ def play_game(white_cmd: str, black_cmd: str, movetime_ms: int, max_plies: int =
|
|
| 98 |
result = board.result()
|
| 99 |
break
|
| 100 |
else:
|
| 101 |
-
result = "1/2-1/2"
|
| 102 |
finally:
|
| 103 |
white.close()
|
| 104 |
black.close()
|
| 105 |
-
return result, moves_uci
|
| 106 |
|
| 107 |
|
| 108 |
def main():
|
|
@@ -112,7 +104,6 @@ def main():
|
|
| 112 |
ap.add_argument("--movetime", type=int, default=300)
|
| 113 |
ap.add_argument("--games", type=int, default=1)
|
| 114 |
args = ap.parse_args()
|
| 115 |
-
|
| 116 |
tally = {"white": 0, "black": 0, "draw": 0}
|
| 117 |
for g in range(args.games):
|
| 118 |
w, b = (args.white, args.black) if g % 2 == 0 else (args.black, args.white)
|
|
@@ -124,7 +115,9 @@ def main():
|
|
| 124 |
tally["black" if g % 2 == 0 else "white"] += 1
|
| 125 |
else:
|
| 126 |
tally["draw"] += 1
|
| 127 |
-
print(
|
|
|
|
|
|
|
| 128 |
|
| 129 |
|
| 130 |
if __name__ == "__main__":
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import argparse
|
| 3 |
import shlex
|
| 4 |
import subprocess
|
| 5 |
import time
|
|
|
|
| 6 |
import chess
|
| 7 |
|
| 8 |
|
| 9 |
class UciEngine:
|
| 10 |
+
|
| 11 |
def __init__(self, command: str, cwd: str | None = None):
|
| 12 |
self.proc = subprocess.Popen(
|
| 13 |
shlex.split(command),
|
|
|
|
| 63 |
self.proc.kill()
|
| 64 |
|
| 65 |
|
| 66 |
+
def play_game(
|
| 67 |
+
white_cmd: str, black_cmd: str, movetime_ms: int, max_plies: int = 300
|
| 68 |
+
) -> tuple[str, list[str]]:
|
| 69 |
white = UciEngine(white_cmd)
|
| 70 |
black = UciEngine(black_cmd)
|
| 71 |
white.new_game()
|
| 72 |
black.new_game()
|
|
|
|
| 73 |
board = chess.Board()
|
| 74 |
moves_uci: list[str] = []
|
| 75 |
result = "*"
|
|
|
|
| 80 |
move = chess.Move.from_uci(mv_uci)
|
| 81 |
if move not in board.legal_moves:
|
| 82 |
result = "0-1" if board.turn == chess.WHITE else "1-0"
|
| 83 |
+
print(
|
| 84 |
+
f"[match] ILLEGAL MOVE {mv_uci} by {('white' if board.turn else 'black')} -- forfeit"
|
| 85 |
+
)
|
| 86 |
break
|
| 87 |
board.push(move)
|
| 88 |
moves_uci.append(mv_uci)
|
|
|
|
| 90 |
result = board.result()
|
| 91 |
break
|
| 92 |
else:
|
| 93 |
+
result = "1/2-1/2"
|
| 94 |
finally:
|
| 95 |
white.close()
|
| 96 |
black.close()
|
| 97 |
+
return (result, moves_uci)
|
| 98 |
|
| 99 |
|
| 100 |
def main():
|
|
|
|
| 104 |
ap.add_argument("--movetime", type=int, default=300)
|
| 105 |
ap.add_argument("--games", type=int, default=1)
|
| 106 |
args = ap.parse_args()
|
|
|
|
| 107 |
tally = {"white": 0, "black": 0, "draw": 0}
|
| 108 |
for g in range(args.games):
|
| 109 |
w, b = (args.white, args.black) if g % 2 == 0 else (args.black, args.white)
|
|
|
|
| 115 |
tally["black" if g % 2 == 0 else "white"] += 1
|
| 116 |
else:
|
| 117 |
tally["draw"] += 1
|
| 118 |
+
print(
|
| 119 |
+
f"[summary] --white wins={tally['white']} --black wins={tally['black']} draws={tally['draw']}"
|
| 120 |
+
)
|
| 121 |
|
| 122 |
|
| 123 |
if __name__ == "__main__":
|
training/pretrain.py
CHANGED
|
@@ -1,39 +1,25 @@
|
|
| 1 |
-
"""Supervised pretraining for ChessMamba: masked-legal next-move cross-entropy +
|
| 2 |
-
game-outcome value regression, over shards produced by data_pipeline.py.
|
| 3 |
-
|
| 4 |
-
Deliberately plain AdamW, not a custom optimizer -- this is a ~3.5-day deadline
|
| 5 |
-
project, not a from-scratch-optimizer research effort like the LM projects.
|
| 6 |
-
"""
|
| 7 |
-
|
| 8 |
from __future__ import annotations
|
| 9 |
-
|
| 10 |
import argparse
|
| 11 |
import glob
|
| 12 |
import os
|
| 13 |
import time
|
| 14 |
-
|
| 15 |
import numpy as np
|
| 16 |
import torch
|
| 17 |
import torch.nn.functional as F
|
| 18 |
-
|
| 19 |
from data_pipeline import SOFT_K
|
| 20 |
from model import ChessMamba, MAX_PLIES
|
| 21 |
|
| 22 |
|
| 23 |
class ShardDataset:
|
| 24 |
-
"""Loads all shards fully into RAM (data is small enough -- tens of GB at most
|
| 25 |
-
against 188GB host RAM) and hands out shuffled batches with per-batch length
|
| 26 |
-
truncation (batches are padded to their own max real length, not the global
|
| 27 |
-
MAX_PLIES, since the scan cost is O(L))."""
|
| 28 |
|
| 29 |
def __init__(self, data_dir: str, val_fraction: float = 0.01):
|
| 30 |
-
# recursive glob: parallel self-play workers each write shard_*.npz into their own
|
| 31 |
-
# subdirectory (avoids filename collisions between workers writing concurrently)
|
| 32 |
paths = sorted(glob.glob(os.path.join(data_dir, "**", "shard_*.npz"), recursive=True))
|
| 33 |
if not paths:
|
| 34 |
raise RuntimeError(f"no shards found in {data_dir}")
|
| 35 |
-
from_ids, to_ids, promo_ids, lengths, target_from_to, target_promo, values = (
|
| 36 |
-
|
|
|
|
|
|
|
| 37 |
for p in paths:
|
| 38 |
d = np.load(p)
|
| 39 |
from_ids.append(d["from_ids"])
|
|
@@ -47,9 +33,6 @@ class ShardDataset:
|
|
| 47 |
soft_move_ids.append(d["soft_move_ids"])
|
| 48 |
soft_weights.append(d["soft_weights"])
|
| 49 |
else:
|
| 50 |
-
# shards written before the soft-target field existed: synthesize the
|
| 51 |
-
# equivalent degenerate one-hot distribution so the loss can treat every
|
| 52 |
-
# example uniformly regardless of which stage produced it.
|
| 53 |
n_ex = len(d["target_from_to"])
|
| 54 |
one_hot_ids = np.full((n_ex, SOFT_K), -1, dtype=np.int16)
|
| 55 |
one_hot_ids[:, 0] = d["target_from_to"]
|
|
@@ -57,7 +40,6 @@ class ShardDataset:
|
|
| 57 |
one_hot_w[:, 0] = 1.0
|
| 58 |
soft_move_ids.append(one_hot_ids)
|
| 59 |
soft_weights.append(one_hot_w)
|
| 60 |
-
|
| 61 |
self.from_ids = np.concatenate(from_ids)
|
| 62 |
self.to_ids = np.concatenate(to_ids)
|
| 63 |
self.promo_ids = np.concatenate(promo_ids)
|
|
@@ -67,20 +49,17 @@ class ShardDataset:
|
|
| 67 |
self.values = np.concatenate(values)
|
| 68 |
self.soft_move_ids = np.concatenate(soft_move_ids)
|
| 69 |
self.soft_weights = np.concatenate(soft_weights)
|
| 70 |
-
|
| 71 |
n = len(self.lengths)
|
| 72 |
rng = np.random.default_rng(0)
|
| 73 |
perm = rng.permutation(n)
|
| 74 |
n_val = max(1, int(n * val_fraction))
|
| 75 |
self.val_idx = perm[:n_val]
|
| 76 |
self.train_idx = perm[n_val:]
|
| 77 |
-
print(
|
|
|
|
|
|
|
| 78 |
|
| 79 |
def get_batch(self, idx: np.ndarray, device: torch.device, fixed_len: int | None = None):
|
| 80 |
-
# fixed_len pins every batch to the same shape (full MAX_PLIES by default) --
|
| 81 |
-
# required under torch.compile, which recompiles (expensive, ~60s) per distinct
|
| 82 |
-
# shape it sees; truncating to each batch's own max real length would otherwise
|
| 83 |
-
# trigger a recompile storm across the many different lengths in the data.
|
| 84 |
max_len = fixed_len if fixed_len is not None else int(self.lengths[idx].max())
|
| 85 |
from_ids = torch.from_numpy(self.from_ids[idx][:, :max_len].astype(np.int64)).to(device)
|
| 86 |
to_ids = torch.from_numpy(self.to_ids[idx][:, :max_len].astype(np.int64)).to(device)
|
|
@@ -91,47 +70,57 @@ class ShardDataset:
|
|
| 91 |
values = torch.from_numpy(self.values[idx].astype(np.float32)).to(device)
|
| 92 |
soft_move_ids = torch.from_numpy(self.soft_move_ids[idx].astype(np.int64)).to(device)
|
| 93 |
soft_weights = torch.from_numpy(self.soft_weights[idx].astype(np.float32)).to(device)
|
| 94 |
-
return
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
|
| 96 |
|
| 97 |
def concat_batches(batch_a, batch_b):
|
| 98 |
-
|
| 99 |
-
same fixed_len, so every tensor pair already shares its non-batch dimensions."""
|
| 100 |
-
return tuple(torch.cat([a, b], dim=0) for a, b in zip(batch_a, batch_b))
|
| 101 |
|
| 102 |
|
| 103 |
def compute_loss(model: ChessMamba, batch, device: torch.device, value_loss_weight: float = 0.5):
|
| 104 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
policy_logits, promo_logits, value = model(from_ids, to_ids, promo_ids, lengths=lengths)
|
| 106 |
-
|
| 107 |
-
# soft cross-entropy over up to SOFT_K candidate moves per example, weighted by
|
| 108 |
-
# soft_weights -- mathematically identical to F.cross_entropy(policy_logits,
|
| 109 |
-
# target_from_to) when the target is a degenerate one-hot (ordinary hard-label
|
| 110 |
-
# data), and a genuine distributional loss when it isn't (self-play fine-tuning
|
| 111 |
-
# targets, e.g. softmax-with-temperature over a search's root candidate scores).
|
| 112 |
-
# Padding slots (id -1) are clamped to a valid index for the gather but contribute
|
| 113 |
-
# nothing since their weight is 0.
|
| 114 |
log_probs = F.log_softmax(policy_logits, dim=-1)
|
| 115 |
-
gathered = log_probs.gather(1, soft_move_ids.clamp(min=0))
|
| 116 |
policy_loss = -(soft_weights * gathered).sum(-1).mean()
|
| 117 |
-
|
| 118 |
promo_mask = target_promo > 0
|
| 119 |
if promo_mask.any():
|
| 120 |
promo_loss = F.cross_entropy(promo_logits[promo_mask], target_promo[promo_mask])
|
| 121 |
else:
|
| 122 |
promo_loss = torch.tensor(0.0, device=device)
|
| 123 |
-
|
| 124 |
value_loss = F.mse_loss(value.squeeze(-1), values)
|
| 125 |
-
|
| 126 |
loss = policy_loss + 0.25 * promo_loss + value_loss_weight * value_loss
|
| 127 |
with torch.no_grad():
|
| 128 |
acc = (policy_logits.argmax(-1) == target_from_to).float().mean()
|
| 129 |
-
return
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
|
|
|
|
|
|
|
|
|
| 135 |
|
| 136 |
|
| 137 |
def main():
|
|
@@ -142,56 +131,75 @@ def main():
|
|
| 142 |
ap.add_argument("--depth", type=int, default=10)
|
| 143 |
ap.add_argument("--state-dim", type=int, default=16)
|
| 144 |
ap.add_argument("--batch-size", type=int, default=1024)
|
| 145 |
-
ap.add_argument("--lr", type=float, default=
|
| 146 |
ap.add_argument("--weight-decay", type=float, default=0.01)
|
| 147 |
-
ap.add_argument(
|
| 148 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 149 |
ap.add_argument("--warmup-steps", type=int, default=500)
|
| 150 |
ap.add_argument("--log-every", type=int, default=50)
|
| 151 |
ap.add_argument("--eval-every", type=int, default=500)
|
| 152 |
ap.add_argument("--ckpt-every", type=int, default=1000)
|
| 153 |
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
| 154 |
-
ap.add_argument(
|
| 155 |
-
|
| 156 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 157 |
ap.add_argument("--no-compile", action="store_true", help="disable torch.compile")
|
| 158 |
-
ap.add_argument(
|
| 159 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 160 |
args = ap.parse_args()
|
| 161 |
-
|
| 162 |
os.makedirs(args.ckpt_dir, exist_ok=True)
|
| 163 |
device = torch.device(args.device)
|
| 164 |
if device.type == "cuda":
|
| 165 |
torch.set_float32_matmul_precision("high")
|
| 166 |
-
|
| 167 |
ds = ShardDataset(args.data_dir)
|
| 168 |
replay_ds = ShardDataset(args.replay_data_dir) if args.replay_data_dir else None
|
| 169 |
-
|
| 170 |
model_cfg = dict(dim=args.dim, depth=args.depth, state_dim=args.state_dim)
|
| 171 |
model = ChessMamba(**model_cfg, use_checkpoint=not args.no_checkpoint).to(device)
|
| 172 |
-
# the parallel scan is many small elementwise/cat ops -- torch.compile fuses them into far
|
| 173 |
-
# fewer kernel launches. Measured impact: 7.8s/step -> 2.1s/step at dim=384/depth=10/batch=512.
|
| 174 |
train_model = model if args.no_compile else torch.compile(model)
|
| 175 |
-
n_params = sum(p.numel() for p in model.parameters())
|
| 176 |
print(f"[model] {n_params:,} params, cfg={model_cfg}")
|
| 177 |
-
|
| 178 |
-
|
| 179 |
-
|
| 180 |
start_step = 0
|
| 181 |
if args.resume and os.path.exists(args.resume):
|
| 182 |
-
# true resume: same phase, same step budget/LR schedule, picking up an interrupted run
|
| 183 |
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
|
| 184 |
model.load_state_dict(ckpt["model"])
|
| 185 |
opt.load_state_dict(ckpt["optimizer"])
|
| 186 |
start_step = ckpt.get("step", 0)
|
| 187 |
print(f"[resume] loaded {args.resume} at step {start_step}")
|
| 188 |
elif args.init_from and os.path.exists(args.init_from):
|
| 189 |
-
# warm-start: new phase (new dataset/step budget/LR schedule), only the weights carry
|
| 190 |
-
# over -- NOT optimizer state or step count, which would otherwise make `range(start_step,
|
| 191 |
-
# args.steps)` empty if this phase's step budget is smaller than the previous phase's.
|
| 192 |
ckpt = torch.load(args.init_from, map_location=device, weights_only=False)
|
| 193 |
model.load_state_dict(ckpt["model"])
|
| 194 |
-
print(
|
|
|
|
|
|
|
| 195 |
|
| 196 |
def lr_at(step):
|
| 197 |
if step < args.warmup_steps:
|
|
@@ -201,62 +209,67 @@ def main():
|
|
| 201 |
|
| 202 |
rng = np.random.default_rng(1234)
|
| 203 |
t0 = time.monotonic()
|
| 204 |
-
|
| 205 |
replay_size = int(args.batch_size * args.replay_fraction) if replay_ds is not None else 0
|
| 206 |
primary_size = args.batch_size - replay_size
|
| 207 |
if replay_ds is not None:
|
| 208 |
-
print(
|
| 209 |
-
|
|
|
|
| 210 |
for step in range(start_step, args.steps):
|
| 211 |
for g in opt.param_groups:
|
| 212 |
g["lr"] = lr_at(step)
|
| 213 |
-
|
| 214 |
idx = rng.choice(ds.train_idx, size=primary_size, replace=False)
|
| 215 |
batch = ds.get_batch(idx, device, fixed_len=MAX_PLIES)
|
| 216 |
if replay_ds is not None:
|
| 217 |
replay_idx = rng.choice(replay_ds.train_idx, size=replay_size, replace=False)
|
| 218 |
replay_batch = replay_ds.get_batch(replay_idx, device, fixed_len=MAX_PLIES)
|
| 219 |
batch = concat_batches(batch, replay_batch)
|
| 220 |
-
|
| 221 |
model.train()
|
| 222 |
-
loss, stats = compute_loss(
|
|
|
|
|
|
|
| 223 |
opt.zero_grad(set_to_none=True)
|
| 224 |
loss.backward()
|
| 225 |
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 226 |
opt.step()
|
| 227 |
-
|
| 228 |
if step % args.log_every == 0:
|
| 229 |
dt = time.monotonic() - t0
|
| 230 |
print(
|
| 231 |
-
f"step {step} loss {loss.item():.4f} policy {stats['policy_loss']:.4f} "
|
| 232 |
-
f"promo {stats['promo_loss']:.4f} value {stats['value_loss']:.4f} "
|
| 233 |
-
f"acc {stats['move_acc']:.3f} lr {lr_at(step):.2e} elapsed {dt:.0f}s"
|
| 234 |
)
|
| 235 |
-
|
| 236 |
if step % args.eval_every == 0 and step > 0:
|
| 237 |
model.eval()
|
| 238 |
with torch.no_grad():
|
| 239 |
-
val_idx = rng.choice(
|
|
|
|
|
|
|
| 240 |
val_batch = ds.get_batch(val_idx, device, fixed_len=MAX_PLIES)
|
| 241 |
-
_, val_stats = compute_loss(
|
| 242 |
-
|
| 243 |
-
|
|
|
|
|
|
|
|
|
|
| 244 |
if step % args.ckpt_every == 0 and step > 0:
|
| 245 |
save_checkpoint(model, opt, step, model_cfg, args.ckpt_dir, "model.pt")
|
| 246 |
print(f" [ckpt] saved step {step}")
|
| 247 |
-
|
| 248 |
save_checkpoint(model, opt, args.steps, model_cfg, args.ckpt_dir, "model.pt")
|
| 249 |
print(f"[done] saved final checkpoint at step {args.steps}")
|
| 250 |
|
| 251 |
|
| 252 |
def save_checkpoint(model, opt, step, model_cfg, ckpt_dir, name):
|
| 253 |
-
"""Write-then-rename so a concurrent reader (e.g. testing the engine mid-training)
|
| 254 |
-
never sees a partially-written file -- os.replace is atomic on the same filesystem."""
|
| 255 |
path = os.path.join(ckpt_dir, name)
|
| 256 |
tmp_path = path + ".tmp"
|
| 257 |
-
torch.save(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 258 |
os.replace(tmp_path, path)
|
| 259 |
-
# keep a numbered snapshot too, so earlier checkpoints stay testable/comparable
|
| 260 |
step_path = os.path.join(ckpt_dir, f"model_step{step}.pt")
|
| 261 |
if not os.path.exists(step_path):
|
| 262 |
import shutil
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import argparse
|
| 3 |
import glob
|
| 4 |
import os
|
| 5 |
import time
|
|
|
|
| 6 |
import numpy as np
|
| 7 |
import torch
|
| 8 |
import torch.nn.functional as F
|
|
|
|
| 9 |
from data_pipeline import SOFT_K
|
| 10 |
from model import ChessMamba, MAX_PLIES
|
| 11 |
|
| 12 |
|
| 13 |
class ShardDataset:
|
|
|
|
|
|
|
|
|
|
|
|
|
| 14 |
|
| 15 |
def __init__(self, data_dir: str, val_fraction: float = 0.01):
|
|
|
|
|
|
|
| 16 |
paths = sorted(glob.glob(os.path.join(data_dir, "**", "shard_*.npz"), recursive=True))
|
| 17 |
if not paths:
|
| 18 |
raise RuntimeError(f"no shards found in {data_dir}")
|
| 19 |
+
from_ids, to_ids, promo_ids, lengths, target_from_to, target_promo, values = (
|
| 20 |
+
[] for _ in range(7)
|
| 21 |
+
)
|
| 22 |
+
soft_move_ids, soft_weights = ([], [])
|
| 23 |
for p in paths:
|
| 24 |
d = np.load(p)
|
| 25 |
from_ids.append(d["from_ids"])
|
|
|
|
| 33 |
soft_move_ids.append(d["soft_move_ids"])
|
| 34 |
soft_weights.append(d["soft_weights"])
|
| 35 |
else:
|
|
|
|
|
|
|
|
|
|
| 36 |
n_ex = len(d["target_from_to"])
|
| 37 |
one_hot_ids = np.full((n_ex, SOFT_K), -1, dtype=np.int16)
|
| 38 |
one_hot_ids[:, 0] = d["target_from_to"]
|
|
|
|
| 40 |
one_hot_w[:, 0] = 1.0
|
| 41 |
soft_move_ids.append(one_hot_ids)
|
| 42 |
soft_weights.append(one_hot_w)
|
|
|
|
| 43 |
self.from_ids = np.concatenate(from_ids)
|
| 44 |
self.to_ids = np.concatenate(to_ids)
|
| 45 |
self.promo_ids = np.concatenate(promo_ids)
|
|
|
|
| 49 |
self.values = np.concatenate(values)
|
| 50 |
self.soft_move_ids = np.concatenate(soft_move_ids)
|
| 51 |
self.soft_weights = np.concatenate(soft_weights)
|
|
|
|
| 52 |
n = len(self.lengths)
|
| 53 |
rng = np.random.default_rng(0)
|
| 54 |
perm = rng.permutation(n)
|
| 55 |
n_val = max(1, int(n * val_fraction))
|
| 56 |
self.val_idx = perm[:n_val]
|
| 57 |
self.train_idx = perm[n_val:]
|
| 58 |
+
print(
|
| 59 |
+
f"[data] loaded {n:,} examples from {len(paths)} shards ({n - n_val:,} train / {n_val:,} val)"
|
| 60 |
+
)
|
| 61 |
|
| 62 |
def get_batch(self, idx: np.ndarray, device: torch.device, fixed_len: int | None = None):
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
max_len = fixed_len if fixed_len is not None else int(self.lengths[idx].max())
|
| 64 |
from_ids = torch.from_numpy(self.from_ids[idx][:, :max_len].astype(np.int64)).to(device)
|
| 65 |
to_ids = torch.from_numpy(self.to_ids[idx][:, :max_len].astype(np.int64)).to(device)
|
|
|
|
| 70 |
values = torch.from_numpy(self.values[idx].astype(np.float32)).to(device)
|
| 71 |
soft_move_ids = torch.from_numpy(self.soft_move_ids[idx].astype(np.int64)).to(device)
|
| 72 |
soft_weights = torch.from_numpy(self.soft_weights[idx].astype(np.float32)).to(device)
|
| 73 |
+
return (
|
| 74 |
+
from_ids,
|
| 75 |
+
to_ids,
|
| 76 |
+
promo_ids,
|
| 77 |
+
lengths,
|
| 78 |
+
target_from_to,
|
| 79 |
+
target_promo,
|
| 80 |
+
values,
|
| 81 |
+
soft_move_ids,
|
| 82 |
+
soft_weights,
|
| 83 |
+
)
|
| 84 |
|
| 85 |
|
| 86 |
def concat_batches(batch_a, batch_b):
|
| 87 |
+
return tuple((torch.cat([a, b], dim=0) for a, b in zip(batch_a, batch_b)))
|
|
|
|
|
|
|
| 88 |
|
| 89 |
|
| 90 |
def compute_loss(model: ChessMamba, batch, device: torch.device, value_loss_weight: float = 0.5):
|
| 91 |
+
(
|
| 92 |
+
from_ids,
|
| 93 |
+
to_ids,
|
| 94 |
+
promo_ids,
|
| 95 |
+
lengths,
|
| 96 |
+
target_from_to,
|
| 97 |
+
target_promo,
|
| 98 |
+
values,
|
| 99 |
+
soft_move_ids,
|
| 100 |
+
soft_weights,
|
| 101 |
+
) = batch
|
| 102 |
policy_logits, promo_logits, value = model(from_ids, to_ids, promo_ids, lengths=lengths)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 103 |
log_probs = F.log_softmax(policy_logits, dim=-1)
|
| 104 |
+
gathered = log_probs.gather(1, soft_move_ids.clamp(min=0))
|
| 105 |
policy_loss = -(soft_weights * gathered).sum(-1).mean()
|
|
|
|
| 106 |
promo_mask = target_promo > 0
|
| 107 |
if promo_mask.any():
|
| 108 |
promo_loss = F.cross_entropy(promo_logits[promo_mask], target_promo[promo_mask])
|
| 109 |
else:
|
| 110 |
promo_loss = torch.tensor(0.0, device=device)
|
|
|
|
| 111 |
value_loss = F.mse_loss(value.squeeze(-1), values)
|
|
|
|
| 112 |
loss = policy_loss + 0.25 * promo_loss + value_loss_weight * value_loss
|
| 113 |
with torch.no_grad():
|
| 114 |
acc = (policy_logits.argmax(-1) == target_from_to).float().mean()
|
| 115 |
+
return (
|
| 116 |
+
loss,
|
| 117 |
+
{
|
| 118 |
+
"policy_loss": policy_loss.item(),
|
| 119 |
+
"promo_loss": promo_loss.item() if promo_loss.requires_grad else float(promo_loss),
|
| 120 |
+
"value_loss": value_loss.item(),
|
| 121 |
+
"move_acc": acc.item(),
|
| 122 |
+
},
|
| 123 |
+
)
|
| 124 |
|
| 125 |
|
| 126 |
def main():
|
|
|
|
| 131 |
ap.add_argument("--depth", type=int, default=10)
|
| 132 |
ap.add_argument("--state-dim", type=int, default=16)
|
| 133 |
ap.add_argument("--batch-size", type=int, default=1024)
|
| 134 |
+
ap.add_argument("--lr", type=float, default=0.0003)
|
| 135 |
ap.add_argument("--weight-decay", type=float, default=0.01)
|
| 136 |
+
ap.add_argument(
|
| 137 |
+
"--value-loss-weight",
|
| 138 |
+
type=float,
|
| 139 |
+
default=0.5,
|
| 140 |
+
help="raised for phase 4: diagnosed bottleneck is the value head's material judgment specifically",
|
| 141 |
+
)
|
| 142 |
+
ap.add_argument("--steps", type=int, default=20000)
|
| 143 |
ap.add_argument("--warmup-steps", type=int, default=500)
|
| 144 |
ap.add_argument("--log-every", type=int, default=50)
|
| 145 |
ap.add_argument("--eval-every", type=int, default=500)
|
| 146 |
ap.add_argument("--ckpt-every", type=int, default=1000)
|
| 147 |
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
| 148 |
+
ap.add_argument(
|
| 149 |
+
"--resume",
|
| 150 |
+
default="",
|
| 151 |
+
help="resume an interrupted run of THIS SAME phase (loads optimizer + step count too)",
|
| 152 |
+
)
|
| 153 |
+
ap.add_argument(
|
| 154 |
+
"--init-from",
|
| 155 |
+
default="",
|
| 156 |
+
help="warm-start a NEW phase from another checkpoint's weights only (fresh optimizer/step/LR schedule)",
|
| 157 |
+
)
|
| 158 |
+
ap.add_argument(
|
| 159 |
+
"--no-checkpoint",
|
| 160 |
+
action="store_true",
|
| 161 |
+
help="disable gradient checkpointing (faster per-step if it fits in memory)",
|
| 162 |
+
)
|
| 163 |
ap.add_argument("--no-compile", action="store_true", help="disable torch.compile")
|
| 164 |
+
ap.add_argument(
|
| 165 |
+
"--replay-data-dir",
|
| 166 |
+
default="",
|
| 167 |
+
help="mix in examples from this dir (e.g. the original supervised data) during fine-tuning, to avoid catastrophic forgetting of the base distribution",
|
| 168 |
+
)
|
| 169 |
+
ap.add_argument(
|
| 170 |
+
"--replay-fraction",
|
| 171 |
+
type=float,
|
| 172 |
+
default=0.4,
|
| 173 |
+
help="fraction of each batch drawn from --replay-data-dir when set",
|
| 174 |
+
)
|
| 175 |
args = ap.parse_args()
|
|
|
|
| 176 |
os.makedirs(args.ckpt_dir, exist_ok=True)
|
| 177 |
device = torch.device(args.device)
|
| 178 |
if device.type == "cuda":
|
| 179 |
torch.set_float32_matmul_precision("high")
|
|
|
|
| 180 |
ds = ShardDataset(args.data_dir)
|
| 181 |
replay_ds = ShardDataset(args.replay_data_dir) if args.replay_data_dir else None
|
|
|
|
| 182 |
model_cfg = dict(dim=args.dim, depth=args.depth, state_dim=args.state_dim)
|
| 183 |
model = ChessMamba(**model_cfg, use_checkpoint=not args.no_checkpoint).to(device)
|
|
|
|
|
|
|
| 184 |
train_model = model if args.no_compile else torch.compile(model)
|
| 185 |
+
n_params = sum((p.numel() for p in model.parameters()))
|
| 186 |
print(f"[model] {n_params:,} params, cfg={model_cfg}")
|
| 187 |
+
opt = torch.optim.AdamW(
|
| 188 |
+
model.parameters(), lr=args.lr, weight_decay=args.weight_decay, betas=(0.9, 0.95)
|
| 189 |
+
)
|
| 190 |
start_step = 0
|
| 191 |
if args.resume and os.path.exists(args.resume):
|
|
|
|
| 192 |
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
|
| 193 |
model.load_state_dict(ckpt["model"])
|
| 194 |
opt.load_state_dict(ckpt["optimizer"])
|
| 195 |
start_step = ckpt.get("step", 0)
|
| 196 |
print(f"[resume] loaded {args.resume} at step {start_step}")
|
| 197 |
elif args.init_from and os.path.exists(args.init_from):
|
|
|
|
|
|
|
|
|
|
| 198 |
ckpt = torch.load(args.init_from, map_location=device, weights_only=False)
|
| 199 |
model.load_state_dict(ckpt["model"])
|
| 200 |
+
print(
|
| 201 |
+
f"[init-from] loaded weights from {args.init_from} (step {ckpt.get('step', '?')} of its own run), starting fresh optimizer/schedule at step 0"
|
| 202 |
+
)
|
| 203 |
|
| 204 |
def lr_at(step):
|
| 205 |
if step < args.warmup_steps:
|
|
|
|
| 209 |
|
| 210 |
rng = np.random.default_rng(1234)
|
| 211 |
t0 = time.monotonic()
|
|
|
|
| 212 |
replay_size = int(args.batch_size * args.replay_fraction) if replay_ds is not None else 0
|
| 213 |
primary_size = args.batch_size - replay_size
|
| 214 |
if replay_ds is not None:
|
| 215 |
+
print(
|
| 216 |
+
f"[replay] mixing {replay_size}/{args.batch_size} examples per batch from {args.replay_data_dir}"
|
| 217 |
+
)
|
| 218 |
for step in range(start_step, args.steps):
|
| 219 |
for g in opt.param_groups:
|
| 220 |
g["lr"] = lr_at(step)
|
|
|
|
| 221 |
idx = rng.choice(ds.train_idx, size=primary_size, replace=False)
|
| 222 |
batch = ds.get_batch(idx, device, fixed_len=MAX_PLIES)
|
| 223 |
if replay_ds is not None:
|
| 224 |
replay_idx = rng.choice(replay_ds.train_idx, size=replay_size, replace=False)
|
| 225 |
replay_batch = replay_ds.get_batch(replay_idx, device, fixed_len=MAX_PLIES)
|
| 226 |
batch = concat_batches(batch, replay_batch)
|
|
|
|
| 227 |
model.train()
|
| 228 |
+
loss, stats = compute_loss(
|
| 229 |
+
train_model, batch, device, value_loss_weight=args.value_loss_weight
|
| 230 |
+
)
|
| 231 |
opt.zero_grad(set_to_none=True)
|
| 232 |
loss.backward()
|
| 233 |
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 234 |
opt.step()
|
|
|
|
| 235 |
if step % args.log_every == 0:
|
| 236 |
dt = time.monotonic() - t0
|
| 237 |
print(
|
| 238 |
+
f"step {step} loss {loss.item():.4f} policy {stats['policy_loss']:.4f} promo {stats['promo_loss']:.4f} value {stats['value_loss']:.4f} acc {stats['move_acc']:.3f} lr {lr_at(step):.2e} elapsed {dt:.0f}s"
|
|
|
|
|
|
|
| 239 |
)
|
|
|
|
| 240 |
if step % args.eval_every == 0 and step > 0:
|
| 241 |
model.eval()
|
| 242 |
with torch.no_grad():
|
| 243 |
+
val_idx = rng.choice(
|
| 244 |
+
ds.val_idx, size=min(args.batch_size, len(ds.val_idx)), replace=False
|
| 245 |
+
)
|
| 246 |
val_batch = ds.get_batch(val_idx, device, fixed_len=MAX_PLIES)
|
| 247 |
+
_, val_stats = compute_loss(
|
| 248 |
+
train_model, val_batch, device, value_loss_weight=args.value_loss_weight
|
| 249 |
+
)
|
| 250 |
+
print(
|
| 251 |
+
f" [val] step {step} policy {val_stats['policy_loss']:.4f} acc {val_stats['move_acc']:.3f}"
|
| 252 |
+
)
|
| 253 |
if step % args.ckpt_every == 0 and step > 0:
|
| 254 |
save_checkpoint(model, opt, step, model_cfg, args.ckpt_dir, "model.pt")
|
| 255 |
print(f" [ckpt] saved step {step}")
|
|
|
|
| 256 |
save_checkpoint(model, opt, args.steps, model_cfg, args.ckpt_dir, "model.pt")
|
| 257 |
print(f"[done] saved final checkpoint at step {args.steps}")
|
| 258 |
|
| 259 |
|
| 260 |
def save_checkpoint(model, opt, step, model_cfg, ckpt_dir, name):
|
|
|
|
|
|
|
| 261 |
path = os.path.join(ckpt_dir, name)
|
| 262 |
tmp_path = path + ".tmp"
|
| 263 |
+
torch.save(
|
| 264 |
+
{
|
| 265 |
+
"model": model.state_dict(),
|
| 266 |
+
"optimizer": opt.state_dict(),
|
| 267 |
+
"step": step,
|
| 268 |
+
"config": model_cfg,
|
| 269 |
+
},
|
| 270 |
+
tmp_path,
|
| 271 |
+
)
|
| 272 |
os.replace(tmp_path, path)
|
|
|
|
| 273 |
step_path = os.path.join(ckpt_dir, f"model_step{step}.pt")
|
| 274 |
if not os.path.exists(step_path):
|
| 275 |
import shutil
|
training/run_selfplay_parallel.py
CHANGED
|
@@ -1,15 +1,4 @@
|
|
| 1 |
-
"""Fan `selfplay_finetune.py` out across multiple CPU processes -- self-play games are
|
| 2 |
-
independent, and single-threaded generation (~1 game at a time) doesn't come close to
|
| 3 |
-
using a 32-core box. Each worker writes to its own subdirectory (pretrain.py's
|
| 4 |
-
ShardDataset globs recursively, so no merging step is needed afterward).
|
| 5 |
-
|
| 6 |
-
Usage:
|
| 7 |
-
python3 run_selfplay_parallel.py --checkpoint /root/chess/ckpt/model.pt \
|
| 8 |
-
--out-dir /root/chess/selfplay_data --games 4000 --workers 28
|
| 9 |
-
"""
|
| 10 |
-
|
| 11 |
from __future__ import annotations
|
| 12 |
-
|
| 13 |
import argparse
|
| 14 |
import os
|
| 15 |
import subprocess
|
|
@@ -25,7 +14,6 @@ def main():
|
|
| 25 |
ap.add_argument("--movetime", type=int, default=200)
|
| 26 |
ap.add_argument("--temperature-plies", type=int, default=10)
|
| 27 |
args = ap.parse_args()
|
| 28 |
-
|
| 29 |
per_worker = max(1, args.games // args.workers)
|
| 30 |
procs = []
|
| 31 |
logs = []
|
|
@@ -35,18 +23,28 @@ def main():
|
|
| 35 |
os.makedirs(args.out_dir, exist_ok=True)
|
| 36 |
log_f = open(log_path, "w")
|
| 37 |
cmd = [
|
| 38 |
-
sys.executable,
|
| 39 |
-
"
|
| 40 |
-
"--
|
| 41 |
-
|
| 42 |
-
"--
|
| 43 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 44 |
]
|
| 45 |
-
p = subprocess.Popen(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
procs.append(p)
|
| 47 |
logs.append(log_f)
|
| 48 |
print(f"[launcher] started worker {i} (pid {p.pid}), {per_worker} games -> {worker_out}")
|
| 49 |
-
|
| 50 |
exit_codes = []
|
| 51 |
for i, p in enumerate(procs):
|
| 52 |
code = p.wait()
|
|
@@ -54,11 +52,14 @@ def main():
|
|
| 54 |
print(f"[launcher] worker {i} exited with code {code}")
|
| 55 |
for f in logs:
|
| 56 |
f.close()
|
| 57 |
-
|
| 58 |
failed = [i for i, c in enumerate(exit_codes) if c != 0]
|
| 59 |
if failed:
|
| 60 |
-
print(
|
| 61 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 62 |
|
| 63 |
|
| 64 |
if __name__ == "__main__":
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import argparse
|
| 3 |
import os
|
| 4 |
import subprocess
|
|
|
|
| 14 |
ap.add_argument("--movetime", type=int, default=200)
|
| 15 |
ap.add_argument("--temperature-plies", type=int, default=10)
|
| 16 |
args = ap.parse_args()
|
|
|
|
| 17 |
per_worker = max(1, args.games // args.workers)
|
| 18 |
procs = []
|
| 19 |
logs = []
|
|
|
|
| 23 |
os.makedirs(args.out_dir, exist_ok=True)
|
| 24 |
log_f = open(log_path, "w")
|
| 25 |
cmd = [
|
| 26 |
+
sys.executable,
|
| 27 |
+
"selfplay_finetune.py",
|
| 28 |
+
"--checkpoint",
|
| 29 |
+
args.checkpoint,
|
| 30 |
+
"--out-dir",
|
| 31 |
+
worker_out,
|
| 32 |
+
"--games",
|
| 33 |
+
str(per_worker),
|
| 34 |
+
"--movetime",
|
| 35 |
+
str(args.movetime),
|
| 36 |
+
"--temperature-plies",
|
| 37 |
+
str(args.temperature_plies),
|
| 38 |
]
|
| 39 |
+
p = subprocess.Popen(
|
| 40 |
+
cmd,
|
| 41 |
+
stdout=log_f,
|
| 42 |
+
stderr=subprocess.STDOUT,
|
| 43 |
+
cwd=os.path.dirname(os.path.abspath(__file__)) or ".",
|
| 44 |
+
)
|
| 45 |
procs.append(p)
|
| 46 |
logs.append(log_f)
|
| 47 |
print(f"[launcher] started worker {i} (pid {p.pid}), {per_worker} games -> {worker_out}")
|
|
|
|
| 48 |
exit_codes = []
|
| 49 |
for i, p in enumerate(procs):
|
| 50 |
code = p.wait()
|
|
|
|
| 52 |
print(f"[launcher] worker {i} exited with code {code}")
|
| 53 |
for f in logs:
|
| 54 |
f.close()
|
|
|
|
| 55 |
failed = [i for i, c in enumerate(exit_codes) if c != 0]
|
| 56 |
if failed:
|
| 57 |
+
print(
|
| 58 |
+
f"[launcher] WARNING: workers failed: {failed} -- check worker*.log in {args.out_dir}"
|
| 59 |
+
)
|
| 60 |
+
print(
|
| 61 |
+
f"[launcher] DONE: {args.workers} workers, {per_worker * args.workers} games total requested"
|
| 62 |
+
)
|
| 63 |
|
| 64 |
|
| 65 |
if __name__ == "__main__":
|
training/selfplay_finetune.py
CHANGED
|
@@ -1,116 +1,95 @@
|
|
| 1 |
-
"""Generates self-play games with the current checkpoint's shallow search and packs
|
| 2 |
-
them into shards in the exact same format data_pipeline.py produces -- so fine-tuning
|
| 3 |
-
on them is just another pretrain.py run with --data-dir pointed here and --resume
|
| 4 |
-
pointed at the pretrained checkpoint. No separate training loop needed.
|
| 5 |
-
|
| 6 |
-
Policy target is a softmax-with-temperature distribution over the search's scored
|
| 7 |
-
root candidates (not just the single chosen move) -- the AlphaZero-style motivation:
|
| 8 |
-
a distributional target is a richer training signal than a hard label, which matters
|
| 9 |
-
more here than it would with millions of games, since the self-play corpus is modest.
|
| 10 |
-
Real AlphaZero uses MCTS visit counts specifically; ours uses negamax score-softmax
|
| 11 |
-
over the same top-K candidates the rest of the search already ranks, which is the
|
| 12 |
-
natural equivalent for a negamax (not MCTS) search -- not visit counts, but still a
|
| 13 |
-
genuine distribution over multiple candidates instead of a single argmax label.
|
| 14 |
-
"""
|
| 15 |
-
|
| 16 |
from __future__ import annotations
|
| 17 |
-
|
| 18 |
import argparse
|
| 19 |
import math
|
| 20 |
import random
|
| 21 |
import time
|
| 22 |
-
|
| 23 |
import torch
|
| 24 |
|
| 25 |
-
# self-play runs many of these processes in parallel across cores (run_selfplay_parallel.py);
|
| 26 |
-
# without this, each process's PyTorch defaults to using every core for its own tiny CPU
|
| 27 |
-
# matmuls, so N workers oversubscribe to N*cores threads instead of using one core each.
|
| 28 |
torch.set_num_threads(1)
|
| 29 |
-
|
| 30 |
import chess
|
| 31 |
-
|
| 32 |
from chess_io import move_to_ids
|
| 33 |
from data_pipeline import ShardWriter
|
| 34 |
from model import MAX_PLIES
|
| 35 |
from search import Searcher
|
| 36 |
|
| 37 |
-
|
| 38 |
-
SOFT_TARGET_TEMPERATURE = 0.5 # lower = closer to a hard argmax label, higher = flatter/softer
|
| 39 |
|
| 40 |
|
| 41 |
def scores_to_distribution(scored: list[tuple]) -> tuple[list[int], list[float]]:
|
| 42 |
-
"""scored: [(move, negamax_score), ...] -> (from_to_ids, softmax weights)."""
|
| 43 |
values = [s for _, s in scored]
|
| 44 |
max_v = max(values)
|
| 45 |
exp_vals = [math.exp((v - max_v) / SOFT_TARGET_TEMPERATURE) for v in values]
|
| 46 |
total = sum(exp_vals) or 1.0
|
| 47 |
ids = [mv.from_square * 64 + mv.to_square for mv, _ in scored]
|
| 48 |
weights = [e / total for e in exp_vals]
|
| 49 |
-
return ids, weights
|
| 50 |
|
| 51 |
|
| 52 |
-
def play_selfplay_game(
|
|
|
|
|
|
|
| 53 |
board = chess.Board()
|
| 54 |
searcher.reset()
|
| 55 |
-
searcher.sync(board)
|
| 56 |
node = searcher.root_node
|
| 57 |
hist_from: list[int] = []
|
| 58 |
hist_to: list[int] = []
|
| 59 |
hist_promo: list[int] = []
|
| 60 |
-
ply_records = []
|
| 61 |
-
|
| 62 |
for ply in range(max_plies):
|
| 63 |
if board.is_game_over():
|
| 64 |
break
|
| 65 |
-
|
| 66 |
priors, _ = searcher._policy_value(board, node)
|
| 67 |
if not priors:
|
| 68 |
break
|
| 69 |
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: searcher.top_k]
|
| 70 |
-
|
| 71 |
-
# score each candidate at the root via one ply of negamax, same mechanism choose_move uses.
|
| 72 |
-
# One shared deadline for the whole ply, not one per candidate -- movetime_ms is meant to
|
| 73 |
-
# bound the total time spent choosing this move, same contract as choose_move's `go movetime`.
|
| 74 |
deadline = time.monotonic() + movetime_ms / 1000.0
|
| 75 |
scored = []
|
| 76 |
for mv, _ in ranked:
|
| 77 |
child = searcher._expand(board, node, mv)
|
| 78 |
-
score = -searcher._negamax(
|
|
|
|
|
|
|
| 79 |
board.pop()
|
| 80 |
scored.append((mv, score))
|
| 81 |
scored.sort(key=lambda x: -x[1])
|
| 82 |
-
|
| 83 |
if ply < temperature_plies:
|
| 84 |
move = random.choice(scored[: min(3, len(scored))])[0]
|
| 85 |
else:
|
| 86 |
move = scored[0][0]
|
| 87 |
-
|
| 88 |
-
side_to_move = board.turn # capture before _expand pushes the move
|
| 89 |
from_to_id, promo_id = move_to_ids(move)
|
| 90 |
soft_ids, soft_weights = scores_to_distribution(scored)
|
| 91 |
window_from = hist_from[-MAX_PLIES:]
|
| 92 |
window_to = hist_to[-MAX_PLIES:]
|
| 93 |
window_promo = hist_promo[-MAX_PLIES:]
|
| 94 |
-
ply_records.append(
|
| 95 |
-
|
| 96 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
hist_from.append(move.from_square)
|
| 98 |
hist_to.append(move.to_square)
|
| 99 |
hist_promo.append(promo_id)
|
| 100 |
-
|
| 101 |
result = board.result() if board.is_game_over() else "1/2-1/2"
|
| 102 |
from data_pipeline import result_value
|
| 103 |
|
| 104 |
examples = [
|
| 105 |
(wf, wt, wp, tft, tp, result_value(result, side), sids, sw)
|
| 106 |
-
for
|
| 107 |
]
|
| 108 |
-
return examples, result
|
| 109 |
|
| 110 |
|
| 111 |
def mirror_square(sq: int) -> int:
|
| 112 |
-
"""Vertical rank-mirror (rank r -> 7-r, file unchanged) = XOR 56, since python-chess
|
| 113 |
-
squares are rank*8+file and XORing the rank's 3 bits with 0b111 gives 7-rank."""
|
| 114 |
return sq ^ 56
|
| 115 |
|
| 116 |
|
|
@@ -119,13 +98,6 @@ def mirror_from_to(from_to_id: int) -> int:
|
|
| 119 |
|
| 120 |
|
| 121 |
def mirror_example(ex):
|
| 122 |
-
"""Chess has a real symmetry: mirror the board vertically and it's an equally legal
|
| 123 |
-
position from the other side's perspective. Doubles self-play data for free -- no
|
| 124 |
-
extra search needed, since it's a deterministic transform of an already-played game.
|
| 125 |
-
Piece types (promo ids), the value target (already side-to-move-relative, not
|
| 126 |
-
color-absolute), and soft-target weights are all unaffected by the mirror; only
|
| 127 |
-
squares need transforming. Padding slots (-1) in soft_ids are left untouched -- XORing
|
| 128 |
-
a sentinel would produce a bogus square id instead of staying a no-op padding marker."""
|
| 129 |
wf, wt, wp, tft, tp, value, soft_ids, soft_w = ex
|
| 130 |
m_wf = [mirror_square(s) for s in wf]
|
| 131 |
m_wt = [mirror_square(s) for s in wt]
|
|
@@ -141,16 +113,15 @@ def main():
|
|
| 141 |
ap.add_argument("--games", type=int, default=2000)
|
| 142 |
ap.add_argument("--movetime", type=int, default=200)
|
| 143 |
ap.add_argument("--temperature-plies", type=int, default=10)
|
| 144 |
-
ap.add_argument(
|
|
|
|
|
|
|
| 145 |
args = ap.parse_args()
|
| 146 |
-
|
| 147 |
searcher = Searcher(checkpoint_path=args.checkpoint)
|
| 148 |
if searcher.model is None:
|
| 149 |
raise RuntimeError(f"no checkpoint found at {args.checkpoint} -- pretrain first")
|
| 150 |
-
|
| 151 |
writer = ShardWriter(args.out_dir)
|
| 152 |
results = {"1-0": 0, "0-1": 0, "1/2-1/2": 0}
|
| 153 |
-
|
| 154 |
for g in range(args.games):
|
| 155 |
examples, result = play_selfplay_game(searcher, args.movetime, args.temperature_plies)
|
| 156 |
for ex in examples:
|
|
@@ -160,7 +131,6 @@ def main():
|
|
| 160 |
results[result] = results.get(result, 0) + 1
|
| 161 |
if (g + 1) % 20 == 0:
|
| 162 |
print(f"[selfplay] {g + 1}/{args.games} games, results so far: {results}")
|
| 163 |
-
|
| 164 |
writer.flush()
|
| 165 |
print(f"[selfplay] DONE {args.games} games, results: {results}")
|
| 166 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
from __future__ import annotations
|
|
|
|
| 2 |
import argparse
|
| 3 |
import math
|
| 4 |
import random
|
| 5 |
import time
|
|
|
|
| 6 |
import torch
|
| 7 |
|
|
|
|
|
|
|
|
|
|
| 8 |
torch.set_num_threads(1)
|
|
|
|
| 9 |
import chess
|
|
|
|
| 10 |
from chess_io import move_to_ids
|
| 11 |
from data_pipeline import ShardWriter
|
| 12 |
from model import MAX_PLIES
|
| 13 |
from search import Searcher
|
| 14 |
|
| 15 |
+
SOFT_TARGET_TEMPERATURE = 0.5
|
|
|
|
| 16 |
|
| 17 |
|
| 18 |
def scores_to_distribution(scored: list[tuple]) -> tuple[list[int], list[float]]:
|
|
|
|
| 19 |
values = [s for _, s in scored]
|
| 20 |
max_v = max(values)
|
| 21 |
exp_vals = [math.exp((v - max_v) / SOFT_TARGET_TEMPERATURE) for v in values]
|
| 22 |
total = sum(exp_vals) or 1.0
|
| 23 |
ids = [mv.from_square * 64 + mv.to_square for mv, _ in scored]
|
| 24 |
weights = [e / total for e in exp_vals]
|
| 25 |
+
return (ids, weights)
|
| 26 |
|
| 27 |
|
| 28 |
+
def play_selfplay_game(
|
| 29 |
+
searcher: Searcher, movetime_ms: int, temperature_plies: int, max_plies: int = 200
|
| 30 |
+
):
|
| 31 |
board = chess.Board()
|
| 32 |
searcher.reset()
|
| 33 |
+
searcher.sync(board)
|
| 34 |
node = searcher.root_node
|
| 35 |
hist_from: list[int] = []
|
| 36 |
hist_to: list[int] = []
|
| 37 |
hist_promo: list[int] = []
|
| 38 |
+
ply_records = []
|
|
|
|
| 39 |
for ply in range(max_plies):
|
| 40 |
if board.is_game_over():
|
| 41 |
break
|
|
|
|
| 42 |
priors, _ = searcher._policy_value(board, node)
|
| 43 |
if not priors:
|
| 44 |
break
|
| 45 |
ranked = sorted(priors.items(), key=lambda kv: -kv[1])[: searcher.top_k]
|
|
|
|
|
|
|
|
|
|
|
|
|
| 46 |
deadline = time.monotonic() + movetime_ms / 1000.0
|
| 47 |
scored = []
|
| 48 |
for mv, _ in ranked:
|
| 49 |
child = searcher._expand(board, node, mv)
|
| 50 |
+
score = -searcher._negamax(
|
| 51 |
+
board, child, searcher.max_depth - 1, -float("inf"), float("inf"), deadline
|
| 52 |
+
)
|
| 53 |
board.pop()
|
| 54 |
scored.append((mv, score))
|
| 55 |
scored.sort(key=lambda x: -x[1])
|
|
|
|
| 56 |
if ply < temperature_plies:
|
| 57 |
move = random.choice(scored[: min(3, len(scored))])[0]
|
| 58 |
else:
|
| 59 |
move = scored[0][0]
|
| 60 |
+
side_to_move = board.turn
|
|
|
|
| 61 |
from_to_id, promo_id = move_to_ids(move)
|
| 62 |
soft_ids, soft_weights = scores_to_distribution(scored)
|
| 63 |
window_from = hist_from[-MAX_PLIES:]
|
| 64 |
window_to = hist_to[-MAX_PLIES:]
|
| 65 |
window_promo = hist_promo[-MAX_PLIES:]
|
| 66 |
+
ply_records.append(
|
| 67 |
+
(
|
| 68 |
+
window_from,
|
| 69 |
+
window_to,
|
| 70 |
+
window_promo,
|
| 71 |
+
from_to_id,
|
| 72 |
+
promo_id,
|
| 73 |
+
side_to_move,
|
| 74 |
+
soft_ids,
|
| 75 |
+
soft_weights,
|
| 76 |
+
)
|
| 77 |
+
)
|
| 78 |
+
node = searcher._expand(board, node, move)
|
| 79 |
hist_from.append(move.from_square)
|
| 80 |
hist_to.append(move.to_square)
|
| 81 |
hist_promo.append(promo_id)
|
|
|
|
| 82 |
result = board.result() if board.is_game_over() else "1/2-1/2"
|
| 83 |
from data_pipeline import result_value
|
| 84 |
|
| 85 |
examples = [
|
| 86 |
(wf, wt, wp, tft, tp, result_value(result, side), sids, sw)
|
| 87 |
+
for wf, wt, wp, tft, tp, side, sids, sw in ply_records
|
| 88 |
]
|
| 89 |
+
return (examples, result)
|
| 90 |
|
| 91 |
|
| 92 |
def mirror_square(sq: int) -> int:
|
|
|
|
|
|
|
| 93 |
return sq ^ 56
|
| 94 |
|
| 95 |
|
|
|
|
| 98 |
|
| 99 |
|
| 100 |
def mirror_example(ex):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 101 |
wf, wt, wp, tft, tp, value, soft_ids, soft_w = ex
|
| 102 |
m_wf = [mirror_square(s) for s in wf]
|
| 103 |
m_wt = [mirror_square(s) for s in wt]
|
|
|
|
| 113 |
ap.add_argument("--games", type=int, default=2000)
|
| 114 |
ap.add_argument("--movetime", type=int, default=200)
|
| 115 |
ap.add_argument("--temperature-plies", type=int, default=10)
|
| 116 |
+
ap.add_argument(
|
| 117 |
+
"--no-mirror-augment", action="store_true", help="disable color-flip data doubling"
|
| 118 |
+
)
|
| 119 |
args = ap.parse_args()
|
|
|
|
| 120 |
searcher = Searcher(checkpoint_path=args.checkpoint)
|
| 121 |
if searcher.model is None:
|
| 122 |
raise RuntimeError(f"no checkpoint found at {args.checkpoint} -- pretrain first")
|
|
|
|
| 123 |
writer = ShardWriter(args.out_dir)
|
| 124 |
results = {"1-0": 0, "0-1": 0, "1/2-1/2": 0}
|
|
|
|
| 125 |
for g in range(args.games):
|
| 126 |
examples, result = play_selfplay_game(searcher, args.movetime, args.temperature_plies)
|
| 127 |
for ex in examples:
|
|
|
|
| 131 |
results[result] = results.get(result, 0) + 1
|
| 132 |
if (g + 1) % 20 == 0:
|
| 133 |
print(f"[selfplay] {g + 1}/{args.games} games, results so far: {results}")
|
|
|
|
| 134 |
writer.flush()
|
| 135 |
print(f"[selfplay] DONE {args.games} games, results: {results}")
|
| 136 |
|