Add Dirichlet root noise + standalone eval CLI, fix self-play stall trap

Trap diagnosis: agent learned to stall (avoid opening expeditions, draw
from discard pile to extend deck) until max_steps timeout, then squeak by
on opponents' negative scores. All eval wins were from timeouts; agent
never won a naturally-terminating game. Self-play reinforced this because
timeout games still got a positive value target.

Fixes (no algorithm change, all MCTS hyperparameters or signal shaping):

- Dirichlet noise at root prior (AlphaZero standard, was missing):
  mcts.pyx `_expand_with_prior` takes `is_root` flag; root expansion
  mixes prior with Dirichlet(α). Callers in interleaved_self_play and
  the internal evaluate_and_backup pass `not item.path`.
- Default config strengthens exploration on the 50-sim batched search:
  c_puct 1.5 -> 3.0, virtual_loss_value 1.0 -> 5.0, plus new
  root_dirichlet_alpha=0.3 / root_dirichlet_epsilon=0.25.
- Self-play timeout signal zeroed: `_finalize_context` sets v_target=0
  if context.state is not terminal. Stops the network from learning
  "stall = positive value".

New standalone evaluator:
- `lost-cities-ismcts eval` subcommand (eval_checkpoint.py): loads a
  checkpoint, runs N games per opponent across a parallel pool, reports
  win/score with 95% CIs plus per-game logging via --verbose. Defaults
  cover heuristic-balanced/aggressive/cautious (rollout policy isn't in
  the training-eval opponent list, so this is the natural way to compare
  the trained policy against its rollout target).

Tests (19) still pass; .so rebuilt.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-05-11 05:17:32 +09:00
co-authored by Claude Opus 4.7
parent 651175e5bd
commit f289997c1c
6 changed files with 380 additions and 11 deletions
+16 -10
View File
@@ -1,8 +1,8 @@
run:
experiment_name: ismcts-default
max_iterations: 100
max_iterations: 500
seed: 1
device: auto
device: cuda
rules:
n_colors: 5
n_ranks: 9
@@ -22,10 +22,14 @@ network:
activation: relu
mcts:
n_simulations: 50
c_puct: 1.5
c_puct: 3.0
max_depth: 200
parallel_simulations: 8
virtual_loss_value: 1.0
parallel_simulations: 64
virtual_loss_value: 5.0
eval_n_simulations: 16
rollout_policy: heuristic_balanced
root_dirichlet_alpha: 0.3
root_dirichlet_epsilon: 0.25
temperature:
training: 1.0
eval: 0.0
@@ -36,15 +40,17 @@ training:
replay_capacity: 100000
interleave_games: 8
interleave_max_batch: 64
num_workers: 8
worker_device: cuda
optimization:
learning_rate: 0.0003
grad_clip: 5.0
checkpoint:
save_every: 10
save_every: 20
save_latest: true
evaluation:
eval_every: 10
games: 20
eval_every: 5
games: 5
opponents: [random, discard-only, heuristic-cautious]
max_steps: 10000
num_workers: 1
max_steps: 500
num_workers: 8
@@ -112,6 +112,16 @@ def main(argv: list[str] | None = None) -> None:
train.add_argument("--wandb-tag", action="append", default=[])
train.add_argument("--wandb-notes")
train.set_defaults(func=train_command)
from .eval_checkpoint import add_eval_args, run_eval
eval_cmd = subparsers.add_parser(
"eval",
help="Evaluate a saved checkpoint vs heuristic bots in parallel.",
)
add_eval_args(eval_cmd)
eval_cmd.set_defaults(func=lambda a: run_eval(a))
args = parser.parse_args(argv)
args.func(args)
@@ -30,6 +30,8 @@ class MctsConfig(StrictModel):
virtual_loss_value: float = 1.0
eval_with_mcts: bool = True
eval_n_simulations: int = 0
root_dirichlet_alpha: float = 0.0
root_dirichlet_epsilon: float = 0.0
@field_validator("n_simulations", "max_depth", "parallel_simulations")
@classmethod
@@ -0,0 +1,323 @@
"""Standalone parallel evaluation of an ISMCTS checkpoint vs heuristic bots.
Reuses the same MCTS / bot stack as the training-loop eval, but does not
interfere with a running training process. Useful for comparing a snapshot
against opponents that are not in `evaluation.opponents` (e.g. heuristic-balanced,
the rollout policy) and for running many more games than per-iter eval typically
allows.
Invoked via ``lost-cities-ismcts eval`` (see ``cli.py``).
"""
from __future__ import annotations
import argparse
import multiprocessing as mp
import os
import random
import sys
import time
from concurrent.futures import ProcessPoolExecutor, as_completed
from dataclasses import dataclass, field
from pathlib import Path
import torch
@dataclass(frozen=True)
class _WorkerJob:
ckpt_path: str
opponent: str
game_indices: tuple[int, ...]
seed: int
device: str
worker_index: int
verbose: bool
@dataclass
class _GameResult:
game_index: int
policy_player: int
score_diff: float
turns: int
policy_turns: int
play_actions: int
timed_out: bool
@dataclass
class _WorkerResult:
worker_index: int
opponent: str
games: list[_GameResult] = field(default_factory=list)
elapsed: float = 0.0
def _run_games(job: _WorkerJob) -> _WorkerResult:
# Inside-worker imports + thread-cap to avoid CPU oversubscription
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 coolrl_lost_cities.games.classic.ismcts.config import IsMctsConfig
from coolrl_lost_cities.games.classic.ismcts.mcts import IsMctsSearcher
from coolrl_lost_cities.games.classic.ismcts.network import AlphaZeroNet
os.environ.setdefault("OMP_NUM_THREADS", "1")
os.environ.setdefault("MKL_NUM_THREADS", "1")
torch.set_num_threads(1)
started = time.perf_counter()
print(
f" [worker {job.worker_index}] start ({len(job.game_indices)} games "
f"vs {job.opponent}, device={job.device})",
flush=True,
)
ckpt = torch.load(job.ckpt_path, map_location="cpu", weights_only=False)
cfg = IsMctsConfig.model_validate(ckpt["config"])
game_config = LostCitiesConfig(**ckpt["game_config"])
probe = GameState.new_game(game_config, seed=cfg.run.seed)
dim = input_dim(probe, cfg.encoding)
device = torch.device(job.device)
net = AlphaZeroNet.from_config(dim, probe.action_size, cfg).to(device)
net.load_state_dict(ckpt["network"])
net.eval()
eval_mcts_cfg = cfg.mcts.model_copy()
if cfg.mcts.eval_n_simulations > 0:
eval_mcts_cfg = eval_mcts_cfg.model_copy(
update={"n_simulations": cfg.mcts.eval_n_simulations}
)
rng = random.Random(job.seed + job.worker_index * 7919)
result = _WorkerResult(worker_index=job.worker_index, opponent=job.opponent)
max_steps = 500
for game_index in job.game_indices:
game_started = time.perf_counter()
policy_player = game_index % 2
opps = [
build_bot(job.opponent, seed=job.seed + game_index),
build_bot(job.opponent, seed=job.seed + game_index + 1),
]
state = GameState.new_game(game_config, seed=job.seed + game_index)
turns = 0
policy_turns = 0
play_actions = 0
timed_out = False
while True:
if state.terminal:
break
if turns >= max_steps:
timed_out = True
break
current = int(state.current_player)
if current == policy_player:
searcher = IsMctsSearcher(
net,
eval_mcts_cfg,
device=device,
encoding=cfg.encoding,
rng=random.Random(rng.randrange(2**31)),
)
visits = searcher.search(state, current)
unified = (
max(visits, key=visits.get) if visits else 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:
state.apply_action(opps[current].act(state))
turns += 1
diff = float(state.score_diff(policy_player))
gr = _GameResult(
game_index=game_index,
policy_player=policy_player,
score_diff=diff,
turns=turns,
policy_turns=policy_turns,
play_actions=play_actions,
timed_out=timed_out,
)
result.games.append(gr)
if job.verbose:
elapsed_g = time.perf_counter() - game_started
pa = play_actions / policy_turns if policy_turns else 0.0
print(
f" [worker {job.worker_index}] game {game_index:3d} "
f"as P{policy_player} | turns={turns:3d} "
f"score={diff:+6.1f} PA={pa:.2f}"
f"{' TIMEOUT' if timed_out else ''} "
f"({elapsed_g:.1f}s)",
flush=True,
)
result.elapsed = time.perf_counter() - started
won = sum(1 for g in result.games if g.score_diff > 0)
print(
f" [worker {job.worker_index}] done {len(result.games)} games "
f"vs {job.opponent} in {result.elapsed:.1f}s "
f"(W={won}/{len(result.games)})",
flush=True,
)
return result
def _split_games(n_games: int, n_workers: int) -> list[tuple[int, ...]]:
n_workers = max(1, min(n_workers, n_games))
base = n_games // n_workers
rem = n_games % n_workers
out: list[tuple[int, ...]] = []
cursor = 0
for i in range(n_workers):
count = base + (1 if i < rem else 0)
out.append(tuple(range(cursor, cursor + count)))
cursor += count
return out
def _summarize(games: list[_GameResult]) -> dict[str, float]:
n = len(games)
if n == 0:
return {}
wins = sum(1 for g in games if g.score_diff > 0)
losses = sum(1 for g in games if g.score_diff < 0)
draws = sum(1 for g in games if g.score_diff == 0)
timeouts = sum(1 for g in games if g.timed_out)
score_diffs = [g.score_diff for g in games]
avg = sum(score_diffs) / n
var = sum((d - avg) ** 2 for d in score_diffs) / n if n > 1 else 0.0
std = var**0.5
total_policy_turns = sum(g.policy_turns for g in games)
total_play_actions = sum(g.play_actions for g in games)
pa = total_play_actions / total_policy_turns if total_policy_turns else 0.0
avg_turns = sum(g.turns for g in games) / n
z = 1.96
# Wilson CI on win rate (half-width only; report as wr ± half)
wr = wins / n
denom = 1 + z * z / n
half = z / denom * ((wr * (1 - wr) / n + z * z / (4 * n * n)) ** 0.5)
score_ci = z * std / (n**0.5) if n > 1 else 0.0
return {
"games": n,
"wins": wins,
"losses": losses,
"draws": draws,
"timeouts": timeouts,
"win_rate": wr,
"win_rate_ci_half": half,
"avg_score_diff": avg,
"score_std": std,
"score_ci_half": score_ci,
"play_action_rate": pa,
"avg_turns": avg_turns,
}
def _default_ckpt() -> Path:
"""Find latest 'ismcts-overnight*' run's latest.pt, or fallback to newest run."""
candidates = sorted(Path("runs").glob("*ismcts-overnight*"))
if candidates:
return candidates[-1] / "latest.pt"
candidates = sorted(Path("runs").iterdir())
if not candidates:
raise SystemExit("no runs/ directory entries")
return candidates[-1] / "latest.pt"
def add_eval_args(parser: argparse.ArgumentParser) -> None:
parser.add_argument(
"--ckpt",
default=None,
help="Path to checkpoint .pt. Default: latest overnight run latest.pt.",
)
parser.add_argument(
"--opponents",
nargs="+",
default=["heuristic-balanced", "heuristic-aggressive", "heuristic-cautious"],
help="Opponent bot names from registry.",
)
parser.add_argument("--games", type=int, default=50, help="Games per opponent (default: 50).")
parser.add_argument(
"--device",
choices=("cpu", "cuda"),
default="cpu",
help="Worker device (default: cpu; cuda may compete with running training).",
)
parser.add_argument(
"--num-workers", type=int, default=8, help="Parallel worker count (default: 8)."
)
parser.add_argument("--seed", type=int, default=99999)
parser.add_argument(
"--verbose",
action="store_true",
help="Print per-game result lines (turns/score/PA) in addition to per-worker summaries.",
)
def run_eval(args: argparse.Namespace) -> None:
ckpt_path = Path(args.ckpt) if args.ckpt else _default_ckpt()
if not ckpt_path.exists():
print(f"checkpoint not found: {ckpt_path}", file=sys.stderr)
raise SystemExit(1)
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
iteration = ckpt.get("iteration", "?")
print(f"checkpoint : {ckpt_path}")
print(f"iteration : {iteration}")
print(f"device : {args.device}")
print(f"workers : {args.num_workers}")
print(f"games/opp : {args.games}")
print(f"opponents : {args.opponents}")
print(f"seed : {args.seed}")
if args.verbose:
print("verbose : True (per-game logging)")
print()
ctx = mp.get_context("spawn")
started_all = time.perf_counter()
for opponent in args.opponents:
opp_started = time.perf_counter()
slices = _split_games(args.games, args.num_workers)
jobs = [
_WorkerJob(
ckpt_path=str(ckpt_path),
opponent=opponent,
game_indices=tuple(slices[i]),
seed=args.seed,
device=args.device,
worker_index=i,
verbose=args.verbose,
)
for i in range(len(slices))
]
all_games: list[_GameResult] = []
with ProcessPoolExecutor(max_workers=len(jobs), mp_context=ctx) as ex:
futures = [ex.submit(_run_games, j) for j in jobs]
for f in as_completed(futures):
res = f.result()
all_games.extend(res.games)
elapsed = time.perf_counter() - opp_started
summary = _summarize(all_games)
n = int(summary["games"])
print()
print(
f"vs {opponent:22s} | W={summary['wins']}/{n} "
f"({summary['win_rate']:.2f} ± {summary['win_rate_ci_half']:.2f}) "
f"| S={summary['avg_score_diff']:+6.1f} ± {summary['score_ci_half']:5.1f} "
f"(σ={summary['score_std']:.1f}) | PA={summary['play_action_rate']:.2f} "
f"| turns={summary['avg_turns']:.0f} "
f"| timeouts={summary['timeouts']} "
f"| elapsed={elapsed:.1f}s"
)
print()
total = time.perf_counter() - started_all
print(f"total elapsed: {total:.1f}s")
@@ -181,6 +181,7 @@ def _evaluate_global_batch(
item.legal_actions,
priors_by_id[id(item)],
values_by_id[id(item)],
not item.path, # is_root for Dirichlet noise
)
job.searcher._backup(item.path, value, item.leaf_player)
@@ -222,6 +223,14 @@ def _finish_decision(
def _finalize_context(context: _GameContext) -> list[ReplaySample]:
# If the game did not terminate naturally (hit max_steps), the score
# reflects an incomplete game — typically a "stall" outcome where both
# sides have under-developed expeditions. Treating that as a real win
# for either player creates a degenerate stall-and-pray learning
# signal. Zero it out so the trajectory is neutral.
if not context.state.terminal:
final_diff0 = 0.0
else:
final_diff0 = float(context.state.score_diff(0))
samples: list[ReplaySample] = []
for decision in context.decisions:
@@ -489,6 +489,7 @@ cdef class IsMctsSearcher:
item.legal_actions,
priors_by_id[id(item)],
values_by_id[id(item)],
not item.path, # is_root: empty path means leaf == root
)
self._backup(item.path, value, item.leaf_player)
@@ -500,14 +501,32 @@ cdef class IsMctsSearcher:
list legal_actions,
object probs,
double network_value,
bint is_root=False,
):
cdef int action
cdef int i
cdef int n_legal
cdef double alpha
cdef double epsilon
cdef object rollout_value
cdef object noise
cdef object np_rng
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
# Dirichlet noise at root (AlphaZero pattern: force exploration of low-prior actions)
alpha = float(self.config.root_dirichlet_alpha)
epsilon = float(self.config.root_dirichlet_epsilon)
if is_root and alpha > 0.0 and epsilon > 0.0:
n_legal = len(legal_actions)
# Use numpy with seed from self.rng for reproducibility under fixed seeds
np_rng = np.random.default_rng(self.rng.randrange(2**31))
noise = np_rng.dirichlet([alpha] * n_legal)
for i in range(n_legal):
action = legal_actions[i]
probs[action] = (1.0 - epsilon) * float(probs[action]) + epsilon * float(noise[i])
for action in legal_actions:
(<_ArrayMap>node.priors).set_float(action, float(probs[action]))
if not (<_ArrayMap>node.visits).has(action):