from __future__ import annotations import math import numpy as np import torch from game import canonical, is_terminal, legal_actions, play, winner from model import MicroZeroNet class NeuralMCTS: def __init__( self, model: MicroZeroNet, simulations: int, exploration: float = 1.35, ) -> None: self.model = model self.simulations = simulations self.exploration = exploration self.priors: dict[tuple, np.ndarray] = {} self.state_visits: dict[tuple, int] = {} self.edge_visits: dict[tuple, int] = {} self.edge_value: dict[tuple, float] = {} @staticmethod def key(board: np.ndarray, player: int) -> tuple: return (*board.tolist(), player) @torch.inference_mode() def expand(self, board: np.ndarray, player: int) -> float: key = self.key(board, player) logits, value = self.model(torch.from_numpy(canonical(board, player))[None]) logits = logits[0].numpy() legal = legal_actions(board) masked = np.full(9, -1e9, dtype=np.float64) masked[legal] = logits[legal] probabilities = np.exp(masked - np.max(masked)) probabilities /= probabilities.sum() self.priors[key] = probabilities self.state_visits[key] = 0 return float(value[0]) def search(self, board: np.ndarray, player: int) -> float: result = winner(board) if result: return 1.0 if result == player else -1.0 if is_terminal(board): return 0.0 state_key = self.key(board, player) if state_key not in self.priors: return self.expand(board, player) best_score = -float("inf") best_action = -1 root_visits = max(1, self.state_visits[state_key]) for action in legal_actions(board): edge = (*state_key, int(action)) visits = self.edge_visits.get(edge, 0) value = self.edge_value.get(edge, 0.0) prior = self.priors[state_key][action] score = value + self.exploration * prior * math.sqrt(root_visits) / ( 1 + visits ) if score > best_score: best_score = score best_action = int(action) child = play(board, best_action, player) value = -self.search(child, -player) edge = (*state_key, best_action) visits = self.edge_visits.get(edge, 0) + 1 previous = self.edge_value.get(edge, 0.0) self.edge_visits[edge] = visits self.edge_value[edge] = previous + (value - previous) / visits self.state_visits[state_key] += 1 return value def policy( self, board: np.ndarray, player: int, *, temperature: float, add_noise: bool, rng: np.random.Generator, ) -> np.ndarray: key = self.key(board, player) if key not in self.priors: self.expand(board, player) if add_noise: legal = legal_actions(board) noise = rng.dirichlet(np.full(len(legal), 0.30)) priors = self.priors[key].copy() priors[legal] = 0.75 * priors[legal] + 0.25 * noise self.priors[key] = priors for _ in range(self.simulations): self.search(board, player) visits = np.zeros(9, dtype=np.float64) for action in legal_actions(board): visits[action] = self.edge_visits.get((*key, int(action)), 0) if temperature <= 1e-6: policy = np.zeros(9, dtype=np.float64) policy[int(np.argmax(visits))] = 1.0 return policy adjusted = visits ** (1 / temperature) if adjusted.sum() == 0: adjusted[legal_actions(board)] = 1 return adjusted / adjusted.sum()