Add multi-process self-play, eval workers, MCTS Cython port

Key changes for ISMCTS speed and correctness:
- Cython port: HeuristicBot helpers (`heuristic_cy.pyx` + new `.pxd`) and
  ISMCTS searcher (`mcts.pyx`) now run as cdef. Both share a fast
  unified-action path through GameState's C interface to avoid Python
  round-trips on hot rollout/tree-walk paths.
- Multi-process self-play and eval: `workers.py`, `eval_worker.py`,
  `interleaved_self_play.py`, plus trainer wiring with ProcessPoolExecutor
  + spawn context. Eval inside `evaluate.py` is parallel per opponent.
- ISMCTS-specific eval (`evaluate.py`) runs MCTS at decision time so the
  metric matches deploy mode; `evaluation.eval_with_mcts` flag preserves
  backwards-compatible policy-only eval when needed.
- Trainer logs progress per phase (self-play start/done, eval per
  opponent), and value loss is now scaled by `value_scale` so policy and
  value losses sit on comparable magnitudes.
- Compact info-set key (`info_set.py`) using packed-struct format and
  child-key reuse during MCTS descent to cut per-step canonicalization.

Tests: 19 ISMCTS suite passing, including parity (Cython-vs-Python
sequential, batched-vs-sequential visit counts, push/pop round-trip).

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