| """
|
| Minimal UCI engine wrapper for the trained ChessResNet model.
|
|
|
| Can be used by chess GUIs (Arena, cutechess, En Croissant, etc.) or
|
| test harnesses that speak the UCI protocol.
|
|
|
| Usage:
|
| python uci_engine.py --ckpt runs/stage1_stockfish_30m/best.pt --device cuda
|
| python uci_engine.py --ckpt runs/stage1_stockfish_30m/best.pt --device cpu
|
| """
|
| import argparse
|
| import sys
|
| from pathlib import Path
|
|
|
| import chess
|
| import torch
|
|
|
| sys.path.insert(0, str(Path(__file__).resolve().parent))
|
| from common_chess import encode_board, move_to_action_id, legal_action_ids
|
| from model import ChessResNet, create_model_from_config
|
|
|
|
|
|
|
| PIECE_VALUES = {
|
| chess.PAWN: 100,
|
| chess.KNIGHT: 320,
|
| chess.BISHOP: 330,
|
| chess.ROOK: 500,
|
| chess.QUEEN: 900,
|
| chess.KING: 0,
|
| }
|
|
|
|
|
| def material_score_for_side(board: chess.Board, side: chess.Color) -> int:
|
| """Return *side*'s material advantage in centipawns (positive = side ahead)."""
|
| score = 0
|
| for piece_type in chess.PIECE_TYPES:
|
| value = PIECE_VALUES[piece_type]
|
| score += len(board.pieces(piece_type, side)) * value
|
| score -= len(board.pieces(piece_type, not side)) * value
|
| return score
|
|
|
|
|
| def move_causes_drawish(board: chess.Board, move: chess.Move) -> bool:
|
| """Check whether *move* immediately leads to a drawish outcome."""
|
| b = board.copy(stack=True)
|
| b.push(move)
|
| if b.is_repetition(3):
|
| return True
|
| if b.can_claim_threefold_repetition():
|
| return True
|
| if b.is_fifty_moves():
|
| return True
|
| if b.can_claim_fifty_moves():
|
| return True
|
| if b.is_stalemate():
|
| return True
|
| if b.is_insufficient_material():
|
| return True
|
| return False
|
|
|
|
|
| def choose_move_with_draw_awareness(
|
| board: chess.Board,
|
| legal_moves: list,
|
| legal_action_ids: list[int],
|
| policy_logits: torch.Tensor,
|
| value_pred,
|
| topk: int = 12,
|
| ) -> chess.Move:
|
| """Draw-aware top-k rerank of legal moves.
|
|
|
| Parameters
|
| ----------
|
| board : current python-chess Board (must retain move stack).
|
| legal_moves : list of chess.Move, parallel to *legal_action_ids*.
|
| legal_action_ids : list of int action IDs parallel to *legal_moves*.
|
| policy_logits : full [N_ACTIONS] torch.Tensor (on any device).
|
| value_pred : scalar value-head output (can be tensor or float).
|
| topk : number of top candidates to consider for reranking.
|
|
|
| Returns a legal ``chess.Move``.
|
| """
|
| side = board.turn
|
| root_value = float(value_pred.squeeze().item() if hasattr(value_pred, "item") else value_pred)
|
| material = material_score_for_side(board, side)
|
|
|
|
|
| if material >= 500:
|
| root_value = max(root_value, 0.50)
|
| if material <= -500:
|
| root_value = min(root_value, -0.50)
|
|
|
|
|
| scored = [
|
| (float(policy_logits[aid].item() if hasattr(policy_logits, "item") else policy_logits[aid]), move)
|
| for aid, move in zip(legal_action_ids, legal_moves)
|
| ]
|
| scored.sort(key=lambda x: x[0], reverse=True)
|
| sorted_moves = [m for _, m in scored]
|
| original_top1 = sorted_moves[0]
|
|
|
|
|
| if root_value > 0.35:
|
| for move in sorted_moves[:topk]:
|
| if not move_causes_drawish(board, move):
|
| return move
|
| return original_top1
|
|
|
|
|
| if root_value < -0.35:
|
| for move in sorted_moves[:topk]:
|
| if move_causes_drawish(board, move):
|
| return move
|
| return original_top1
|
|
|
|
|
| return original_top1
|
|
|
|
|
| class UCIEngine:
|
| """Minimal UCI chess engine using a trained ChessResNet model."""
|
|
|
| def __init__(self, ckpt_path: str, device: str = "cuda"):
|
| self.device = device if torch.cuda.is_available() and device == "cuda" else "cpu"
|
| self.board = chess.Board()
|
| self.model = self._load_model(ckpt_path)
|
| self.model.eval()
|
| self._stop_requested = False
|
|
|
| def _load_model(self, ckpt_path: str) -> ChessResNet:
|
| ckpt = torch.load(ckpt_path, map_location=self.device, weights_only=True)
|
| model_config = ckpt.get("model_config", {})
|
| if not model_config:
|
| model_config = {
|
| "channels": ckpt.get("args", {}).get("channels", 256),
|
| "blocks": ckpt.get("args", {}).get("blocks", 20),
|
| "num_actions": 20480,
|
| }
|
| model = create_model_from_config(model_config)
|
| model.load_state_dict(ckpt["model"])
|
| model.to(self.device)
|
| return model
|
|
|
| def uci_new_game(self):
|
| """Reset board for a new game."""
|
| self.board.reset()
|
| self._stop_requested = False
|
|
|
| def set_position(self, fen: str | None = None, moves: list[str] | None = None):
|
| """Set up position from FEN and optional move list."""
|
| if fen:
|
| self.board.set_fen(fen)
|
| else:
|
| self.board.reset()
|
| if moves:
|
| for m in moves:
|
| self.board.push(chess.Move.from_uci(m))
|
|
|
| def get_best_move(self, movetime_ms: int = 1000) -> tuple[str, float]:
|
| """
|
| Return (bestmove_uci, top_logit) by evaluating the current board.
|
|
|
| Only considers legal moves. This is a single-forward-pass evaluator;
|
| it does NOT do MCTS or search.
|
| """
|
|
|
| planes = encode_board(self.board)
|
| inp = torch.from_numpy(planes).unsqueeze(0).float().to(self.device)
|
|
|
| with torch.no_grad():
|
| with torch.autocast(device_type=self.device, enabled=(self.device == "cuda")):
|
| policy_logits, value = self.model(inp)
|
|
|
| policy_logits = policy_logits.squeeze(0)
|
|
|
|
|
| action_ids, moves = legal_action_ids(self.board)
|
|
|
| if not moves:
|
| return "0000", float("-inf")
|
|
|
|
|
| best_move_obj = choose_move_with_draw_awareness(
|
| self.board, moves, action_ids, policy_logits, value, topk=12
|
| )
|
| best_move = best_move_obj.uci()
|
| best_logit = float(policy_logits[move_to_action_id(best_move_obj, self.board.turn)])
|
|
|
| return best_move, best_logit
|
|
|
| def handle_go(self, tokens: list[str]):
|
| """Process 'go' command and output bestmove."""
|
| movetime_ms = 1000
|
| if "movetime" in tokens:
|
| idx = tokens.index("movetime") + 1
|
| if idx < len(tokens):
|
| movetime_ms = int(tokens[idx])
|
|
|
| best_move, _ = self.get_best_move(movetime_ms)
|
| print(f"bestmove {best_move}", flush=True)
|
|
|
| def handle_position(self, tokens: list[str]):
|
| """Process 'position' command."""
|
| fen = None
|
| moves = []
|
|
|
| if "startpos" in tokens:
|
| pass
|
| elif "fen" in tokens:
|
|
|
| idx = tokens.index("fen") + 1
|
| fen_parts = []
|
| while idx < len(tokens) and tokens[idx] != "moves":
|
| fen_parts.append(tokens[idx])
|
| idx += 1
|
| fen = " ".join(fen_parts)
|
|
|
| if "moves" in tokens:
|
| idx = tokens.index("moves") + 1
|
| moves = tokens[idx:]
|
|
|
| self.set_position(fen, moves)
|
|
|
| def run(self):
|
| """Main UCI loop: read commands from stdin, respond to stdout."""
|
| while True:
|
| line = sys.stdin.readline()
|
| if not line:
|
| break
|
| line = line.strip()
|
| if not line:
|
| continue
|
|
|
| parts = line.split()
|
| cmd = parts[0]
|
|
|
| if cmd == "uci":
|
| print("id name JoeyStage1StockfishDistill", flush=True)
|
| print("id author Joey", flush=True)
|
| print("uciok", flush=True)
|
| elif cmd == "isready":
|
| print("readyok", flush=True)
|
| elif cmd == "ucinewgame":
|
| self.uci_new_game()
|
| elif cmd == "position":
|
| self.handle_position(parts[1:])
|
| elif cmd == "go":
|
| self.handle_go(parts[1:])
|
| elif cmd == "stop":
|
| self._stop_requested = True
|
| elif cmd == "quit":
|
| break
|
|
|
|
|
| def main():
|
| parser = argparse.ArgumentParser(description="UCI chess engine")
|
| parser.add_argument("--ckpt", default='', help="Path to checkpoint .pt file")
|
| parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
|
| args = parser.parse_args()
|
|
|
| engine = UCIEngine(args.ckpt, args.device)
|
| engine.run()
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|