Parallax-Chess-Preview / mcts_engine.py
ParallaxOpen's picture
Upload mcts_engine.py with huggingface_hub
25e96d1 verified
Raw History Blame
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
@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)