diff --git a/src/coolrl_lost_cities/games/classic/ismcts/config.py b/src/coolrl_lost_cities/games/classic/ismcts/config.py index 9b9cfcd..dfce91c 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/config.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/config.py @@ -38,6 +38,12 @@ class MctsConfig(StrictModel): # bad backup permanently kills an action. Setting q_scale=100 normalizes Q # to ~[-1, 1] (consistent with AlphaZero's convention). q_scale: float = 100.0 + # Opponent-aware search: when set, the search tree treats the opponent + # seat as a fixed external policy (heuristic bot) instead of expanding it + # with the network's priors/value. Used during mixed-opponent self-play + # so root visit distributions reflect the *actual* opponent the trainee + # faces. Bot name is taken from training.mixed_opponent_bot. + opponent_aware_search: bool = False @field_validator("n_simulations", "max_depth", "parallel_simulations") @classmethod @@ -90,6 +96,15 @@ class TrainingConfig(StrictModel): md_target_alpha_start: float = 0.3 md_target_alpha_end: float = 0.8 md_target_alpha_iters: int = 500 + # Mixed-opponent self-play: a fraction of games per iteration are played + # against a fixed external bot instead of the current network. Only the + # trainee's decisions are stored as policy targets; opponent moves are + # taken by `mixed_opponent_bot.act(state)`. Combined with + # mcts.opponent_aware_search, the MCTS tree models the opponent as that + # same bot so root-visit distributions reflect the real opponent. + # Set fraction=0 to disable (pure self-play). + mixed_opponent_fraction: float = 0.0 + mixed_opponent_bot: str = "heuristic-balanced" @field_validator( "games_per_iter", 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 index 583e5b4..a8044a6 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/interleaved_self_play.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/interleaved_self_play.py @@ -6,6 +6,7 @@ from dataclasses import dataclass, field import numpy as np import torch +from coolrl_lost_cities.games.classic.bots.registry import build_bot from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig @@ -34,6 +35,11 @@ class _GameContext: game_index: int decisions: list[_PendingDecision] = field(default_factory=list) steps: int = 0 + # Mixed-opponent setup: when traverser_seat is not None, only that seat + # uses MCTS+network; the other seat is played by `opponent_bot`. None for + # pure self-play games (both seats use MCTS). + traverser_seat: int | None = None + opponent_bot: object | None = None @dataclass @@ -61,15 +67,28 @@ def play_self_play_iteration( active: list[_GameContext] = [] started = 0 target_games = training_config.games_per_iter + mixed_fraction = float(training_config.mixed_opponent_fraction) def fill_active() -> None: nonlocal started while len(active) < training_config.interleave_games and started < target_games: + traverser_seat: int | None = None + opponent_bot = None + if mixed_fraction > 0.0 and rng.random() < mixed_fraction: + # Alternate trainee seat so MCTS sees both first- and + # second-player perspectives equally. + traverser_seat = started % 2 + opponent_bot = build_bot( + training_config.mixed_opponent_bot, + seed=rng.randrange(2**31), + ) active.append( _GameContext( state=GameState.new_game(game_config, seed=rng.randrange(2**31)), rng=random.Random(rng.randrange(2**31)), game_index=started, + traverser_seat=traverser_seat, + opponent_bot=opponent_bot, ) ) started += 1 @@ -83,6 +102,19 @@ def play_self_play_iteration( completed.append(_finalize_context(context)) continue player = int(context.state.current_player) + # Mixed-opponent: if it's the opponent's turn in a mixed game, + # let the heuristic bot move directly (no MCTS, no sample). + if ( + context.traverser_seat is not None + and context.opponent_bot is not None + and player != context.traverser_seat + ): + phase_action = context.opponent_bot.act(context.state) + unified = context.state.to_unified_action(phase_action) + context.state.apply_unified_action(unified) + context.steps += 1 + still_active.append(context) + continue searcher = IsMctsSearcher( network, mcts_config, @@ -90,6 +122,18 @@ def play_self_play_iteration( encoding=encoding, rng=random.Random(context.rng.randrange(2**31)), ) + # Pass opponent bot into the searcher for opponent-aware + # determinization (search models the real opponent the trainee + # faces, not a self-play mirror). + if ( + mcts_config.opponent_aware_search + and context.traverser_seat is not None + and context.opponent_bot is not None + ): + searcher.set_opponent_bot( + context.opponent_bot, + traverser_seat=context.traverser_seat, + ) jobs.append( _SearchJob( context=context, diff --git a/src/coolrl_lost_cities/games/classic/ismcts/mcts.py b/src/coolrl_lost_cities/games/classic/ismcts/mcts.py index 2b516df..c8bdda1 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/mcts.py +++ b/src/coolrl_lost_cities/games/classic/ismcts/mcts.py @@ -89,6 +89,13 @@ class IsMctsSearcher: self._rollout_bot = ( HeuristicBot() if config.rollout_policy == "heuristic_balanced" else None ) + # Opponent-aware search: see mcts.pyx for the rationale. + self._opponent_bot: object | None = None + self._traverser_seat: int = -1 + + def set_opponent_bot(self, bot: object, *, traverser_seat: int) -> None: + self._opponent_bot = bot + self._traverser_seat = int(traverser_seat) def search( self, @@ -133,34 +140,45 @@ class IsMctsSearcher: # 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 + opponent_aware = self._opponent_bot is not None + trav_seat = self._traverser_seat while True: player = int(state.current_player) if state.terminal or depth >= self.config.max_depth: + leaf_seat = trav_seat if opponent_aware else player return PendingSimulation( path=path, leaf_state=state, leaf_node=None, - leaf_player=player, + leaf_player=leaf_seat, info_state=None, legal_mask=None, legal_actions=[], - terminal_value=float(state.score_diff(player)), + terminal_value=float(state.score_diff(leaf_seat)), ) + if opponent_aware and player != trav_seat: + phase_action = self._opponent_bot.act(state) + unified = state.to_unified_action(phase_action) + state.apply_unified_action(unified) + cached_key = None + depth += 1 + continue 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 + leaf_seat = trav_seat if opponent_aware else player return PendingSimulation( path=path, leaf_state=state, leaf_node=node, - leaf_player=player, + leaf_player=leaf_seat, info_state=None, legal_mask=None, legal_actions=[], - terminal_value=float(state.score_diff(player)), + terminal_value=float(state.score_diff(leaf_seat)), ) return PendingSimulation( path=path, diff --git a/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx b/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx index 28c4041..8185552 100644 --- a/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx +++ b/src/coolrl_lost_cities/games/classic/ismcts/mcts.pyx @@ -296,6 +296,12 @@ cdef class IsMctsSearcher: cdef public MctsTree tree cdef HeuristicBot _rollout_bot cdef int action_size + # Opponent-aware search: when set, the search treats one seat as a fixed + # external policy (heuristic bot). Opponent moves are applied directly + # without entering the tree, and all values are taken from the + # traverser's perspective. None for standard symmetric self-play search. + cdef public object _opponent_bot + cdef public int _traverser_seat def __init__( self, @@ -318,6 +324,12 @@ cdef class IsMctsSearcher: self._rollout_bot = ( PyHeuristicBot() if config.rollout_policy == "heuristic_balanced" else None ) + self._opponent_bot = None + self._traverser_seat = -1 + + def set_opponent_bot(self, object bot, *, int traverser_seat): + self._opponent_bot = bot + self._traverser_seat = traverser_seat cdef inline int _from_unified_action_c(self, GameState state, int action_id) noexcept: cdef int card_action_size = 2 * state.hand_size @@ -390,19 +402,41 @@ cdef class IsMctsSearcher: cdef int actions[MAX_ACTIONS] cdef int action_count cdef int i + cdef bint opponent_aware = self._opponent_bot is not None + cdef int leaf_player_seat + cdef int trav_seat = self._traverser_seat + cdef object phase_action + cdef int unified_action while True: player = state.current_player if state.terminal or depth >= int(self.config.max_depth): + if opponent_aware: + leaf_player_seat = trav_seat + else: + leaf_player_seat = player return PendingSimulation( path=path, leaf_state=state, leaf_node=None, - leaf_player=player, + leaf_player=leaf_player_seat, info_state=None, legal_mask=None, legal_actions=[], - terminal_value=float(state.total_scores[player] - state.total_scores[1 - player]), + terminal_value=float( + state.total_scores[leaf_player_seat] + - state.total_scores[1 - leaf_player_seat] + ), ) + # Opponent-aware: if it's the opponent's turn, let the heuristic + # bot move directly instead of expanding the tree. + if opponent_aware and player != trav_seat: + phase_action = self._opponent_bot.act(state) + unified_action = state.to_unified_action(phase_action) + local_action = self._from_unified_action_c(state, unified_action) + state._push_action_c(local_action) + cached_key = None + depth += 1 + continue if cached_key is None: key = canonical_info_set_key(state, player) else: @@ -413,15 +447,22 @@ cdef class IsMctsSearcher: legal_actions = [actions[i] for i in range(action_count)] if not legal_actions: node.terminal = True + if opponent_aware: + leaf_player_seat = trav_seat + else: + leaf_player_seat = player return PendingSimulation( path=path, leaf_state=state, leaf_node=node, - leaf_player=player, + leaf_player=leaf_player_seat, info_state=None, legal_mask=None, legal_actions=[], - terminal_value=float(state.total_scores[player] - state.total_scores[1 - player]), + terminal_value=float( + state.total_scores[leaf_player_seat] + - state.total_scores[1 - leaf_player_seat] + ), ) return PendingSimulation( path=path,