Joeyfully commited on
Commit
ba69de3
Β·
verified Β·
1 Parent(s): bd3ca24

Upload 5 files

Browse files
Files changed (5) hide show
  1. best.pt +3 -0
  2. common_chess.py +124 -0
  3. engine.sh +3 -0
  4. model.py +176 -0
  5. uci_engine.py +267 -0
best.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fa1c91b6e90fa7fe706b664c5ca0a6dc1533e5d48346e62b87d91da51cc4ad7a
3
+ size 348079598
common_chess.py ADDED
@@ -0,0 +1,124 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Common chess encoding utilities shared across training, evaluation, and UCI engine.
3
+
4
+ Must match the encoding used in `scripts/generate_data/label_stockfish_dataset.py`.
5
+ """
6
+ import chess
7
+ import numpy as np
8
+
9
+ NUM_ACTIONS = 64 * 64 * 5 # 20480: canonical 64Γ—64 squares Γ— 5 promotion types
10
+ NUM_BOARD_PLANES = 18 # 6 own pieces + 6 opponent + 4 castling + 1 ep + 1 side
11
+
12
+
13
+ def encode_board(board: chess.Board) -> np.ndarray:
14
+ """Encode a chess.Board into 18 canonical planes, shape (18, 8, 8), uint8.
15
+
16
+ If black is to move, squares are mirrored vertically so that the side
17
+ to move is always at the bottom (ranks 0-3 own, ranks 4-7 opponent).
18
+
19
+ Planes:
20
+ 0-5 : own pawn, knight, bishop, rook, queen, king
21
+ 6-11 : opponent pawn, knight, bishop, rook, queen, king
22
+ 12 : own kingside castling right (all-1 plane)
23
+ 13 : own queenside castling right (all-1 plane)
24
+ 14 : opponent kingside castling right (all-1 plane)
25
+ 15 : opponent queenside castling right (all-1 plane)
26
+ 16 : en-passant target square (1-hot)
27
+ 17 : original side-to-move (1 = white, 0 = black)
28
+ """
29
+ planes = np.zeros((18, 8, 8), dtype=np.uint8)
30
+ turn = board.turn
31
+ mirror = not turn # mirror squares if black to move
32
+
33
+ # Piece planes (own = 0-5, opponent = 6-11)
34
+ for sq in chess.SQUARES:
35
+ piece = board.piece_at(sq)
36
+ if piece is None:
37
+ continue
38
+ csq = chess.square_mirror(sq) if mirror else sq
39
+ row, col = divmod(csq, 8)
40
+ if piece.color == turn:
41
+ idx = piece.piece_type - 1 # 0-5 own
42
+ else:
43
+ idx = piece.piece_type - 1 + 6 # 6-11 opponent
44
+ planes[idx, row, col] = 1
45
+
46
+ # Castling rights (own / opponent perspective)
47
+ if board.has_kingside_castling_rights(turn):
48
+ planes[12, :, :] = 1
49
+ if board.has_queenside_castling_rights(turn):
50
+ planes[13, :, :] = 1
51
+ if board.has_kingside_castling_rights(not turn):
52
+ planes[14, :, :] = 1
53
+ if board.has_queenside_castling_rights(not turn):
54
+ planes[15, :, :] = 1
55
+
56
+ # En-passant
57
+ ep = board.ep_square
58
+ if ep is not None:
59
+ cep = chess.square_mirror(ep) if mirror else ep
60
+ row, col = divmod(cep, 8)
61
+ planes[16, row, col] = 1
62
+
63
+ # Original side-to-move indicator
64
+ if turn == chess.WHITE:
65
+ planes[17, :, :] = 1
66
+
67
+ return planes
68
+
69
+
70
+ def move_to_action_id(move: chess.Move, turn: chess.Color) -> int:
71
+ """Convert a chess.Move to canonical action id [0, 20480).
72
+
73
+ If black to move, from_square and to_square are mirrored vertically
74
+ so the encoding is invariant under board orientation.
75
+ """
76
+ if turn == chess.BLACK:
77
+ from_sq = chess.square_mirror(move.from_square)
78
+ to_sq = chess.square_mirror(move.to_square)
79
+ else:
80
+ from_sq = move.from_square
81
+ to_sq = move.to_square
82
+
83
+ promo = move.promotion
84
+ if promo is None:
85
+ pid = 0
86
+ elif promo == chess.QUEEN:
87
+ pid = 1
88
+ elif promo == chess.ROOK:
89
+ pid = 2
90
+ elif promo == chess.BISHOP:
91
+ pid = 3
92
+ elif promo == chess.KNIGHT:
93
+ pid = 4
94
+ else:
95
+ pid = 0 # should never happen
96
+
97
+ return (from_sq * 64 + to_sq) * 5 + pid
98
+
99
+
100
+ def action_id_to_move(action_id: int, turn: chess.Color) -> chess.Move:
101
+ """Inverse of move_to_action_id."""
102
+ pid = action_id % 5
103
+ raw = action_id // 5
104
+ to_sq = raw % 64
105
+ from_sq = raw // 64
106
+
107
+ if turn == chess.BLACK:
108
+ from_sq = chess.square_mirror(from_sq)
109
+ to_sq = chess.square_mirror(to_sq)
110
+
111
+ promo_map = {0: None, 1: chess.QUEEN, 2: chess.ROOK,
112
+ 3: chess.BISHOP, 4: chess.KNIGHT}
113
+ return chess.Move(from_sq, to_sq, promotion=promo_map[pid])
114
+
115
+
116
+ def legal_action_ids(board: chess.Board):
117
+ """Return (action_ids, moves) for all legal moves on *board*.
118
+
119
+ action_ids: list of int length len(moves)
120
+ moves: list of chess.Move
121
+ """
122
+ moves = list(board.legal_moves)
123
+ action_ids = [move_to_action_id(m, board.turn) for m in moves]
124
+ return action_ids, moves
engine.sh ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ cd "$(dirname "$0")"
3
+ exec /root/miniconda3/envs/lichessbot/bin/python -u uci_engine.py --ckpt best.pt --device cpu
model.py ADDED
@@ -0,0 +1,176 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ ChessResNet: a ResNet-style policy-value network for chess.
3
+
4
+ Architecture:
5
+ Stem: Conv3x3 18β†’channels, GroupNorm, GELU
6
+ Tower: N residual blocks (Conv3x3→GN→GELU→Conv3x3→GN→+→GELU)
7
+ Policy head: spatial Conv1x1 β†’ 320 channels β†’ reshape to [B, 20480]
8
+ Value head: Conv1x1 β†’ 32 β†’ Flatten β†’ Linear 256 β†’ Linear 1 β†’ tanh
9
+
10
+ Default config (channels=256, blocks=24) yields ~29M parameters.
11
+ """
12
+ import math
13
+ from typing import Optional
14
+
15
+ import torch
16
+ from torch import nn
17
+
18
+
19
+ class ResidualBlock(nn.Module):
20
+ """Pre-activation residual block with GroupNorm."""
21
+
22
+ def __init__(self, channels: int, norm_groups: int = 32):
23
+ super().__init__()
24
+ self.conv1 = nn.Conv2d(channels, channels, 3, padding=1, bias=False)
25
+ self.norm1 = nn.GroupNorm(norm_groups, channels)
26
+ self.act1 = nn.GELU()
27
+ self.conv2 = nn.Conv2d(channels, channels, 3, padding=1, bias=False)
28
+ self.norm2 = nn.GroupNorm(norm_groups, channels)
29
+ self.act2 = nn.GELU()
30
+
31
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
32
+ residual = x
33
+ out = self.conv1(x)
34
+ out = self.norm1(out)
35
+ out = self.act1(out)
36
+ out = self.conv2(out)
37
+ out = self.norm2(out)
38
+ out = out + residual
39
+ out = self.act2(out)
40
+ return out
41
+
42
+
43
+ class ChessResNet(nn.Module):
44
+ """Policy-value ResNet for chess.
45
+
46
+ Args:
47
+ channels: Number of filters in the residual tower (default: 256).
48
+ blocks: Number of residual blocks (default: 20).
49
+ num_actions: Size of action space (default: 20480).
50
+ norm_groups: Number of groups for GroupNorm (default: 32).
51
+ """
52
+
53
+ def __init__(
54
+ self,
55
+ channels: int = 256,
56
+ blocks: int = 24,
57
+ num_actions: int = 20480,
58
+ norm_groups: int = 32,
59
+ ):
60
+ super().__init__()
61
+ self.channels = channels
62
+ self.blocks = blocks
63
+ self.num_actions = num_actions
64
+
65
+ # ---- Stem ----
66
+ self.stem = nn.Sequential(
67
+ nn.Conv2d(18, channels, 3, padding=1, bias=False),
68
+ nn.GroupNorm(norm_groups, channels),
69
+ nn.GELU(),
70
+ )
71
+
72
+ # ---- Residual tower ----
73
+ tower = []
74
+ for _ in range(blocks):
75
+ tower.append(ResidualBlock(channels, norm_groups))
76
+ self.tower = nn.Sequential(*tower)
77
+
78
+ # ---- Policy head (spatial): linear logits, no activation ----
79
+ # 320 = 64 destination squares Γ— 5 promotion types
80
+ self.policy_head = nn.Conv2d(channels, 320, 1, bias=True)
81
+
82
+ # ---- Value head ----
83
+ self.value_head = nn.Sequential(
84
+ nn.Conv2d(channels, 32, 1, bias=False),
85
+ nn.GroupNorm(8, 32),
86
+ nn.GELU(),
87
+ nn.Flatten(),
88
+ nn.Linear(32 * 8 * 8, 256),
89
+ nn.GELU(),
90
+ nn.Linear(256, 1),
91
+ nn.Tanh(),
92
+ )
93
+
94
+ self._init_weights()
95
+
96
+ def _init_weights(self):
97
+ """Initialize weights with scaled normal for stability."""
98
+ for m in self.modules():
99
+ if isinstance(m, nn.Conv2d):
100
+ nn.init.kaiming_normal_(m.weight, mode='fan_out',
101
+ nonlinearity='relu')
102
+ elif isinstance(m, nn.Linear):
103
+ nn.init.trunc_normal_(m.weight, std=0.02)
104
+ if m.bias is not None:
105
+ nn.init.zeros_(m.bias)
106
+
107
+ def forward(self, boards: torch.Tensor):
108
+ """
109
+ Args:
110
+ boards: [B, 18, 8, 8] float tensor (values 0.0 or 1.0)
111
+
112
+ Returns:
113
+ policy_logits: [B, 20480] raw logits for all action IDs
114
+ value: [B] tanh-squashed scalar [-1, 1]
115
+ """
116
+ x = self.stem(boards)
117
+ x = self.tower(x)
118
+
119
+ # Policy head: [B, 320, 8, 8] (NCHW) -> [B, 20480]
120
+ # Action ID encoding:
121
+ # action_id = ((from_sq * 64) + to_sq) * 5 + promo_id
122
+ # Spatial (h,w) = from_square = h*8 + w
123
+ # Channel c = to_sq * 5 + promo_id (0-319)
124
+ pol = self.policy_head(x) # [B, 320, 8, 8] NCHW
125
+ B = pol.shape[0]
126
+ pol = pol.permute(0, 2, 3, 1) # [B, 8, 8, 320] NHWC
127
+ policy_logits = pol.reshape(B, 8 * 8 * 320) # [B, 20480]
128
+ # After permute+reshape:
129
+ # flat_idx = (h*8+w) * 320 + c
130
+ # = from_sq * 320 + to_sq * 5 + promo_id
131
+ # = ((from_sq * 64) + to_sq) * 5 + promo_id βœ“
132
+
133
+ # Value head
134
+ value = self.value_head(x).squeeze(-1) # [B]
135
+
136
+ return policy_logits, value
137
+
138
+
139
+ def count_parameters(model: nn.Module) -> int:
140
+ """Return total number of trainable parameters."""
141
+ return sum(p.numel() for p in model.parameters() if p.requires_grad)
142
+
143
+
144
+ def get_model_config(model: ChessResNet) -> dict:
145
+ """Return model hyperparameters for checkpoint saving."""
146
+ return dict(
147
+ channels=model.channels,
148
+ blocks=model.blocks,
149
+ num_actions=model.num_actions,
150
+ )
151
+
152
+
153
+ def create_model_from_config(config: dict) -> ChessResNet:
154
+ """Create a model from a config dict (as stored in checkpoints)."""
155
+ return ChessResNet(
156
+ channels=config.get("channels", 256),
157
+ blocks=config.get("blocks", 24),
158
+ num_actions=config.get("num_actions", 20480),
159
+ )
160
+
161
+
162
+ if __name__ == "__main__":
163
+ m = ChessResNet(channels=256, blocks=24)
164
+ n_params = count_parameters(m)
165
+ print(f"ChessResNet(channels=256, blocks=24): {n_params:,} parameters")
166
+ # ~29.0M expected
167
+
168
+ m = ChessResNet(channels=128, blocks=4)
169
+ n_params = count_parameters(m)
170
+ print(f"ChessResNet(channels=128, blocks=4): {n_params:,} parameters")
171
+
172
+ # Test forward
173
+ x = torch.randn(4, 18, 8, 8)
174
+ pol, val = m(x)
175
+ print(f"Policy logits shape: {pol.shape} (expected [4, 20480])")
176
+ print(f"Value shape: {val.shape} (expected [4])")
uci_engine.py ADDED
@@ -0,0 +1,267 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Minimal UCI engine wrapper for the trained ChessResNet model.
3
+
4
+ Can be used by chess GUIs (Arena, cutechess, En Croissant, etc.) or
5
+ test harnesses that speak the UCI protocol.
6
+
7
+ Usage:
8
+ python uci_engine.py --ckpt runs/stage1_stockfish_30m/best.pt --device cuda
9
+ python uci_engine.py --ckpt runs/stage1_stockfish_30m/best.pt --device cpu
10
+ """
11
+ import argparse
12
+ import sys
13
+ from pathlib import Path
14
+
15
+ import chess
16
+ import torch
17
+
18
+ sys.path.insert(0, str(Path(__file__).resolve().parent))
19
+ from common_chess import encode_board, move_to_action_id, legal_action_ids
20
+ from model import ChessResNet, create_model_from_config
21
+
22
+ # ── Draw-aware move selection helpers ──────────────────────────────────────
23
+
24
+ PIECE_VALUES = {
25
+ chess.PAWN: 100,
26
+ chess.KNIGHT: 320,
27
+ chess.BISHOP: 330,
28
+ chess.ROOK: 500,
29
+ chess.QUEEN: 900,
30
+ chess.KING: 0,
31
+ }
32
+
33
+
34
+ def material_score_for_side(board: chess.Board, side: chess.Color) -> int:
35
+ """Return *side*'s material advantage in centipawns (positive = side ahead)."""
36
+ score = 0
37
+ for piece_type in chess.PIECE_TYPES:
38
+ value = PIECE_VALUES[piece_type]
39
+ score += len(board.pieces(piece_type, side)) * value
40
+ score -= len(board.pieces(piece_type, not side)) * value
41
+ return score
42
+
43
+
44
+ def move_causes_drawish(board: chess.Board, move: chess.Move) -> bool:
45
+ """Check whether *move* immediately leads to a drawish outcome."""
46
+ b = board.copy(stack=True)
47
+ b.push(move)
48
+ if b.is_repetition(3):
49
+ return True
50
+ if b.can_claim_threefold_repetition():
51
+ return True
52
+ if b.is_fifty_moves():
53
+ return True
54
+ if b.can_claim_fifty_moves():
55
+ return True
56
+ if b.is_stalemate():
57
+ return True
58
+ if b.is_insufficient_material():
59
+ return True
60
+ return False
61
+
62
+
63
+ def choose_move_with_draw_awareness(
64
+ board: chess.Board,
65
+ legal_moves: list,
66
+ legal_action_ids: list[int],
67
+ policy_logits: torch.Tensor,
68
+ value_pred,
69
+ topk: int = 12,
70
+ ) -> chess.Move:
71
+ """Draw-aware top-k rerank of legal moves.
72
+
73
+ Parameters
74
+ ----------
75
+ board : current python-chess Board (must retain move stack).
76
+ legal_moves : list of chess.Move, parallel to *legal_action_ids*.
77
+ legal_action_ids : list of int action IDs parallel to *legal_moves*.
78
+ policy_logits : full [N_ACTIONS] torch.Tensor (on any device).
79
+ value_pred : scalar value-head output (can be tensor or float).
80
+ topk : number of top candidates to consider for reranking.
81
+
82
+ Returns a legal ``chess.Move``.
83
+ """
84
+ side = board.turn
85
+ root_value = float(value_pred.squeeze().item() if hasattr(value_pred, "item") else value_pred)
86
+ material = material_score_for_side(board, side)
87
+
88
+ # Material fallback: override value signal when material gap is large
89
+ if material >= 500:
90
+ root_value = max(root_value, 0.50)
91
+ if material <= -500:
92
+ root_value = min(root_value, -0.50)
93
+
94
+ # Sort legal moves by policy logit descending
95
+ scored = [
96
+ (float(policy_logits[aid].item() if hasattr(policy_logits, "item") else policy_logits[aid]), move)
97
+ for aid, move in zip(legal_action_ids, legal_moves)
98
+ ]
99
+ scored.sort(key=lambda x: x[0], reverse=True)
100
+ sorted_moves = [m for _, m in scored]
101
+ original_top1 = sorted_moves[0]
102
+
103
+ # ── Advantage: avoid draws ────────────────────────────────────────
104
+ if root_value > 0.35:
105
+ for move in sorted_moves[:topk]:
106
+ if not move_causes_drawish(board, move):
107
+ return move
108
+ return original_top1 # fallback – every top-k move is drawish
109
+
110
+ # ── Disadvantage: prefer draws ────────────────────────────────────
111
+ if root_value < -0.35:
112
+ for move in sorted_moves[:topk]:
113
+ if move_causes_drawish(board, move):
114
+ return move
115
+ return original_top1 # fallback – no drawish move in top-k
116
+
117
+ # ── Near-equality: stick with policy top1 ─────────────────────────
118
+ return original_top1
119
+
120
+
121
+ class UCIEngine:
122
+ """Minimal UCI chess engine using a trained ChessResNet model."""
123
+
124
+ def __init__(self, ckpt_path: str, device: str = "cuda"):
125
+ self.device = device if torch.cuda.is_available() and device == "cuda" else "cpu"
126
+ self.board = chess.Board()
127
+ self.model = self._load_model(ckpt_path)
128
+ self.model.eval()
129
+ self._stop_requested = False
130
+
131
+ def _load_model(self, ckpt_path: str) -> ChessResNet:
132
+ ckpt = torch.load(ckpt_path, map_location=self.device, weights_only=True)
133
+ model_config = ckpt.get("model_config", {})
134
+ if not model_config:
135
+ model_config = {
136
+ "channels": ckpt.get("args", {}).get("channels", 256),
137
+ "blocks": ckpt.get("args", {}).get("blocks", 20),
138
+ "num_actions": 20480,
139
+ }
140
+ model = create_model_from_config(model_config)
141
+ model.load_state_dict(ckpt["model"])
142
+ model.to(self.device)
143
+ return model
144
+
145
+ def uci_new_game(self):
146
+ """Reset board for a new game."""
147
+ self.board.reset()
148
+ self._stop_requested = False
149
+
150
+ def set_position(self, fen: str | None = None, moves: list[str] | None = None):
151
+ """Set up position from FEN and optional move list."""
152
+ if fen:
153
+ self.board.set_fen(fen)
154
+ else:
155
+ self.board.reset()
156
+ if moves:
157
+ for m in moves:
158
+ self.board.push(chess.Move.from_uci(m))
159
+
160
+ def get_best_move(self, movetime_ms: int = 1000) -> tuple[str, float]:
161
+ """
162
+ Return (bestmove_uci, top_logit) by evaluating the current board.
163
+
164
+ Only considers legal moves. This is a single-forward-pass evaluator;
165
+ it does NOT do MCTS or search.
166
+ """
167
+ # Encode board
168
+ planes = encode_board(self.board)
169
+ inp = torch.from_numpy(planes).unsqueeze(0).float().to(self.device) # [1,18,8,8]
170
+
171
+ with torch.no_grad():
172
+ with torch.autocast(device_type=self.device, enabled=(self.device == "cuda")):
173
+ policy_logits, value = self.model(inp)
174
+
175
+ policy_logits = policy_logits.squeeze(0) # [20480]
176
+
177
+ # Get legal moves and their action IDs
178
+ action_ids, moves = legal_action_ids(self.board)
179
+
180
+ if not moves:
181
+ return "0000", float("-inf")
182
+
183
+ # Draw-aware top-k rerank (avoids threefold-repetition, 50-move, etc.)
184
+ best_move_obj = choose_move_with_draw_awareness(
185
+ self.board, moves, action_ids, policy_logits, value, topk=12
186
+ )
187
+ best_move = best_move_obj.uci()
188
+ best_logit = float(policy_logits[move_to_action_id(best_move_obj, self.board.turn)])
189
+
190
+ return best_move, best_logit
191
+
192
+ def handle_go(self, tokens: list[str]):
193
+ """Process 'go' command and output bestmove."""
194
+ movetime_ms = 1000
195
+ if "movetime" in tokens:
196
+ idx = tokens.index("movetime") + 1
197
+ if idx < len(tokens):
198
+ movetime_ms = int(tokens[idx])
199
+
200
+ best_move, _ = self.get_best_move(movetime_ms)
201
+ print(f"bestmove {best_move}", flush=True)
202
+
203
+ def handle_position(self, tokens: list[str]):
204
+ """Process 'position' command."""
205
+ fen = None
206
+ moves = []
207
+
208
+ if "startpos" in tokens:
209
+ pass # use starting position (board.reset() already done or standard)
210
+ elif "fen" in tokens:
211
+ # Collect FEN string up to "moves" keyword
212
+ idx = tokens.index("fen") + 1
213
+ fen_parts = []
214
+ while idx < len(tokens) and tokens[idx] != "moves":
215
+ fen_parts.append(tokens[idx])
216
+ idx += 1
217
+ fen = " ".join(fen_parts)
218
+
219
+ if "moves" in tokens:
220
+ idx = tokens.index("moves") + 1
221
+ moves = tokens[idx:]
222
+
223
+ self.set_position(fen, moves)
224
+
225
+ def run(self):
226
+ """Main UCI loop: read commands from stdin, respond to stdout."""
227
+ while True:
228
+ line = sys.stdin.readline()
229
+ if not line:
230
+ break
231
+ line = line.strip()
232
+ if not line:
233
+ continue
234
+
235
+ parts = line.split()
236
+ cmd = parts[0]
237
+
238
+ if cmd == "uci":
239
+ print("id name JoeyStage1StockfishDistill", flush=True)
240
+ print("id author Joey", flush=True)
241
+ print("uciok", flush=True)
242
+ elif cmd == "isready":
243
+ print("readyok", flush=True)
244
+ elif cmd == "ucinewgame":
245
+ self.uci_new_game()
246
+ elif cmd == "position":
247
+ self.handle_position(parts[1:])
248
+ elif cmd == "go":
249
+ self.handle_go(parts[1:])
250
+ elif cmd == "stop":
251
+ self._stop_requested = True
252
+ elif cmd == "quit":
253
+ break
254
+
255
+
256
+ def main():
257
+ parser = argparse.ArgumentParser(description="UCI chess engine")
258
+ parser.add_argument("--ckpt", default='', help="Path to checkpoint .pt file")
259
+ parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
260
+ args = parser.parse_args()
261
+
262
+ engine = UCIEngine(args.ckpt, args.device)
263
+ engine.run()
264
+
265
+
266
+ if __name__ == "__main__":
267
+ main()