Add mixed-opponent self-play with opponent-aware MCTS
C13/C14 cycles: heuristic-balanced bot plays a configurable fraction of self-play games (training.mixed_opponent_fraction). Trainee turns are stored as policy samples; opponent turns are taken by the bot directly and not stored. When mcts.opponent_aware_search is set, the MCTS tree also treats the opponent seat as that bot — opponent moves are applied without expanding into the search tree, and all values are taken from the traverser's perspective. This was Codex's top recommendation for breaking the symmetric self-play weak fixed point. Empirical: opponent-aware mixed self-play does NOT lift win-rate above BC pretrain (vs heuristic-cautious 100-game eval): C13 (mixed=0.5, no KL): 0/100 — catastrophic forgetting C14 (mixed=0.2, KL beta=1): 17/100 — preserved BC, no improvement BC pretrain baseline: 21/100 Combined with C10-C12 results, BC remains the ceiling under our compute budget (1 GPU + 50 sims + 768x4 net). Code is left in place as configurable dials for future runs with more compute.
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 = (
|
||||
<HeuristicBot>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,
|
||||
|
||||
Reference in New Issue
Block a user