TobiasLogic commited on
Commit
380c43e
·
verified ·
1 Parent(s): 9f1fb01

Upload folder using huggingface_hub

Browse files
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("CHESSMAMBA_CKPT", os.path.join(os.path.dirname(__file__), "ckpt", "model.pt"))
23
- POLICY_ONLY = os.environ.get("CHESSMAMBA_POLICY_ONLY", "") == "1" # diagnostic: skip negamax, play raw top policy move
 
 
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 # window of recent half-moves fed to the model
23
 
24
 
25
  class RMSNorm(nn.Module):
26
- def __init__(self, dim: int, eps: float = 1e-5):
 
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.in_proj = nn.Linear(dim, inner_dim * 2, bias=False) # -> (x, gate)
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) # each (B, L, inner_dim)
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) # (B, L, inner_dim), > 0
105
-
106
- A = -torch.exp(self.A_log) # (inner_dim, state_dim), negative
107
-
108
- # discretize for every timestep at once: A_bar_t = exp(delta_t * A), Bx_t = delta_t * B_t * x_t
109
- A_bar = torch.exp(delta.unsqueeze(-1) * A.view(1, 1, self.inner_dim, self.state_dim)) # (B,L,inner_dim,state_dim)
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(self, x_t: torch.Tensor, h_prev: torch.Tensor | None) -> tuple[torch.Tensor, torch.Tensor]:
119
- """Single-timestep update for live/incremental play: given the token at the
120
- current step and the carried recurrent state, returns (output, new_state) --
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) # (B, inner_dim) each
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) # (B, inner_dim)
138
-
139
- A = -torch.exp(self.A_log) # (inner_dim, state_dim)
140
- A_bar = torch.exp(delta.unsqueeze(-1) * A.unsqueeze(0)) # (B, inner_dim, state_dim)
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(self, x_t: torch.Tensor, h_prev: torch.Tensor | None) -> tuple[torch.Tensor, torch.Tensor]:
 
 
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) # +1 for the <start> position
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(self, from_ids: torch.Tensor, to_ids: torch.Tensor, promo_ids: torch.Tensor) -> torch.Tensor:
204
- """from_ids/to_ids/promo_ids: (B, L) move history, L can be 0 (empty game)."""
 
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(self, from_ids: torch.Tensor, to_ids: torch.Tensor, promo_ids: torch.Tensor, lengths: torch.Tensor | None = None):
216
- """Returns (policy_logits (B,4096), promo_logits (B,5), value (B,1)).
217
-
218
- `lengths` (B,), if given, is the real (unpadded) move-history length per example --
219
- needed so a batch can mix examples of different history lengths: sequences are
220
- right-padded to a common L_max, and since the scan is causal, padding after an
221
- example's real length never affects the output gathered at that example's own
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] # (B, dim)
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) # noqa: E731
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(self, from_list: list[int], to_list: list[int], promo_list: list[int], device: torch.device | str = "cpu"):
284
- """Rebuild state from a (already-windowed, len <= max_plies) move list -- exactly
285
- matching what a windowed `forward()` call would see. Used once per real move
286
- played (cheap: len(list) O(1) steps), not per search node."""
 
 
 
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 # raised from 3 now that each node is O(1) instead of O(window length)
45
  DEFAULT_MOVETIME_S = 3.0
46
- DEFAULT_MAX_QDEPTH = 6 # quiescence extension depth beyond the main search horizon
47
- DEFAULT_QS_TOP_K = 6 # capture/promotion candidates examined per quiescence node
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 # combined non-king material below this = "endgame"
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 # (state, outputs), see model.py's incremental interface
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(from_list, to_list, promo_list, device=self.device)
 
 
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
- _MATERIAL_VALUE.get(p.piece_type, 0)
138
- for p in board.piece_map().values()
139
- if p.piece_type != chess.KING
 
 
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(self, board: chess.Board, node, alpha: float, beta: float, deadline: float, qdepth: int = 0) -> float:
176
- """Extends search through captures/promotions only, past the main search's fixed
177
- horizon, so a leaf doesn't get statically evaluated mid-capture-sequence (the
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(self, board: chess.Board, node, depth: int, alpha: float, beta: float, deadline: float) -> float:
 
 
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 = int(go_tokens[go_tokens.index(time_key) + 1]) if time_key in go_tokens else 0
252
- own_inc_ms = int(go_tokens[go_tokens.index(inc_key) + 1]) if inc_key in go_tokens else 0
 
 
 
 
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 # low-material positions are where a small model's static eval is weakest
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 = 200_000
33
- SOFT_K = 10 # candidate slots for distributional self-play targets (see ShardWriter.add)
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(self, window_from, window_to, window_promo, target_from_to, target_promo, value, soft_move_ids=None, soft_weights=None):
146
- """soft_move_ids/soft_weights (optional): up to SOFT_K (from_to_id, weight) pairs
147
- giving a distribution over candidate moves, e.g. a search's score distribution
148
- over root candidates -- richer self-play training signal than a single hard
149
- label. Omit for ordinary hard-label examples (data_pipeline.py's default path);
150
- target_from_to is then stored as a degenerate one-hot distribution instead, so
151
- pretrain.py's loss can treat every example uniformly."""
 
 
 
 
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=500_000)
198
- ap.add_argument("--target-examples", type=int, default=0, help="0 = no cap, stop on target-games instead")
 
 
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(white_cmd: str, black_cmd: str, movetime_ms: int, max_plies: int = 300) -> tuple[str, list[str]]:
 
 
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(f"[match] ILLEGAL MOVE {mv_uci} by {'white' if board.turn else 'black'} -- forfeit")
 
 
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" # ply cap reached, call it a draw for calibration purposes
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(f"[summary] --white wins={tally['white']} --black wins={tally['black']} draws={tally['draw']}")
 
 
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 = ([] for _ in range(7))
36
- soft_move_ids, soft_weights = [], []
 
 
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(f"[data] loaded {n:,} examples from {len(paths)} shards ({n - n_val:,} train / {n_val:,} val)")
 
 
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 from_ids, to_ids, promo_ids, lengths, target_from_to, target_promo, values, soft_move_ids, soft_weights
 
 
 
 
 
 
 
 
 
 
95
 
96
 
97
  def concat_batches(batch_a, batch_b):
98
- """Concatenate two get_batch() tuples along the batch dim -- both were built with the
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
- from_ids, to_ids, promo_ids, lengths, target_from_to, target_promo, values, soft_move_ids, soft_weights = batch
 
 
 
 
 
 
 
 
 
 
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)) # (B, SOFT_K)
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 loss, {
130
- "policy_loss": policy_loss.item(),
131
- "promo_loss": promo_loss.item() if promo_loss.requires_grad else float(promo_loss),
132
- "value_loss": value_loss.item(),
133
- "move_acc": acc.item(),
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=3e-4)
146
  ap.add_argument("--weight-decay", type=float, default=0.01)
147
- ap.add_argument("--value-loss-weight", type=float, default=0.5, help="raised for phase 4: diagnosed bottleneck is the value head's material judgment specifically")
148
- ap.add_argument("--steps", type=int, default=20_000)
 
 
 
 
 
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("--resume", default="", help="resume an interrupted run of THIS SAME phase (loads optimizer + step count too)")
155
- ap.add_argument("--init-from", default="", help="warm-start a NEW phase from another checkpoint's weights only (fresh optimizer/step/LR schedule)")
156
- ap.add_argument("--no-checkpoint", action="store_true", help="disable gradient checkpointing (faster per-step if it fits in memory)")
 
 
 
 
 
 
 
 
 
 
 
 
157
  ap.add_argument("--no-compile", action="store_true", help="disable torch.compile")
158
- ap.add_argument("--replay-data-dir", default="", help="mix in examples from this dir (e.g. the original supervised data) during fine-tuning, to avoid catastrophic forgetting of the base distribution")
159
- ap.add_argument("--replay-fraction", type=float, default=0.4, help="fraction of each batch drawn from --replay-data-dir when set")
 
 
 
 
 
 
 
 
 
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
- opt = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay, betas=(0.9, 0.95))
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(f"[init-from] loaded weights from {args.init_from} (step {ckpt.get('step', '?')} of its own run), starting fresh optimizer/schedule at step 0")
 
 
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(f"[replay] mixing {replay_size}/{args.batch_size} examples per batch from {args.replay_data_dir}")
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(train_model, batch, device, value_loss_weight=args.value_loss_weight)
 
 
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(ds.val_idx, size=min(args.batch_size, len(ds.val_idx)), replace=False)
 
 
240
  val_batch = ds.get_batch(val_idx, device, fixed_len=MAX_PLIES)
241
- _, val_stats = compute_loss(train_model, val_batch, device, value_loss_weight=args.value_loss_weight)
242
- print(f" [val] step {step} policy {val_stats['policy_loss']:.4f} acc {val_stats['move_acc']:.3f}")
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({"model": model.state_dict(), "optimizer": opt.state_dict(), "step": step, "config": model_cfg}, tmp_path)
 
 
 
 
 
 
 
 
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, "selfplay_finetune.py",
39
- "--checkpoint", args.checkpoint,
40
- "--out-dir", worker_out,
41
- "--games", str(per_worker),
42
- "--movetime", str(args.movetime),
43
- "--temperature-plies", str(args.temperature_plies),
 
 
 
 
 
 
44
  ]
45
- p = subprocess.Popen(cmd, stdout=log_f, stderr=subprocess.STDOUT, cwd=os.path.dirname(os.path.abspath(__file__)) or ".")
 
 
 
 
 
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(f"[launcher] WARNING: workers failed: {failed} -- check worker*.log in {args.out_dir}")
61
- print(f"[launcher] DONE: {args.workers} workers, {per_worker * args.workers} games total requested")
 
 
 
 
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(searcher: Searcher, movetime_ms: int, temperature_plies: int, max_plies: int = 200):
 
 
53
  board = chess.Board()
54
  searcher.reset()
55
- searcher.sync(board) # empty board -> root node = initial state
56
  node = searcher.root_node
57
  hist_from: list[int] = []
58
  hist_to: list[int] = []
59
  hist_promo: list[int] = []
60
- ply_records = [] # (window_from, window_to, window_promo, target_from_to, target_promo, side_to_move, soft_ids, soft_weights)
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(board, child, searcher.max_depth - 1, -float("inf"), float("inf"), deadline)
 
 
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((window_from, window_to, window_promo, from_to_id, promo_id, side_to_move, soft_ids, soft_weights))
95
-
96
- node = searcher._expand(board, node, move) # advances board+state for real (no pop -- keeping this move)
 
 
 
 
 
 
 
 
 
 
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 (wf, wt, wp, tft, tp, side, sids, sw) in ply_records
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("--no-mirror-augment", action="store_true", help="disable color-flip data doubling")
 
 
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