Add multi-process self-play, eval workers, MCTS Cython port
Key changes for ISMCTS speed and correctness: - Cython port: HeuristicBot helpers (`heuristic_cy.pyx` + new `.pxd`) and ISMCTS searcher (`mcts.pyx`) now run as cdef. Both share a fast unified-action path through GameState's C interface to avoid Python round-trips on hot rollout/tree-walk paths. - Multi-process self-play and eval: `workers.py`, `eval_worker.py`, `interleaved_self_play.py`, plus trainer wiring with ProcessPoolExecutor + spawn context. Eval inside `evaluate.py` is parallel per opponent. - ISMCTS-specific eval (`evaluate.py`) runs MCTS at decision time so the metric matches deploy mode; `evaluation.eval_with_mcts` flag preserves backwards-compatible policy-only eval when needed. - Trainer logs progress per phase (self-play start/done, eval per opponent), and value loss is now scaled by `value_scale` so policy and value losses sit on comparable magnitudes. - Compact info-set key (`info_set.py`) using packed-struct format and child-key reuse during MCTS descent to cut per-step canonicalization. Tests: 19 ISMCTS suite passing, including parity (Cython-vs-Python sequential, batched-vs-sequential visit counts, push/pop round-trip). Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -24,6 +24,8 @@ mcts:
|
||||
n_simulations: 50
|
||||
c_puct: 1.5
|
||||
max_depth: 200
|
||||
parallel_simulations: 8
|
||||
virtual_loss_value: 1.0
|
||||
temperature:
|
||||
training: 1.0
|
||||
eval: 0.0
|
||||
@@ -32,6 +34,8 @@ training:
|
||||
gradient_steps_per_iter: 10
|
||||
batch_size: 128
|
||||
replay_capacity: 100000
|
||||
interleave_games: 8
|
||||
interleave_max_batch: 64
|
||||
optimization:
|
||||
learning_rate: 0.0003
|
||||
grad_clip: 5.0
|
||||
|
||||
@@ -24,6 +24,8 @@ mcts:
|
||||
n_simulations: 50
|
||||
c_puct: 1.5
|
||||
max_depth: 100
|
||||
parallel_simulations: 8
|
||||
virtual_loss_value: 1.0
|
||||
temperature:
|
||||
training: 1.0
|
||||
eval: 0.0
|
||||
@@ -32,6 +34,8 @@ training:
|
||||
gradient_steps_per_iter: 10
|
||||
batch_size: 128
|
||||
replay_capacity: 50000
|
||||
interleave_games: 8
|
||||
interleave_max_batch: 64
|
||||
optimization:
|
||||
learning_rate: 0.001
|
||||
grad_clip: 5.0
|
||||
|
||||
@@ -59,6 +59,14 @@ include = ["coolrl_lost_cities*"]
|
||||
"*.pxd",
|
||||
"*.pyx",
|
||||
]
|
||||
"coolrl_lost_cities.games.classic.bots" = [
|
||||
"*.pxd",
|
||||
"*.pyx",
|
||||
]
|
||||
"coolrl_lost_cities.games.classic.ismcts" = [
|
||||
"*.pxd",
|
||||
"*.pyx",
|
||||
]
|
||||
|
||||
[tool.ruff]
|
||||
line-length = 100
|
||||
|
||||
@@ -30,6 +30,10 @@ extensions = cythonize(
|
||||
"coolrl_lost_cities.games.classic.bots.heuristic_cy",
|
||||
["src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pyx"],
|
||||
),
|
||||
Extension(
|
||||
"coolrl_lost_cities.games.classic.ismcts.mcts",
|
||||
["src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx"],
|
||||
),
|
||||
],
|
||||
language_level=3,
|
||||
compiler_directives={
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
from ..game cimport GameState
|
||||
|
||||
|
||||
cdef class _CachedState:
|
||||
cdef GameState _state
|
||||
cdef public object config
|
||||
cdef public object hands
|
||||
cdef public object expeditions
|
||||
cdef public object discards
|
||||
cdef public object deck
|
||||
cdef int hand_encoded[2][16]
|
||||
cdef int hand_size[2]
|
||||
cdef int expedition_top[2][8]
|
||||
cdef int expedition_count[2][8]
|
||||
cdef int expedition_handshakes[2][8]
|
||||
cdef int expedition_numeric_sum[2][8]
|
||||
cdef int expedition_last_numeric[2][8]
|
||||
cdef int discard_top[8]
|
||||
cdef int discard_count[8]
|
||||
cdef int deck_remaining
|
||||
cdef int total_scores[2]
|
||||
cdef int current_player
|
||||
cdef int phase
|
||||
cdef int turn_count
|
||||
cdef int n_colors
|
||||
cdef int n_ranks
|
||||
cdef int min_rank
|
||||
cdef int hand_capacity
|
||||
cdef int bonus_threshold
|
||||
cdef int bonus_amount
|
||||
cdef int expedition_penalty
|
||||
cdef void _build(self, GameState state) except *
|
||||
cpdef list legal_card_mask(self)
|
||||
cpdef list legal_draw_mask(self)
|
||||
cpdef bint can_play_card(self, int player, object card)
|
||||
cpdef bint can_play_encoded(self, int player, int card)
|
||||
cpdef bint has_numeric(self, int player, int color)
|
||||
cpdef int score_diff(self, int player)
|
||||
|
||||
|
||||
cdef class HeuristicBot:
|
||||
cdef public object params
|
||||
cdef double color_commit_cache[2][8]
|
||||
cdef unsigned char color_commit_valid[2][8]
|
||||
cdef signed char playability_cache[2][8][17]
|
||||
cdef void _reset_caches(self) noexcept
|
||||
cpdef int act_cython(self, GameState state) except -1
|
||||
cdef int _card_color_c(self, _CachedState state, int card) noexcept
|
||||
cdef int _card_rank_c(self, _CachedState state, int card) noexcept
|
||||
cdef int _num_c(self, _CachedState state, int card) noexcept
|
||||
cdef int _play_action_c(self, int slot) noexcept
|
||||
cdef int _discard_action_c(self, int slot) noexcept
|
||||
cdef bint _legal_card_action_c(self, _CachedState state, int action) noexcept
|
||||
cdef bint _legal_draw_action_c(self, _CachedState state, int action) noexcept
|
||||
cdef bint _can_play_card_c(self, _CachedState state, int player, int card) noexcept
|
||||
cdef bint _has_numeric_c(self, _CachedState state, int player, int color) noexcept
|
||||
cdef int _opened_colors_c(self, _CachedState state, int player) noexcept
|
||||
cdef int _first_legal_card_c(self, _CachedState state) noexcept
|
||||
cdef int _first_legal_draw_c(self, _CachedState state) noexcept
|
||||
cdef int _act_card_c(self, _CachedState state, object derived) except -1
|
||||
cdef int _act_draw_c(self, _CachedState state, object derived) except -1
|
||||
cdef int _best_handshake_play_c(
|
||||
self, _CachedState state, int player, object derived, int deck_left
|
||||
) except -2
|
||||
cdef int _best_number_play_c(
|
||||
self, _CachedState state, int player, object derived, int deck_left
|
||||
) except -2
|
||||
cdef double _started_expedition_play_value_c(
|
||||
self, _CachedState state, int player, int card, object derived, int deck_left
|
||||
) except *
|
||||
cdef bint _should_open_expedition_c(
|
||||
self, _CachedState state, int player, int color, int opening_card, object derived, int deck_left
|
||||
) except *
|
||||
cdef double _opening_plan_value_c(
|
||||
self, _CachedState state, int player, int color, int opening_card, object derived, int deck_left
|
||||
) except *
|
||||
cdef double _open_expedition_value_c(
|
||||
self, _CachedState state, int player, int color, int opening_card, object derived, int deck_left
|
||||
) except *
|
||||
cdef int _best_forced_open_c(
|
||||
self, _CachedState state, int player, object derived, int deck_left
|
||||
) except -2
|
||||
cdef int _best_discard_c(self, _CachedState state, int player, object derived) except -2
|
||||
cdef double _visible_draw_value_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *
|
||||
cdef double _visible_open_support_value_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *
|
||||
cdef bint _visible_number_can_help_open_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *
|
||||
cdef double _deck_draw_value_c(self, _CachedState state, object derived) except *
|
||||
cdef double _card_value_for_me_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *
|
||||
cdef double _card_value_for_opponent_c(
|
||||
self, _CachedState state, int opponent, int card, object derived
|
||||
) except *
|
||||
cdef double _color_commitment_c(
|
||||
self, _CachedState state, int player, int color, object derived
|
||||
) except *
|
||||
cdef double _public_color_commitment_for_opponent_c(
|
||||
self, _CachedState state, int opponent, int color, object derived
|
||||
) except *
|
||||
cdef double _bonus_potential_c(
|
||||
self,
|
||||
_CachedState state,
|
||||
int player,
|
||||
int color,
|
||||
int extra_cards,
|
||||
object derived,
|
||||
int committed_cards,
|
||||
int exclude_card,
|
||||
) except *
|
||||
cdef double _new_color_open_penalty_c(self, int opened_colors) noexcept
|
||||
cdef double _late_penalty_c(self, object derived, int deck_left) except *
|
||||
@@ -5,6 +5,10 @@ import logging
|
||||
from dataclasses import dataclass
|
||||
from functools import lru_cache
|
||||
|
||||
from libc.string cimport memset
|
||||
|
||||
from ..game cimport GameState as CGameState
|
||||
|
||||
from ..game import Card, GameState, LostCitiesConfig
|
||||
from ..policy import LostCitiesPolicy, PolicyInput
|
||||
from .base import first_legal, legal_from_obs
|
||||
@@ -18,6 +22,11 @@ except ImportError as exc: # pragma: no cover
|
||||
PLAY_OR_DISCARD_ACTIONS_PER_SLOT = 2
|
||||
DRAW_FROM_DECK_ACTION = 0
|
||||
|
||||
DEF MAX_HAND_SIZE = 16
|
||||
DEF MAX_COLORS = 8
|
||||
DEF MAX_RANKS = 16
|
||||
DEF MAX_ACTIONS = 64
|
||||
|
||||
|
||||
def play_action(slot: int) -> int:
|
||||
return PLAY_OR_DISCARD_ACTIONS_PER_SLOT * slot
|
||||
@@ -34,35 +43,87 @@ def draw_from_discard_action(color: int) -> int:
|
||||
LOGGER = logging.getLogger("coolrl_lost_cities.games.classic.bots.heuristic")
|
||||
|
||||
|
||||
class _CachedState:
|
||||
def __init__(self, state):
|
||||
self._state = state
|
||||
cdef class _CachedState:
|
||||
def __init__(self, CGameState state):
|
||||
self._build(state)
|
||||
self.config = state.config
|
||||
self.current_player = state.current_player
|
||||
self.phase = state.phase
|
||||
self.turn_count = state.turn_count
|
||||
self.hands = state.hands
|
||||
self.expeditions = state.expeditions
|
||||
self.discards = state.discards
|
||||
self.deck = state.deck
|
||||
|
||||
def legal_card_mask(self):
|
||||
cdef void _build(self, CGameState state) except *:
|
||||
cdef int player
|
||||
cdef int slot
|
||||
cdef int color
|
||||
cdef int idx
|
||||
cdef int length
|
||||
if state.hand_size > MAX_HAND_SIZE:
|
||||
raise ValueError("hand_size exceeds HeuristicBot fixed hand buffer")
|
||||
if state.n_colors > MAX_COLORS:
|
||||
raise ValueError("n_colors exceeds HeuristicBot fixed color buffer")
|
||||
if state.n_ranks > MAX_RANKS:
|
||||
raise ValueError("n_ranks exceeds HeuristicBot fixed rank buffer")
|
||||
self._state = state
|
||||
self.n_colors = state.n_colors
|
||||
self.n_ranks = state.n_ranks
|
||||
self.min_rank = state.min_rank
|
||||
self.hand_capacity = state.hand_size
|
||||
self.bonus_threshold = state.bonus_threshold
|
||||
self.bonus_amount = state.bonus_amount
|
||||
self.expedition_penalty = state.expedition_penalty
|
||||
self.current_player = state.current_player
|
||||
self.phase = state.phase_id
|
||||
self.turn_count = state.turn_count
|
||||
self.deck_remaining = state.deck_len
|
||||
self.total_scores[0] = state.total_scores[0]
|
||||
self.total_scores[1] = state.total_scores[1]
|
||||
memset(&self.hand_encoded[0][0], 0, sizeof(self.hand_encoded))
|
||||
memset(&self.expedition_top[0][0], 0, sizeof(self.expedition_top))
|
||||
memset(&self.expedition_count[0][0], 0, sizeof(self.expedition_count))
|
||||
memset(&self.expedition_handshakes[0][0], 0, sizeof(self.expedition_handshakes))
|
||||
memset(&self.expedition_numeric_sum[0][0], 0, sizeof(self.expedition_numeric_sum))
|
||||
memset(&self.expedition_last_numeric[0][0], 0, sizeof(self.expedition_last_numeric))
|
||||
memset(&self.discard_top[0], 0, sizeof(self.discard_top))
|
||||
memset(&self.discard_count[0], 0, sizeof(self.discard_count))
|
||||
for player in range(2):
|
||||
self.hand_size[player] = state.hand_lens[player]
|
||||
for slot in range(state.hand_lens[player]):
|
||||
self.hand_encoded[player][slot] = state.hand_cards[state._hand_index(player, slot)]
|
||||
for color in range(state.n_colors):
|
||||
idx = state._expedition_len_index(player, color)
|
||||
length = state.expedition_lens[idx]
|
||||
self.expedition_count[player][color] = length
|
||||
self.expedition_handshakes[player][color] = state.handshake_counts[idx]
|
||||
self.expedition_numeric_sum[player][color] = state.numeric_sums[idx]
|
||||
self.expedition_last_numeric[player][color] = state.last_numeric_ranks[idx]
|
||||
if length > 0:
|
||||
self.expedition_top[player][color] = state.expedition_cards[
|
||||
state._expedition_index(player, color, length - 1)
|
||||
]
|
||||
for color in range(state.n_colors):
|
||||
length = state.discard_lens[color]
|
||||
self.discard_count[color] = length
|
||||
if length > 0:
|
||||
self.discard_top[color] = state.discard_cards[state._discard_index(color, length - 1)]
|
||||
|
||||
cpdef list legal_card_mask(self):
|
||||
return self._state.legal_card_mask()
|
||||
|
||||
def legal_draw_mask(self):
|
||||
cpdef list legal_draw_mask(self):
|
||||
return self._state.legal_draw_mask()
|
||||
|
||||
def can_play_card(self, player, card):
|
||||
cpdef bint can_play_card(self, int player, object card):
|
||||
return self._state.can_play_card(player, card)
|
||||
|
||||
def has_numeric(self, player, color):
|
||||
return self._state.has_numeric(player, color)
|
||||
cpdef bint can_play_encoded(self, int player, int card):
|
||||
return self._state._can_play_encoded_card_c(player, card)
|
||||
|
||||
def score_diff(self, player):
|
||||
return self._state.score_diff(player)
|
||||
cpdef bint has_numeric(self, int player, int color):
|
||||
return self._state.last_numeric_rank(player, color) > 0
|
||||
|
||||
def __getattr__(self, name):
|
||||
return getattr(self._state, name)
|
||||
cpdef int score_diff(self, int player):
|
||||
return self.total_scores[player] - self.total_scores[1 - player]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -178,29 +239,756 @@ def derive_heuristic_config(
|
||||
)
|
||||
|
||||
|
||||
class HeuristicBot(LostCitiesPolicy):
|
||||
cdef class HeuristicBot:
|
||||
def __init__(self, params: HeuristicParams | None = None):
|
||||
self.params = params or HeuristicParams()
|
||||
self._reset_caches()
|
||||
|
||||
cdef void _reset_caches(self) noexcept:
|
||||
memset(&self.color_commit_cache[0][0], 0, sizeof(self.color_commit_cache))
|
||||
memset(&self.color_commit_valid[0][0], 0, sizeof(self.color_commit_valid))
|
||||
memset(&self.playability_cache[0][0][0], -1, sizeof(self.playability_cache))
|
||||
|
||||
cpdef int act_cython(self, CGameState state) except -1:
|
||||
cdef _CachedState cstate = _CachedState(state)
|
||||
cdef object derived
|
||||
self._reset_caches()
|
||||
derived = self._derived(state)
|
||||
if cstate.phase == 0:
|
||||
return self._act_card_c(cstate, derived)
|
||||
return self._act_draw_c(cstate, derived)
|
||||
|
||||
def act(self, obs_or_state: PolicyInput) -> int:
|
||||
if not isinstance(obs_or_state, GameState) and not hasattr(obs_or_state, "legal_mask"):
|
||||
LOGGER.debug(
|
||||
"HeuristicBot fallback to first legal: input_type=%s",
|
||||
type(obs_or_state).__name__,
|
||||
)
|
||||
return first_legal(legal_from_obs(obs_or_state))
|
||||
|
||||
LOGGER.debug(
|
||||
"HeuristicBot heuristic path: player=%s phase=%s turn=%s",
|
||||
obs_or_state.current_player,
|
||||
obs_or_state.phase,
|
||||
obs_or_state.turn_count,
|
||||
)
|
||||
state = _CachedState(obs_or_state)
|
||||
if state.phase == "card":
|
||||
return self._act_card(state)
|
||||
if isinstance(obs_or_state, GameState):
|
||||
return self.act_cython(obs_or_state)
|
||||
|
||||
return self._act_draw(state)
|
||||
if obs_or_state.phase == "card":
|
||||
return self._act_card(obs_or_state)
|
||||
|
||||
return self._act_draw(obs_or_state)
|
||||
|
||||
cdef inline int _card_color_c(self, _CachedState state, int card) noexcept:
|
||||
return card // (state.n_ranks + 1)
|
||||
|
||||
cdef inline int _card_rank_c(self, _CachedState state, int card) noexcept:
|
||||
return card % (state.n_ranks + 1)
|
||||
|
||||
cdef inline int _num_c(self, _CachedState state, int card) noexcept:
|
||||
cdef int rank = self._card_rank_c(state, card)
|
||||
if rank == 0:
|
||||
return 0
|
||||
return state.min_rank + rank - 1
|
||||
|
||||
cdef inline int _play_action_c(self, int slot) noexcept:
|
||||
return 2 * slot
|
||||
|
||||
cdef inline int _discard_action_c(self, int slot) noexcept:
|
||||
return 2 * slot + 1
|
||||
|
||||
cdef inline bint _legal_card_action_c(self, _CachedState state, int action) noexcept:
|
||||
cdef int slot = action // 2
|
||||
if action < 0 or action >= 2 * state.hand_capacity:
|
||||
return False
|
||||
if slot >= state.hand_size[state.current_player]:
|
||||
return False
|
||||
if action % 2 == 1:
|
||||
return True
|
||||
return self._can_play_card_c(
|
||||
state,
|
||||
state.current_player,
|
||||
state.hand_encoded[state.current_player][slot],
|
||||
)
|
||||
|
||||
cdef inline bint _legal_draw_action_c(self, _CachedState state, int action) noexcept:
|
||||
cdef int color
|
||||
if action == 0:
|
||||
return state.deck_remaining > 0
|
||||
color = action - 1
|
||||
if color < 0 or color >= state.n_colors:
|
||||
return False
|
||||
return (
|
||||
state.discard_count[color] > 0
|
||||
and (state._state.pending_discarded_color < 0 or color != state._state.pending_discarded_color)
|
||||
)
|
||||
|
||||
cdef bint _can_play_card_c(self, _CachedState state, int player, int card) noexcept:
|
||||
cdef int color = self._card_color_c(state, card)
|
||||
cdef int rank = self._card_rank_c(state, card)
|
||||
cdef signed char cached
|
||||
if color < 0 or color >= state.n_colors or rank < 0 or rank > state.n_ranks:
|
||||
return False
|
||||
cached = self.playability_cache[player][color][rank]
|
||||
if cached >= 0:
|
||||
return cached == 1
|
||||
if rank == 0:
|
||||
cached = 1 if state.expedition_last_numeric[player][color] == 0 else 0
|
||||
else:
|
||||
cached = 1 if rank > state.expedition_last_numeric[player][color] else 0
|
||||
self.playability_cache[player][color][rank] = cached
|
||||
return cached == 1
|
||||
|
||||
cdef bint _has_numeric_c(self, _CachedState state, int player, int color) noexcept:
|
||||
return state.expedition_last_numeric[player][color] > 0
|
||||
|
||||
cdef int _opened_colors_c(self, _CachedState state, int player) noexcept:
|
||||
cdef int color
|
||||
cdef int count = 0
|
||||
for color in range(state.n_colors):
|
||||
if state.expedition_count[player][color] > 0:
|
||||
count += 1
|
||||
return count
|
||||
|
||||
cdef int _first_legal_card_c(self, _CachedState state) noexcept:
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int player = state.current_player
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if self._can_play_card_c(state, player, card):
|
||||
return 2 * slot
|
||||
return 2 * slot + 1
|
||||
return 0
|
||||
|
||||
cdef int _first_legal_draw_c(self, _CachedState state) noexcept:
|
||||
cdef int color
|
||||
if self._legal_draw_action_c(state, 0):
|
||||
return 0
|
||||
for color in range(state.n_colors):
|
||||
if self._legal_draw_action_c(state, 1 + color):
|
||||
return 1 + color
|
||||
return 0
|
||||
|
||||
cdef int _act_card_c(self, _CachedState state, object derived) except -1:
|
||||
cdef int player = state.current_player
|
||||
cdef int action
|
||||
action = self._best_handshake_play_c(state, player, derived, state.deck_remaining)
|
||||
if action >= 0:
|
||||
return action
|
||||
action = self._best_number_play_c(state, player, derived, state.deck_remaining)
|
||||
if action >= 0:
|
||||
return action
|
||||
if self._opened_colors_c(state, player) == 0:
|
||||
action = self._best_forced_open_c(state, player, derived, state.deck_remaining)
|
||||
if action >= 0:
|
||||
return action
|
||||
action = self._best_discard_c(state, player, derived)
|
||||
if action >= 0:
|
||||
return action
|
||||
return self._first_legal_card_c(state)
|
||||
|
||||
cdef int _act_draw_c(self, _CachedState state, object derived) except -1:
|
||||
cdef int player = state.current_player
|
||||
cdef int color
|
||||
cdef int action
|
||||
cdef int best_action = -1
|
||||
cdef int best_tie = -1
|
||||
cdef double value
|
||||
cdef double best_value = -1.0e100
|
||||
if self._legal_draw_action_c(state, 0):
|
||||
best_action = 0
|
||||
best_tie = 1
|
||||
best_value = self._deck_draw_value_c(state, derived)
|
||||
for color in range(state.n_colors):
|
||||
action = 1 + color
|
||||
if not self._legal_draw_action_c(state, action):
|
||||
continue
|
||||
value = self._visible_draw_value_c(state, player, state.discard_top[color], derived)
|
||||
if (
|
||||
value > best_value
|
||||
or (
|
||||
value == best_value
|
||||
and (0 > best_tie or (best_tie == 0 and action > best_action))
|
||||
)
|
||||
):
|
||||
best_value = value
|
||||
best_tie = 0
|
||||
best_action = action
|
||||
if best_action >= 0:
|
||||
return best_action
|
||||
return self._first_legal_draw_c(state)
|
||||
|
||||
cdef int _best_handshake_play_c(
|
||||
self, _CachedState state, int player, object derived, int deck_left
|
||||
) except -2:
|
||||
cdef int slot
|
||||
cdef int other_slot
|
||||
cdef int card
|
||||
cdef int other
|
||||
cdef int color
|
||||
cdef int number_count
|
||||
cdef int number_sum
|
||||
cdef double value
|
||||
cdef double best_value = -1.0e100
|
||||
cdef int best_action = -1
|
||||
if state._state.n_handshakes <= 0:
|
||||
return -1
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if not self._legal_card_action_c(state, 2 * slot) or self._card_rank_c(state, card) != 0:
|
||||
continue
|
||||
color = self._card_color_c(state, card)
|
||||
if state.expedition_last_numeric[player][color] > 0:
|
||||
continue
|
||||
number_count = 0
|
||||
number_sum = 0
|
||||
for other_slot in range(state.hand_size[player]):
|
||||
if other_slot == slot:
|
||||
continue
|
||||
other = state.hand_encoded[player][other_slot]
|
||||
if (
|
||||
self._card_color_c(state, other) == color
|
||||
and self._card_rank_c(state, other) != 0
|
||||
and self._can_play_card_c(state, player, other)
|
||||
):
|
||||
number_count += 1
|
||||
number_sum += self._num_c(state, other)
|
||||
if number_count < derived.min_handshake_numeric_cards:
|
||||
continue
|
||||
if number_sum < derived.open_target_sum * self.params.handshake_target_multiplier:
|
||||
continue
|
||||
if deck_left <= derived.late_open_block_threshold:
|
||||
continue
|
||||
value = number_sum + 2.0 * number_count
|
||||
value += self._bonus_potential_c(state, player, color, 0, derived, 1, card)
|
||||
value -= self._late_penalty_c(derived, deck_left)
|
||||
if value > best_value or (value == best_value and 2 * slot > best_action):
|
||||
best_value = value
|
||||
best_action = 2 * slot
|
||||
return best_action
|
||||
|
||||
cdef int _best_number_play_c(
|
||||
self, _CachedState state, int player, object derived, int deck_left
|
||||
) except -2:
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int color
|
||||
cdef double value
|
||||
cdef double best_value = -1.0e100
|
||||
cdef int best_action = -1
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if not self._legal_card_action_c(state, 2 * slot) or self._card_rank_c(state, card) == 0:
|
||||
continue
|
||||
color = self._card_color_c(state, card)
|
||||
if state.expedition_count[player][color] > 0:
|
||||
value = self._started_expedition_play_value_c(
|
||||
state, player, card, derived, deck_left
|
||||
)
|
||||
if value > best_value or (value == best_value and 2 * slot > best_action):
|
||||
best_value = value
|
||||
best_action = 2 * slot
|
||||
continue
|
||||
if self._should_open_expedition_c(state, player, color, card, derived, deck_left):
|
||||
value = self._open_expedition_value_c(state, player, color, card, derived, deck_left)
|
||||
if value > best_value or (value == best_value and 2 * slot > best_action):
|
||||
best_value = value
|
||||
best_action = 2 * slot
|
||||
return best_action
|
||||
|
||||
cdef double _started_expedition_play_value_c(
|
||||
self, _CachedState state, int player, int card, object derived, int deck_left
|
||||
) except *:
|
||||
cdef int color = self._card_color_c(state, card)
|
||||
cdef int numeric_value = self._num_c(state, card)
|
||||
cdef int slot
|
||||
cdef int followup
|
||||
cdef int projected_sum = state.expedition_numeric_sum[player][color] + numeric_value
|
||||
cdef double value = 0.0
|
||||
for slot in range(state.hand_size[player]):
|
||||
followup = state.hand_encoded[player][slot]
|
||||
if (
|
||||
followup != card
|
||||
and self._card_color_c(state, followup) == color
|
||||
and self._card_rank_c(state, followup) != 0
|
||||
and self._card_rank_c(state, followup) > self._card_rank_c(state, card)
|
||||
):
|
||||
projected_sum += self._num_c(state, followup)
|
||||
value += self.params.started_expedition_play_bonus
|
||||
value += self.params.started_expedition_followup_bonus
|
||||
value += <double>(state._state.min_rank + state._state.n_ranks - numeric_value)
|
||||
if deck_left <= derived.late_deck_threshold:
|
||||
value += 2.0 * numeric_value
|
||||
elif deck_left <= derived.mid_deck_threshold:
|
||||
value += 0.8 * numeric_value
|
||||
if projected_sum < derived.open_target_sum:
|
||||
value -= 6.0
|
||||
value += 3.0 * state.expedition_handshakes[player][color]
|
||||
value += self._bonus_potential_c(state, player, color, 0, derived, 1, card)
|
||||
return value
|
||||
|
||||
cdef bint _should_open_expedition_c(
|
||||
self, _CachedState state, int player, int color, int opening_card, object derived, int deck_left
|
||||
) except *:
|
||||
if deck_left <= derived.late_open_block_threshold:
|
||||
return False
|
||||
return self._opening_plan_value_c(state, player, color, opening_card, derived, deck_left) > 0.0
|
||||
|
||||
cdef double _opening_plan_value_c(
|
||||
self, _CachedState state, int player, int color, int opening_card, object derived, int deck_left
|
||||
) except *:
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int rank
|
||||
cdef int numbers = 0
|
||||
cdef int handshakes = 0
|
||||
cdef int number_sum = 0
|
||||
cdef int high_count = 0
|
||||
cdef int opened_colors = self._opened_colors_c(state, player)
|
||||
cdef int opening_value = self._num_c(state, opening_card)
|
||||
cdef double new_color_penalty = self._new_color_open_penalty_c(opened_colors)
|
||||
cdef bint strong_open
|
||||
cdef bint speculative_open
|
||||
cdef bint single_late_open
|
||||
cdef bint exceptional_open
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if self._card_color_c(state, card) != color:
|
||||
continue
|
||||
rank = self._card_rank_c(state, card)
|
||||
if rank == 0:
|
||||
handshakes += 1
|
||||
elif rank >= self._card_rank_c(state, opening_card):
|
||||
numbers += 1
|
||||
number_sum += self._num_c(state, card)
|
||||
if rank >= derived.middle_rank:
|
||||
high_count += 1
|
||||
strong_open = (
|
||||
numbers >= derived.min_open_cards
|
||||
and number_sum >= derived.open_target_sum
|
||||
and (high_count > 0 or number_sum >= 0.85 * derived.max_color_sum)
|
||||
)
|
||||
speculative_open = (
|
||||
opened_colors <= 2
|
||||
and numbers >= 2
|
||||
and number_sum >= 0.65 * derived.open_target_sum
|
||||
and high_count > 0
|
||||
)
|
||||
single_late_open = (
|
||||
deck_left <= derived.mid_deck_threshold and numbers >= 1 and opening_value >= 8
|
||||
)
|
||||
exceptional_open = (
|
||||
numbers >= derived.min_open_cards + 1
|
||||
and number_sum >= max(float(derived.break_even_sum), derived.open_target_sum * 1.4)
|
||||
and high_count >= 2
|
||||
and deck_left > derived.mid_deck_threshold
|
||||
)
|
||||
if opened_colors == 3:
|
||||
speculative_open = False
|
||||
if opened_colors >= 4:
|
||||
strong_open = False
|
||||
speculative_open = False
|
||||
single_late_open = False
|
||||
if strong_open:
|
||||
return 6.0 + 0.25 * number_sum + 0.8 * numbers + 0.5 * handshakes - new_color_penalty
|
||||
if speculative_open:
|
||||
return 3.0 + 0.18 * number_sum + 0.7 * numbers + 0.4 * handshakes - new_color_penalty
|
||||
if opened_colors == 3:
|
||||
return 0.0
|
||||
if exceptional_open:
|
||||
return 10.0 + 0.3 * number_sum + 1.0 * numbers + 0.7 * high_count - new_color_penalty
|
||||
if single_late_open:
|
||||
return 1.5 + 0.2 * opening_value - new_color_penalty
|
||||
return 0.0
|
||||
|
||||
cdef double _open_expedition_value_c(
|
||||
self, _CachedState state, int player, int color, int opening_card, object derived, int deck_left
|
||||
) except *:
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int rank
|
||||
cdef int numbers = 0
|
||||
cdef int handshakes = 0
|
||||
cdef int number_sum = 0
|
||||
cdef double value
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if self._card_color_c(state, card) != color:
|
||||
continue
|
||||
rank = self._card_rank_c(state, card)
|
||||
if rank == 0:
|
||||
handshakes += 1
|
||||
elif rank >= self._card_rank_c(state, opening_card):
|
||||
numbers += 1
|
||||
number_sum += self._num_c(state, card)
|
||||
value = number_sum + 2.0 * numbers + 1.5 * handshakes
|
||||
value += self._opening_plan_value_c(state, player, color, opening_card, derived, deck_left)
|
||||
value += state.expedition_penalty
|
||||
value += max(0.0, float(derived.middle_rank - self._card_rank_c(state, opening_card)))
|
||||
value += self._bonus_potential_c(state, player, color, 0, derived, 1, opening_card)
|
||||
value -= self._late_penalty_c(derived, deck_left)
|
||||
return value
|
||||
|
||||
cdef int _best_forced_open_c(
|
||||
self, _CachedState state, int player, object derived, int deck_left
|
||||
) except -2:
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int color
|
||||
cdef int other
|
||||
cdef int number_count
|
||||
cdef int number_sum
|
||||
cdef double opening_value
|
||||
cdef double forced_value
|
||||
cdef double best_value = -1.0e100
|
||||
cdef int best_action = -1
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if not self._legal_card_action_c(state, 2 * slot) or self._card_rank_c(state, card) == 0:
|
||||
continue
|
||||
color = self._card_color_c(state, card)
|
||||
if state.expedition_count[player][color] > 0:
|
||||
continue
|
||||
opening_value = self._opening_plan_value_c(state, player, color, card, derived, deck_left)
|
||||
number_count = 0
|
||||
number_sum = 0
|
||||
for other_slot in range(state.hand_size[player]):
|
||||
other = state.hand_encoded[player][other_slot]
|
||||
if self._card_color_c(state, other) == color and self._card_rank_c(state, other) != 0:
|
||||
number_count += 1
|
||||
number_sum += self._num_c(state, other)
|
||||
if (
|
||||
opening_value <= 0.0
|
||||
and number_count < 2
|
||||
and number_sum < 0.5 * derived.open_target_sum
|
||||
and deck_left > derived.mid_deck_threshold
|
||||
):
|
||||
continue
|
||||
forced_value = opening_value + 0.2 * number_sum
|
||||
forced_value += <double>(state._state.min_rank + state._state.n_ranks - self._num_c(state, card))
|
||||
if forced_value > best_value or (forced_value == best_value and 2 * slot > best_action):
|
||||
best_value = forced_value
|
||||
best_action = 2 * slot
|
||||
return best_action
|
||||
|
||||
cdef int _best_discard_c(self, _CachedState state, int player, object derived) except -2:
|
||||
cdef int opponent = 1 - player
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef double my_value
|
||||
cdef double opponent_value
|
||||
cdef double score
|
||||
cdef double best_score = -1.0e100
|
||||
cdef int best_action = -1
|
||||
for slot in range(state.hand_size[player]):
|
||||
if not self._legal_card_action_c(state, 2 * slot + 1):
|
||||
continue
|
||||
card = state.hand_encoded[player][slot]
|
||||
my_value = self._card_value_for_me_c(state, player, card, derived)
|
||||
opponent_value = self._card_value_for_opponent_c(state, opponent, card, derived)
|
||||
score = -my_value - self.params.gift_penalty_weight * opponent_value
|
||||
if not self._can_play_card_c(state, player, card):
|
||||
score += self.params.unusable_discard_bonus
|
||||
if not self._can_play_card_c(state, opponent, card):
|
||||
score += self.params.discard_safety_bonus
|
||||
if self._card_rank_c(state, card) == 0 and self._can_play_card_c(state, player, card):
|
||||
score -= 4.0
|
||||
if score > best_score or (score == best_score and 2 * slot + 1 > best_action):
|
||||
best_score = score
|
||||
best_action = 2 * slot + 1
|
||||
return best_action
|
||||
|
||||
cdef double _visible_draw_value_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *:
|
||||
cdef int color = self._card_color_c(state, card)
|
||||
cdef int opponent = 1 - player
|
||||
cdef int opened_colors = self._opened_colors_c(state, player)
|
||||
cdef bint is_unopened_color = state.expedition_count[player][color] == 0
|
||||
cdef double commitment = self._color_commitment_c(state, player, color, derived)
|
||||
cdef double opponent_value = self._card_value_for_opponent_c(state, opponent, card, derived)
|
||||
cdef int score_diff = state.total_scores[player] - state.total_scores[opponent]
|
||||
cdef double value = self.params.deny_opponent_weight * opponent_value
|
||||
cdef double support
|
||||
cdef bint exceptional_support = False
|
||||
cdef int slot
|
||||
cdef int other
|
||||
cdef int number_count
|
||||
cdef int number_sum
|
||||
cdef double required_sum
|
||||
if score_diff <= 0:
|
||||
value += self.params.losing_visible_draw_bonus
|
||||
if is_unopened_color:
|
||||
if opened_colors >= 4:
|
||||
value -= self.params.unopened_draw_penalty_four_open
|
||||
elif opened_colors >= 3:
|
||||
value -= self.params.unopened_draw_penalty_three_open
|
||||
if self._card_rank_c(state, card) == 0:
|
||||
if self._has_numeric_c(state, player, color):
|
||||
return value - self.params.dead_visible_draw_penalty
|
||||
if state.expedition_count[player][color] == 0:
|
||||
number_count = 0
|
||||
number_sum = 0
|
||||
for slot in range(state.hand_size[player]):
|
||||
other = state.hand_encoded[player][slot]
|
||||
if (
|
||||
self._card_color_c(state, other) == color
|
||||
and self._card_rank_c(state, other) != 0
|
||||
and self._can_play_card_c(state, player, other)
|
||||
):
|
||||
number_count += 1
|
||||
number_sum += self._num_c(state, other)
|
||||
required_sum = derived.open_target_sum * self.params.handshake_target_multiplier
|
||||
if number_count < derived.min_handshake_numeric_cards or number_sum < required_sum:
|
||||
support = self._visible_open_support_value_c(state, player, card, derived)
|
||||
exceptional_support = support >= 6.0
|
||||
if (
|
||||
is_unopened_color
|
||||
and opened_colors >= 4
|
||||
and not exceptional_support
|
||||
and opponent_value < self.params.strong_deny_threshold
|
||||
and score_diff > -15
|
||||
):
|
||||
return -8.0
|
||||
return value + support - 0.5
|
||||
return value + 6.0 + commitment
|
||||
if self._can_play_card_c(state, player, card):
|
||||
value += self._num_c(state, card)
|
||||
value += 0.7 * commitment
|
||||
if state.expedition_count[player][color] > 0:
|
||||
value += 5.0
|
||||
else:
|
||||
support = self._visible_open_support_value_c(state, player, card, derived)
|
||||
exceptional_support = support >= 6.0
|
||||
value += support
|
||||
else:
|
||||
value -= self.params.dead_visible_draw_penalty
|
||||
if state.expedition_count[player][color] == 0:
|
||||
support = self._visible_open_support_value_c(state, player, card, derived)
|
||||
exceptional_support = support >= 6.0
|
||||
value += support
|
||||
if (
|
||||
is_unopened_color
|
||||
and opened_colors >= 4
|
||||
and not exceptional_support
|
||||
and opponent_value < self.params.strong_deny_threshold
|
||||
and score_diff > -15
|
||||
):
|
||||
return -8.0
|
||||
value += self._bonus_potential_c(state, player, color, 1, derived, 0, -1)
|
||||
return value
|
||||
|
||||
cdef double _visible_open_support_value_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *:
|
||||
cdef int color = self._card_color_c(state, card)
|
||||
cdef int opened_colors = self._opened_colors_c(state, player)
|
||||
cdef int slot
|
||||
cdef int other
|
||||
cdef int rank
|
||||
cdef int future_numbers = 0
|
||||
cdef int same_color_handshakes = 0
|
||||
cdef double value = 0.0
|
||||
for slot in range(state.hand_size[player]):
|
||||
other = state.hand_encoded[player][slot]
|
||||
if self._card_color_c(state, other) != color:
|
||||
continue
|
||||
rank = self._card_rank_c(state, other)
|
||||
if rank == 0:
|
||||
same_color_handshakes += 1
|
||||
elif other != card and rank >= self._card_rank_c(state, card):
|
||||
future_numbers += 1
|
||||
value += 0.8 * future_numbers
|
||||
value += 1.0 * same_color_handshakes
|
||||
if self._card_rank_c(state, card) <= derived.middle_rank:
|
||||
value += self.params.speculative_visible_draw_bonus
|
||||
if self._visible_number_can_help_open_c(state, player, card, derived):
|
||||
value += 4.0
|
||||
elif opened_colors <= 2 and (future_numbers > 0 or same_color_handshakes > 0):
|
||||
value += self.params.speculative_visible_draw_bonus
|
||||
if opened_colors <= 2:
|
||||
value += 0.25 * self._opening_plan_value_c(
|
||||
state, player, color, card, derived, state.deck_remaining
|
||||
)
|
||||
elif opened_colors == 3:
|
||||
value += 0.1 * max(
|
||||
0.0,
|
||||
self._opening_plan_value_c(state, player, color, card, derived, state.deck_remaining),
|
||||
)
|
||||
return value
|
||||
|
||||
cdef bint _visible_number_can_help_open_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *:
|
||||
cdef int color = self._card_color_c(state, card)
|
||||
cdef int slot
|
||||
cdef int other
|
||||
cdef int count = 1
|
||||
cdef int number_sum = self._num_c(state, card)
|
||||
cdef bint has_high = self._card_rank_c(state, card) >= derived.middle_rank
|
||||
for slot in range(state.hand_size[player]):
|
||||
other = state.hand_encoded[player][slot]
|
||||
if (
|
||||
self._card_color_c(state, other) == color
|
||||
and self._card_rank_c(state, other) != 0
|
||||
and self._card_rank_c(state, other) >= self._card_rank_c(state, card)
|
||||
):
|
||||
count += 1
|
||||
number_sum += self._num_c(state, other)
|
||||
if self._card_rank_c(state, other) >= derived.middle_rank:
|
||||
has_high = True
|
||||
if count < derived.min_open_cards:
|
||||
return False
|
||||
if number_sum < derived.open_target_sum:
|
||||
return False
|
||||
return has_high
|
||||
|
||||
cdef double _deck_draw_value_c(self, _CachedState state, object derived) except *:
|
||||
cdef int score_diff = state.total_scores[state.current_player] - state.total_scores[1 - state.current_player]
|
||||
cdef double value
|
||||
if state.deck_remaining > derived.mid_deck_threshold:
|
||||
value = self.params.deck_draw_early_value
|
||||
elif state.deck_remaining > derived.late_deck_threshold:
|
||||
value = self.params.deck_draw_mid_value
|
||||
else:
|
||||
value = self.params.deck_draw_late_value
|
||||
if score_diff > 0:
|
||||
value += self.params.winning_deck_bonus
|
||||
else:
|
||||
value -= self.params.losing_deck_penalty
|
||||
return value
|
||||
|
||||
cdef double _card_value_for_me_c(
|
||||
self, _CachedState state, int player, int card, object derived
|
||||
) except *:
|
||||
cdef int color = self._card_color_c(state, card)
|
||||
cdef int rank = self._card_rank_c(state, card)
|
||||
cdef double commitment
|
||||
cdef int numeric_value
|
||||
cdef double value
|
||||
if not self._can_play_card_c(state, player, card):
|
||||
return 0.0
|
||||
commitment = self._color_commitment_c(state, player, color, derived)
|
||||
if rank == 0:
|
||||
return 7.0 + 1.2 * commitment
|
||||
numeric_value = self._num_c(state, card)
|
||||
value = 0.8 * numeric_value + self.params.commitment_weight * commitment
|
||||
if state.expedition_count[player][color] > 0:
|
||||
value += self.params.started_expedition_play_bonus
|
||||
value += self.params.started_expedition_followup_bonus
|
||||
if commitment >= 6.0 and rank <= derived.middle_rank:
|
||||
value += self.params.low_card_sequence_bonus
|
||||
return value
|
||||
|
||||
cdef double _card_value_for_opponent_c(
|
||||
self, _CachedState state, int opponent, int card, object derived
|
||||
) except *:
|
||||
cdef int rank = self._card_rank_c(state, card)
|
||||
cdef double interest
|
||||
if not self._can_play_card_c(state, opponent, card):
|
||||
return 0.0
|
||||
interest = self._public_color_commitment_for_opponent_c(
|
||||
state, opponent, self._card_color_c(state, card), derived
|
||||
)
|
||||
if rank == 0:
|
||||
return 8.0 + 1.5 * interest
|
||||
return self._num_c(state, card) * (0.4 + 0.25 * interest)
|
||||
|
||||
cdef double _color_commitment_c(
|
||||
self, _CachedState state, int player, int color, object derived
|
||||
) except *:
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int rank
|
||||
cdef int playable_numbers = 0
|
||||
cdef int playable_handshakes = 0
|
||||
cdef int playable_sum = 0
|
||||
cdef double value
|
||||
if self.color_commit_valid[player][color] != 0:
|
||||
return self.color_commit_cache[player][color]
|
||||
value = 0.0
|
||||
if state.expedition_count[player][color] > 0:
|
||||
value += 5.0
|
||||
value += 2.0 * state.expedition_handshakes[player][color]
|
||||
value += 0.25 * state.expedition_numeric_sum[player][color]
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if self._card_color_c(state, card) != color or not self._can_play_card_c(state, player, card):
|
||||
continue
|
||||
rank = self._card_rank_c(state, card)
|
||||
if rank == 0:
|
||||
playable_handshakes += 1
|
||||
else:
|
||||
playable_numbers += 1
|
||||
playable_sum += self._num_c(state, card)
|
||||
value += 1.2 * playable_numbers
|
||||
value += 1.5 * playable_handshakes
|
||||
value += 0.15 * playable_sum
|
||||
value += 0.05 * self._bonus_potential_c(state, player, color, 0, derived, 0, -1)
|
||||
self.color_commit_cache[player][color] = value
|
||||
self.color_commit_valid[player][color] = 1
|
||||
return value
|
||||
|
||||
cdef double _public_color_commitment_for_opponent_c(
|
||||
self, _CachedState state, int opponent, int color, object derived
|
||||
) except *:
|
||||
cdef int top_card
|
||||
cdef double value = 0.0
|
||||
if state.expedition_count[opponent][color] > 0:
|
||||
value += 5.0
|
||||
value += 2.0 * state.expedition_handshakes[opponent][color]
|
||||
value += 0.25 * state.expedition_numeric_sum[opponent][color]
|
||||
if state.expedition_last_numeric[opponent][color] > 0:
|
||||
value += 0.4 * (state.min_rank + state.expedition_last_numeric[opponent][color] - 1)
|
||||
if state.discard_count[color] > 0:
|
||||
top_card = state.discard_top[color]
|
||||
if self._can_play_card_c(state, opponent, top_card):
|
||||
if self._card_rank_c(state, top_card) == 0:
|
||||
value += 1.5
|
||||
else:
|
||||
value += 1.0 + 0.1 * self._num_c(state, top_card)
|
||||
if derived.bonus_possible and state.expedition_count[opponent][color] + 1 >= state.bonus_threshold:
|
||||
value += 0.2 * state.bonus_amount
|
||||
return value
|
||||
|
||||
cdef double _bonus_potential_c(
|
||||
self,
|
||||
_CachedState state,
|
||||
int player,
|
||||
int color,
|
||||
int extra_cards,
|
||||
object derived,
|
||||
int committed_cards,
|
||||
int exclude_card,
|
||||
) except *:
|
||||
cdef int need
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int playable_count = 0
|
||||
if not derived.bonus_possible:
|
||||
return 0.0
|
||||
need = state.bonus_threshold - (state.expedition_count[player][color] + committed_cards)
|
||||
if need <= 0:
|
||||
return <double>state.bonus_amount
|
||||
for slot in range(state.hand_size[player]):
|
||||
card = state.hand_encoded[player][slot]
|
||||
if (
|
||||
card != exclude_card
|
||||
and self._card_color_c(state, card) == color
|
||||
and self._can_play_card_c(state, player, card)
|
||||
):
|
||||
playable_count += 1
|
||||
if playable_count + extra_cards >= need:
|
||||
return 0.4 * state.bonus_amount
|
||||
return 0.0
|
||||
|
||||
cdef double _new_color_open_penalty_c(self, int opened_colors) noexcept:
|
||||
if opened_colors <= 1:
|
||||
return 0.0
|
||||
if opened_colors == 2:
|
||||
return 6.0
|
||||
if opened_colors == 3:
|
||||
return 14.0
|
||||
return 28.0
|
||||
|
||||
cdef double _late_penalty_c(self, object derived, int deck_left) except *:
|
||||
if deck_left <= derived.late_deck_threshold:
|
||||
return 15.0
|
||||
if deck_left <= derived.mid_deck_threshold:
|
||||
return 8.0
|
||||
return 0.0
|
||||
|
||||
def _act_card(self, state: GameState) -> int:
|
||||
player = state.current_player
|
||||
|
||||
@@ -25,14 +25,26 @@ class MctsConfig(StrictModel):
|
||||
c_puct: float = 1.5
|
||||
max_depth: int = 200
|
||||
use_rollout_value: bool = True
|
||||
rollout_policy: str = "random"
|
||||
parallel_simulations: int = 8
|
||||
virtual_loss_value: float = 1.0
|
||||
eval_with_mcts: bool = True
|
||||
eval_n_simulations: int = 0
|
||||
|
||||
@field_validator("n_simulations", "max_depth")
|
||||
@field_validator("n_simulations", "max_depth", "parallel_simulations")
|
||||
@classmethod
|
||||
def _positive_int(cls, value: int) -> int:
|
||||
if value <= 0:
|
||||
raise ValueError("must be positive")
|
||||
return value
|
||||
|
||||
@field_validator("rollout_policy")
|
||||
@classmethod
|
||||
def _rollout_policy(cls, value: str) -> str:
|
||||
if value not in {"random", "heuristic_balanced"}:
|
||||
raise ValueError("rollout_policy must be 'random' or 'heuristic_balanced'")
|
||||
return value
|
||||
|
||||
|
||||
class TemperatureConfig(StrictModel):
|
||||
training: float = 1.0
|
||||
@@ -44,8 +56,20 @@ class TrainingConfig(StrictModel):
|
||||
gradient_steps_per_iter: int = 10
|
||||
batch_size: int = 128
|
||||
replay_capacity: int = 100_000
|
||||
interleave_games: int = 8
|
||||
interleave_max_batch: int = 64
|
||||
num_workers: int = 1
|
||||
worker_device: str = "cpu"
|
||||
|
||||
@field_validator("games_per_iter", "gradient_steps_per_iter", "batch_size", "replay_capacity")
|
||||
@field_validator(
|
||||
"games_per_iter",
|
||||
"gradient_steps_per_iter",
|
||||
"batch_size",
|
||||
"replay_capacity",
|
||||
"interleave_games",
|
||||
"interleave_max_batch",
|
||||
"num_workers",
|
||||
)
|
||||
@classmethod
|
||||
def _positive_int(cls, value: int) -> int:
|
||||
if value <= 0:
|
||||
|
||||
@@ -0,0 +1,125 @@
|
||||
"""Multi-process eval workers for ISMCTS — slice games across processes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.bots.registry import build_bot
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim
|
||||
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||
|
||||
from .config import IsMctsConfig, config_from_dict
|
||||
from .mcts import IsMctsSearcher
|
||||
from .network import AlphaZeroNet
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class EvalWorkerBatch:
|
||||
worker_index: int
|
||||
config: dict[str, Any]
|
||||
game_config: dict[str, Any]
|
||||
network_state: dict[str, Any]
|
||||
mcts_config: dict[str, Any]
|
||||
opponent: str
|
||||
game_indices: list[int]
|
||||
seed: int
|
||||
device: str
|
||||
max_steps: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class EvalWorkerResult:
|
||||
worker_index: int
|
||||
score_diffs: list[float]
|
||||
wins0: int
|
||||
wins1: int
|
||||
draws: int
|
||||
policy_turns: int
|
||||
play_actions: int
|
||||
timeouts: int
|
||||
|
||||
|
||||
def run_eval_worker(batch: EvalWorkerBatch) -> EvalWorkerResult:
|
||||
import os
|
||||
|
||||
os.environ.setdefault("OMP_NUM_THREADS", "1")
|
||||
os.environ.setdefault("MKL_NUM_THREADS", "1")
|
||||
torch.set_num_threads(1)
|
||||
cfg: IsMctsConfig = config_from_dict(batch.config)
|
||||
game_config = LostCitiesConfig(**batch.game_config)
|
||||
device = torch.device(batch.device)
|
||||
probe = GameState.new_game(game_config, seed=batch.seed)
|
||||
in_dim = input_dim(probe, cfg.encoding)
|
||||
network = AlphaZeroNet.from_config(in_dim, probe.action_size, cfg).to(device)
|
||||
network.load_state_dict(batch.network_state)
|
||||
network.eval()
|
||||
from .config import MctsConfig
|
||||
|
||||
mcts_config = MctsConfig.model_validate(batch.mcts_config)
|
||||
|
||||
rng = random.Random(batch.seed + batch.worker_index * 7919)
|
||||
score_diffs: list[float] = []
|
||||
wins0 = wins1 = draws = 0
|
||||
policy_turns = 0
|
||||
play_actions = 0
|
||||
timeouts = 0
|
||||
for game_index in batch.game_indices:
|
||||
policy_player = game_index % 2
|
||||
opponents = [
|
||||
build_bot(batch.opponent, seed=batch.seed + game_index),
|
||||
build_bot(batch.opponent, seed=batch.seed + game_index + 1),
|
||||
]
|
||||
state = GameState.new_game(game_config, seed=batch.seed + game_index)
|
||||
steps = 0
|
||||
terminated = False
|
||||
while steps < batch.max_steps:
|
||||
if state.terminal:
|
||||
terminated = True
|
||||
break
|
||||
current = int(state.current_player)
|
||||
if current == policy_player:
|
||||
searcher = IsMctsSearcher(
|
||||
network,
|
||||
mcts_config,
|
||||
device=device,
|
||||
encoding=cfg.encoding,
|
||||
rng=random.Random(rng.randrange(2**31)),
|
||||
)
|
||||
visits = searcher.search(state, current)
|
||||
if visits:
|
||||
unified = max(visits, key=visits.get)
|
||||
else:
|
||||
unified = state.unified_legal_actions()[0]
|
||||
if state.phase == "card":
|
||||
policy_turns += 1
|
||||
if unified % 2 == 0:
|
||||
play_actions += 1
|
||||
state.apply_unified_action(unified)
|
||||
else:
|
||||
action = opponents[current].act(state)
|
||||
state.apply_action(action)
|
||||
steps += 1
|
||||
if not terminated:
|
||||
timeouts += 1
|
||||
diff = float(state.score_diff(policy_player))
|
||||
score_diffs.append(diff)
|
||||
if diff > 0:
|
||||
wins0 += 1
|
||||
elif diff < 0:
|
||||
wins1 += 1
|
||||
else:
|
||||
draws += 1
|
||||
return EvalWorkerResult(
|
||||
worker_index=batch.worker_index,
|
||||
score_diffs=score_diffs,
|
||||
wins0=wins0,
|
||||
wins1=wins1,
|
||||
draws=draws,
|
||||
policy_turns=policy_turns,
|
||||
play_actions=play_actions,
|
||||
timeouts=timeouts,
|
||||
)
|
||||
@@ -0,0 +1,190 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import multiprocessing as mp
|
||||
import random
|
||||
import time
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.bots.registry import build_bot
|
||||
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||
|
||||
from .config import IsMctsConfig, MctsConfig
|
||||
from .eval_worker import EvalWorkerBatch, run_eval_worker
|
||||
from .mcts import IsMctsSearcher
|
||||
from .network import AlphaZeroNet
|
||||
|
||||
|
||||
def evaluate_with_mcts(
|
||||
network: AlphaZeroNet,
|
||||
game_config: LostCitiesConfig,
|
||||
mcts_config: MctsConfig,
|
||||
*,
|
||||
games: int,
|
||||
seed: int,
|
||||
opponent: str,
|
||||
device: torch.device | str = "cpu",
|
||||
encoding=None,
|
||||
max_steps: int = 10_000,
|
||||
config: IsMctsConfig | None = None,
|
||||
num_workers: int = 1,
|
||||
) -> dict[str, float | int]:
|
||||
started = time.perf_counter()
|
||||
if num_workers > 1 and config is not None and games > 1:
|
||||
return _evaluate_parallel(
|
||||
network,
|
||||
game_config,
|
||||
mcts_config,
|
||||
config=config,
|
||||
games=games,
|
||||
seed=seed,
|
||||
opponent=opponent,
|
||||
num_workers=num_workers,
|
||||
max_steps=max_steps,
|
||||
started=started,
|
||||
)
|
||||
rng = random.Random(seed)
|
||||
score_diffs: list[float] = []
|
||||
wins0 = wins1 = draws = 0
|
||||
policy_turns = 0
|
||||
play_actions = 0
|
||||
timeouts = 0
|
||||
|
||||
network.eval()
|
||||
for game_index in range(games):
|
||||
policy_player = game_index % 2
|
||||
opponents = [
|
||||
build_bot(opponent, seed=seed + game_index),
|
||||
build_bot(opponent, seed=seed + game_index + 1),
|
||||
]
|
||||
state = GameState.new_game(game_config, seed=seed + game_index)
|
||||
steps = 0
|
||||
terminated = False
|
||||
while steps < max_steps:
|
||||
if state.terminal:
|
||||
terminated = True
|
||||
break
|
||||
current = int(state.current_player)
|
||||
if current == policy_player:
|
||||
searcher = IsMctsSearcher(
|
||||
network,
|
||||
mcts_config,
|
||||
device=device,
|
||||
encoding=encoding,
|
||||
rng=random.Random(rng.randrange(2**31)),
|
||||
)
|
||||
visits = searcher.search(state, current)
|
||||
if visits:
|
||||
unified = max(visits, key=visits.get)
|
||||
else:
|
||||
unified = state.unified_legal_actions()[0]
|
||||
if state.phase == "card":
|
||||
policy_turns += 1
|
||||
if unified % 2 == 0:
|
||||
play_actions += 1
|
||||
state.apply_unified_action(unified)
|
||||
else:
|
||||
action = opponents[current].act(state)
|
||||
state.apply_action(action)
|
||||
steps += 1
|
||||
if not terminated:
|
||||
timeouts += 1
|
||||
diff = float(state.score_diff(policy_player))
|
||||
score_diffs.append(diff)
|
||||
if diff > 0:
|
||||
wins0 += 1
|
||||
elif diff < 0:
|
||||
wins1 += 1
|
||||
else:
|
||||
draws += 1
|
||||
|
||||
n = len(score_diffs)
|
||||
avg_diff = sum(score_diffs) / n if n else 0.0
|
||||
return {
|
||||
"games": n,
|
||||
"win_rate0": wins0 / n if n else 0.0,
|
||||
"win_rate1": wins1 / n if n else 0.0,
|
||||
"wins0": wins0,
|
||||
"wins1": wins1,
|
||||
"draws": draws,
|
||||
"avg_score_diff0": avg_diff,
|
||||
"policy_turns": policy_turns,
|
||||
"play_action_rate": play_actions / policy_turns if policy_turns else 0.0,
|
||||
"max_step_timeouts": timeouts,
|
||||
"elapsed_seconds": time.perf_counter() - started,
|
||||
}
|
||||
|
||||
|
||||
def _evaluate_parallel(
|
||||
network: AlphaZeroNet,
|
||||
game_config: LostCitiesConfig,
|
||||
mcts_config: MctsConfig,
|
||||
*,
|
||||
config: IsMctsConfig,
|
||||
games: int,
|
||||
seed: int,
|
||||
opponent: str,
|
||||
num_workers: int,
|
||||
max_steps: int,
|
||||
started: float,
|
||||
) -> dict[str, float | int]:
|
||||
effective_workers = min(num_workers, games)
|
||||
base = games // effective_workers
|
||||
rem = games % effective_workers
|
||||
counts = [base + (1 if i < rem else 0) for i in range(effective_workers)]
|
||||
indices_per_worker: list[list[int]] = []
|
||||
cursor = 0
|
||||
for c in counts:
|
||||
indices_per_worker.append(list(range(cursor, cursor + c)))
|
||||
cursor += c
|
||||
cpu_state = {name: tensor.detach().cpu() for name, tensor in network.state_dict().items()}
|
||||
config_dict = config.to_dict()
|
||||
game_snapshot = game_config.to_snapshot()
|
||||
mcts_dict = mcts_config.model_dump(mode="json")
|
||||
worker_device = str(config.training.worker_device)
|
||||
batches = [
|
||||
EvalWorkerBatch(
|
||||
worker_index=i,
|
||||
config=config_dict,
|
||||
game_config=game_snapshot,
|
||||
network_state=cpu_state,
|
||||
mcts_config=mcts_dict,
|
||||
opponent=opponent,
|
||||
game_indices=indices_per_worker[i],
|
||||
seed=seed,
|
||||
device=worker_device,
|
||||
max_steps=max_steps,
|
||||
)
|
||||
for i in range(effective_workers)
|
||||
]
|
||||
ctx = mp.get_context("spawn")
|
||||
score_diffs: list[float] = []
|
||||
wins0 = wins1 = draws = 0
|
||||
policy_turns = 0
|
||||
play_actions = 0
|
||||
timeouts = 0
|
||||
with ProcessPoolExecutor(max_workers=effective_workers, mp_context=ctx) as executor:
|
||||
for res in executor.map(run_eval_worker, batches):
|
||||
score_diffs.extend(res.score_diffs)
|
||||
wins0 += res.wins0
|
||||
wins1 += res.wins1
|
||||
draws += res.draws
|
||||
policy_turns += res.policy_turns
|
||||
play_actions += res.play_actions
|
||||
timeouts += res.timeouts
|
||||
n = len(score_diffs)
|
||||
avg_diff = sum(score_diffs) / n if n else 0.0
|
||||
return {
|
||||
"games": n,
|
||||
"win_rate0": wins0 / n if n else 0.0,
|
||||
"win_rate1": wins1 / n if n else 0.0,
|
||||
"wins0": wins0,
|
||||
"wins1": wins1,
|
||||
"draws": draws,
|
||||
"avg_score_diff0": avg_diff,
|
||||
"policy_turns": policy_turns,
|
||||
"play_action_rate": play_actions / policy_turns if policy_turns else 0.0,
|
||||
"max_step_timeouts": timeouts,
|
||||
"elapsed_seconds": time.perf_counter() - started,
|
||||
}
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import struct
|
||||
from collections import Counter
|
||||
|
||||
from coolrl_lost_cities.games.classic.game import Card, GameState, build_deck
|
||||
@@ -14,30 +14,73 @@ def _sorted_cards(cards: list[Card]) -> list[tuple[int, int]]:
|
||||
return sorted(_card_tuple(card) for card in cards)
|
||||
|
||||
|
||||
# Phase encoding: "card" -> 0, "draw" -> 1, anything else -> 2.
|
||||
_PHASE_TO_INT = {"card": 0, "draw": 1}
|
||||
|
||||
|
||||
def canonical_info_set_key(state: GameState, player: int) -> bytes:
|
||||
"""Deterministic key for observable information from ``player``'s POV."""
|
||||
"""Deterministic key for observable information from ``player``'s POV.
|
||||
|
||||
Packed binary representation (big-endian) covering the same fields as
|
||||
the previous JSON encoding. Faster to compute and produces a more
|
||||
compact key while remaining a stable, hashable ``bytes`` value.
|
||||
"""
|
||||
p = int(player)
|
||||
payload = {
|
||||
"config": state.config.to_snapshot(),
|
||||
"player": p,
|
||||
"current_player": int(state.current_player),
|
||||
"phase": state.phase,
|
||||
"pending_discarded_color": (
|
||||
None if state.pending_discarded_color < 0 else int(state.pending_discarded_color)
|
||||
),
|
||||
"turn_count": int(state.turn_count),
|
||||
"terminal": bool(state.terminal),
|
||||
"deck_size": len(state.deck),
|
||||
"hand": _sorted_cards(state.hands[p]),
|
||||
"hand_size_opp": len(state.hands[1 - p]),
|
||||
"expeditions": [
|
||||
[[_card_tuple(card) for card in expedition] for expedition in player_expeditions]
|
||||
for player_expeditions in state.expeditions
|
||||
],
|
||||
"discards": [[_card_tuple(card) for card in discard] for discard in state.discards],
|
||||
"legal_mask": list(map(bool, state.unified_legal_mask())),
|
||||
}
|
||||
return json.dumps(payload, sort_keys=True, separators=(",", ":")).encode("utf-8")
|
||||
cfg = state.config
|
||||
parts: list[bytes] = []
|
||||
# Header: rule constants that pin down the action/observation shape.
|
||||
parts.append(
|
||||
struct.pack(
|
||||
">BBBBBhhBBBBB",
|
||||
int(cfg.n_colors) & 0xFF,
|
||||
int(cfg.n_ranks) & 0xFF,
|
||||
int(cfg.min_rank) & 0xFF,
|
||||
int(cfg.n_handshakes) & 0xFF,
|
||||
int(cfg.hand_size) & 0xFF,
|
||||
int(cfg.expedition_penalty),
|
||||
int(cfg.bonus_amount),
|
||||
int(cfg.bonus_threshold) & 0xFF,
|
||||
p & 0xFF,
|
||||
int(state.current_player) & 0xFF,
|
||||
_PHASE_TO_INT.get(state.phase, 2) & 0xFF,
|
||||
(1 if state.terminal else 0) & 0xFF,
|
||||
)
|
||||
)
|
||||
# Variable scalars.
|
||||
pending_color = -1 if state.pending_discarded_color < 0 else int(state.pending_discarded_color)
|
||||
parts.append(
|
||||
struct.pack(
|
||||
">bHHH",
|
||||
pending_color,
|
||||
int(state.turn_count) & 0xFFFF,
|
||||
len(state.deck) & 0xFFFF,
|
||||
len(state.hands[1 - p]) & 0xFFFF,
|
||||
)
|
||||
)
|
||||
# Sorted hand for the POV player. Cards encoded as (color, rank).
|
||||
hand = state.hands[p]
|
||||
parts.append(struct.pack(">H", len(hand)))
|
||||
if hand:
|
||||
sorted_pairs = sorted((int(c.color), int(c.rank)) for c in hand)
|
||||
parts.append(b"".join(struct.pack(">BB", c, r) for c, r in sorted_pairs))
|
||||
# Expeditions per player/color (ordered, since order matters for legality).
|
||||
expeditions = state.expeditions
|
||||
for player_expeditions in expeditions:
|
||||
for expedition in player_expeditions:
|
||||
parts.append(struct.pack(">H", len(expedition)))
|
||||
if expedition:
|
||||
parts.append(
|
||||
b"".join(struct.pack(">BB", int(c.color), int(c.rank)) for c in expedition)
|
||||
)
|
||||
# Discards per color.
|
||||
for discard in state.discards:
|
||||
parts.append(struct.pack(">H", len(discard)))
|
||||
if discard:
|
||||
parts.append(b"".join(struct.pack(">BB", int(c.color), int(c.rank)) for c in discard))
|
||||
# Legal mask (packed as raw bytes from the underlying list).
|
||||
mask = state.unified_legal_mask()
|
||||
parts.append(bytes(1 if bool(b) else 0 for b in mask))
|
||||
return b"".join(parts)
|
||||
|
||||
|
||||
def visible_cards(state: GameState, player: int) -> list[Card]:
|
||||
|
||||
@@ -0,0 +1,240 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import random
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
|
||||
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||
|
||||
from .config import MctsConfig, TrainingConfig
|
||||
from .info_set import canonical_info_set_key
|
||||
from .mcts import IsMctsSearcher, PendingSimulation
|
||||
from .network import AlphaZeroNet
|
||||
from .replay_buffer import ReplaySample
|
||||
from .self_play import select_from_distribution, visit_distribution
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PendingDecision:
|
||||
info_state: np.ndarray
|
||||
legal_mask: np.ndarray
|
||||
pi_target: np.ndarray
|
||||
player: int
|
||||
prior: np.ndarray
|
||||
game_index: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class _GameContext:
|
||||
state: GameState
|
||||
rng: random.Random
|
||||
game_index: int
|
||||
decisions: list[_PendingDecision] = field(default_factory=list)
|
||||
steps: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SearchJob:
|
||||
context: _GameContext
|
||||
searcher: IsMctsSearcher
|
||||
traverser: int
|
||||
remaining: int
|
||||
|
||||
|
||||
def play_self_play_iteration(
|
||||
network: AlphaZeroNet,
|
||||
mcts_config: MctsConfig,
|
||||
training_config: TrainingConfig,
|
||||
game_config: LostCitiesConfig,
|
||||
rng: random.Random,
|
||||
*,
|
||||
device: torch.device | str = "cpu",
|
||||
encoding=None,
|
||||
temperature: float = 1.0,
|
||||
max_steps: int = 10_000,
|
||||
) -> list[ReplaySample]:
|
||||
device = torch.device(device)
|
||||
completed: list[list[ReplaySample]] = []
|
||||
active: list[_GameContext] = []
|
||||
started = 0
|
||||
target_games = training_config.games_per_iter
|
||||
|
||||
def fill_active() -> None:
|
||||
nonlocal started
|
||||
while len(active) < training_config.interleave_games and started < target_games:
|
||||
active.append(
|
||||
_GameContext(
|
||||
state=GameState.new_game(game_config, seed=rng.randrange(2**31)),
|
||||
rng=random.Random(rng.randrange(2**31)),
|
||||
game_index=started,
|
||||
)
|
||||
)
|
||||
started += 1
|
||||
|
||||
fill_active()
|
||||
while active:
|
||||
jobs: list[_SearchJob] = []
|
||||
still_active: list[_GameContext] = []
|
||||
for context in active:
|
||||
if context.state.terminal or context.steps >= max_steps:
|
||||
completed.append(_finalize_context(context))
|
||||
continue
|
||||
player = int(context.state.current_player)
|
||||
searcher = IsMctsSearcher(
|
||||
network,
|
||||
mcts_config,
|
||||
device=device,
|
||||
encoding=encoding,
|
||||
rng=random.Random(context.rng.randrange(2**31)),
|
||||
)
|
||||
jobs.append(
|
||||
_SearchJob(
|
||||
context=context,
|
||||
searcher=searcher,
|
||||
traverser=player,
|
||||
remaining=mcts_config.n_simulations,
|
||||
)
|
||||
)
|
||||
still_active.append(context)
|
||||
|
||||
active = still_active
|
||||
if jobs:
|
||||
_run_search_jobs(network, jobs, training_config.interleave_max_batch, device)
|
||||
for job in jobs:
|
||||
_finish_decision(job, mcts_config, encoding, temperature)
|
||||
|
||||
fill_active()
|
||||
|
||||
samples: list[ReplaySample] = []
|
||||
for game_samples in completed:
|
||||
samples.extend(game_samples)
|
||||
return samples
|
||||
|
||||
|
||||
def _run_search_jobs(
|
||||
network: AlphaZeroNet,
|
||||
jobs: list[_SearchJob],
|
||||
max_batch: int,
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
while any(job.remaining > 0 for job in jobs):
|
||||
pending: list[tuple[_SearchJob, PendingSimulation]] = []
|
||||
for job in jobs:
|
||||
job_quota = min(
|
||||
job.remaining,
|
||||
job.searcher.config.parallel_simulations,
|
||||
max_batch - len(pending),
|
||||
)
|
||||
if job_quota <= 0:
|
||||
break
|
||||
job_pending = job.searcher.prepare_simulation_batch(
|
||||
job.context.state,
|
||||
job.traverser,
|
||||
job_quota,
|
||||
)
|
||||
pending.extend((job, item) for item in job_pending)
|
||||
job.remaining -= len(job_pending)
|
||||
if len(pending) >= max_batch:
|
||||
break
|
||||
if not pending:
|
||||
break
|
||||
_evaluate_global_batch(network, pending, device)
|
||||
|
||||
|
||||
def _evaluate_global_batch(
|
||||
network: AlphaZeroNet,
|
||||
pending: list[tuple[_SearchJob, PendingSimulation]],
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
network_pending = [(job, item) for job, item in pending if item.terminal_value is None]
|
||||
values_by_id: dict[int, float] = {}
|
||||
priors_by_id: dict[int, np.ndarray] = {}
|
||||
if network_pending:
|
||||
infos = np.stack(
|
||||
[item.info_state for _job, item in network_pending if item.info_state is not None]
|
||||
)
|
||||
masks = np.stack(
|
||||
[item.legal_mask for _job, item in network_pending if item.legal_mask is not None]
|
||||
)
|
||||
with torch.inference_mode():
|
||||
x = torch.as_tensor(infos, dtype=torch.float32, device=device)
|
||||
mask = torch.as_tensor(masks, dtype=torch.bool, device=device)
|
||||
probs = network.policy_distribution(x, mask).detach().cpu().numpy()
|
||||
_logits, values = network(x, mask)
|
||||
values_np = values.detach().cpu().numpy()
|
||||
for index, (_job, item) in enumerate(network_pending):
|
||||
priors_by_id[id(item)] = probs[index]
|
||||
values_by_id[id(item)] = float(values_np[index])
|
||||
|
||||
for job, item in pending:
|
||||
if item.terminal_value is not None:
|
||||
value = item.terminal_value
|
||||
else:
|
||||
assert item.leaf_node is not None
|
||||
value = job.searcher._expand_with_prior(
|
||||
item.leaf_node,
|
||||
item.leaf_state,
|
||||
item.leaf_player,
|
||||
item.legal_actions,
|
||||
priors_by_id[id(item)],
|
||||
values_by_id[id(item)],
|
||||
)
|
||||
job.searcher._backup(item.path, value, item.leaf_player)
|
||||
|
||||
|
||||
def _finish_decision(
|
||||
job: _SearchJob,
|
||||
mcts_config: MctsConfig,
|
||||
encoding,
|
||||
temperature: float,
|
||||
) -> None:
|
||||
context = job.context
|
||||
state = context.state
|
||||
player = int(state.current_player)
|
||||
legal_mask = np.asarray(state.unified_legal_mask(), dtype=bool)
|
||||
info = encode_info_state(state, player, encoding)
|
||||
root_key = canonical_info_set_key(state, player)
|
||||
root = job.searcher.tree.get_or_create(root_key, player=player, terminal=state.terminal)
|
||||
visits = {action: root.visits.get(action, 0) for action in state.unified_legal_actions()}
|
||||
pi = visit_distribution(visits, state.action_size, temperature=temperature)
|
||||
if pi.sum() <= 0:
|
||||
legal_actions = np.flatnonzero(legal_mask)
|
||||
pi[legal_actions] = 1.0 / len(legal_actions)
|
||||
prior = np.zeros(state.action_size, dtype=np.float32)
|
||||
for action in state.unified_legal_actions():
|
||||
prior[action] = float(root.priors.get(action, 0.0))
|
||||
context.decisions.append(
|
||||
_PendingDecision(
|
||||
info_state=info.astype(np.float32),
|
||||
legal_mask=legal_mask.astype(bool),
|
||||
pi_target=pi.astype(np.float32),
|
||||
player=player,
|
||||
prior=prior,
|
||||
game_index=context.game_index,
|
||||
)
|
||||
)
|
||||
action = select_from_distribution(pi, context.rng)
|
||||
state.apply_unified_action(action)
|
||||
context.steps += 1
|
||||
|
||||
|
||||
def _finalize_context(context: _GameContext) -> list[ReplaySample]:
|
||||
final_diff0 = float(context.state.score_diff(0))
|
||||
samples: list[ReplaySample] = []
|
||||
for decision in context.decisions:
|
||||
value = final_diff0 if decision.player == 0 else -final_diff0
|
||||
samples.append(
|
||||
ReplaySample(
|
||||
info_state=decision.info_state,
|
||||
legal_mask=decision.legal_mask,
|
||||
pi_target=decision.pi_target,
|
||||
v_target=value,
|
||||
player=decision.player,
|
||||
prior=decision.prior,
|
||||
game_index=decision.game_index,
|
||||
)
|
||||
)
|
||||
return samples
|
||||
@@ -7,6 +7,7 @@ from dataclasses import dataclass, field
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.bots.heuristic import HeuristicBot
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
|
||||
from coolrl_lost_cities.games.classic.game import GameState
|
||||
|
||||
@@ -23,6 +24,7 @@ class MctsNode:
|
||||
priors: dict[int, float] = field(default_factory=dict)
|
||||
visits: dict[int, int] = field(default_factory=dict)
|
||||
value_sum: dict[int, float] = field(default_factory=dict)
|
||||
virtual_visits: dict[int, int] = field(default_factory=dict)
|
||||
children: dict[int, bytes] = field(default_factory=dict)
|
||||
terminal: bool = False
|
||||
|
||||
@@ -36,6 +38,26 @@ class MctsNode:
|
||||
return self.value_sum.get(action, 0.0) / n
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchPathEntry:
|
||||
node: MctsNode
|
||||
action: int
|
||||
parent_player: int
|
||||
child_player: int
|
||||
|
||||
|
||||
@dataclass
|
||||
class PendingSimulation:
|
||||
path: list[SearchPathEntry]
|
||||
leaf_state: GameState
|
||||
leaf_node: MctsNode | None
|
||||
leaf_player: int
|
||||
info_state: np.ndarray | None
|
||||
legal_mask: np.ndarray | None
|
||||
legal_actions: list[int]
|
||||
terminal_value: float | None = None
|
||||
|
||||
|
||||
class MctsTree:
|
||||
def __init__(self) -> None:
|
||||
self.nodes: dict[bytes, MctsNode] = {}
|
||||
@@ -64,6 +86,9 @@ class IsMctsSearcher:
|
||||
self.encoding = encoding
|
||||
self.rng = rng or random.Random()
|
||||
self.tree = MctsTree()
|
||||
self._rollout_bot = (
|
||||
HeuristicBot() if config.rollout_policy == "heuristic_balanced" else None
|
||||
)
|
||||
|
||||
def search(
|
||||
self,
|
||||
@@ -76,71 +101,209 @@ class IsMctsSearcher:
|
||||
root_key, player=state.current_player, terminal=state.terminal
|
||||
)
|
||||
sims = int(n_sims or self.config.n_simulations)
|
||||
for _ in range(sims):
|
||||
det = sample_determinization(state, traverser, self.rng)
|
||||
self._simulate(det, depth=0)
|
||||
completed = 0
|
||||
while completed < sims:
|
||||
batch_size = min(self.config.parallel_simulations, sims - completed)
|
||||
pending = self.prepare_simulation_batch(state, traverser, batch_size)
|
||||
if not pending:
|
||||
break
|
||||
self.evaluate_and_backup(pending)
|
||||
completed += len(pending)
|
||||
legal = state.unified_legal_actions()
|
||||
return {action: root.visits.get(action, 0) for action in legal}
|
||||
|
||||
def _simulate(self, state: GameState, *, depth: int) -> float:
|
||||
player = int(state.current_player)
|
||||
if state.terminal or depth >= self.config.max_depth:
|
||||
return float(state.score_diff(player))
|
||||
def prepare_simulation_batch(
|
||||
self,
|
||||
root_state: GameState,
|
||||
traverser: int,
|
||||
max_simulations: int,
|
||||
) -> list[PendingSimulation]:
|
||||
pending: list[PendingSimulation] = []
|
||||
for _ in range(max_simulations):
|
||||
item = self.prepare_simulation(root_state, traverser)
|
||||
pending.append(item)
|
||||
if item.terminal_value is None and item.leaf_node is not None and not item.path:
|
||||
break
|
||||
return pending
|
||||
|
||||
key = canonical_info_set_key(state, player)
|
||||
node = self.tree.get_or_create(key, player=player, terminal=state.terminal)
|
||||
if not node.is_expanded():
|
||||
value = self._expand_and_evaluate(node, state, player)
|
||||
return value
|
||||
def prepare_simulation(self, root_state: GameState, traverser: int) -> PendingSimulation:
|
||||
state = sample_determinization(root_state, traverser, self.rng)
|
||||
path: list[SearchPathEntry] = []
|
||||
depth = 0
|
||||
# Cache the info-set key for the current node so we don't recompute it
|
||||
# after applying an action (the child's key becomes the next iter's key).
|
||||
cached_key: bytes | None = None
|
||||
while True:
|
||||
player = int(state.current_player)
|
||||
if state.terminal or depth >= self.config.max_depth:
|
||||
return PendingSimulation(
|
||||
path=path,
|
||||
leaf_state=state,
|
||||
leaf_node=None,
|
||||
leaf_player=player,
|
||||
info_state=None,
|
||||
legal_mask=None,
|
||||
legal_actions=[],
|
||||
terminal_value=float(state.score_diff(player)),
|
||||
)
|
||||
key = cached_key if cached_key is not None else canonical_info_set_key(state, player)
|
||||
node = self.tree.get_or_create(key, player=player, terminal=state.terminal)
|
||||
if not node.is_expanded():
|
||||
legal_actions = state.unified_legal_actions()
|
||||
if not legal_actions:
|
||||
node.terminal = True
|
||||
return PendingSimulation(
|
||||
path=path,
|
||||
leaf_state=state,
|
||||
leaf_node=node,
|
||||
leaf_player=player,
|
||||
info_state=None,
|
||||
legal_mask=None,
|
||||
legal_actions=[],
|
||||
terminal_value=float(state.score_diff(player)),
|
||||
)
|
||||
return PendingSimulation(
|
||||
path=path,
|
||||
leaf_state=state,
|
||||
leaf_node=node,
|
||||
leaf_player=player,
|
||||
info_state=encode_info_state(state, player, self.encoding),
|
||||
legal_mask=np.asarray(state.unified_legal_mask(), dtype=bool),
|
||||
legal_actions=legal_actions,
|
||||
)
|
||||
|
||||
action = self._select_action(node, state.unified_legal_actions())
|
||||
child = state.clone()
|
||||
child.apply_unified_action(action)
|
||||
child_value = self._simulate(child, depth=depth + 1)
|
||||
value = child_value if child.current_player == player else -child_value
|
||||
node.visits[action] = node.visits.get(action, 0) + 1
|
||||
node.value_sum[action] = node.value_sum.get(action, 0.0) + value
|
||||
child_key = canonical_info_set_key(child, child.current_player)
|
||||
node.children[action] = child_key
|
||||
return value
|
||||
# Reuse the same legal_actions list for selection (avoids one
|
||||
# extra unified_legal_actions() call inside _select_action).
|
||||
legal_actions = state.unified_legal_actions()
|
||||
action = self._select_action(node, legal_actions)
|
||||
node.virtual_visits[action] = node.virtual_visits.get(action, 0) + 1
|
||||
child = state.clone()
|
||||
child.apply_unified_action(action)
|
||||
child_player = int(child.current_player)
|
||||
child_key = canonical_info_set_key(child, child_player)
|
||||
node.children[action] = child_key
|
||||
path.append(
|
||||
SearchPathEntry(
|
||||
node=node,
|
||||
action=action,
|
||||
parent_player=player,
|
||||
child_player=child_player,
|
||||
)
|
||||
)
|
||||
state = child
|
||||
cached_key = child_key
|
||||
depth += 1
|
||||
|
||||
def _expand_and_evaluate(self, node: MctsNode, state: GameState, player: int) -> float:
|
||||
def evaluate_and_backup(self, pending: list[PendingSimulation]) -> None:
|
||||
network_pending = [item for item in pending if item.terminal_value is None]
|
||||
values_by_id: dict[int, float] = {}
|
||||
priors_by_id: dict[int, np.ndarray] = {}
|
||||
if network_pending:
|
||||
infos = np.stack(
|
||||
[item.info_state for item in network_pending if item.info_state is not None]
|
||||
)
|
||||
masks = np.stack(
|
||||
[item.legal_mask for item in network_pending if item.legal_mask is not None]
|
||||
)
|
||||
with torch.inference_mode():
|
||||
x = torch.as_tensor(infos, dtype=torch.float32, device=self.device)
|
||||
mask = torch.as_tensor(masks, dtype=torch.bool, device=self.device)
|
||||
probs = self.network.policy_distribution(x, mask).detach().cpu().numpy()
|
||||
_logits, network_values = self.network(x, mask)
|
||||
network_values_np = network_values.detach().cpu().numpy()
|
||||
for index, item in enumerate(network_pending):
|
||||
priors_by_id[id(item)] = probs[index]
|
||||
values_by_id[id(item)] = float(network_values_np[index])
|
||||
|
||||
for item in pending:
|
||||
if item.terminal_value is not None:
|
||||
value = item.terminal_value
|
||||
else:
|
||||
assert item.leaf_node is not None
|
||||
value = self._expand_with_prior(
|
||||
item.leaf_node,
|
||||
item.leaf_state,
|
||||
item.leaf_player,
|
||||
item.legal_actions,
|
||||
priors_by_id[id(item)],
|
||||
values_by_id[id(item)],
|
||||
)
|
||||
self._backup(item.path, value, item.leaf_player)
|
||||
|
||||
def _expand_with_prior(
|
||||
self,
|
||||
node: MctsNode,
|
||||
state: GameState,
|
||||
player: int,
|
||||
legal_actions: list[int],
|
||||
probs: np.ndarray,
|
||||
network_value: float,
|
||||
) -> float:
|
||||
legal_actions = state.unified_legal_actions()
|
||||
if not legal_actions:
|
||||
node.terminal = True
|
||||
return float(state.score_diff(player))
|
||||
info = encode_info_state(state, player, self.encoding)
|
||||
legal_mask = np.asarray(state.unified_legal_mask(), dtype=bool)
|
||||
with torch.inference_mode():
|
||||
x = torch.as_tensor(info[None, :], dtype=torch.float32, device=self.device)
|
||||
mask = torch.as_tensor(legal_mask[None, :], dtype=torch.bool, device=self.device)
|
||||
probs = self.network.policy_distribution(x, mask).squeeze(0).detach().cpu().numpy()
|
||||
_logits, network_value = self.network(x, mask)
|
||||
for action in legal_actions:
|
||||
node.priors[action] = float(probs[action])
|
||||
node.visits.setdefault(action, 0)
|
||||
node.value_sum.setdefault(action, 0.0)
|
||||
node.virtual_visits.setdefault(action, 0)
|
||||
rollout_value = (
|
||||
self._rollout_value(state, player) if self.config.use_rollout_value else None
|
||||
)
|
||||
if rollout_value is None:
|
||||
return float(network_value.item())
|
||||
return float(network_value)
|
||||
return rollout_value
|
||||
|
||||
def _select_action(self, node: MctsNode, legal_actions: list[int]) -> int:
|
||||
total_visits = sum(node.visits.get(action, 0) for action in legal_actions)
|
||||
total_visits = sum(
|
||||
node.visits.get(action, 0) + node.virtual_visits.get(action, 0)
|
||||
for action in legal_actions
|
||||
)
|
||||
sqrt_total = math.sqrt(max(1, total_visits))
|
||||
best_score = -float("inf")
|
||||
best_action = legal_actions[0]
|
||||
for action in legal_actions:
|
||||
n = node.visits.get(action, 0)
|
||||
virtual = node.virtual_visits.get(action, 0)
|
||||
n_eff = n + virtual
|
||||
prior = node.priors.get(action, 0.0)
|
||||
score = node.q(action) + self.config.c_puct * prior * sqrt_total / (1 + n)
|
||||
if n_eff <= 0:
|
||||
q_eff = 0.0
|
||||
else:
|
||||
q_eff = (
|
||||
node.value_sum.get(action, 0.0) - virtual * self.config.virtual_loss_value
|
||||
) / n_eff
|
||||
score = q_eff + self.config.c_puct * prior * sqrt_total / (1 + n_eff)
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_action = action
|
||||
return int(best_action)
|
||||
|
||||
def _backup(
|
||||
self,
|
||||
path: list[SearchPathEntry],
|
||||
leaf_value: float,
|
||||
leaf_player: int,
|
||||
) -> None:
|
||||
value = float(leaf_value)
|
||||
value_player = int(leaf_player)
|
||||
for entry in reversed(path):
|
||||
parent_value = value if value_player == entry.parent_player else -value
|
||||
current_virtual = entry.node.virtual_visits.get(entry.action, 0)
|
||||
entry.node.virtual_visits[entry.action] = max(0, current_virtual - 1)
|
||||
entry.node.visits[entry.action] = entry.node.visits.get(entry.action, 0) + 1
|
||||
entry.node.value_sum[entry.action] = (
|
||||
entry.node.value_sum.get(entry.action, 0.0) + parent_value
|
||||
)
|
||||
value = parent_value
|
||||
value_player = entry.parent_player
|
||||
|
||||
def _release_virtual_path(self, path: list[SearchPathEntry]) -> None:
|
||||
for entry in path:
|
||||
current_virtual = entry.node.virtual_visits.get(entry.action, 0)
|
||||
entry.node.virtual_visits[entry.action] = max(0, current_virtual - 1)
|
||||
|
||||
def _rollout_value(self, state: GameState, player: int) -> float | None:
|
||||
rollout = state.clone()
|
||||
steps = 0
|
||||
@@ -148,7 +311,13 @@ class IsMctsSearcher:
|
||||
legal = rollout.unified_legal_actions()
|
||||
if not legal:
|
||||
break
|
||||
action = self.rng.choice(legal)
|
||||
if self._rollout_bot is not None:
|
||||
phase_action = self._rollout_bot.act(rollout)
|
||||
action = rollout.to_unified_action(phase_action)
|
||||
if action not in legal:
|
||||
action = self.rng.choice(legal)
|
||||
else:
|
||||
action = self.rng.choice(legal)
|
||||
rollout.apply_unified_action(action)
|
||||
steps += 1
|
||||
return float(rollout.score_diff(player))
|
||||
|
||||
@@ -0,0 +1,621 @@
|
||||
# cython: language_level=3, boundscheck=False, wraparound=False, cdivision=True, initializedcheck=False
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.bots.heuristic_cy cimport HeuristicBot
|
||||
from coolrl_lost_cities.games.classic.bots.heuristic import HeuristicBot as PyHeuristicBot
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
|
||||
from coolrl_lost_cities.games.classic.game cimport GameState
|
||||
|
||||
from .determinization import sample_determinization
|
||||
from .info_set import canonical_info_set_key
|
||||
|
||||
|
||||
DEF MAX_ACTIONS = 64
|
||||
DEF DEFAULT_ACTION_SIZE = 64
|
||||
|
||||
|
||||
cdef class _ArrayMap:
|
||||
cdef MctsNode node
|
||||
cdef int kind
|
||||
cdef int action_size
|
||||
cdef bint is_int
|
||||
|
||||
def __init__(self, MctsNode node, int kind, bint is_int=False):
|
||||
self.node = node
|
||||
self.kind = kind
|
||||
self.action_size = node.action_size
|
||||
self.is_int = is_int
|
||||
|
||||
cdef inline void _check(self, int action) except *:
|
||||
if action < 0 or action >= self.action_size:
|
||||
raise KeyError(action)
|
||||
|
||||
cdef inline bint has(self, int action) noexcept:
|
||||
return 0 <= action < self.action_size and self.node.active_present[action] != 0
|
||||
|
||||
cdef inline long get_int(self, int action, long default_value=0) noexcept:
|
||||
if 0 <= action < self.action_size and self.node.active_present[action] != 0:
|
||||
if self.kind == 1:
|
||||
return self.node.visits_arr[action]
|
||||
if self.kind == 3:
|
||||
return self.node.virtual_visits_arr[action]
|
||||
return default_value
|
||||
|
||||
cdef inline double get_float(self, int action, double default_value=0.0) noexcept:
|
||||
if 0 <= action < self.action_size and self.node.active_present[action] != 0:
|
||||
if self.kind == 0:
|
||||
return self.node.priors_arr[action]
|
||||
if self.kind == 2:
|
||||
return self.node.value_sum_arr[action]
|
||||
return default_value
|
||||
|
||||
cdef inline void _mark_active(self, int action) noexcept:
|
||||
if self.node.active_present[action] == 0:
|
||||
self.node.active_present[action] = 1
|
||||
self.node.active_actions[self.node.n_active] = action
|
||||
self.node.n_active += 1
|
||||
|
||||
cdef inline void set_int(self, int action, long value) except *:
|
||||
self._check(action)
|
||||
self._mark_active(action)
|
||||
if self.kind == 1:
|
||||
self.node.visits_arr[action] = <int>value
|
||||
elif self.kind == 3:
|
||||
self.node.virtual_visits_arr[action] = <int>value
|
||||
else:
|
||||
raise TypeError("integer write to float node map")
|
||||
|
||||
cdef inline void set_float(self, int action, double value) except *:
|
||||
self._check(action)
|
||||
self._mark_active(action)
|
||||
if self.kind == 0:
|
||||
self.node.priors_arr[action] = value
|
||||
elif self.kind == 2:
|
||||
self.node.value_sum_arr[action] = value
|
||||
else:
|
||||
raise TypeError("float write to integer node map")
|
||||
|
||||
def get(self, action, default=None):
|
||||
cdef int a = int(action)
|
||||
if self.has(a):
|
||||
if self.is_int:
|
||||
return int(self.get_int(a, 0))
|
||||
return float(self.get_float(a, 0.0))
|
||||
return default
|
||||
|
||||
def setdefault(self, action, default=None):
|
||||
cdef int a = int(action)
|
||||
if self.has(a):
|
||||
if self.is_int:
|
||||
return int(self.get_int(a, 0))
|
||||
return float(self.get_float(a, 0.0))
|
||||
if default is None:
|
||||
default = 0 if self.is_int else 0.0
|
||||
if self.is_int:
|
||||
self.set_int(a, int(default))
|
||||
return int(default)
|
||||
self.set_float(a, float(default))
|
||||
return float(default)
|
||||
|
||||
def __getitem__(self, action):
|
||||
cdef int a = int(action)
|
||||
self._check(a)
|
||||
if self.node.active_present[a] == 0:
|
||||
raise KeyError(action)
|
||||
if self.is_int:
|
||||
return int(self.get_int(a, 0))
|
||||
return float(self.get_float(a, 0.0))
|
||||
|
||||
def __setitem__(self, action, value):
|
||||
cdef int a = int(action)
|
||||
if self.is_int:
|
||||
self.set_int(a, int(value))
|
||||
else:
|
||||
self.set_float(a, float(value))
|
||||
|
||||
def __contains__(self, action):
|
||||
return self.has(int(action))
|
||||
|
||||
def __bool__(self):
|
||||
return self.node.n_active > 0
|
||||
|
||||
def __len__(self):
|
||||
return self.node.n_active
|
||||
|
||||
def items(self):
|
||||
cdef int i
|
||||
result = []
|
||||
cdef int action
|
||||
for i in range(self.node.n_active):
|
||||
action = self.node.active_actions[i]
|
||||
if self.is_int:
|
||||
result.append((action, int(self.get_int(action, 0))))
|
||||
else:
|
||||
result.append((action, float(self.get_float(action, 0.0))))
|
||||
return result
|
||||
|
||||
def keys(self):
|
||||
cdef int i
|
||||
return [self.node.active_actions[i] for i in range(self.node.n_active)]
|
||||
|
||||
def values(self):
|
||||
cdef int i
|
||||
result = []
|
||||
cdef int action
|
||||
for i in range(self.node.n_active):
|
||||
action = self.node.active_actions[i]
|
||||
if self.is_int:
|
||||
result.append(int(self.get_int(action, 0)))
|
||||
else:
|
||||
result.append(float(self.get_float(action, 0.0)))
|
||||
return result
|
||||
|
||||
def __repr__(self):
|
||||
return repr(dict(self.items()))
|
||||
|
||||
|
||||
cdef class MctsNode:
|
||||
cdef public bytes info_set_key
|
||||
cdef public int player
|
||||
cdef public object priors
|
||||
cdef public object visits
|
||||
cdef public object value_sum
|
||||
cdef public object virtual_visits
|
||||
cdef public dict children
|
||||
cdef public bint terminal
|
||||
cdef public bint expanded
|
||||
cdef int action_size
|
||||
cdef int visits_arr[MAX_ACTIONS]
|
||||
cdef double value_sum_arr[MAX_ACTIONS]
|
||||
cdef double priors_arr[MAX_ACTIONS]
|
||||
cdef int virtual_visits_arr[MAX_ACTIONS]
|
||||
cdef int active_actions[MAX_ACTIONS]
|
||||
cdef unsigned char active_present[MAX_ACTIONS]
|
||||
cdef int n_active
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bytes info_set_key,
|
||||
int player,
|
||||
object priors=None,
|
||||
object visits=None,
|
||||
object value_sum=None,
|
||||
object virtual_visits=None,
|
||||
object children=None,
|
||||
bint terminal=False,
|
||||
int action_size=DEFAULT_ACTION_SIZE,
|
||||
):
|
||||
if action_size > MAX_ACTIONS:
|
||||
raise ValueError("action_size exceeds fixed MCTS action buffer")
|
||||
self.info_set_key = info_set_key
|
||||
self.player = player
|
||||
self.terminal = terminal
|
||||
self.expanded = terminal
|
||||
self.action_size = action_size
|
||||
self.n_active = 0
|
||||
self.priors = _ArrayMap(self, 0, False)
|
||||
self.visits = _ArrayMap(self, 1, True)
|
||||
self.value_sum = _ArrayMap(self, 2, False)
|
||||
self.virtual_visits = _ArrayMap(self, 3, True)
|
||||
self.children = {} if children is None else dict(children)
|
||||
if priors is not None:
|
||||
for action, value in dict(priors).items():
|
||||
self.priors[action] = value
|
||||
if visits is not None:
|
||||
for action, value in dict(visits).items():
|
||||
self.visits[action] = value
|
||||
if value_sum is not None:
|
||||
for action, value in dict(value_sum).items():
|
||||
self.value_sum[action] = value
|
||||
if virtual_visits is not None:
|
||||
for action, value in dict(virtual_visits).items():
|
||||
self.virtual_visits[action] = value
|
||||
|
||||
cpdef bint is_expanded(self):
|
||||
return self.terminal or self.expanded or self.n_active > 0
|
||||
|
||||
cpdef double q(self, int action):
|
||||
cdef long n = (<_ArrayMap>self.visits).get_int(action, 0)
|
||||
if n <= 0:
|
||||
return 0.0
|
||||
return (<_ArrayMap>self.value_sum).get_float(action, 0.0) / n
|
||||
|
||||
|
||||
cdef class SearchPathEntry:
|
||||
cdef public MctsNode node
|
||||
cdef public int action
|
||||
cdef public int parent_player
|
||||
cdef public int child_player
|
||||
|
||||
def __init__(self, MctsNode node, int action, int parent_player, int child_player):
|
||||
self.node = node
|
||||
self.action = action
|
||||
self.parent_player = parent_player
|
||||
self.child_player = child_player
|
||||
|
||||
|
||||
cdef class PendingSimulation:
|
||||
cdef public list path
|
||||
cdef public GameState leaf_state
|
||||
cdef public object leaf_node
|
||||
cdef public int leaf_player
|
||||
cdef public object info_state
|
||||
cdef public object legal_mask
|
||||
cdef public list legal_actions
|
||||
cdef public object terminal_value
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
list path,
|
||||
GameState leaf_state,
|
||||
object leaf_node,
|
||||
int leaf_player,
|
||||
object info_state,
|
||||
object legal_mask,
|
||||
list legal_actions,
|
||||
object terminal_value=None,
|
||||
):
|
||||
self.path = path
|
||||
self.leaf_state = leaf_state
|
||||
self.leaf_node = leaf_node
|
||||
self.leaf_player = leaf_player
|
||||
self.info_state = info_state
|
||||
self.legal_mask = legal_mask
|
||||
self.legal_actions = legal_actions
|
||||
self.terminal_value = terminal_value
|
||||
|
||||
|
||||
cdef class MctsTree:
|
||||
cdef public dict nodes
|
||||
cdef int action_size
|
||||
|
||||
def __init__(self, int action_size=DEFAULT_ACTION_SIZE):
|
||||
self.nodes = {}
|
||||
self.action_size = action_size
|
||||
|
||||
def get_or_create(self, bytes key, *, int player, bint terminal=False):
|
||||
cdef MctsNode node = self.nodes.get(key)
|
||||
if node is None:
|
||||
node = MctsNode(key, player=player, terminal=terminal, action_size=self.action_size)
|
||||
self.nodes[key] = node
|
||||
return node
|
||||
|
||||
|
||||
cdef class IsMctsSearcher:
|
||||
cdef public object network
|
||||
cdef public object config
|
||||
cdef public object device
|
||||
cdef public object encoding
|
||||
cdef public object rng
|
||||
cdef public MctsTree tree
|
||||
cdef HeuristicBot _rollout_bot
|
||||
cdef int action_size
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
object network,
|
||||
object config,
|
||||
*,
|
||||
object device="cpu",
|
||||
object encoding=None,
|
||||
object rng=None,
|
||||
):
|
||||
self.network = network
|
||||
self.config = config
|
||||
self.device = torch.device(device)
|
||||
self.encoding = encoding
|
||||
self.rng = rng or random.Random()
|
||||
self.action_size = int(getattr(network, "action_size", DEFAULT_ACTION_SIZE))
|
||||
if self.action_size > MAX_ACTIONS:
|
||||
raise ValueError("action_size exceeds fixed MCTS action buffer")
|
||||
self.tree = MctsTree(self.action_size)
|
||||
self._rollout_bot = (
|
||||
<HeuristicBot>PyHeuristicBot() if config.rollout_policy == "heuristic_balanced" else None
|
||||
)
|
||||
|
||||
cdef inline int _from_unified_action_c(self, GameState state, int action_id) noexcept:
|
||||
cdef int card_action_size = 2 * state.hand_size
|
||||
if state.phase_id == 0:
|
||||
return action_id
|
||||
return action_id - card_action_size
|
||||
|
||||
cdef list _unified_legal_actions_list_c(self, GameState state):
|
||||
cdef int actions[MAX_ACTIONS]
|
||||
cdef int count = state._unified_legal_actions_c(actions)
|
||||
cdef int i
|
||||
return [actions[i] for i in range(count)]
|
||||
|
||||
cpdef dict search(self, GameState state, int traverser, object n_sims=None):
|
||||
cdef bytes root_key = canonical_info_set_key(state, state.current_player)
|
||||
cdef MctsNode root = self.tree.get_or_create(
|
||||
root_key, player=state.current_player, terminal=state.terminal
|
||||
)
|
||||
cdef int sims = int(n_sims or self.config.n_simulations)
|
||||
cdef int completed = 0
|
||||
cdef list pending
|
||||
cdef list legal
|
||||
cdef int action
|
||||
cdef dict result
|
||||
while completed < sims:
|
||||
pending = self.prepare_simulation_batch(state, traverser, 1)
|
||||
if not pending:
|
||||
break
|
||||
self.evaluate_and_backup(pending)
|
||||
completed += len(pending)
|
||||
legal = state.unified_legal_actions()
|
||||
result = {}
|
||||
for action in legal:
|
||||
result[action] = (<_ArrayMap>root.visits).get_int(action, 0)
|
||||
return result
|
||||
|
||||
cpdef list prepare_simulation_batch(
|
||||
self,
|
||||
GameState root_state,
|
||||
int traverser,
|
||||
int max_simulations,
|
||||
):
|
||||
cdef list pending = []
|
||||
cdef PendingSimulation item
|
||||
cdef int i
|
||||
for i in range(max_simulations):
|
||||
item = self.prepare_simulation(root_state, traverser)
|
||||
pending.append(item)
|
||||
if item.terminal_value is None and item.leaf_node is not None and not item.path:
|
||||
break
|
||||
return pending
|
||||
|
||||
cpdef PendingSimulation prepare_simulation(self, GameState root_state, int traverser):
|
||||
cdef GameState state = sample_determinization(root_state, traverser, self.rng)
|
||||
cdef list path = []
|
||||
cdef int depth = 0
|
||||
cdef object cached_key = None
|
||||
cdef int player
|
||||
cdef bytes key
|
||||
cdef MctsNode node
|
||||
cdef list legal_actions
|
||||
cdef int action
|
||||
cdef int local_action
|
||||
cdef int child_player
|
||||
cdef bytes child_key
|
||||
cdef int actions[MAX_ACTIONS]
|
||||
cdef int action_count
|
||||
cdef int i
|
||||
while True:
|
||||
player = state.current_player
|
||||
if state.terminal or depth >= int(self.config.max_depth):
|
||||
return PendingSimulation(
|
||||
path=path,
|
||||
leaf_state=state,
|
||||
leaf_node=None,
|
||||
leaf_player=player,
|
||||
info_state=None,
|
||||
legal_mask=None,
|
||||
legal_actions=[],
|
||||
terminal_value=float(state.total_scores[player] - state.total_scores[1 - player]),
|
||||
)
|
||||
if cached_key is None:
|
||||
key = canonical_info_set_key(state, player)
|
||||
else:
|
||||
key = cached_key
|
||||
node = self.tree.get_or_create(key, player=player, terminal=state.terminal)
|
||||
if not node.is_expanded():
|
||||
action_count = state._unified_legal_actions_c(actions)
|
||||
legal_actions = [actions[i] for i in range(action_count)]
|
||||
if not legal_actions:
|
||||
node.terminal = True
|
||||
return PendingSimulation(
|
||||
path=path,
|
||||
leaf_state=state,
|
||||
leaf_node=node,
|
||||
leaf_player=player,
|
||||
info_state=None,
|
||||
legal_mask=None,
|
||||
legal_actions=[],
|
||||
terminal_value=float(state.total_scores[player] - state.total_scores[1 - player]),
|
||||
)
|
||||
return PendingSimulation(
|
||||
path=path,
|
||||
leaf_state=state,
|
||||
leaf_node=node,
|
||||
leaf_player=player,
|
||||
info_state=encode_info_state(state, player, self.encoding),
|
||||
legal_mask=np.asarray(state.unified_legal_mask(), dtype=bool),
|
||||
legal_actions=legal_actions,
|
||||
)
|
||||
|
||||
action_count = state._unified_legal_actions_c(actions)
|
||||
legal_actions = [actions[i] for i in range(action_count)]
|
||||
action = self._select_action(node, legal_actions)
|
||||
(<_ArrayMap>node.virtual_visits).set_int(
|
||||
action, (<_ArrayMap>node.virtual_visits).get_int(action, 0) + 1
|
||||
)
|
||||
local_action = self._from_unified_action_c(state, action)
|
||||
state._push_action_c(local_action)
|
||||
child_player = state.current_player
|
||||
child_key = canonical_info_set_key(state, child_player)
|
||||
node.children[action] = child_key
|
||||
path.append(
|
||||
SearchPathEntry(
|
||||
node=node,
|
||||
action=action,
|
||||
parent_player=player,
|
||||
child_player=child_player,
|
||||
)
|
||||
)
|
||||
cached_key = child_key
|
||||
depth += 1
|
||||
|
||||
cpdef evaluate_and_backup(self, list pending):
|
||||
cdef list network_pending = [item for item in pending if item.terminal_value is None]
|
||||
cdef dict values_by_id = {}
|
||||
cdef dict priors_by_id = {}
|
||||
cdef object infos
|
||||
cdef object masks
|
||||
cdef object x
|
||||
cdef object mask
|
||||
cdef object probs
|
||||
cdef object network_values
|
||||
cdef object network_values_np
|
||||
cdef int index
|
||||
cdef PendingSimulation item
|
||||
cdef double value
|
||||
if network_pending:
|
||||
infos = np.stack([item.info_state for item in network_pending if item.info_state is not None])
|
||||
masks = np.stack([item.legal_mask for item in network_pending if item.legal_mask is not None])
|
||||
with torch.inference_mode():
|
||||
x = torch.as_tensor(infos, dtype=torch.float32, device=self.device)
|
||||
mask = torch.as_tensor(masks, dtype=torch.bool, device=self.device)
|
||||
probs = self.network.policy_distribution(x, mask).detach().cpu().numpy()
|
||||
_logits, network_values = self.network(x, mask)
|
||||
network_values_np = network_values.detach().cpu().numpy()
|
||||
for index, item in enumerate(network_pending):
|
||||
priors_by_id[id(item)] = probs[index]
|
||||
values_by_id[id(item)] = float(network_values_np[index])
|
||||
|
||||
for item in pending:
|
||||
if item.terminal_value is not None:
|
||||
value = item.terminal_value
|
||||
else:
|
||||
value = self._expand_with_prior(
|
||||
item.leaf_node,
|
||||
item.leaf_state,
|
||||
item.leaf_player,
|
||||
item.legal_actions,
|
||||
priors_by_id[id(item)],
|
||||
values_by_id[id(item)],
|
||||
)
|
||||
self._backup(item.path, value, item.leaf_player)
|
||||
|
||||
cpdef double _expand_with_prior(
|
||||
self,
|
||||
MctsNode node,
|
||||
GameState state,
|
||||
int player,
|
||||
list legal_actions,
|
||||
object probs,
|
||||
double network_value,
|
||||
):
|
||||
cdef int action
|
||||
cdef object rollout_value
|
||||
legal_actions = self._unified_legal_actions_list_c(state)
|
||||
if not legal_actions:
|
||||
node.terminal = True
|
||||
return float(state.total_scores[player] - state.total_scores[1 - player])
|
||||
node.expanded = True
|
||||
for action in legal_actions:
|
||||
(<_ArrayMap>node.priors).set_float(action, float(probs[action]))
|
||||
if not (<_ArrayMap>node.visits).has(action):
|
||||
(<_ArrayMap>node.visits).set_int(action, 0)
|
||||
if not (<_ArrayMap>node.value_sum).has(action):
|
||||
(<_ArrayMap>node.value_sum).set_float(action, 0.0)
|
||||
if not (<_ArrayMap>node.virtual_visits).has(action):
|
||||
(<_ArrayMap>node.virtual_visits).set_int(action, 0)
|
||||
rollout_value = self._rollout_value(state, player) if self.config.use_rollout_value else None
|
||||
if rollout_value is None:
|
||||
return float(network_value)
|
||||
return float(rollout_value)
|
||||
|
||||
cpdef int _select_action(self, MctsNode node, list legal_actions):
|
||||
cdef int total_visits = 0
|
||||
cdef int action
|
||||
cdef long n
|
||||
cdef long virtual
|
||||
cdef long n_eff
|
||||
cdef double sqrt_total
|
||||
cdef double prior
|
||||
cdef double q_eff
|
||||
cdef double score
|
||||
cdef double best_score = -float("inf")
|
||||
cdef int best_action = int(legal_actions[0])
|
||||
cdef _ArrayMap visits = <_ArrayMap>node.visits
|
||||
cdef _ArrayMap virtual_visits = <_ArrayMap>node.virtual_visits
|
||||
cdef _ArrayMap priors = <_ArrayMap>node.priors
|
||||
cdef _ArrayMap value_sum = <_ArrayMap>node.value_sum
|
||||
for action in legal_actions:
|
||||
total_visits += visits.get_int(action, 0) + virtual_visits.get_int(action, 0)
|
||||
sqrt_total = math.sqrt(max(1, total_visits))
|
||||
for action in legal_actions:
|
||||
n = visits.get_int(action, 0)
|
||||
virtual = virtual_visits.get_int(action, 0)
|
||||
n_eff = n + virtual
|
||||
prior = priors.get_float(action, 0.0)
|
||||
if n_eff <= 0:
|
||||
q_eff = 0.0
|
||||
else:
|
||||
q_eff = (
|
||||
value_sum.get_float(action, 0.0)
|
||||
- virtual * float(self.config.virtual_loss_value)
|
||||
) / n_eff
|
||||
score = q_eff + float(self.config.c_puct) * prior * sqrt_total / (1 + n_eff)
|
||||
if score > best_score:
|
||||
best_score = score
|
||||
best_action = action
|
||||
return int(best_action)
|
||||
|
||||
cpdef _backup(self, list path, double leaf_value, int leaf_player):
|
||||
cdef double value = float(leaf_value)
|
||||
cdef int value_player = int(leaf_player)
|
||||
cdef SearchPathEntry entry
|
||||
cdef double parent_value
|
||||
cdef long current_virtual
|
||||
cdef _ArrayMap visits
|
||||
cdef _ArrayMap virtual_visits
|
||||
cdef _ArrayMap value_sum
|
||||
for entry in reversed(path):
|
||||
parent_value = value if value_player == entry.parent_player else -value
|
||||
virtual_visits = <_ArrayMap>entry.node.virtual_visits
|
||||
visits = <_ArrayMap>entry.node.visits
|
||||
value_sum = <_ArrayMap>entry.node.value_sum
|
||||
current_virtual = virtual_visits.get_int(entry.action, 0)
|
||||
virtual_visits.set_int(entry.action, max(0, current_virtual - 1))
|
||||
visits.set_int(entry.action, visits.get_int(entry.action, 0) + 1)
|
||||
value_sum.set_float(
|
||||
entry.action,
|
||||
value_sum.get_float(entry.action, 0.0) + parent_value,
|
||||
)
|
||||
value = parent_value
|
||||
value_player = entry.parent_player
|
||||
|
||||
cpdef _release_virtual_path(self, list path):
|
||||
cdef SearchPathEntry entry
|
||||
cdef _ArrayMap virtual_visits
|
||||
cdef long current_virtual
|
||||
for entry in path:
|
||||
virtual_visits = <_ArrayMap>entry.node.virtual_visits
|
||||
current_virtual = virtual_visits.get_int(entry.action, 0)
|
||||
virtual_visits.set_int(entry.action, max(0, current_virtual - 1))
|
||||
|
||||
cpdef object _rollout_value(self, GameState state, int player):
|
||||
cdef int steps = 0
|
||||
cdef int actions[MAX_ACTIONS]
|
||||
cdef int count
|
||||
cdef int unified_action
|
||||
cdef int action
|
||||
cdef int max_depth = int(self.config.max_depth)
|
||||
while not state.terminal and steps < max_depth:
|
||||
if self._rollout_bot is not None:
|
||||
action = self._rollout_bot.act_cython(state)
|
||||
if not state._is_legal_action_c(action):
|
||||
count = state._unified_legal_actions_c(actions)
|
||||
if count <= 0:
|
||||
break
|
||||
unified_action = actions[self.rng.randrange(count)]
|
||||
action = self._from_unified_action_c(state, unified_action)
|
||||
else:
|
||||
count = state._unified_legal_actions_c(actions)
|
||||
if count <= 0:
|
||||
break
|
||||
unified_action = actions[self.rng.randrange(count)]
|
||||
action = self._from_unified_action_c(state, unified_action)
|
||||
state._push_action_c(action)
|
||||
steps += 1
|
||||
while steps > 0:
|
||||
state._pop_action_c()
|
||||
steps -= 1
|
||||
return float(state.total_scores[player] - state.total_scores[1 - player])
|
||||
@@ -15,6 +15,7 @@ class ReplaySample:
|
||||
v_target: float
|
||||
player: int
|
||||
prior: np.ndarray | None = None
|
||||
game_index: int | None = None
|
||||
|
||||
|
||||
class ReplayBuffer:
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import random
|
||||
import time
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
@@ -15,9 +17,11 @@ from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy
|
||||
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||
|
||||
from .config import IsMctsConfig
|
||||
from .evaluate import evaluate_with_mcts
|
||||
from .interleaved_self_play import play_self_play_iteration
|
||||
from .network import AlphaZeroLogitsView, AlphaZeroNet
|
||||
from .replay_buffer import ReplayBuffer, ReplaySample
|
||||
from .self_play import play_self_play_game
|
||||
from .workers import SelfPlayWorkerBatch, run_self_play_worker
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -112,14 +116,19 @@ class IsMctsTrainer:
|
||||
return metrics
|
||||
|
||||
def run_iteration(self, iteration: int) -> IterationMetrics:
|
||||
print(
|
||||
f"[iter {iteration}] self-play start (workers={self.config.training.num_workers})",
|
||||
flush=True,
|
||||
)
|
||||
self.network.eval()
|
||||
sp_started = time.perf_counter()
|
||||
added = 0
|
||||
iteration_samples: list[ReplaySample] = []
|
||||
for _ in range(self.config.training.games_per_iter):
|
||||
samples = play_self_play_game(
|
||||
if self.config.training.num_workers > 1:
|
||||
iteration_samples = self._run_self_play_parallel(iteration)
|
||||
else:
|
||||
iteration_samples = play_self_play_iteration(
|
||||
self.network,
|
||||
self.config.mcts,
|
||||
self.config.training,
|
||||
self.game_config,
|
||||
self.rng,
|
||||
device=self.device,
|
||||
@@ -127,10 +136,13 @@ class IsMctsTrainer:
|
||||
temperature=self.config.temperature.training,
|
||||
max_steps=self.config.evaluation.max_steps,
|
||||
)
|
||||
self.buffer.add(samples)
|
||||
iteration_samples.extend(samples)
|
||||
added += len(samples)
|
||||
self.buffer.add(iteration_samples)
|
||||
added = len(iteration_samples)
|
||||
self_play_seconds = time.perf_counter() - sp_started
|
||||
print(
|
||||
f"[iter {iteration}] self-play done in {self_play_seconds:.1f}s, {added} samples",
|
||||
flush=True,
|
||||
)
|
||||
mcts_metrics = self._compute_mcts_metrics(iteration_samples)
|
||||
|
||||
train_started = time.perf_counter()
|
||||
@@ -140,7 +152,12 @@ class IsMctsTrainer:
|
||||
losses.append(self._train_batch(batch))
|
||||
train_seconds = time.perf_counter() - train_started
|
||||
loss_arr = np.asarray(losses, dtype=np.float64)
|
||||
print(f"[iter {iteration}] train done in {train_seconds:.1f}s, eval starting", flush=True)
|
||||
eval_started = time.perf_counter()
|
||||
eval_metrics = self._evaluate(iteration)
|
||||
print(
|
||||
f"[iter {iteration}] eval done in {time.perf_counter() - eval_started:.1f}s", flush=True
|
||||
)
|
||||
return IterationMetrics(
|
||||
iteration=iteration,
|
||||
samples_added=added,
|
||||
@@ -154,6 +171,69 @@ class IsMctsTrainer:
|
||||
mcts_metrics=mcts_metrics,
|
||||
)
|
||||
|
||||
def _run_self_play_parallel(self, iteration: int) -> list[ReplaySample]:
|
||||
training_cfg = self.config.training
|
||||
num_workers = max(1, int(training_cfg.num_workers))
|
||||
total_games = int(training_cfg.games_per_iter)
|
||||
if total_games <= 0:
|
||||
return []
|
||||
# Split games across workers as evenly as possible.
|
||||
effective_workers = min(num_workers, total_games)
|
||||
base = total_games // effective_workers
|
||||
remainder = total_games % effective_workers
|
||||
per_worker = [base + (1 if i < remainder else 0) for i in range(effective_workers)]
|
||||
# Move network state dict to CPU for cross-process transfer.
|
||||
cpu_state = {
|
||||
name: tensor.detach().cpu() for name, tensor in self.network.state_dict().items()
|
||||
}
|
||||
config_dict = self.config.to_dict()
|
||||
game_snapshot = self.game_config.to_snapshot()
|
||||
max_steps = self.config.evaluation.max_steps
|
||||
temperature = self.config.temperature.training
|
||||
worker_device = str(self.config.training.worker_device)
|
||||
batches: list[SelfPlayWorkerBatch] = []
|
||||
for worker_index in range(effective_workers):
|
||||
seed = self.rng.randrange(2**31)
|
||||
batches.append(
|
||||
SelfPlayWorkerBatch(
|
||||
worker_index=worker_index,
|
||||
games_for_worker=per_worker[worker_index],
|
||||
base_seed=seed,
|
||||
config=config_dict,
|
||||
game_config=game_snapshot,
|
||||
network_state=cpu_state,
|
||||
temperature=temperature,
|
||||
max_steps=max_steps,
|
||||
device=worker_device,
|
||||
)
|
||||
)
|
||||
samples: list[ReplaySample] = []
|
||||
ctx = mp.get_context("spawn")
|
||||
print(
|
||||
f" spawning {effective_workers} workers for {total_games} games "
|
||||
f"(per_worker={per_worker})...",
|
||||
flush=True,
|
||||
)
|
||||
spawn_started = time.perf_counter()
|
||||
with ProcessPoolExecutor(max_workers=effective_workers, mp_context=ctx) as executor:
|
||||
futures = [executor.submit(run_self_play_worker, batch) for batch in batches]
|
||||
print(
|
||||
f" workers submitted in {time.perf_counter() - spawn_started:.1f}s, waiting for results...",
|
||||
flush=True,
|
||||
)
|
||||
results = []
|
||||
for future in futures:
|
||||
res = future.result()
|
||||
results.append(res)
|
||||
print(
|
||||
f" worker {res.worker_index} done ({len(res.samples)} samples, "
|
||||
f"elapsed {time.perf_counter() - spawn_started:.1f}s)",
|
||||
flush=True,
|
||||
)
|
||||
for result in sorted(results, key=lambda item: item.worker_index):
|
||||
samples.extend(result.samples)
|
||||
return samples
|
||||
|
||||
def _train_batch(self, batch: list[ReplaySample]) -> tuple[float, float, float]:
|
||||
self.network.train()
|
||||
info = torch.as_tensor(
|
||||
@@ -179,7 +259,8 @@ class IsMctsTrainer:
|
||||
logits, value_pred = self.network(info, legal)
|
||||
log_probs = torch.log_softmax(logits, dim=-1)
|
||||
policy_loss = -(pi * log_probs).sum(dim=-1).mean()
|
||||
value_loss = nn.functional.mse_loss(value_pred, value_target)
|
||||
v_scale = float(self.network.value_scale)
|
||||
value_loss = nn.functional.mse_loss(value_pred / v_scale, value_target / v_scale)
|
||||
loss = policy_loss + value_loss
|
||||
self.optimizer.zero_grad(set_to_none=True)
|
||||
loss.backward()
|
||||
@@ -204,22 +285,52 @@ class IsMctsTrainer:
|
||||
return {}
|
||||
self.network.eval()
|
||||
results: dict[str, float | int] = {}
|
||||
logits_view = AlphaZeroLogitsView(self.network)
|
||||
for opponent in opponents:
|
||||
result = evaluate_strategy_network(
|
||||
logits_view,
|
||||
self.game_config,
|
||||
games=self.config.evaluation.games,
|
||||
seed=self.config.run.seed + iteration * 1000,
|
||||
opponent=opponent,
|
||||
device=self.device,
|
||||
encoding=self.config.encoding,
|
||||
max_steps=self.config.evaluation.max_steps,
|
||||
batch_size=self.config.evaluation.batch_size,
|
||||
)
|
||||
print(f" eval vs {opponent}...", flush=True)
|
||||
opp_started = time.perf_counter()
|
||||
if self.config.mcts.eval_with_mcts:
|
||||
eval_mcts_cfg = self.config.mcts.model_copy()
|
||||
if self.config.mcts.eval_n_simulations > 0:
|
||||
eval_mcts_cfg = eval_mcts_cfg.model_copy(
|
||||
update={"n_simulations": self.config.mcts.eval_n_simulations}
|
||||
)
|
||||
result = evaluate_with_mcts(
|
||||
self.network,
|
||||
self.game_config,
|
||||
eval_mcts_cfg,
|
||||
games=self.config.evaluation.games,
|
||||
seed=self.config.run.seed + iteration * 1000,
|
||||
opponent=opponent,
|
||||
device=self.device,
|
||||
encoding=self.config.encoding,
|
||||
max_steps=self.config.evaluation.max_steps,
|
||||
config=self.config,
|
||||
num_workers=self.config.training.num_workers,
|
||||
)
|
||||
else:
|
||||
logits_view = AlphaZeroLogitsView(self.network)
|
||||
result = evaluate_strategy_network(
|
||||
logits_view,
|
||||
self.game_config,
|
||||
games=self.config.evaluation.games,
|
||||
seed=self.config.run.seed + iteration * 1000,
|
||||
opponent=opponent,
|
||||
device=self.device,
|
||||
encoding=self.config.encoding,
|
||||
max_steps=self.config.evaluation.max_steps,
|
||||
batch_size=self.config.evaluation.batch_size,
|
||||
)
|
||||
key = opponent.replace("-", "_")
|
||||
for metric_key, value in result.items():
|
||||
results[f"eval/{key}/{metric_key}"] = value
|
||||
par = result.get("play_action_rate", 0.0)
|
||||
sd = result.get("avg_score_diff0", 0.0)
|
||||
wr = result.get("win_rate0", 0.0)
|
||||
print(
|
||||
f" eval vs {opponent} done in {time.perf_counter() - opp_started:.1f}s "
|
||||
f"PA={par:.2f} W={wr:.2f} S={sd:.1f}",
|
||||
flush=True,
|
||||
)
|
||||
return results
|
||||
|
||||
def _compute_mcts_metrics(self, samples: list[ReplaySample]) -> dict[str, float]:
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Multi-process self-play workers for ISMCTS.
|
||||
|
||||
Mirrors the Deep CFR pattern: a ``ProcessPoolExecutor`` (spawn context)
|
||||
runs N worker processes, each receiving the current network state dict and
|
||||
a slice of the iteration's self-play games. Workers run network inference
|
||||
on CPU by default (small policy/value MLP, GPU contention is the bottleneck
|
||||
when sharing a single device across many workers).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim
|
||||
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||
|
||||
from .config import IsMctsConfig, config_from_dict
|
||||
from .interleaved_self_play import play_self_play_iteration
|
||||
from .network import AlphaZeroNet
|
||||
from .replay_buffer import ReplaySample
|
||||
|
||||
_TORCH_THREADS_CONFIGURED = False
|
||||
|
||||
|
||||
def _configure_worker_torch_threads() -> None:
|
||||
global _TORCH_THREADS_CONFIGURED
|
||||
if _TORCH_THREADS_CONFIGURED:
|
||||
return
|
||||
os.environ.setdefault("OMP_NUM_THREADS", "1")
|
||||
os.environ.setdefault("MKL_NUM_THREADS", "1")
|
||||
torch.set_num_threads(1)
|
||||
if hasattr(torch, "set_num_interop_threads"):
|
||||
try:
|
||||
torch.set_num_interop_threads(1)
|
||||
except RuntimeError:
|
||||
pass
|
||||
_TORCH_THREADS_CONFIGURED = True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SelfPlayWorkerBatch:
|
||||
worker_index: int
|
||||
games_for_worker: int
|
||||
base_seed: int
|
||||
config: dict[str, Any]
|
||||
game_config: dict[str, Any]
|
||||
network_state: dict[str, Any]
|
||||
temperature: float
|
||||
max_steps: int
|
||||
device: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class SelfPlayWorkerResult:
|
||||
worker_index: int
|
||||
samples: list[ReplaySample]
|
||||
|
||||
|
||||
def run_self_play_worker(batch: SelfPlayWorkerBatch) -> SelfPlayWorkerResult:
|
||||
import time as _time
|
||||
|
||||
_t0 = _time.perf_counter()
|
||||
print(f" [worker {batch.worker_index}] starting ({batch.games_for_worker} games)", flush=True)
|
||||
_configure_worker_torch_threads()
|
||||
cfg: IsMctsConfig = config_from_dict(batch.config)
|
||||
game_config = LostCitiesConfig(**batch.game_config)
|
||||
device = torch.device(batch.device)
|
||||
probe = GameState.new_game(game_config, seed=batch.base_seed)
|
||||
in_dim = input_dim(probe, cfg.encoding)
|
||||
action_size = probe.action_size
|
||||
network = AlphaZeroNet.from_config(in_dim, action_size, cfg).to(device)
|
||||
network.load_state_dict(batch.network_state)
|
||||
network.eval()
|
||||
print(
|
||||
f" [worker {batch.worker_index}] init done in {_time.perf_counter() - _t0:.1f}s, self-play start",
|
||||
flush=True,
|
||||
)
|
||||
# Build a per-worker TrainingConfig with the worker's game count.
|
||||
worker_training = cfg.training.model_copy(
|
||||
update={"games_per_iter": int(batch.games_for_worker)}
|
||||
)
|
||||
rng = random.Random(batch.base_seed)
|
||||
_sp_t0 = _time.perf_counter()
|
||||
samples = play_self_play_iteration(
|
||||
network,
|
||||
cfg.mcts,
|
||||
worker_training,
|
||||
game_config,
|
||||
rng,
|
||||
device=device,
|
||||
encoding=cfg.encoding,
|
||||
temperature=batch.temperature,
|
||||
max_steps=batch.max_steps,
|
||||
)
|
||||
print(
|
||||
f" [worker {batch.worker_index}] self-play done in {_time.perf_counter() - _sp_t0:.1f}s ({len(samples)} samples)",
|
||||
flush=True,
|
||||
)
|
||||
return SelfPlayWorkerResult(worker_index=batch.worker_index, samples=samples)
|
||||
@@ -1,22 +1,55 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import random
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim
|
||||
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
|
||||
|
||||
from coolrl_lost_cities.games.classic.bots.heuristic import HeuristicBot
|
||||
from coolrl_lost_cities.games.classic.bots.heuristic_py import (
|
||||
HeuristicBot as PythonHeuristicBot,
|
||||
)
|
||||
from coolrl_lost_cities.games.classic.ismcts.config import IsMctsConfig, MctsConfig
|
||||
from coolrl_lost_cities.games.classic.ismcts.determinization import sample_determinization
|
||||
from coolrl_lost_cities.games.classic.ismcts.info_set import canonical_info_set_key
|
||||
from coolrl_lost_cities.games.classic.ismcts.mcts import IsMctsSearcher
|
||||
from coolrl_lost_cities.games.classic.ismcts.interleaved_self_play import (
|
||||
play_self_play_iteration,
|
||||
)
|
||||
from coolrl_lost_cities.games.classic.ismcts.mcts import IsMctsSearcher, MctsNode
|
||||
from coolrl_lost_cities.games.classic.ismcts.network import AlphaZeroLogitsView, AlphaZeroNet
|
||||
from coolrl_lost_cities.games.classic.ismcts.replay_buffer import ReplayBuffer, ReplaySample
|
||||
from coolrl_lost_cities.games.classic.ismcts.self_play import play_self_play_game
|
||||
from coolrl_lost_cities.games.classic.ismcts.trainer import IsMctsTrainer
|
||||
|
||||
|
||||
def _python_mcts_searcher():
|
||||
module_name = "coolrl_lost_cities.games.classic.ismcts._mcts_python_baseline"
|
||||
existing = sys.modules.get(module_name)
|
||||
if existing is not None:
|
||||
return existing.IsMctsSearcher
|
||||
path = (
|
||||
Path(__file__).parents[4]
|
||||
/ "src"
|
||||
/ "coolrl_lost_cities"
|
||||
/ "games"
|
||||
/ "classic"
|
||||
/ "ismcts"
|
||||
/ "mcts.py"
|
||||
)
|
||||
spec = importlib.util.spec_from_file_location(module_name, path)
|
||||
assert spec is not None
|
||||
assert spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
sys.modules[module_name] = module
|
||||
spec.loader.exec_module(module)
|
||||
return module.IsMctsSearcher
|
||||
|
||||
|
||||
def mini_config(seed: int = 1) -> LostCitiesConfig:
|
||||
return LostCitiesConfig(
|
||||
n_colors=3,
|
||||
@@ -97,6 +130,143 @@ def test_mcts_prior_drives_visits() -> None:
|
||||
assert visits[favored] == max(visits.values())
|
||||
|
||||
|
||||
def test_search_correctness_vs_sequential() -> None:
|
||||
for n_sims in (8, 16, 64):
|
||||
state = GameState.new_game(mini_config(), seed=17)
|
||||
dim = input_dim(state)
|
||||
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=1)
|
||||
config = MctsConfig(
|
||||
n_simulations=n_sims,
|
||||
parallel_simulations=1,
|
||||
use_rollout_value=False,
|
||||
)
|
||||
left = IsMctsSearcher(net, config, rng=random.Random(18))
|
||||
right = IsMctsSearcher(net, config, rng=random.Random(18))
|
||||
assert left.search(state, state.current_player) == right.search(state, state.current_player)
|
||||
|
||||
|
||||
def test_cython_sequential_matches_python_sequential_visit_counts() -> None:
|
||||
PythonIsMctsSearcher = _python_mcts_searcher()
|
||||
for n_sims in (8, 32, 128):
|
||||
state = GameState.new_game(mini_config(), seed=23)
|
||||
dim = input_dim(state)
|
||||
torch.manual_seed(24)
|
||||
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=1)
|
||||
config = MctsConfig(
|
||||
n_simulations=n_sims,
|
||||
parallel_simulations=1,
|
||||
use_rollout_value=False,
|
||||
)
|
||||
python_searcher = PythonIsMctsSearcher(net, config, rng=random.Random(25))
|
||||
cython_searcher = IsMctsSearcher(net, config, rng=random.Random(25))
|
||||
|
||||
assert cython_searcher.search(state, state.current_player) == python_searcher.search(
|
||||
state, state.current_player
|
||||
)
|
||||
|
||||
|
||||
def test_search_visit_counts_match_with_parallel_simulations() -> None:
|
||||
for n_sims in (8, 32, 128):
|
||||
state = GameState.new_game(mini_config(), seed=26)
|
||||
dim = input_dim(state)
|
||||
torch.manual_seed(27)
|
||||
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=1)
|
||||
sequential = IsMctsSearcher(
|
||||
net,
|
||||
MctsConfig(n_simulations=n_sims, parallel_simulations=1, use_rollout_value=False),
|
||||
rng=random.Random(28),
|
||||
)
|
||||
batched = IsMctsSearcher(
|
||||
net,
|
||||
MctsConfig(n_simulations=n_sims, parallel_simulations=8, use_rollout_value=False),
|
||||
rng=random.Random(28),
|
||||
)
|
||||
|
||||
assert batched.search(state, state.current_player) == sequential.search(
|
||||
state, state.current_player
|
||||
)
|
||||
|
||||
|
||||
def test_search_with_virtual_loss_diversity() -> None:
|
||||
state = GameState.new_game(mini_config(), seed=19)
|
||||
dim = input_dim(state)
|
||||
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=0)
|
||||
for param in net.parameters():
|
||||
param.data.zero_()
|
||||
searcher = IsMctsSearcher(
|
||||
net,
|
||||
MctsConfig(n_simulations=64, parallel_simulations=4, virtual_loss_value=1.0),
|
||||
rng=random.Random(20),
|
||||
)
|
||||
first = searcher.prepare_simulation_batch(state, state.current_player, 1)
|
||||
searcher.evaluate_and_backup(first)
|
||||
|
||||
pending = searcher.prepare_simulation_batch(state, state.current_player, 4)
|
||||
first_actions = [item.path[0].action for item in pending if item.path]
|
||||
assert len(set(first_actions)) >= 2
|
||||
|
||||
|
||||
def test_heuristic_cython_fast_path_matches_python_for_random_states() -> None:
|
||||
configs = [mini_config(seed=31), LostCitiesConfig(seed=32)]
|
||||
py_bot = PythonHeuristicBot()
|
||||
cy_bot = HeuristicBot()
|
||||
|
||||
for config in configs:
|
||||
rng = random.Random(33)
|
||||
checked = 0
|
||||
attempts = 0
|
||||
while checked < 100 and attempts < 1000:
|
||||
attempts += 1
|
||||
state = GameState.new_game(config, seed=rng.randrange(2**31))
|
||||
for _ in range(rng.randrange(40)):
|
||||
if state.terminal:
|
||||
break
|
||||
legal = state.unified_legal_actions()
|
||||
if not legal:
|
||||
break
|
||||
state.apply_unified_action(rng.choice(legal))
|
||||
if state.terminal or not state.unified_legal_actions():
|
||||
continue
|
||||
|
||||
assert cy_bot.act_cython(state) == py_bot.act(state)
|
||||
checked += 1
|
||||
|
||||
assert checked == 100
|
||||
|
||||
|
||||
def test_game_state_push_pop_unified_round_trip_snapshot() -> None:
|
||||
rng = random.Random(34)
|
||||
for config in (mini_config(seed=35), LostCitiesConfig(seed=36)):
|
||||
state = GameState.new_game(config, seed=37)
|
||||
for _ in range(100):
|
||||
if state.terminal:
|
||||
break
|
||||
before = state.to_snapshot()
|
||||
unified = rng.choice(state.unified_legal_actions())
|
||||
local = state.from_unified_action(unified)
|
||||
state.push_action(local)
|
||||
state.pop_action()
|
||||
assert state.to_snapshot() == before
|
||||
state.apply_unified_action(unified)
|
||||
|
||||
|
||||
def test_mcts_node_c_array_maps_are_dict_like() -> None:
|
||||
node = MctsNode(b"root", player=0, action_size=16)
|
||||
node.priors[3] = 0.25
|
||||
node.visits.setdefault(3, 0)
|
||||
node.value_sum[3] = 1.5
|
||||
node.virtual_visits[3] = 2
|
||||
node.visits[3] = node.visits.get(3, 0) + 4
|
||||
|
||||
assert bool(node.priors)
|
||||
assert node.priors.get(3, 0.0) == 0.25
|
||||
assert node.visits.get(3, 0) == 4
|
||||
assert node.value_sum[3] == 1.5
|
||||
assert node.virtual_visits[3] == 2
|
||||
assert 3 in node.visits
|
||||
assert dict(node.visits.items()) == {3: 4}
|
||||
|
||||
|
||||
def test_replay_buffer_capacity_and_sample() -> None:
|
||||
sample = ReplaySample(
|
||||
info_state=np.zeros(4, dtype=np.float32),
|
||||
@@ -127,6 +297,29 @@ def test_self_play_game_returns_signed_targets() -> None:
|
||||
assert all(sample.prior is not None for sample in samples)
|
||||
|
||||
|
||||
def test_interleaved_self_play_yields_complete_games() -> None:
|
||||
config = mini_config()
|
||||
state = GameState.new_game(config, seed=21)
|
||||
net = AlphaZeroNet(input_dim(state), state.action_size, hidden_size=8, num_layers=1)
|
||||
ismcts_config = IsMctsConfig.model_validate(
|
||||
{
|
||||
"mcts": {"n_simulations": 2, "parallel_simulations": 2},
|
||||
"training": {"games_per_iter": 4, "interleave_games": 4, "interleave_max_batch": 16},
|
||||
}
|
||||
)
|
||||
samples = play_self_play_iteration(
|
||||
net,
|
||||
ismcts_config.mcts,
|
||||
ismcts_config.training,
|
||||
config,
|
||||
random.Random(22),
|
||||
max_steps=80,
|
||||
)
|
||||
assert samples
|
||||
assert {sample.game_index for sample in samples} == {0, 1, 2, 3}
|
||||
assert all(np.isfinite(sample.v_target) for sample in samples)
|
||||
|
||||
|
||||
def test_trainer_one_iteration_smoke(tmp_path) -> None:
|
||||
config = IsMctsConfig.model_validate(
|
||||
{
|
||||
@@ -185,9 +378,9 @@ def test_trainer_emits_full_eval_metrics(tmp_path) -> None:
|
||||
run_dir=tmp_path,
|
||||
)
|
||||
metrics = trainer.train()[0].to_dict()
|
||||
assert "eval/random/avg_opened_colors" in metrics
|
||||
assert "eval/random/bad_open_rate" in metrics
|
||||
assert "eval/random/per_game_negative_expeditions" in metrics
|
||||
assert "eval/random/avg_score_diff0" in metrics
|
||||
assert "eval/random/play_action_rate" in metrics
|
||||
assert "eval/random/win_rate0" in metrics
|
||||
|
||||
|
||||
def test_trainer_emits_mcts_metrics(tmp_path) -> None:
|
||||
@@ -221,3 +414,39 @@ def test_trainer_emits_mcts_metrics(tmp_path) -> None:
|
||||
):
|
||||
assert key in metrics
|
||||
assert np.isfinite(metrics[key])
|
||||
|
||||
|
||||
def test_smoke_iter_with_batching(tmp_path) -> None:
|
||||
config = IsMctsConfig.model_validate(
|
||||
{
|
||||
"run": {"max_iterations": 1, "seed": 14, "device": "cpu"},
|
||||
"rules": {
|
||||
"n_colors": 3,
|
||||
"n_ranks": 5,
|
||||
"n_handshakes": 1,
|
||||
"hand_size": 4,
|
||||
"bonus_threshold": 4,
|
||||
},
|
||||
"network": {"hidden_size": 16, "num_layers": 1},
|
||||
"mcts": {"n_simulations": 4, "parallel_simulations": 4},
|
||||
"training": {
|
||||
"games_per_iter": 1,
|
||||
"gradient_steps_per_iter": 1,
|
||||
"batch_size": 8,
|
||||
"interleave_games": 4,
|
||||
"interleave_max_batch": 16,
|
||||
},
|
||||
"checkpoint": {"save_every": 0},
|
||||
"evaluation": {"eval_every": 0, "num_workers": 1, "max_steps": 80},
|
||||
}
|
||||
)
|
||||
trainer = IsMctsTrainer(
|
||||
config,
|
||||
config.rules.to_lost_cities_config(seed=config.run.seed),
|
||||
run_dir=tmp_path,
|
||||
)
|
||||
metrics = trainer.train()[0].to_dict()
|
||||
assert metrics["samples/added"] > 0
|
||||
assert "mcts/avg_visit_entropy" in metrics
|
||||
assert "mcts/value_prediction_error" in metrics
|
||||
assert "mcts/policy_mcts_kl" in metrics
|
||||
|
||||
Reference in New Issue
Block a user