Upload 5 files
Browse files- best.pt +3 -0
- common_chess.py +124 -0
- engine.sh +3 -0
- model.py +176 -0
- 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()
|