Download uci_engine.py from Joeyfully/ChessResNet-30M: direct link, hf CLI and curl.
- Browser
- Download file 9.7 kB
-
https://huggingface.co/Joeyfully/ChessResNet-30M/resolve/main/uci_engine.py
- Command line
-
hf download hf://Joeyfully/ChessResNet-30M/uci_engine.py
-
curl -L -o uci_engine.py https://huggingface.co/Joeyfully/ChessResNet-30M/resolve/main/uci_engine.py
9.7 kB
| """ | |
| 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() | |