diff --git a/configs/ismcts/default.yaml b/configs/ismcts/default.yaml index 5885b38..d76fc7b 100644 --- a/configs/ismcts/default.yaml +++ b/configs/ismcts/default.yaml @@ -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 diff --git a/configs/ismcts/mini.yaml b/configs/ismcts/mini.yaml index 4844ff4..095e6cc 100644 --- a/configs/ismcts/mini.yaml +++ b/configs/ismcts/mini.yaml @@ -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 diff --git a/pyproject.toml b/pyproject.toml index cf705a5..3491e1e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 diff --git a/setup.py b/setup.py index 158247a..6852ce1 100644 --- a/setup.py +++ b/setup.py @@ -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={ diff --git a/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pxd b/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pxd new file mode 100644 index 0000000..d9b8372 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pxd @@ -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 * diff --git a/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pyx b/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pyx index 1f38e0a..05326ef 100644 --- a/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pyx +++ b/src/coolrl_lost_cities/games/classic/bots/heuristic_cy.pyx @@ -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 += (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 += (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 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 diff --git a/src/coolrl_lost_cities/games/classic/ismcts/config.py b/src/coolrl_lost_cities/games/classic/ismcts/config.py index d5c6fb4..0b849fa 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/config.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/config.py @@ -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: diff --git a/src/coolrl_lost_cities/games/classic/ismcts/eval_worker.py b/src/coolrl_lost_cities/games/classic/ismcts/eval_worker.py new file mode 100644 index 0000000..f4a3276 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/ismcts/eval_worker.py @@ -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, + ) diff --git a/src/coolrl_lost_cities/games/classic/ismcts/evaluate.py b/src/coolrl_lost_cities/games/classic/ismcts/evaluate.py new file mode 100644 index 0000000..a392330 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/ismcts/evaluate.py @@ -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, + } diff --git a/src/coolrl_lost_cities/games/classic/ismcts/info_set.py b/src/coolrl_lost_cities/games/classic/ismcts/info_set.py index 008dd25..b988426 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/info_set.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/info_set.py @@ -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]: diff --git a/src/coolrl_lost_cities/games/classic/ismcts/interleaved_self_play.py b/src/coolrl_lost_cities/games/classic/ismcts/interleaved_self_play.py new file mode 100644 index 0000000..3e4b763 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/ismcts/interleaved_self_play.py @@ -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 diff --git a/src/coolrl_lost_cities/games/classic/ismcts/mcts.py b/src/coolrl_lost_cities/games/classic/ismcts/mcts.py index 6ae2f2a..749606a 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/mcts.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/mcts.py @@ -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)) diff --git a/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx b/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx new file mode 100644 index 0000000..f142a70 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx @@ -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] = value + elif self.kind == 3: + self.node.virtual_visits_arr[action] = 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 = ( + 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]) diff --git a/src/coolrl_lost_cities/games/classic/ismcts/replay_buffer.py b/src/coolrl_lost_cities/games/classic/ismcts/replay_buffer.py index a55aa33..d254b5e 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/replay_buffer.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/replay_buffer.py @@ -15,6 +15,7 @@ class ReplaySample: v_target: float player: int prior: np.ndarray | None = None + game_index: int | None = None class ReplayBuffer: diff --git a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py index ed52112..2f263b7 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/trainer.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/trainer.py @@ -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]: diff --git a/src/coolrl_lost_cities/games/classic/ismcts/workers.py b/src/coolrl_lost_cities/games/classic/ismcts/workers.py new file mode 100644 index 0000000..28d2ccc --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/ismcts/workers.py @@ -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) diff --git a/tests/games/classic/ismcts/test_ismcts.py b/tests/games/classic/ismcts/test_ismcts.py index 644ccea..a8b19da 100644 --- a/tests/games/classic/ismcts/test_ismcts.py +++ b/tests/games/classic/ismcts/test_ismcts.py @@ -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