chess-position-evaluator / inference.py
engdarwish's picture
Fix escaped newline corruption from previous commit
1cd8369 verified
Raw
History Blame Contribute Delete
6.5 kB
"""
Chess Position Evaluator (V7 HULK) - Inference Example
Author: Ahmed Darwish (@engdarwish)
https://huggingface.co/engdarwish/chess-position-evaluator
Minimal, runnable example: download the weights, rebuild the exact
SE-ResNet-20 dual-head architecture, encode a FEN position, and get
a value estimate + ranked legal moves.
Install:
pip install torch python-chess huggingface_hub numpy
"""
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import chess
from huggingface_hub import hf_hub_download
REPO_ID = "engdarwish/chess-position-evaluator"
MOVE_SPACE = 64 * 64 * 5
PROMO_MAP = {None: 0, chess.QUEEN: 1, chess.ROOK: 2, chess.BISHOP: 3, chess.KNIGHT: 4}
PROMO_MAP_INV = {v: k for k, v in PROMO_MAP.items() if k is not None}
def uci_to_idx(uci_str: str) -> int:
move = chess.Move.from_uci(uci_str)
promo = PROMO_MAP.get(move.promotion, 0)
return (move.from_square * 64 + move.to_square) * 5 + promo
def idx_to_uci(idx: int) -> str:
promo_type = idx % 5
rest = idx // 5
to_sq = rest % 64
from_sq = rest // 64
promo_piece = PROMO_MAP_INV.get(promo_type)
return chess.Move(from_sq, to_sq, promotion=promo_piece).uci()
PIECE_IDX = {
chess.PAWN: 0, chess.KNIGHT: 1, chess.BISHOP: 2,
chess.ROOK: 3, chess.QUEEN: 4, chess.KING: 5,
}
def fen_to_tensor(fen: str) -> np.ndarray:
board = chess.Board(fen)
t = np.zeros((18, 8, 8), dtype=np.float32)
for sq in chess.SQUARES:
piece = board.piece_at(sq)
if piece:
r = 7 - (sq >> 3)
c = sq & 7
ch = PIECE_IDX[piece.piece_type]
t[ch if piece.color else ch + 6, r, c] = 1.0
if board.turn == chess.WHITE:
t[12] = 1.0
t[13] = float(board.has_kingside_castling_rights(chess.WHITE))
t[14] = float(board.has_queenside_castling_rights(chess.WHITE))
t[15] = float(board.has_kingside_castling_rights(chess.BLACK))
t[16] = float(board.has_queenside_castling_rights(chess.BLACK))
if board.ep_square is not None:
t[17, 7 - (board.ep_square >> 3), board.ep_square & 7] = 1.0
return t
class SEBlock(nn.Module):
def __init__(self, channels, reduction=16):
super().__init__()
self.pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Sequential(
nn.Linear(channels, channels // reduction, bias=False),
nn.ReLU(inplace=True),
nn.Linear(channels // reduction, channels, bias=False),
nn.Sigmoid(),
)
def forward(self, x):
b, c, _, _ = x.size()
y = self.pool(x).view(b, c)
y = self.fc(y).view(b, c, 1, 1)
return x * y
class SEResBlock(nn.Module):
def __init__(self, channels):
super().__init__()
self.conv1 = nn.Conv2d(channels, channels, 3, padding=1, bias=False)
self.bn1 = nn.BatchNorm2d(channels)
self.conv2 = nn.Conv2d(channels, channels, 3, padding=1, bias=False)
self.bn2 = nn.BatchNorm2d(channels)
self.se = SEBlock(channels)
def forward(self, x):
r = x
out = F.relu(self.bn1(self.conv1(x)), inplace=True)
out = self.bn2(self.conv2(out))
out = self.se(out)
return F.relu(out + r, inplace=True)
class ChessPositionEvaluator(nn.Module):
"""
Chess Position Evaluator - V7 HULK (SE-ResNet-20, Dual Head)
Input : (B, 18, 8, 8)
Value : scalar in [-1, 1] (tanh, from White's perspective)
Policy : logits over 20,480 moves (includes underpromotion)
"""
def __init__(self, in_channels=18, num_filters=256, num_res_blocks=20, policy_size=MOVE_SPACE):
super().__init__()
self.input_block = nn.Sequential(
nn.Conv2d(in_channels, num_filters, 3, padding=1, bias=False),
nn.BatchNorm2d(num_filters),
nn.ReLU(inplace=True),
)
self.tower = nn.Sequential(*[SEResBlock(num_filters) for _ in range(num_res_blocks)])
self.value_head = nn.Sequential(
nn.Conv2d(num_filters, 32, 1, bias=False),
nn.BatchNorm2d(32), nn.ReLU(inplace=True),
nn.Flatten(),
nn.Linear(32 * 8 * 8, 256), nn.ReLU(inplace=True),
nn.Dropout(0.3),
nn.Linear(256, 1), nn.Tanh(),
)
self.policy_head = nn.Sequential(
nn.Conv2d(num_filters, 32, 1, bias=False),
nn.BatchNorm2d(32), nn.ReLU(inplace=True),
nn.Flatten(),
nn.Linear(32 * 8 * 8, 1024), nn.ReLU(inplace=True),
nn.Dropout(0.3),
nn.Linear(1024, policy_size),
)
def forward(self, x):
x = self.input_block(x)
x = self.tower(x)
value = self.value_head(x).squeeze(-1)
policy = self.policy_head(x)
return value, policy
def load_model(device: str = "cpu") -> ChessPositionEvaluator:
weights_path = hf_hub_download(repo_id=REPO_ID, filename="model_weights.pt")
model = ChessPositionEvaluator().to(device)
state_dict = torch.load(weights_path, map_location=device)
model.load_state_dict(state_dict)
model.eval()
return model
@torch.no_grad()
def get_best_move(model: ChessPositionEvaluator, fen: str, top_k: int = 5, device: str = "cpu"):
"""Returns (value, [(uci_move, probability), ...]) using legal-move masking."""
board = chess.Board(fen)
t = fen_to_tensor(fen)
inp = torch.tensor(t, dtype=torch.float32).unsqueeze(0).to(device)
value, policy_logits = model(inp)
logits = policy_logits[0].float().cpu()
mask = torch.full((MOVE_SPACE,), float("-inf"))
legal_moves = list(board.legal_moves)
for move in legal_moves:
idx = uci_to_idx(move.uci())
if 0 <= idx < MOVE_SPACE:
mask[idx] = 0.0
probs = torch.softmax(logits + mask, dim=0)
top_indices = torch.topk(probs, min(top_k, len(legal_moves))).indices
return value.item(), [(idx_to_uci(i.item()), probs[i].item()) for i in top_indices]
if __name__ == "__main__":
model = load_model()
print(f"Parameters: {sum(p.numel() for p in model.parameters()):,}")
positions = [
("Starting position", chess.STARTING_FEN),
("After 1.e4", "rnbqkbnr/pppppppp/8/8/4P3/8/PPPP1PPP/RNBQKBNR b KQkq e3 0 1"),
]
for name, fen in positions:
val, moves = get_best_move(model, fen, top_k=3)
print(f"\n{name}")
print(f" Value: {val:.4f}")
for uci, prob in moves:
print(f" {uci} ({prob * 100:.1f}%)")