Deep CFR checkpoint eval CLI 추가

This commit is contained in:
2026-05-06 23:52:50 +09:00
parent 65e7abefad
commit fc5b3d9bc8
8 changed files with 354 additions and 2 deletions
+1
View File
@@ -6,6 +6,7 @@ build/
dist/ dist/
wheels/ wheels/
*.egg-info *.egg-info
runs/
# Cython-generated sources # Cython-generated sources
src/coolrl_lost_cities/games/classic/game.c src/coolrl_lost_cities/games/classic/game.c
+1
View File
@@ -22,6 +22,7 @@ gui = [
lost-cities-classic = "coolrl_lost_cities.games.classic:main" lost-cities-classic = "coolrl_lost_cities.games.classic:main"
lost-cities-eval = "coolrl_lost_cities.games.classic.evaluation: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-classic-gui = "coolrl_lost_cities.games.classic.pygame_pvp:main"
lost-cities-deep-cfr = "coolrl_lost_cities.games.classic.deep_cfr.cli:main"
[dependency-groups] [dependency-groups]
dev = [ dev = [
@@ -0,0 +1,17 @@
from __future__ import annotations
from pathlib import Path
from typing import Any
import torch
def save_checkpoint(path: str | Path, payload: dict[str, Any]) -> Path:
output = Path(path)
output.parent.mkdir(parents=True, exist_ok=True)
torch.save(payload, output)
return output
def load_checkpoint(path: str | Path, *, device: torch.device | str = "cpu") -> dict[str, Any]:
return torch.load(path, map_location=device)
@@ -0,0 +1,93 @@
from __future__ import annotations
import argparse
import json
from dataclasses import replace
from pathlib import Path
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, config_from_dict
from coolrl_lost_cities.games.classic.deep_cfr.evaluate import (
evaluate_strategy_network,
load_strategy_policy_from_checkpoint,
)
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
from coolrl_lost_cities.games.classic.game import classic_config
def _load_config(path: str | None) -> DeepCFRConfig:
if path is None:
return DeepCFRConfig()
return config_from_dict(json.loads(Path(path).read_text(encoding="utf-8")))
def train_command(args: argparse.Namespace) -> None:
config = _load_config(args.config)
overrides = {}
for key in (
"iterations",
"traversals_per_iteration",
"checkpoint_dir",
"eval_every",
"eval_games",
"seed",
):
value = getattr(args, key)
if value is not None:
overrides[key] = value
if args.no_save:
overrides["save_every_iteration"] = False
config = replace(config, **overrides)
trainer = DeepCFRTrainer(config, classic_config(seed=config.seed), device=args.device)
if args.resume:
trainer.load_checkpoint(args.resume)
metrics = trainer.train()
for item in metrics:
print(json.dumps(item.__dict__, sort_keys=True))
def eval_command(args: argparse.Namespace) -> None:
policy, game_config = load_strategy_policy_from_checkpoint(args.checkpoint, device=args.device)
result = evaluate_strategy_network(
policy.strategy_network,
game_config,
games=args.games,
seed=args.seed,
opponent=args.opponent,
device=args.device,
max_steps=args.max_steps,
)
print(json.dumps(result, indent=2, sort_keys=True))
def main(argv: list[str] | None = None) -> None:
parser = argparse.ArgumentParser(description="Lost Cities classic Deep CFR tools.")
subparsers = parser.add_subparsers(dest="command", required=True)
train = subparsers.add_parser("train")
train.add_argument("--config")
train.add_argument("--iterations", type=int)
train.add_argument("--traversals-per-iteration", type=int)
train.add_argument("--checkpoint-dir")
train.add_argument("--resume")
train.add_argument("--device", default="cpu")
train.add_argument("--eval-every", type=int)
train.add_argument("--eval-games", type=int)
train.add_argument("--seed", type=int)
train.add_argument("--no-save", action="store_true")
train.set_defaults(func=train_command)
evaluate = subparsers.add_parser("eval")
evaluate.add_argument("--checkpoint", required=True)
evaluate.add_argument("--opponent", default="random")
evaluate.add_argument("--games", type=int, default=10)
evaluate.add_argument("--seed", type=int, default=1)
evaluate.add_argument("--max-steps", type=int, default=10_000)
evaluate.add_argument("--device", default="cpu")
evaluate.set_defaults(func=eval_command)
args = parser.parse_args(argv)
args.func(args)
if __name__ == "__main__":
main()
@@ -1,6 +1,8 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
@dataclass(frozen=True) @dataclass(frozen=True)
@@ -28,3 +30,23 @@ class DeepCFRConfig:
hidden_size: int = 64 hidden_size: int = 64
learning_rate: float = 1.0e-3 learning_rate: float = 1.0e-3
seed: int = 1 seed: int = 1
checkpoint_dir: str = "runs/deep_cfr/default"
save_every_iteration: bool = True
eval_every: int = 0
eval_games: int = 10
eval_opponents: tuple[str, ...] = ("random",)
eval_max_steps: int = 10_000
def to_dict(self) -> dict[str, Any]:
return asdict(self)
@property
def checkpoint_path(self) -> Path:
return Path(self.checkpoint_dir)
def config_from_dict(data: dict[str, Any]) -> DeepCFRConfig:
values = dict(data)
if "eval_opponents" in values:
values["eval_opponents"] = tuple(values["eval_opponents"])
return DeepCFRConfig(**values)
@@ -0,0 +1,104 @@
from __future__ import annotations
from pathlib import Path
import numpy as np
import torch
from coolrl_lost_cities.games.classic.bots import build_bot
from coolrl_lost_cities.games.classic.deep_cfr.config import config_from_dict
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
from coolrl_lost_cities.games.classic.evaluation import evaluate_policy
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
from coolrl_lost_cities.games.classic.policy import LostCitiesPolicy, PolicyInput
class StrategyNetPolicy(LostCitiesPolicy):
def __init__(
self,
strategy_network: torch.nn.Module,
*,
device: torch.device | str = "cpu",
sample: bool = False,
seed: int | None = None,
) -> None:
self.strategy_network = strategy_network
self.device = torch.device(device)
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)
legal_actions = np.flatnonzero(legal)
if len(legal_actions) == 0:
raise RuntimeError("no legal action available")
return int(legal_actions[0])
state = obs_or_state
legal = np.asarray(state.unified_legal_mask(), dtype=bool)
legal_actions = np.flatnonzero(legal)
if len(legal_actions) == 0:
raise RuntimeError("no legal action available")
info = encode_info_state(state, state.current_player)
with torch.inference_mode():
x = torch.as_tensor(info, dtype=torch.float32, device=self.device).unsqueeze(0)
logits = self.strategy_network(x).squeeze(0).detach().cpu().numpy()
masked = np.where(legal, logits, -np.inf)
if self.sample:
stable = masked[legal_actions] - np.max(masked[legal_actions])
probs = np.exp(stable)
probs = probs / probs.sum()
unified = int(self.rng.choice(legal_actions, p=probs))
else:
unified = int(np.argmax(masked))
return state.from_unified_action(unified)
def evaluate_strategy_network(
strategy_network: torch.nn.Module,
config: LostCitiesConfig,
*,
games: int,
seed: int,
opponent: str = "random",
device: torch.device | str = "cpu",
max_steps: int = 10_000,
) -> dict[str, float | int]:
strategy_network.eval()
def make_strategy(seed_value: int | None = None) -> StrategyNetPolicy:
return StrategyNetPolicy(strategy_network, device=device, seed=seed_value)
def opponent_factory(seed_value=None):
return build_bot(opponent, seed=seed_value)
result = evaluate_policy(
make_strategy,
opponent_factory,
config,
games=games,
seed=seed,
max_steps=max_steps,
)
return result.to_dict()
def load_strategy_policy_from_checkpoint(
checkpoint_path: str | Path,
*,
device: torch.device | str = "cpu",
sample: bool = False,
seed: int | None = None,
) -> tuple[StrategyNetPolicy, LostCitiesConfig]:
payload = torch.load(checkpoint_path, map_location="cpu")
cfg = config_from_dict(payload["config"])
game_config = LostCitiesConfig(**payload["game_config"])
network = DeepCFRMLP(
int(payload["input_dim"]),
int(payload["action_size"]),
cfg.hidden_size,
).to(device)
network.load_state_dict(payload["strategy_network"])
network.eval()
return StrategyNetPolicy(network, device=device, sample=sample, seed=seed), game_config
@@ -1,13 +1,19 @@
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path
import numpy as np import numpy as np
import torch import torch
from torch import nn from torch import nn
from coolrl_lost_cities.games.classic.deep_cfr.checkpoints import (
load_checkpoint,
save_checkpoint,
)
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig
from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim
from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy_network
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
from coolrl_lost_cities.games.classic.deep_cfr.traverser import DeepCFRTraverser, TraversalStats from coolrl_lost_cities.games.classic.deep_cfr.traverser import DeepCFRTraverser, TraversalStats
@@ -26,6 +32,7 @@ class IterationMetrics:
traversal_depth_cutoffs: int traversal_depth_cutoffs: int
traversal_node_limit_cutoffs: int traversal_node_limit_cutoffs: int
traversal_max_depth_reached: int traversal_max_depth_reached: int
eval_metrics: dict[str, float | int]
class DeepCFRTrainer: class DeepCFRTrainer:
@@ -62,8 +69,44 @@ class DeepCFRTrainer:
self.advantage_memory = ReservoirMemory(self.config.advantage_memory_capacity) self.advantage_memory = ReservoirMemory(self.config.advantage_memory_capacity)
self.strategy_memory = ReservoirMemory(self.config.strategy_memory_capacity) self.strategy_memory = ReservoirMemory(self.config.strategy_memory_capacity)
self.rng = np.random.default_rng(self.config.seed + 101) self.rng = np.random.default_rng(self.config.seed + 101)
self.iteration = 0
def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict:
return {
"config": self.config.to_dict(),
"game_config": self.game_config.to_snapshot(),
"iteration": self.iteration,
"input_dim": self.input_dim,
"action_size": self.action_size,
"advantage_networks": [network.state_dict() for network in self.advantage_networks],
"strategy_network": self.strategy_network.state_dict(),
"advantage_optimizers": [
optimizer.state_dict() for optimizer in self.advantage_optimizers
],
"strategy_optimizer": self.strategy_optimizer.state_dict(),
"metrics": None if metrics is None else metrics.__dict__,
}
def save_checkpoint(self, path: str | Path, metrics: IterationMetrics | None = None) -> Path:
return save_checkpoint(path, self.checkpoint_payload(metrics))
def load_checkpoint(self, path: str | Path) -> None:
payload = load_checkpoint(path, device=self.device)
self.iteration = int(payload.get("iteration", 0))
for network, state_dict in zip(
self.advantage_networks, payload["advantage_networks"], strict=True
):
network.load_state_dict(state_dict)
self.strategy_network.load_state_dict(payload["strategy_network"])
for optimizer, state_dict in zip(
self.advantage_optimizers, payload.get("advantage_optimizers", []), strict=False
):
optimizer.load_state_dict(state_dict)
if "strategy_optimizer" in payload:
self.strategy_optimizer.load_state_dict(payload["strategy_optimizer"])
def run_iteration(self, iteration: int) -> IterationMetrics: def run_iteration(self, iteration: int) -> IterationMetrics:
self.iteration = iteration
total_stats = TraversalStats() total_stats = TraversalStats()
traverser = DeepCFRTraverser( traverser = DeepCFRTraverser(
self.advantage_networks, self.advantage_networks,
@@ -97,6 +140,7 @@ class DeepCFRTrainer:
advantage_loss = self._train_advantage_networks() advantage_loss = self._train_advantage_networks()
strategy_loss = self._train_strategy_network() strategy_loss = self._train_strategy_network()
eval_metrics = self._evaluate(iteration)
return IterationMetrics( return IterationMetrics(
iteration=iteration, iteration=iteration,
advantage_samples=len(self.advantage_memory), advantage_samples=len(self.advantage_memory),
@@ -108,10 +152,39 @@ class DeepCFRTrainer:
traversal_depth_cutoffs=total_stats.depth_cutoffs, traversal_depth_cutoffs=total_stats.depth_cutoffs,
traversal_node_limit_cutoffs=total_stats.node_limit_cutoffs, traversal_node_limit_cutoffs=total_stats.node_limit_cutoffs,
traversal_max_depth_reached=total_stats.max_depth_reached, traversal_max_depth_reached=total_stats.max_depth_reached,
eval_metrics=eval_metrics,
) )
def train(self) -> list[IterationMetrics]: def train(self) -> list[IterationMetrics]:
return [self.run_iteration(iteration) for iteration in range(1, self.config.iterations + 1)] metrics: list[IterationMetrics] = []
start = self.iteration + 1
stop = self.iteration + self.config.iterations
for iteration in range(start, stop + 1):
item = self.run_iteration(iteration)
metrics.append(item)
if self.config.save_every_iteration:
checkpoint_dir = self.config.checkpoint_path
self.save_checkpoint(checkpoint_dir / f"iteration_{iteration:05d}.pt", item)
self.save_checkpoint(checkpoint_dir / "latest.pt", item)
return metrics
def _evaluate(self, iteration: int) -> dict[str, float | int]:
if self.config.eval_every <= 0 or iteration % self.config.eval_every != 0:
return {}
results: dict[str, float | int] = {}
for opponent in self.config.eval_opponents:
result = evaluate_strategy_network(
self.strategy_network,
self.game_config,
games=self.config.eval_games,
seed=self.config.seed + iteration * 1000,
opponent=opponent,
device=self.device,
max_steps=self.config.eval_max_steps,
)
for key, value in result.items():
results[f"eval_{opponent}_{key}"] = value
return results
def _train_advantage_networks(self) -> float: def _train_advantage_networks(self) -> float:
losses: list[float] = [] losses: list[float] = []
@@ -21,6 +21,7 @@ def test_deep_cfr_trainer_smoke_run() -> None:
batch_size=2, batch_size=2,
hidden_size=16, hidden_size=16,
seed=23, seed=23,
save_every_iteration=False,
), ),
LostCitiesConfig(seed=23), LostCitiesConfig(seed=23),
) )
@@ -46,6 +47,7 @@ def test_deep_cfr_recursive_traverser_restores_state_and_collects_samples() -> N
batch_size=2, batch_size=2,
hidden_size=16, hidden_size=16,
seed=29, seed=29,
save_every_iteration=False,
), ),
LostCitiesConfig(seed=29), LostCitiesConfig(seed=29),
) )
@@ -94,6 +96,7 @@ def test_deep_cfr_traverser_supports_outcome_sampling_and_rollout_cutoffs() -> N
batch_size=2, batch_size=2,
hidden_size=16, hidden_size=16,
seed=31, seed=31,
save_every_iteration=False,
), ),
LostCitiesConfig(seed=31), LostCitiesConfig(seed=31),
) )
@@ -149,3 +152,41 @@ def test_reservoir_memory_caps_samples_and_filters_player_batches() -> None:
player_one = memory.sample(8, rng, player=1) player_one = memory.sample(8, rng, player=1)
assert player_one assert player_one
assert all(sample.player == 1 for sample in player_one) assert all(sample.player == 1 for sample in player_one)
def test_deep_cfr_trainer_saves_loads_and_evaluates_checkpoint(tmp_path) -> None:
checkpoint_dir = tmp_path / "deep_cfr"
trainer = DeepCFRTrainer(
DeepCFRConfig(
iterations=1,
traversals_per_iteration=1,
max_traversal_depth=2,
max_nodes_per_traversal=32,
batch_size=2,
hidden_size=16,
seed=41,
checkpoint_dir=str(checkpoint_dir),
save_every_iteration=True,
eval_every=1,
eval_games=2,
eval_opponents=("random",),
),
LostCitiesConfig(seed=41),
)
metrics = trainer.train()
latest = checkpoint_dir / "latest.pt"
restored = DeepCFRTrainer(
DeepCFRConfig(
hidden_size=16,
seed=41,
checkpoint_dir=str(checkpoint_dir),
save_every_iteration=False,
),
LostCitiesConfig(seed=41),
)
restored.load_checkpoint(latest)
assert latest.exists()
assert restored.iteration == 1
assert "eval_random_games" in metrics[0].eval_metrics