ParallaxOpen commited on
Commit
25e96d1
·
verified ·
1 Parent(s): e861cd4

Upload mcts_engine.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. mcts_engine.py +214 -0
mcts_engine.py ADDED
@@ -0,0 +1,214 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Parallax-Chess V3: MCTS Neural Search Engine.
3
+ Uses policy + value network with Monte Carlo Tree Search.
4
+ """
5
+ import math, time
6
+ from pathlib import Path
7
+ import torch
8
+ import chess
9
+
10
+ import sys
11
+ sys.path.insert(0, str(Path(__file__).parent.parent / "scripts"))
12
+ sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
13
+ sys.path.insert(0, str(Path(__file__).parent.parent.parent))
14
+
15
+ from train_v3 import ParallaxChessV3, move_to_index
16
+ from centauri.models.base import SmallLMConfig
17
+ from chess_encoding import VOCAB_SIZE
18
+
19
+
20
+ class MCTSNode:
21
+ __slots__ = ['board', 'parent', 'move', 'children', 'visits', 'value_sum', 'prior']
22
+
23
+ def __init__(self, board, parent=None, move=None, prior=0.0):
24
+ self.board = board
25
+ self.parent = parent
26
+ self.move = move
27
+ self.prior = prior
28
+ self.children = {}
29
+ self.visits = 0
30
+ self.value_sum = 0.0
31
+
32
+ @property
33
+ def q(self):
34
+ return self.value_sum / self.visits if self.visits > 0 else 0.0
35
+
36
+ def ucb(self, child, c_puct=1.5):
37
+ return child.q + c_puct * child.prior * math.sqrt(self.visits) / (1 + child.visits)
38
+
39
+ def best_child(self, c_puct=1.5):
40
+ return max(self.children.values(), key=lambda c: self.ucb(c, c_puct))
41
+
42
+
43
+ class ParallaxChessMCTS:
44
+ """MCTS chess engine with neural network guidance."""
45
+
46
+ def __init__(self, model_path=None, device=None):
47
+ self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu")
48
+
49
+ if model_path is None:
50
+ # Find latest checkpoint
51
+ ckpt_dir = Path(__file__).parent.parent / "checkpoints" / "parallax_chess_v3"
52
+ candidates = sorted(ckpt_dir.glob("step_*.pt"), reverse=True)
53
+ if not candidates:
54
+ candidates = list(ckpt_dir.glob("final.pt"))
55
+ model_path = candidates[0]
56
+
57
+ cfg = SmallLMConfig(
58
+ vocab_size=VOCAB_SIZE, d_model=512, n_heads=8, n_kv_heads=4,
59
+ n_layers=8, intermediate_size=2048, max_seq_len=128,
60
+ norm_type="rms", rope_type="neox", n_experts=0,
61
+ n_loops=1, loop_mode="per_layer",
62
+ )
63
+ self.model = ParallaxChessV3(cfg).to(self.device)
64
+ ckpt = torch.load(str(model_path), map_location=self.device)
65
+ self.model.load_state_dict(ckpt["model"])
66
+ self.model.eval()
67
+ self.n_params = sum(p.numel() for p in self.model.parameters())
68
+
69
+ def _predict(self, board):
70
+ """Get policy and value for a position."""
71
+ from chess_encoding import encode_board
72
+ board_tokens = encode_board(board)
73
+ inp = torch.tensor([board_tokens], dtype=torch.long).to(self.device)
74
+ with torch.no_grad():
75
+ out = self.model(inp)
76
+ logits = out["logits"]
77
+ value = out["value"].item()
78
+ return logits, value
79
+
80
+ def _get_policy_probs(self, logits, board):
81
+ """Get normalized move probabilities for legal moves only."""
82
+ probs = torch.softmax(logits, dim=-1)
83
+ legal_moves = list(board.legal_moves)
84
+ move_priors = {}
85
+ for move in legal_moves:
86
+ idx = move_to_index(move)
87
+ move_priors[move] = probs[0, idx].item()
88
+ total = sum(move_priors.values()) + 1e-8
89
+ for m in move_priors:
90
+ move_priors[m] /= total
91
+ return move_priors
92
+
93
+ def search(self, board, n_simulations=200, c_puct=1.5, verbose=False):
94
+ """Run MCTS from current position."""
95
+ root = MCTSNode(board.copy())
96
+
97
+ # Expand root
98
+ logits, value = self._predict(board)
99
+ move_priors = self._get_policy_probs(logits, board)
100
+
101
+ for move, prior in move_priors.items():
102
+ child_board = board.copy()
103
+ child_board.push(move)
104
+ root.children[move] = MCTSNode(child_board, parent=root, move=move, prior=prior)
105
+ root.visits = 1
106
+
107
+ for sim in range(n_simulations):
108
+ node = root
109
+
110
+ # Selection: traverse tree using UCB
111
+ while node.children and not node.board.is_game_over():
112
+ node = node.best_child(c_puct)
113
+
114
+ # Evaluation
115
+ if node.board.is_game_over():
116
+ result = node.board.result()
117
+ if result == "1-0":
118
+ value = 1.0
119
+ elif result == "0-1":
120
+ value = -1.0
121
+ else:
122
+ value = 0.0
123
+ # Flip perspective (we evaluate from the perspective of the side to move)
124
+ if node.board.turn == chess.BLACK:
125
+ value = -value
126
+ elif node.visits == 0:
127
+ # First visit: expand with neural network
128
+ logits, value = self._predict(node.board)
129
+ move_priors = self._get_policy_probs(logits, node.board)
130
+ for move, prior in move_priors.items():
131
+ child_board = node.board.copy()
132
+ child_board.push(move)
133
+ node.children[move] = MCTSNode(child_board, parent=node, move=move, prior=prior)
134
+ else:
135
+ # Already expanded: use value head
136
+ _, value = self._predict(node.board)
137
+
138
+ # Backpropagation
139
+ while node is not None:
140
+ node.visits += 1
141
+ node.value_sum += value
142
+ value = -value # Flip for opponent
143
+ node = node.parent
144
+
145
+ if verbose:
146
+ print("MCTS stats: %d simulations" % n_simulations)
147
+ sorted_children = sorted(root.children.values(), key=lambda c: c.visits, reverse=True)
148
+ for child in sorted_children[:5]:
149
+ print(" %s: visits=%d Q=%.3f prior=%.3f" % (
150
+ child.move.uci(), child.visits, child.q, child.prior))
151
+
152
+ # Return most visited move
153
+ best = max(root.children.values(), key=lambda c: c.visits)
154
+ return best.move
155
+
156
+ def choose_move(self, board, n_simulations=200):
157
+ """Choose best move for a position."""
158
+ return self.search(board, n_simulations=n_simulations)
159
+
160
+ def evaluate(self, board):
161
+ """Evaluate position (centipawns, White perspective)."""
162
+ _, value = self._predict(board)
163
+ return value * 1000
164
+
165
+
166
+ def play_game(engine, stockfish_path, sf_depth=10, engine_color=chess.WHITE, verbose=True):
167
+ """Play one game against Stockfish."""
168
+ board = chess.Board()
169
+ sf = chess.engine.SimpleEngine.popen_uci(str(stockfish_path))
170
+ sf.configure({"Threads": 1, "Hash": 64})
171
+
172
+ moves = 0
173
+ while not board.is_game_over() and moves < 200:
174
+ is_engine = (board.turn == engine_color)
175
+ if is_engine:
176
+ t0 = time.time()
177
+ move = engine.choose_move(board, n_simulations=100)
178
+ elapsed = time.time() - t0
179
+ if verbose:
180
+ print("Engine: %s (%.1fs)" % (move.uci(), elapsed))
181
+ else:
182
+ result = sf.play(board, chess.engine.Limit(depth=sf_depth))
183
+ move = result.move
184
+ if verbose:
185
+ print("Stockfish: %s" % move.uci())
186
+
187
+ board.push(move)
188
+ moves += 1
189
+
190
+ sf.quit()
191
+ result = board.result()
192
+ if verbose:
193
+ print("Result: %s in %d moves" % (result, moves))
194
+ return result
195
+
196
+
197
+ if __name__ == "__main__":
198
+ import argparse
199
+ parser = argparse.ArgumentParser()
200
+ parser.add_argument("--simulations", type=int, default=200)
201
+ parser.add_argument("--stockfish", default=str(Path(__file__).parent.parent / "stockfish.exe"))
202
+ args = parser.parse_args()
203
+
204
+ print("Loading model...")
205
+ engine = ParallaxChessMCTS()
206
+
207
+ # Quick test
208
+ board = chess.Board()
209
+ print("\nStarting position:")
210
+ move = engine.choose_move(board, n_simulations=args.simulations)
211
+ print("Best move: %s" % move.uci())
212
+
213
+ eval_cp = engine.evaluate(board)
214
+ print("Eval: %.0f cp" % eval_cp)