""" 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 # ── Draw-aware move selection helpers ────────────────────────────────────── 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) # Material fallback: override value signal when material gap is large if material >= 500: root_value = max(root_value, 0.50) if material <= -500: root_value = min(root_value, -0.50) # Sort legal moves by policy logit descending 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] # ── Advantage: avoid draws ──────────────────────────────────────── if root_value > 0.35: for move in sorted_moves[:topk]: if not move_causes_drawish(board, move): return move return original_top1 # fallback – every top-k move is drawish # ── Disadvantage: prefer draws ──────────────────────────────────── if root_value < -0.35: for move in sorted_moves[:topk]: if move_causes_drawish(board, move): return move return original_top1 # fallback – no drawish move in top-k # ── Near-equality: stick with policy 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. """ # Encode board planes = encode_board(self.board) inp = torch.from_numpy(planes).unsqueeze(0).float().to(self.device) # [1,18,8,8] 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) # [20480] # Get legal moves and their action IDs action_ids, moves = legal_action_ids(self.board) if not moves: return "0000", float("-inf") # Draw-aware top-k rerank (avoids threefold-repetition, 50-move, etc.) 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 # use starting position (board.reset() already done or standard) elif "fen" in tokens: # Collect FEN string up to "moves" keyword 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()