Download mcts_engine.py from ParallaxOpen/Parallax-Chess-Preview: direct link, hf CLI and curl.
- Browser
- Download file 7.84 kB
-
https://huggingface.co/ParallaxOpen/Parallax-Chess-Preview/resolve/679e8f3a4cae3226cd44db0be88c42d7e4b5ad88/mcts_engine.py
- Command line
-
hf download hf://ParallaxOpen/Parallax-Chess-Preview@679e8f3a4cae3226cd44db0be88c42d7e4b5ad88/mcts_engine.py
-
curl -L -o mcts_engine.py https://huggingface.co/ParallaxOpen/Parallax-Chess-Preview/resolve/679e8f3a4cae3226cd44db0be88c42d7e4b5ad88/mcts_engine.py
7.84 kB
| #!/usr/bin/env python3 | |
| """Parallax-Chess V3: MCTS Neural Search Engine. | |
| Uses policy + value network with Monte Carlo Tree Search. | |
| """ | |
| import math, time | |
| from pathlib import Path | |
| import torch | |
| import chess | |
| import sys | |
| sys.path.insert(0, str(Path(__file__).parent.parent / "scripts")) | |
| sys.path.insert(0, str(Path(__file__).parent.parent / "src")) | |
| sys.path.insert(0, str(Path(__file__).parent.parent.parent)) | |
| from train_v3 import ParallaxChessV3, move_to_index | |
| from centauri.models.base import SmallLMConfig | |
| from chess_encoding import VOCAB_SIZE | |
| class MCTSNode: | |
| __slots__ = ['board', 'parent', 'move', 'children', 'visits', 'value_sum', 'prior'] | |
| def __init__(self, board, parent=None, move=None, prior=0.0): | |
| self.board = board | |
| self.parent = parent | |
| self.move = move | |
| self.prior = prior | |
| self.children = {} | |
| self.visits = 0 | |
| self.value_sum = 0.0 | |
| def q(self): | |
| return self.value_sum / self.visits if self.visits > 0 else 0.0 | |
| def ucb(self, child, c_puct=1.5): | |
| return child.q + c_puct * child.prior * math.sqrt(self.visits) / (1 + child.visits) | |
| def best_child(self, c_puct=1.5): | |
| return max(self.children.values(), key=lambda c: self.ucb(c, c_puct)) | |
| class ParallaxChessMCTS: | |
| """MCTS chess engine with neural network guidance.""" | |
| def __init__(self, model_path=None, device=None): | |
| self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| if model_path is None: | |
| # Find latest checkpoint | |
| ckpt_dir = Path(__file__).parent.parent / "checkpoints" / "parallax_chess_v3" | |
| candidates = sorted(ckpt_dir.glob("step_*.pt"), reverse=True) | |
| if not candidates: | |
| candidates = list(ckpt_dir.glob("final.pt")) | |
| model_path = candidates[0] | |
| cfg = SmallLMConfig( | |
| vocab_size=VOCAB_SIZE, d_model=512, n_heads=8, n_kv_heads=4, | |
| n_layers=8, intermediate_size=2048, max_seq_len=128, | |
| norm_type="rms", rope_type="neox", n_experts=0, | |
| n_loops=1, loop_mode="per_layer", | |
| ) | |
| self.model = ParallaxChessV3(cfg).to(self.device) | |
| ckpt = torch.load(str(model_path), map_location=self.device) | |
| self.model.load_state_dict(ckpt["model"]) | |
| self.model.eval() | |
| self.n_params = sum(p.numel() for p in self.model.parameters()) | |
| def _predict(self, board): | |
| """Get policy and value for a position.""" | |
| from chess_encoding import encode_board | |
| board_tokens = encode_board(board) | |
| inp = torch.tensor([board_tokens], dtype=torch.long).to(self.device) | |
| with torch.no_grad(): | |
| out = self.model(inp) | |
| logits = out["logits"] | |
| value = out["value"].item() | |
| return logits, value | |
| def _get_policy_probs(self, logits, board): | |
| """Get normalized move probabilities for legal moves only.""" | |
| probs = torch.softmax(logits, dim=-1) | |
| legal_moves = list(board.legal_moves) | |
| move_priors = {} | |
| for move in legal_moves: | |
| idx = move_to_index(move) | |
| move_priors[move] = probs[0, idx].item() | |
| total = sum(move_priors.values()) + 1e-8 | |
| for m in move_priors: | |
| move_priors[m] /= total | |
| return move_priors | |
| def search(self, board, n_simulations=200, c_puct=1.5, verbose=False): | |
| """Run MCTS from current position.""" | |
| root = MCTSNode(board.copy()) | |
| # Expand root | |
| logits, value = self._predict(board) | |
| move_priors = self._get_policy_probs(logits, board) | |
| for move, prior in move_priors.items(): | |
| child_board = board.copy() | |
| child_board.push(move) | |
| root.children[move] = MCTSNode(child_board, parent=root, move=move, prior=prior) | |
| root.visits = 1 | |
| for sim in range(n_simulations): | |
| node = root | |
| # Selection: traverse tree using UCB | |
| while node.children and not node.board.is_game_over(): | |
| node = node.best_child(c_puct) | |
| # Evaluation | |
| if node.board.is_game_over(): | |
| result = node.board.result() | |
| if result == "1-0": | |
| value = 1.0 | |
| elif result == "0-1": | |
| value = -1.0 | |
| else: | |
| value = 0.0 | |
| # Flip perspective (we evaluate from the perspective of the side to move) | |
| if node.board.turn == chess.BLACK: | |
| value = -value | |
| elif node.visits == 0: | |
| # First visit: expand with neural network | |
| logits, value = self._predict(node.board) | |
| move_priors = self._get_policy_probs(logits, node.board) | |
| for move, prior in move_priors.items(): | |
| child_board = node.board.copy() | |
| child_board.push(move) | |
| node.children[move] = MCTSNode(child_board, parent=node, move=move, prior=prior) | |
| else: | |
| # Already expanded: use value head | |
| _, value = self._predict(node.board) | |
| # Backpropagation | |
| while node is not None: | |
| node.visits += 1 | |
| node.value_sum += value | |
| value = -value # Flip for opponent | |
| node = node.parent | |
| if verbose: | |
| print("MCTS stats: %d simulations" % n_simulations) | |
| sorted_children = sorted(root.children.values(), key=lambda c: c.visits, reverse=True) | |
| for child in sorted_children[:5]: | |
| print(" %s: visits=%d Q=%.3f prior=%.3f" % ( | |
| child.move.uci(), child.visits, child.q, child.prior)) | |
| # Return most visited move | |
| best = max(root.children.values(), key=lambda c: c.visits) | |
| return best.move | |
| def choose_move(self, board, n_simulations=200): | |
| """Choose best move for a position.""" | |
| return self.search(board, n_simulations=n_simulations) | |
| def evaluate(self, board): | |
| """Evaluate position (centipawns, White perspective).""" | |
| _, value = self._predict(board) | |
| return value * 1000 | |
| def play_game(engine, stockfish_path, sf_depth=10, engine_color=chess.WHITE, verbose=True): | |
| """Play one game against Stockfish.""" | |
| board = chess.Board() | |
| sf = chess.engine.SimpleEngine.popen_uci(str(stockfish_path)) | |
| sf.configure({"Threads": 1, "Hash": 64}) | |
| moves = 0 | |
| while not board.is_game_over() and moves < 200: | |
| is_engine = (board.turn == engine_color) | |
| if is_engine: | |
| t0 = time.time() | |
| move = engine.choose_move(board, n_simulations=100) | |
| elapsed = time.time() - t0 | |
| if verbose: | |
| print("Engine: %s (%.1fs)" % (move.uci(), elapsed)) | |
| else: | |
| result = sf.play(board, chess.engine.Limit(depth=sf_depth)) | |
| move = result.move | |
| if verbose: | |
| print("Stockfish: %s" % move.uci()) | |
| board.push(move) | |
| moves += 1 | |
| sf.quit() | |
| result = board.result() | |
| if verbose: | |
| print("Result: %s in %d moves" % (result, moves)) | |
| return result | |
| if __name__ == "__main__": | |
| import argparse | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--simulations", type=int, default=200) | |
| parser.add_argument("--stockfish", default=str(Path(__file__).parent.parent / "stockfish.exe")) | |
| args = parser.parse_args() | |
| print("Loading model...") | |
| engine = ParallaxChessMCTS() | |
| # Quick test | |
| board = chess.Board() | |
| print("\nStarting position:") | |
| move = engine.choose_move(board, n_simulations=args.simulations) | |
| print("Best move: %s" % move.uci()) | |
| eval_cp = engine.evaluate(board) | |
| print("Eval: %.0f cp" % eval_cp) | |