Add SO-ISMCTS mini trainer

Implements a proof-of-concept single-observer IS-MCTS trainer with AlphaZero-style policy/value network, determinization, replay, self-play, CLI configs, and focused tests. Mini acceptance run reaches positive random eval while keeping play_action_rate above the Deep CFR trap threshold.

Tests: uv run python -m pytest tests/games/classic/ismcts/ -x; uv run python -m pytest tests/games/classic/test_deep_cfr_trainer.py -x; uv run lost-cities-ismcts train --config configs/ismcts/mini.yaml
This commit is contained in:
2026-05-10 22:46:22 +09:00
parent 5acda3f272
commit e69f3165b6
15 changed files with 1205 additions and 0 deletions
+46
View File
@@ -0,0 +1,46 @@
run:
experiment_name: ismcts-default
max_iterations: 100
seed: 1
device: auto
rules:
n_colors: 5
n_ranks: 9
min_rank: 2
n_handshakes: 3
hand_size: 8
expedition_penalty: -20
bonus_threshold: 8
bonus_amount: 20
encoding:
derived_playability: true
slot_aware_playability: true
network:
kind: mlp
hidden_size: 512
num_layers: 3
activation: relu
mcts:
n_simulations: 50
c_puct: 1.5
max_depth: 200
temperature:
training: 1.0
eval: 0.0
training:
games_per_iter: 10
gradient_steps_per_iter: 10
batch_size: 128
replay_capacity: 100000
optimization:
learning_rate: 0.0003
grad_clip: 5.0
checkpoint:
save_every: 10
save_latest: true
evaluation:
eval_every: 10
games: 20
opponents: [random, discard-only, heuristic-cautious]
max_steps: 10000
num_workers: 1
+46
View File
@@ -0,0 +1,46 @@
run:
experiment_name: ismcts-mini
max_iterations: 50
seed: 1
device: cpu
rules:
n_colors: 3
n_ranks: 5
min_rank: 2
n_handshakes: 1
hand_size: 4
expedition_penalty: -20
bonus_threshold: 4
bonus_amount: 20
encoding:
derived_playability: true
slot_aware_playability: true
network:
kind: mlp
hidden_size: 128
num_layers: 2
activation: relu
mcts:
n_simulations: 50
c_puct: 1.5
max_depth: 100
temperature:
training: 1.0
eval: 0.0
training:
games_per_iter: 10
gradient_steps_per_iter: 10
batch_size: 128
replay_capacity: 50000
optimization:
learning_rate: 0.001
grad_clip: 5.0
checkpoint:
save_every: 0
save_latest: true
evaluation:
eval_every: 5
games: 20
opponents: [random, discard-only, heuristic-cautious]
max_steps: 300
num_workers: 1
+1
View File
@@ -29,6 +29,7 @@ lost-cities-classic = "coolrl_lost_cities.games.classic:main"
lost-cities-eval = "coolrl_lost_cities.games.classic.evaluation:main"
lost-cities-classic-gui = "coolrl_lost_cities.games.classic.pygame_pvp:main"
lost-cities-deep-cfr = "coolrl_lost_cities.games.classic.deep_cfr.cli:main"
lost-cities-ismcts = "coolrl_lost_cities.games.classic.ismcts.cli:main"
[dependency-groups]
dev = [
@@ -0,0 +1 @@
"""Single-observer IS-MCTS AlphaZero-style training for Lost Cities."""
@@ -0,0 +1,96 @@
from __future__ import annotations
import argparse
import re
from datetime import datetime
from pathlib import Path
from typing import Any
import yaml
from .config import IsMctsConfig, load_config
from .trainer import IsMctsTrainer
_SLUG_RE = re.compile(r"[^a-z0-9]+")
def _kebab_slug(value: str) -> str:
return _SLUG_RE.sub("-", value.strip().lower()).strip("-") or "run"
def _deep_update(base: dict[str, Any], patch: dict[str, Any]) -> None:
for key, value in patch.items():
if isinstance(value, dict) and isinstance(base.get(key), dict):
_deep_update(base[key], value)
else:
base[key] = value
def _set_path_override(overrides: dict[str, Any], assignment: str) -> None:
if "=" not in assignment:
raise ValueError(f"config override must be PATH=VALUE: {assignment}")
path, raw_value = assignment.split("=", 1)
keys = path.split(".")
value = yaml.safe_load(raw_value)
cursor = overrides
for key in keys[:-1]:
cursor = cursor.setdefault(key, {})
cursor[keys[-1]] = value
def _with_overrides(config: IsMctsConfig, assignments: list[str]) -> IsMctsConfig:
overrides: dict[str, Any] = {}
for assignment in assignments:
_set_path_override(overrides, assignment)
data = config.model_dump(mode="python")
_deep_update(data, overrides)
return IsMctsConfig.model_validate(data)
def _resolve_run_dir(config: IsMctsConfig, *, keep: bool) -> Path:
parent = Path("runs") if keep else Path("runs/tmp")
timestamp = datetime.now().strftime("%Y-%m-%d_%H%M%S")
return parent / f"{timestamp}_{_kebab_slug(config.run.experiment_name)}"
def train_command(args: argparse.Namespace) -> None:
config = load_config(args.config) if args.config else IsMctsConfig()
config = _with_overrides(config, args.config_overrides)
run_dir = _resolve_run_dir(config, keep=args.keep)
trainer = IsMctsTrainer(
config,
config.rules.to_lost_cities_config(seed=config.run.seed),
run_dir=run_dir,
device=config.run.device,
)
trainer.train()
def main(argv: list[str] | None = None) -> None:
parser = argparse.ArgumentParser(description="Lost Cities SO-ISMCTS tools.")
subparsers = parser.add_subparsers(dest="command", required=True)
train = subparsers.add_parser("train")
train.add_argument("--config")
train.add_argument("--keep", action="store_true")
train.add_argument(
"--set",
action="append",
default=[],
dest="config_overrides",
metavar="PATH=VALUE",
)
train.add_argument("--wandb", action="store_true", help="Accepted for CLI parity; ignored.")
train.add_argument("--wandb-project", default="coolrl-lost-cities")
train.add_argument("--wandb-name")
train.add_argument("--wandb-mode", choices=("online", "offline", "disabled"), default="online")
train.add_argument("--wandb-group")
train.add_argument("--wandb-job-type")
train.add_argument("--wandb-tag", action="append", default=[])
train.add_argument("--wandb-notes")
train.set_defaults(func=train_command)
args = parser.parse_args(argv)
args.func(args)
if __name__ == "__main__":
main()
@@ -0,0 +1,85 @@
from __future__ import annotations
import json
from collections.abc import Mapping
from pathlib import Path
from typing import Any
import yaml
from pydantic import Field, field_validator
from coolrl_lost_cities.games.classic.deep_cfr.config import (
CheckpointConfig,
EncodingConfig,
EvaluationConfig,
NetworkConfig,
OptimizationConfig,
RulesConfig,
RunConfig,
StrictModel,
)
class MctsConfig(StrictModel):
n_simulations: int = 50
c_puct: float = 1.5
max_depth: int = 200
use_rollout_value: bool = True
@field_validator("n_simulations", "max_depth")
@classmethod
def _positive_int(cls, value: int) -> int:
if value <= 0:
raise ValueError("must be positive")
return value
class TemperatureConfig(StrictModel):
training: float = 1.0
eval: float = 0.0
class TrainingConfig(StrictModel):
games_per_iter: int = 10
gradient_steps_per_iter: int = 10
batch_size: int = 128
replay_capacity: int = 100_000
@field_validator("games_per_iter", "gradient_steps_per_iter", "batch_size", "replay_capacity")
@classmethod
def _positive_int(cls, value: int) -> int:
if value <= 0:
raise ValueError("must be positive")
return value
class IsMctsConfig(StrictModel):
run: RunConfig = Field(default_factory=lambda: RunConfig(experiment_name="ismcts"))
rules: RulesConfig = Field(default_factory=RulesConfig)
encoding: EncodingConfig = Field(default_factory=EncodingConfig)
network: NetworkConfig = Field(
default_factory=lambda: NetworkConfig(hidden_size=512, num_layers=3)
)
mcts: MctsConfig = Field(default_factory=MctsConfig)
temperature: TemperatureConfig = Field(default_factory=TemperatureConfig)
training: TrainingConfig = Field(default_factory=TrainingConfig)
optimization: OptimizationConfig = Field(default_factory=OptimizationConfig)
checkpoint: CheckpointConfig = Field(default_factory=CheckpointConfig)
evaluation: EvaluationConfig = Field(default_factory=EvaluationConfig)
def to_dict(self) -> dict[str, Any]:
return self.model_dump(mode="json")
def config_from_dict(data: Mapping[str, Any]) -> IsMctsConfig:
return IsMctsConfig.model_validate(data)
def load_config(path: str | Path) -> IsMctsConfig:
config_path = Path(path)
text = config_path.read_text(encoding="utf-8")
if config_path.suffix.lower() in {".yaml", ".yml"}:
data = yaml.safe_load(text) or {}
else:
data = json.loads(text)
return config_from_dict(data)
@@ -0,0 +1,28 @@
from __future__ import annotations
import random
from coolrl_lost_cities.games.classic.game import GameState
from .info_set import unseen_cards
def sample_determinization(state: GameState, player: int, rng: random.Random) -> GameState:
"""Sample a concrete state uniformly from ``player``'s current information set."""
p = int(player)
opponent = 1 - p
snapshot = state.to_snapshot()
unseen = unseen_cards(state, p)
rng.shuffle(unseen)
opponent_hand_size = len(state.hands[opponent])
deck_size = len(state.deck)
if len(unseen) != opponent_hand_size + deck_size:
raise ValueError(
"information set card count mismatch: "
f"unseen={len(unseen)} opponent_hand={opponent_hand_size} deck={deck_size}"
)
snapshot["hands"][opponent] = [card.to_snapshot() for card in unseen[:opponent_hand_size]]
snapshot["deck"] = [card.to_snapshot() for card in unseen[opponent_hand_size:]]
det = GameState.from_snapshot(snapshot)
det.validate_invariants()
return det
@@ -0,0 +1,61 @@
from __future__ import annotations
import json
from collections import Counter
from coolrl_lost_cities.games.classic.game import Card, GameState, build_deck
def _card_tuple(card: Card) -> tuple[int, int]:
return int(card.color), int(card.rank)
def _sorted_cards(cards: list[Card]) -> list[tuple[int, int]]:
return sorted(_card_tuple(card) for card in cards)
def canonical_info_set_key(state: GameState, player: int) -> bytes:
"""Deterministic key for observable information from ``player``'s POV."""
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")
def visible_cards(state: GameState, player: int) -> list[Card]:
cards: list[Card] = []
cards.extend(state.hands[int(player)])
for player_expeditions in state.expeditions:
for expedition in player_expeditions:
cards.extend(expedition)
for discard in state.discards:
cards.extend(discard)
return cards
def unseen_cards(state: GameState, player: int) -> list[Card]:
remaining = Counter(_card_tuple(card) for card in build_deck(state.config))
for card in visible_cards(state, player):
remaining[_card_tuple(card)] -= 1
cards: list[Card] = []
for (color, rank), count in remaining.items():
cards.extend(Card(color, rank) for _ in range(count))
return cards
@@ -0,0 +1,154 @@
from __future__ import annotations
import math
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
from .config import MctsConfig
from .determinization import sample_determinization
from .info_set import canonical_info_set_key
from .network import AlphaZeroNet
@dataclass
class MctsNode:
info_set_key: bytes
player: int
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)
children: dict[int, bytes] = field(default_factory=dict)
terminal: bool = False
def is_expanded(self) -> bool:
return self.terminal or bool(self.priors)
def q(self, action: int) -> float:
n = self.visits.get(action, 0)
if n <= 0:
return 0.0
return self.value_sum.get(action, 0.0) / n
class MctsTree:
def __init__(self) -> None:
self.nodes: dict[bytes, MctsNode] = {}
def get_or_create(self, key: bytes, *, player: int, terminal: bool = False) -> MctsNode:
node = self.nodes.get(key)
if node is None:
node = MctsNode(key, player=player, terminal=terminal)
self.nodes[key] = node
return node
class IsMctsSearcher:
def __init__(
self,
network: AlphaZeroNet,
config: MctsConfig,
*,
device: torch.device | str = "cpu",
encoding=None,
rng: random.Random | None = None,
) -> None:
self.network = network
self.config = config
self.device = torch.device(device)
self.encoding = encoding
self.rng = rng or random.Random()
self.tree = MctsTree()
def search(
self,
state: GameState,
traverser: int,
n_sims: int | None = None,
) -> dict[int, int]:
root_key = canonical_info_set_key(state, state.current_player)
root = self.tree.get_or_create(
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)
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))
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
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
def _expand_and_evaluate(self, node: MctsNode, state: GameState, player: int) -> 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)
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 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)
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)
prior = node.priors.get(action, 0.0)
score = node.q(action) + self.config.c_puct * prior * sqrt_total / (1 + n)
if score > best_score:
best_score = score
best_action = action
return int(best_action)
def _rollout_value(self, state: GameState, player: int) -> float | None:
rollout = state.clone()
steps = 0
while not rollout.terminal and steps < self.config.max_depth:
legal = rollout.unified_legal_actions()
if not legal:
break
action = self.rng.choice(legal)
rollout.apply_unified_action(action)
steps += 1
return float(rollout.score_diff(player))
@@ -0,0 +1,72 @@
from __future__ import annotations
import torch
from torch import nn
from coolrl_lost_cities.games.classic.deep_cfr.networks import _activation
from .config import IsMctsConfig
class AlphaZeroNet(nn.Module):
def __init__(
self,
input_dim: int,
action_size: int,
hidden_size: int = 512,
*,
num_layers: int = 3,
activation: str = "relu",
value_scale: float = 100.0,
) -> None:
super().__init__()
self.input_dim = int(input_dim)
self.action_size = int(action_size)
self.value_scale = float(value_scale)
layers: list[nn.Module] = []
last_dim = self.input_dim
for _ in range(max(0, int(num_layers))):
layers.append(nn.Linear(last_dim, hidden_size))
layers.append(_activation(activation))
last_dim = hidden_size
self.backbone = nn.Sequential(*layers)
self.policy_head = nn.Linear(last_dim, self.action_size)
self.value_head = nn.Linear(last_dim, 1)
@classmethod
def from_config(
cls,
input_dim: int,
action_size: int,
config: IsMctsConfig | object,
) -> AlphaZeroNet:
network_config = config.network if hasattr(config, "network") else config
return cls(
input_dim,
action_size,
network_config.hidden_size,
num_layers=network_config.num_layers,
activation=network_config.activation,
)
def forward(
self,
info_state_tensor: torch.Tensor,
legal_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
hidden = self.backbone(info_state_tensor)
logits = self.policy_head(hidden)
value = torch.tanh(self.value_head(hidden)).squeeze(-1) * self.value_scale
if legal_mask is not None:
logits = logits.masked_fill(~legal_mask.bool(), torch.finfo(logits.dtype).min)
return logits, value
def policy_distribution(
self,
info_state_tensor: torch.Tensor,
legal_mask: torch.Tensor,
) -> torch.Tensor:
logits, _value = self.forward(info_state_tensor, legal_mask)
probs = torch.softmax(logits, dim=-1).masked_fill(~legal_mask.bool(), 0.0)
normalizer = probs.sum(dim=-1, keepdim=True).clamp_min(1.0e-12)
return probs / normalizer
@@ -0,0 +1,45 @@
from __future__ import annotations
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
from coolrl_lost_cities.games.classic.policy import LostCitiesPolicy, PolicyInput
from .network import AlphaZeroNet
class AlphaZeroPolicy(LostCitiesPolicy):
def __init__(
self,
network: AlphaZeroNet,
*,
device: torch.device | str = "cpu",
encoding=None,
sample: bool = False,
seed: int | None = None,
) -> None:
self.network = network
self.device = torch.device(device)
self.encoding = encoding
self.sample = sample
self.rng = np.random.default_rng(seed)
def act(self, obs_or_state: PolicyInput) -> int:
if not isinstance(obs_or_state, GameState):
legal = np.asarray(obs_or_state["legal_mask"], dtype=bool)
return int(np.flatnonzero(legal)[0])
state = obs_or_state
legal = np.asarray(state.unified_legal_mask(), dtype=bool)
info = encode_info_state(state, state.current_player, self.encoding)
with torch.inference_mode():
x = torch.as_tensor(info[None, :], dtype=torch.float32, device=self.device)
mask = torch.as_tensor(legal[None, :], dtype=torch.bool, device=self.device)
probs = self.network.policy_distribution(x, mask).squeeze(0).cpu().numpy()
legal_actions = np.flatnonzero(legal)
if self.sample:
unified = int(self.rng.choice(legal_actions, p=probs[legal_actions]))
else:
unified = int(legal_actions[int(np.argmax(probs[legal_actions]))])
return state.from_unified_action(unified)
@@ -0,0 +1,37 @@
from __future__ import annotations
from collections import deque
from collections.abc import Iterable
from dataclasses import dataclass
import numpy as np
@dataclass
class ReplaySample:
info_state: np.ndarray
legal_mask: np.ndarray
pi_target: np.ndarray
v_target: float
player: int
class ReplayBuffer:
def __init__(self, capacity: int, *, seed: int | None = None) -> None:
self.capacity = int(capacity)
self._items: deque[ReplaySample] = deque(maxlen=self.capacity)
self.rng = np.random.default_rng(seed)
def __len__(self) -> int:
return len(self._items)
def add(self, samples: Iterable[ReplaySample]) -> None:
self._items.extend(samples)
def sample(self, batch_size: int) -> list[ReplaySample]:
if not self._items:
raise ValueError("cannot sample from an empty replay buffer")
size = min(int(batch_size), len(self._items))
indices = self.rng.choice(len(self._items), size=size, replace=False)
items = list(self._items)
return [items[int(index)] for index in indices]
@@ -0,0 +1,99 @@
from __future__ import annotations
import random
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
from .mcts import IsMctsSearcher
from .network import AlphaZeroNet
from .replay_buffer import ReplaySample
def visit_distribution(
visit_counts: dict[int, int],
action_size: int,
*,
temperature: float,
) -> np.ndarray:
pi = np.zeros(action_size, dtype=np.float32)
if not visit_counts:
return pi
actions = np.asarray(list(visit_counts), dtype=np.int64)
counts = np.asarray([visit_counts[int(action)] for action in actions], dtype=np.float64)
if temperature <= 0.0:
best = int(actions[int(np.argmax(counts))])
pi[best] = 1.0
return pi
adjusted = np.power(np.maximum(counts, 1.0e-12), 1.0 / temperature)
adjusted /= adjusted.sum()
pi[actions] = adjusted.astype(np.float32)
return pi
def select_from_distribution(pi: np.ndarray, rng: random.Random) -> int:
total = float(pi.sum())
if total <= 0.0:
raise RuntimeError("empty action distribution")
threshold = rng.random() * total
cumsum = 0.0
for action, prob in enumerate(pi):
cumsum += float(prob)
if cumsum >= threshold:
return action
return int(len(pi) - 1)
def play_self_play_game(
network: AlphaZeroNet,
mcts_config: MctsConfig,
game_config: LostCitiesConfig,
rng: random.Random,
*,
device: torch.device | str = "cpu",
encoding=None,
temperature: float = 1.0,
max_steps: int = 10_000,
) -> list[ReplaySample]:
state = GameState.new_game(game_config, seed=rng.randrange(2**31))
pending: list[tuple[np.ndarray, np.ndarray, np.ndarray, int]] = []
steps = 0
while not state.terminal and steps < max_steps:
player = int(state.current_player)
searcher = IsMctsSearcher(
network,
mcts_config,
device=device,
encoding=encoding,
rng=random.Random(rng.randrange(2**31)),
)
visits = searcher.search(state, player)
legal_mask = np.asarray(state.unified_legal_mask(), dtype=bool)
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)
info = encode_info_state(state, player, encoding)
pending.append((info.astype(np.float32), legal_mask, pi, player))
action = select_from_distribution(pi, rng)
state.apply_unified_action(action)
steps += 1
final_diff0 = float(state.score_diff(0))
samples: list[ReplaySample] = []
for info, legal_mask, pi, player in pending:
value = final_diff0 if player == 0 else -final_diff0
samples.append(
ReplaySample(
info_state=info,
legal_mask=legal_mask.astype(bool),
pi_target=pi.astype(np.float32),
v_target=value,
player=player,
)
)
return samples
@@ -0,0 +1,290 @@
from __future__ import annotations
import json
import random
import time
from dataclasses import dataclass
from pathlib import Path
import numpy as np
import torch
from torch import nn
from coolrl_lost_cities.games.classic.bots 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
from .mcts import IsMctsSearcher
from .network import AlphaZeroNet
from .replay_buffer import ReplayBuffer, ReplaySample
from .self_play import play_self_play_game
@dataclass
class IterationMetrics:
iteration: int
samples_added: int
replay_size: int
policy_loss: float
value_loss: float
total_loss: float
self_play_seconds: float
train_seconds: float
eval_metrics: dict[str, float | int]
def to_dict(self) -> dict[str, float | int]:
data: dict[str, float | int] = {
"iteration": self.iteration,
"samples/added": self.samples_added,
"memory/replay": self.replay_size,
"loss/policy": self.policy_loss,
"loss/value": self.value_loss,
"loss/total": self.total_loss,
"time/self_play_seconds": self.self_play_seconds,
"time/train_seconds": self.train_seconds,
}
data.update(self.eval_metrics)
return data
class IsMctsTrainer:
def __init__(
self,
config: IsMctsConfig,
game_config: LostCitiesConfig,
*,
run_dir: str | Path,
device: torch.device | str = "cpu",
) -> None:
self.config = config
self.game_config = game_config
self.run_dir = Path(run_dir)
self.device = self._resolve_device(device)
probe = GameState.new_game(game_config, seed=config.run.seed)
self.input_dim = input_dim(probe, config.encoding)
self.action_size = probe.action_size
self.network = AlphaZeroNet.from_config(self.input_dim, self.action_size, config).to(
self.device
)
self.optimizer = torch.optim.AdamW(
self.network.parameters(),
lr=config.optimization.learning_rate,
weight_decay=max(float(config.optimization.weight_decay), 1.0e-4),
)
self.buffer = ReplayBuffer(config.training.replay_capacity, seed=config.run.seed)
self.metrics_path = self.run_dir / "metrics.jsonl"
self.rng = random.Random(config.run.seed)
def _resolve_device(self, device: torch.device | str) -> torch.device:
token = str(device)
if token == "auto":
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
return torch.device(token)
def train(self) -> list[IterationMetrics]:
self.run_dir.mkdir(parents=True, exist_ok=True)
(self.run_dir / "config.json").write_text(
json.dumps(self.config.to_dict(), indent=2, sort_keys=True),
encoding="utf-8",
)
if self.metrics_path.exists():
self.metrics_path.unlink()
metrics: list[IterationMetrics] = []
max_iterations = self.config.run.max_iterations or 1
started = time.perf_counter()
for iteration in range(1, max_iterations + 1):
if self._time_limit_reached(started):
break
item = self.run_iteration(iteration)
metrics.append(item)
self._append_metrics(item)
self._save_checkpoints(iteration, item)
print(json.dumps(item.to_dict(), sort_keys=True))
return metrics
def run_iteration(self, iteration: int) -> IterationMetrics:
self.network.eval()
sp_started = time.perf_counter()
added = 0
for _ in range(self.config.training.games_per_iter):
samples = play_self_play_game(
self.network,
self.config.mcts,
self.game_config,
self.rng,
device=self.device,
encoding=self.config.encoding,
temperature=self.config.temperature.training,
max_steps=self.config.evaluation.max_steps,
)
self.buffer.add(samples)
added += len(samples)
self_play_seconds = time.perf_counter() - sp_started
train_started = time.perf_counter()
losses = []
for _ in range(self.config.training.gradient_steps_per_iter):
batch = self.buffer.sample(self.config.training.batch_size)
losses.append(self._train_batch(batch))
train_seconds = time.perf_counter() - train_started
loss_arr = np.asarray(losses, dtype=np.float64)
eval_metrics = self._evaluate(iteration)
return IterationMetrics(
iteration=iteration,
samples_added=added,
replay_size=len(self.buffer),
policy_loss=float(loss_arr[:, 0].mean()) if len(loss_arr) else 0.0,
value_loss=float(loss_arr[:, 1].mean()) if len(loss_arr) else 0.0,
total_loss=float(loss_arr[:, 2].mean()) if len(loss_arr) else 0.0,
self_play_seconds=self_play_seconds,
train_seconds=train_seconds,
eval_metrics=eval_metrics,
)
def _train_batch(self, batch: list[ReplaySample]) -> tuple[float, float, float]:
self.network.train()
info = torch.as_tensor(
np.stack([sample.info_state for sample in batch]),
dtype=torch.float32,
device=self.device,
)
legal = torch.as_tensor(
np.stack([sample.legal_mask for sample in batch]),
dtype=torch.bool,
device=self.device,
)
pi = torch.as_tensor(
np.stack([sample.pi_target for sample in batch]),
dtype=torch.float32,
device=self.device,
)
value_target = torch.as_tensor(
[sample.v_target for sample in batch],
dtype=torch.float32,
device=self.device,
)
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)
loss = policy_loss + value_loss
self.optimizer.zero_grad(set_to_none=True)
loss.backward()
if self.config.optimization.grad_clip > 0:
torch.nn.utils.clip_grad_norm_(
self.network.parameters(),
self.config.optimization.grad_clip,
)
self.optimizer.step()
return float(policy_loss.item()), float(value_loss.item()), float(loss.item())
def _evaluate(self, iteration: int) -> dict[str, float | int]:
opponents = self.config.evaluation.opponents_for_iteration(iteration)
if not opponents:
return {}
self.network.eval()
results: dict[str, float | int] = {}
for opponent in opponents:
result = evaluate_policy(
self.network,
self.game_config,
opponent=opponent,
games=self.config.evaluation.games,
seed=self.config.run.seed + iteration * 1000,
device=self.device,
encoding=self.config.encoding,
max_steps=self.config.evaluation.max_steps,
mcts_config=self.config.mcts,
)
key = opponent.replace("-", "_")
for metric_key, value in result.items():
results[f"eval/{key}/{metric_key}"] = value
return results
def _append_metrics(self, metrics: IterationMetrics) -> None:
with self.metrics_path.open("a", encoding="utf-8") as handle:
handle.write(json.dumps(metrics.to_dict(), sort_keys=True) + "\n")
def _save_checkpoints(self, iteration: int, metrics: IterationMetrics) -> None:
payload = {
"config": self.config.to_dict(),
"game_config": self.game_config.to_snapshot(),
"iteration": iteration,
"network": self.network.state_dict(),
"optimizer": self.optimizer.state_dict(),
"metrics": metrics.to_dict(),
}
if self.config.checkpoint.save_latest:
torch.save(payload, self.run_dir / "latest.pt")
if (
self.config.checkpoint.save_every > 0
and iteration % self.config.checkpoint.save_every == 0
):
torch.save(payload, self.run_dir / f"iteration_{iteration:05d}.pt")
def _time_limit_reached(self, started: float) -> bool:
if self.config.run.max_minutes is None:
return False
return (time.perf_counter() - started) / 60.0 >= self.config.run.max_minutes
def evaluate_policy(
network: AlphaZeroNet,
config: LostCitiesConfig,
*,
opponent: str,
games: int,
seed: int,
device: torch.device | str,
encoding=None,
max_steps: int = 10_000,
mcts_config=None,
) -> dict[str, float | int]:
rng = random.Random(seed)
score_diffs: list[int] = []
wins = losses = draws = 0
policy_actions = play_actions = 0
for game_index in range(games):
policy_player = game_index % 2
policies = [
build_bot(opponent, seed=seed + game_index),
build_bot(opponent, seed=seed + game_index),
]
state = GameState.new_game(config, seed=seed + game_index)
for _ in range(max_steps):
if state.terminal:
break
current = int(state.current_player)
if current == policy_player:
searcher = IsMctsSearcher(
network,
mcts_config or IsMctsConfig().mcts,
device=device,
encoding=encoding,
rng=random.Random(rng.randrange(2**31)),
)
visits = searcher.search(state, current)
unified = max(visits, key=visits.get)
action = state.from_unified_action(unified)
else:
action = policies[current].act(state)
if current == policy_player and state.phase == "card":
policy_actions += 1
if action % 2 == 0:
play_actions += 1
state.apply_action(action)
diff = state.score_diff(policy_player)
score_diffs.append(diff)
wins += int(diff > 0)
losses += int(diff < 0)
draws += int(diff == 0)
return {
"games": games,
"wins0": wins,
"wins1": losses,
"draws": draws,
"win_rate0": wins / max(1, games),
"avg_score_diff0": float(np.mean(score_diffs)) if score_diffs else 0.0,
"play_action_rate": play_actions / max(1, policy_actions),
}
+144
View File
@@ -0,0 +1,144 @@
from __future__ import annotations
import random
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.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.network import 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 mini_config(seed: int = 1) -> LostCitiesConfig:
return LostCitiesConfig(
n_colors=3,
n_ranks=5,
n_handshakes=1,
hand_size=4,
bonus_threshold=4,
seed=seed,
)
def test_canonical_key_ignores_opponent_hand_identity() -> None:
config = mini_config()
state = GameState.new_game(config, seed=3)
snap = state.to_snapshot()
snap["hands"][1] = list(reversed(snap["hands"][1]))
other = GameState.from_snapshot(snap)
assert canonical_info_set_key(state, 0) == canonical_info_set_key(other, 0)
def test_determinization_consistent_with_info_set() -> None:
config = mini_config()
state = GameState.new_game(config, seed=4)
rng = random.Random(5)
key = canonical_info_set_key(state, 0)
for _ in range(100):
det = sample_determinization(state, 0, rng)
assert canonical_info_set_key(det, 0) == key
assert len(det.hands[1]) == config.hand_size
all_cards = det.hands[0] + det.hands[1] + det.deck
for player_expeditions in det.expeditions:
for expedition in player_expeditions:
all_cards.extend(expedition)
for discard in det.discards:
all_cards.extend(discard)
assert len(all_cards) == config.deck_size
def test_network_shapes_and_mask() -> None:
state = GameState.new_game(mini_config(), seed=6)
dim = input_dim(state)
net = AlphaZeroNet(dim, state.action_size, hidden_size=16, num_layers=1)
x = torch.as_tensor(encode_info_state(state, 0)[None, :], dtype=torch.float32)
mask = torch.as_tensor(np.asarray(state.unified_legal_mask(), dtype=bool)[None, :])
logits, value = net(x, mask)
probs = net.policy_distribution(x, mask)
assert logits.shape == (1, state.action_size)
assert value.shape == (1,)
assert torch.allclose(probs.sum(dim=-1), torch.ones(1))
assert torch.all(probs[~mask] == 0)
def test_mcts_prior_drives_visits() -> None:
state = GameState.new_game(mini_config(), seed=7)
dim = input_dim(state)
net = AlphaZeroNet(dim, state.action_size, hidden_size=8, num_layers=0)
for param in net.parameters():
param.data.zero_()
legal = state.unified_legal_actions()
favored = legal[0]
net.policy_head.bias.data[favored] = 5.0
searcher = IsMctsSearcher(
net,
MctsConfig(n_simulations=12, c_puct=2.0, use_rollout_value=False),
rng=random.Random(8),
)
visits = searcher.search(state, state.current_player)
assert visits[favored] == max(visits.values())
def test_replay_buffer_capacity_and_sample() -> None:
sample = ReplaySample(
info_state=np.zeros(4, dtype=np.float32),
legal_mask=np.ones(3, dtype=bool),
pi_target=np.array([1.0, 0.0, 0.0], dtype=np.float32),
v_target=1.0,
player=0,
)
buffer = ReplayBuffer(2, seed=1)
buffer.add([sample, sample, sample])
assert len(buffer) == 2
assert len(buffer.sample(2)) == 2
def test_self_play_game_returns_signed_targets() -> None:
config = mini_config()
state = GameState.new_game(config, seed=9)
net = AlphaZeroNet(input_dim(state), state.action_size, hidden_size=8, num_layers=1)
samples = play_self_play_game(
net,
MctsConfig(n_simulations=2),
config,
random.Random(10),
)
assert samples
assert {sample.player for sample in samples} <= {0, 1}
assert all(sample.pi_target.sum() > 0 for sample in samples)
def test_trainer_one_iteration_smoke(tmp_path) -> None:
config = IsMctsConfig.model_validate(
{
"run": {"max_iterations": 1, "seed": 11, "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": 2},
"training": {"games_per_iter": 1, "gradient_steps_per_iter": 1, "batch_size": 8},
"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()
assert len(metrics) == 1
assert (tmp_path / "metrics.jsonl").exists()