#!/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 @property 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)