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
+233 -4
View File
@@ -1,22 +1,55 @@
from __future__ import annotations
import importlib.util
import random
import sys
from pathlib import Path
import numpy as np
import torch
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
from coolrl_lost_cities.games.classic.bots.heuristic import HeuristicBot
from coolrl_lost_cities.games.classic.bots.heuristic_py import (
HeuristicBot as PythonHeuristicBot,
)
from coolrl_lost_cities.games.classic.ismcts.config import IsMctsConfig, MctsConfig
from coolrl_lost_cities.games.classic.ismcts.determinization import sample_determinization
from coolrl_lost_cities.games.classic.ismcts.info_set import canonical_info_set_key
from coolrl_lost_cities.games.classic.ismcts.mcts import IsMctsSearcher
from coolrl_lost_cities.games.classic.ismcts.interleaved_self_play import (
play_self_play_iteration,
)
from coolrl_lost_cities.games.classic.ismcts.mcts import IsMctsSearcher, MctsNode
from coolrl_lost_cities.games.classic.ismcts.network import AlphaZeroLogitsView, AlphaZeroNet
from coolrl_lost_cities.games.classic.ismcts.replay_buffer import ReplayBuffer, ReplaySample
from coolrl_lost_cities.games.classic.ismcts.self_play import play_self_play_game
from coolrl_lost_cities.games.classic.ismcts.trainer import IsMctsTrainer
def _python_mcts_searcher():
module_name = "coolrl_lost_cities.games.classic.ismcts._mcts_python_baseline"
existing = sys.modules.get(module_name)
if existing is not None:
return existing.IsMctsSearcher
path = (
Path(__file__).parents[4]
/ "src"
/ "coolrl_lost_cities"
/ "games"
/ "classic"
/ "ismcts"
/ "mcts.py"
)
spec = importlib.util.spec_from_file_location(module_name, path)
assert spec is not None
assert spec.loader is not None
module = importlib.util.module_from_spec(spec)
sys.modules[module_name] = module
spec.loader.exec_module(module)
return module.IsMctsSearcher
def mini_config(seed: int = 1) -> LostCitiesConfig:
return LostCitiesConfig(
n_colors=3,
@@ -97,6 +130,143 @@ def test_mcts_prior_drives_visits() -> None:
assert visits[favored] == max(visits.values())
def test_search_correctness_vs_sequential() -> None:
for n_sims in (8, 16, 64):
state = GameState.new_game(mini_config(), seed=17)
dim = input_dim(state)
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=1)
config = MctsConfig(
n_simulations=n_sims,
parallel_simulations=1,
use_rollout_value=False,
)
left = IsMctsSearcher(net, config, rng=random.Random(18))
right = IsMctsSearcher(net, config, rng=random.Random(18))
assert left.search(state, state.current_player) == right.search(state, state.current_player)
def test_cython_sequential_matches_python_sequential_visit_counts() -> None:
PythonIsMctsSearcher = _python_mcts_searcher()
for n_sims in (8, 32, 128):
state = GameState.new_game(mini_config(), seed=23)
dim = input_dim(state)
torch.manual_seed(24)
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=1)
config = MctsConfig(
n_simulations=n_sims,
parallel_simulations=1,
use_rollout_value=False,
)
python_searcher = PythonIsMctsSearcher(net, config, rng=random.Random(25))
cython_searcher = IsMctsSearcher(net, config, rng=random.Random(25))
assert cython_searcher.search(state, state.current_player) == python_searcher.search(
state, state.current_player
)
def test_search_visit_counts_match_with_parallel_simulations() -> None:
for n_sims in (8, 32, 128):
state = GameState.new_game(mini_config(), seed=26)
dim = input_dim(state)
torch.manual_seed(27)
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=1)
sequential = IsMctsSearcher(
net,
MctsConfig(n_simulations=n_sims, parallel_simulations=1, use_rollout_value=False),
rng=random.Random(28),
)
batched = IsMctsSearcher(
net,
MctsConfig(n_simulations=n_sims, parallel_simulations=8, use_rollout_value=False),
rng=random.Random(28),
)
assert batched.search(state, state.current_player) == sequential.search(
state, state.current_player
)
def test_search_with_virtual_loss_diversity() -> None:
state = GameState.new_game(mini_config(), seed=19)
dim = input_dim(state)
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=0)
for param in net.parameters():
param.data.zero_()
searcher = IsMctsSearcher(
net,
MctsConfig(n_simulations=64, parallel_simulations=4, virtual_loss_value=1.0),
rng=random.Random(20),
)
first = searcher.prepare_simulation_batch(state, state.current_player, 1)
searcher.evaluate_and_backup(first)
pending = searcher.prepare_simulation_batch(state, state.current_player, 4)
first_actions = [item.path[0].action for item in pending if item.path]
assert len(set(first_actions)) >= 2
def test_heuristic_cython_fast_path_matches_python_for_random_states() -> None:
configs = [mini_config(seed=31), LostCitiesConfig(seed=32)]
py_bot = PythonHeuristicBot()
cy_bot = HeuristicBot()
for config in configs:
rng = random.Random(33)
checked = 0
attempts = 0
while checked < 100 and attempts < 1000:
attempts += 1
state = GameState.new_game(config, seed=rng.randrange(2**31))
for _ in range(rng.randrange(40)):
if state.terminal:
break
legal = state.unified_legal_actions()
if not legal:
break
state.apply_unified_action(rng.choice(legal))
if state.terminal or not state.unified_legal_actions():
continue
assert cy_bot.act_cython(state) == py_bot.act(state)
checked += 1
assert checked == 100
def test_game_state_push_pop_unified_round_trip_snapshot() -> None:
rng = random.Random(34)
for config in (mini_config(seed=35), LostCitiesConfig(seed=36)):
state = GameState.new_game(config, seed=37)
for _ in range(100):
if state.terminal:
break
before = state.to_snapshot()
unified = rng.choice(state.unified_legal_actions())
local = state.from_unified_action(unified)
state.push_action(local)
state.pop_action()
assert state.to_snapshot() == before
state.apply_unified_action(unified)
def test_mcts_node_c_array_maps_are_dict_like() -> None:
node = MctsNode(b"root", player=0, action_size=16)
node.priors[3] = 0.25
node.visits.setdefault(3, 0)
node.value_sum[3] = 1.5
node.virtual_visits[3] = 2
node.visits[3] = node.visits.get(3, 0) + 4
assert bool(node.priors)
assert node.priors.get(3, 0.0) == 0.25
assert node.visits.get(3, 0) == 4
assert node.value_sum[3] == 1.5
assert node.virtual_visits[3] == 2
assert 3 in node.visits
assert dict(node.visits.items()) == {3: 4}
def test_replay_buffer_capacity_and_sample() -> None:
sample = ReplaySample(
info_state=np.zeros(4, dtype=np.float32),
@@ -127,6 +297,29 @@ def test_self_play_game_returns_signed_targets() -> None:
assert all(sample.prior is not None for sample in samples)
def test_interleaved_self_play_yields_complete_games() -> None:
config = mini_config()
state = GameState.new_game(config, seed=21)
net = AlphaZeroNet(input_dim(state), state.action_size, hidden_size=8, num_layers=1)
ismcts_config = IsMctsConfig.model_validate(
{
"mcts": {"n_simulations": 2, "parallel_simulations": 2},
"training": {"games_per_iter": 4, "interleave_games": 4, "interleave_max_batch": 16},
}
)
samples = play_self_play_iteration(
net,
ismcts_config.mcts,
ismcts_config.training,
config,
random.Random(22),
max_steps=80,
)
assert samples
assert {sample.game_index for sample in samples} == {0, 1, 2, 3}
assert all(np.isfinite(sample.v_target) for sample in samples)
def test_trainer_one_iteration_smoke(tmp_path) -> None:
config = IsMctsConfig.model_validate(
{
@@ -185,9 +378,9 @@ def test_trainer_emits_full_eval_metrics(tmp_path) -> None:
run_dir=tmp_path,
)
metrics = trainer.train()[0].to_dict()
assert "eval/random/avg_opened_colors" in metrics
assert "eval/random/bad_open_rate" in metrics
assert "eval/random/per_game_negative_expeditions" in metrics
assert "eval/random/avg_score_diff0" in metrics
assert "eval/random/play_action_rate" in metrics
assert "eval/random/win_rate0" in metrics
def test_trainer_emits_mcts_metrics(tmp_path) -> None:
@@ -221,3 +414,39 @@ def test_trainer_emits_mcts_metrics(tmp_path) -> None:
):
assert key in metrics
assert np.isfinite(metrics[key])
def test_smoke_iter_with_batching(tmp_path) -> None:
config = IsMctsConfig.model_validate(
{
"run": {"max_iterations": 1, "seed": 14, "device": "cpu"},
"rules": {
"n_colors": 3,
"n_ranks": 5,
"n_handshakes": 1,
"hand_size": 4,
"bonus_threshold": 4,
},
"network": {"hidden_size": 16, "num_layers": 1},
"mcts": {"n_simulations": 4, "parallel_simulations": 4},
"training": {
"games_per_iter": 1,
"gradient_steps_per_iter": 1,
"batch_size": 8,
"interleave_games": 4,
"interleave_max_batch": 16,
},
"checkpoint": {"save_every": 0},
"evaluation": {"eval_every": 0, "num_workers": 1, "max_steps": 80},
}
)
trainer = IsMctsTrainer(
config,
config.rules.to_lost_cities_config(seed=config.run.seed),
run_dir=tmp_path,
)
metrics = trainer.train()[0].to_dict()
assert metrics["samples/added"] > 0
assert "mcts/avg_visit_entropy" in metrics
assert "mcts/value_prediction_error" in metrics
assert "mcts/policy_mcts_kl" in metrics