microzero / source /game.py
ARotting's picture
Publish 7.6K parameter neural MCTS self-play agent
c80bb82 verified
Raw
History Blame Contribute Delete
2.34 kB
from __future__ import annotations
from functools import cache
import numpy as np
WIN_LINES = [
(0, 1, 2),
(3, 4, 5),
(6, 7, 8),
(0, 3, 6),
(1, 4, 7),
(2, 5, 8),
(0, 4, 8),
(2, 4, 6),
]
def winner(board: np.ndarray) -> int:
for left, middle, right in WIN_LINES:
value = int(board[left])
if value and value == board[middle] == board[right]:
return value
return 0
def is_terminal(board: np.ndarray) -> bool:
return winner(board) != 0 or not np.any(board == 0)
def legal_actions(board: np.ndarray) -> np.ndarray:
return np.flatnonzero(board == 0)
def play(board: np.ndarray, action: int, player: int) -> np.ndarray:
if board[action] != 0:
raise ValueError(f"Cell {action} is occupied.")
next_board = board.copy()
next_board[action] = player
return next_board
def canonical(board: np.ndarray, player: int) -> np.ndarray:
return board.astype(np.float32) * player
def symmetries(
state: np.ndarray,
policy: np.ndarray,
) -> list[tuple[np.ndarray, np.ndarray]]:
board = state.reshape(3, 3)
probabilities = policy.reshape(3, 3)
transformed = []
for rotations in range(4):
rotated_board = np.rot90(board, rotations)
rotated_policy = np.rot90(probabilities, rotations)
transformed.append(
(rotated_board.reshape(-1).copy(), rotated_policy.reshape(-1).copy())
)
transformed.append(
(
np.fliplr(rotated_board).reshape(-1).copy(),
np.fliplr(rotated_policy).reshape(-1).copy(),
)
)
return transformed
@cache
def minimax_value(board_key: tuple[int, ...], player: int) -> int:
board = np.asarray(board_key, dtype=np.int8)
result = winner(board)
if result:
return 1 if result == player else -1
legal = legal_actions(board)
if not len(legal):
return 0
return max(
-minimax_value(tuple(play(board, int(action), player)), -player)
for action in legal
)
def minimax_action(board: np.ndarray, player: int) -> int:
legal = legal_actions(board)
values = [
-minimax_value(tuple(play(board, int(action), player)), -player)
for action in legal
]
best = max(values)
return int(legal[values.index(best)])