diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/benchmark.py b/src/coolrl_lost_cities/games/classic/deep_cfr/benchmark.py new file mode 100644 index 0000000..9fe1080 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/benchmark.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +import time +from dataclasses import replace + +from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig +from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer +from coolrl_lost_cities.games.classic.game import classic_config + + +def benchmark_traversal( + config: DeepCFRConfig | None = None, + *, + num_workers: int = 0, +) -> dict[str, float | int]: + base = config or DeepCFRConfig( + iterations=1, + traversals_per_iteration=8, + max_traversal_depth=4, + save_every_iteration=False, + ) + cfg = replace(base, num_workers=num_workers, save_every_iteration=False, eval_every=0) + trainer = DeepCFRTrainer(cfg, classic_config(seed=cfg.seed)) + started = time.perf_counter() + metrics = trainer.run_iteration(1) + elapsed = time.perf_counter() - started + return { + "num_workers": num_workers, + "elapsed_seconds": elapsed, + "traversal_nodes": metrics.traversal_nodes, + "nodes_per_second": metrics.traversal_nodes / max(elapsed, 1.0e-12), + "advantage_samples": metrics.advantage_samples, + "strategy_samples": metrics.strategy_samples, + } diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py index a1e10e6..5328504 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py @@ -5,6 +5,7 @@ import json from dataclasses import replace from pathlib import Path +from coolrl_lost_cities.games.classic.deep_cfr.benchmark import benchmark_traversal 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, @@ -59,6 +60,19 @@ def eval_command(args: argparse.Namespace) -> None: print(json.dumps(result, indent=2, sort_keys=True)) +def benchmark_command(args: argparse.Namespace) -> None: + result = benchmark_traversal( + DeepCFRConfig( + traversals_per_iteration=args.traversals, + max_traversal_depth=args.depth, + seed=args.seed, + save_every_iteration=False, + ), + num_workers=args.workers, + ) + 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) @@ -85,6 +99,13 @@ def main(argv: list[str] | None = None) -> None: evaluate.add_argument("--device", default="cpu") evaluate.set_defaults(func=eval_command) + benchmark = subparsers.add_parser("benchmark") + benchmark.add_argument("--workers", type=int, default=0) + benchmark.add_argument("--traversals", type=int, default=8) + benchmark.add_argument("--depth", type=int, default=4) + benchmark.add_argument("--seed", type=int, default=1) + benchmark.set_defaults(func=benchmark_command) + args = parser.parse_args(argv) args.func(args) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py index 22572e3..7ae77e8 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -36,6 +36,8 @@ class DeepCFRConfig: eval_games: int = 10 eval_opponents: tuple[str, ...] = ("random",) eval_max_steps: int = 10_000 + num_workers: int = 0 + traversal_worker_chunk_size: int = 4 def to_dict(self) -> dict[str, Any]: return asdict(self) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py index 39bf1c9..bcde4f0 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -1,7 +1,9 @@ from __future__ import annotations import json +import multiprocessing as mp import time +from concurrent.futures import ProcessPoolExecutor, as_completed from dataclasses import dataclass from pathlib import Path @@ -19,6 +21,10 @@ from coolrl_lost_cities.games.classic.deep_cfr.evaluate import evaluate_strategy 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.traverser import DeepCFRTraverser, TraversalStats +from coolrl_lost_cities.games.classic.deep_cfr.workers import ( + TraversalWorkerBatch, + run_traversal_worker_batch, +) from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig @@ -129,6 +135,29 @@ class DeepCFRTrainer: def run_iteration(self, iteration: int) -> IterationMetrics: self.iteration = iteration + if self.config.num_workers > 1: + total_stats = self._run_traversals_parallel(iteration) + else: + total_stats = self._run_traversals_single_process(iteration) + + advantage_loss = self._train_advantage_networks() + strategy_loss = self._train_strategy_network() + eval_metrics = self._evaluate(iteration) + return IterationMetrics( + iteration=iteration, + advantage_samples=len(self.advantage_memory), + strategy_samples=len(self.strategy_memory), + advantage_loss=advantage_loss, + strategy_loss=strategy_loss, + traversal_nodes=total_stats.nodes, + traversal_terminals=total_stats.terminals, + traversal_depth_cutoffs=total_stats.depth_cutoffs, + traversal_node_limit_cutoffs=total_stats.node_limit_cutoffs, + traversal_max_depth_reached=total_stats.max_depth_reached, + eval_metrics=eval_metrics, + ) + + def _run_traversals_single_process(self, iteration: int) -> TraversalStats: total_stats = TraversalStats() traverser = DeepCFRTraverser( self.advantage_networks, @@ -159,23 +188,56 @@ class DeepCFRTrainer: state = GameState.new_game(self.game_config, seed=seed) _, stats = traverser.traverse(state, player, iteration) total_stats.accumulate(stats) + return total_stats - advantage_loss = self._train_advantage_networks() - strategy_loss = self._train_strategy_network() - eval_metrics = self._evaluate(iteration) - return IterationMetrics( - iteration=iteration, - advantage_samples=len(self.advantage_memory), - strategy_samples=len(self.strategy_memory), - advantage_loss=advantage_loss, - strategy_loss=strategy_loss, - traversal_nodes=total_stats.nodes, - traversal_terminals=total_stats.terminals, - traversal_depth_cutoffs=total_stats.depth_cutoffs, - traversal_node_limit_cutoffs=total_stats.node_limit_cutoffs, - traversal_max_depth_reached=total_stats.max_depth_reached, - eval_metrics=eval_metrics, - ) + def _run_traversals_parallel(self, iteration: int) -> TraversalStats: + batches = self._worker_batches(iteration) + total_stats = TraversalStats() + if not batches: + return total_stats + max_workers = min(self.config.num_workers, len(batches)) + with ProcessPoolExecutor( + max_workers=max_workers, + mp_context=mp.get_context("spawn"), + ) as executor: + futures = [executor.submit(run_traversal_worker_batch, batch) for batch in batches] + for future in as_completed(futures): + result = future.result() + total_stats.accumulate(result.stats) + self.advantage_memory.extend(result.advantage_samples, self.rng) + self.strategy_memory.extend(result.strategy_samples, self.rng) + return total_stats + + def _worker_batches(self, iteration: int) -> list[TraversalWorkerBatch]: + batches: list[TraversalWorkerBatch] = [] + network_payloads = [ + {name: value.detach().cpu() for name, value in network.state_dict().items()} + for network in self.advantage_networks + ] + chunk_size = max(1, self.config.traversal_worker_chunk_size) + batch_index = 0 + for player in range(2): + seeds = [ + self.config.seed + iteration * 10_000 + index * 10 + player + for index in range(self.config.traversals_per_iteration) + ] + for start in range(0, len(seeds), chunk_size): + chunk = seeds[start : start + chunk_size] + batches.append( + TraversalWorkerBatch( + player=player, + iteration=iteration, + seeds=chunk, + config=self.config.to_dict(), + game_config=self.game_config.to_snapshot(), + input_dim=self.input_dim, + action_size=self.action_size, + advantage_networks=network_payloads, + worker_seed=self.config.seed + iteration * 1_000_003 + batch_index, + ) + ) + batch_index += 1 + return batches def train(self) -> list[IterationMetrics]: self._start_run_logging() diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py new file mode 100644 index 0000000..f750488 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -0,0 +1,82 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import numpy as np +import torch + +from coolrl_lost_cities.games.classic.deep_cfr.config import config_from_dict +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.traverser import DeepCFRTraverser, TraversalStats +from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig + + +@dataclass(frozen=True) +class TraversalWorkerBatch: + player: int + iteration: int + seeds: list[int] + config: dict[str, Any] + game_config: dict[str, Any] + input_dim: int + action_size: int + advantage_networks: list[dict[str, Any]] + worker_seed: int + + +@dataclass(frozen=True) +class TraversalWorkerResult: + player: int + stats: TraversalStats + advantage_samples: list[TrainingSample] + strategy_samples: list[TrainingSample] + traversals: int + + +def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerResult: + cfg = config_from_dict(batch.config) + device = torch.device("cpu") + networks = [ + DeepCFRMLP(batch.input_dim, batch.action_size, cfg.hidden_size).to(device) for _ in range(2) + ] + for network, state_dict in zip(networks, batch.advantage_networks, strict=True): + network.load_state_dict(state_dict) + network.eval() + advantage_memory = ReservoirMemory() + strategy_memory = ReservoirMemory() + traverser = DeepCFRTraverser( + networks, + advantage_memory, + strategy_memory, + device=device, + action_size=batch.action_size, + epsilon=cfg.regret_matching_epsilon, + strategy_sample_interval=cfg.strategy_sample_interval, + store_strategy_on_traverser_nodes=cfg.store_strategy_on_traverser_nodes, + store_strategy_on_opponent_nodes=cfg.store_strategy_on_opponent_nodes, + max_depth=cfg.max_traversal_depth, + max_nodes=cfg.max_nodes_per_traversal, + outcome_sampling_epsilon=cfg.outcome_sampling_epsilon, + outcome_sampling_value_clip=cfg.outcome_sampling_value_clip, + outcome_unsampled_regret=cfg.outcome_unsampled_regret, + cutoff_value_mode=cfg.cutoff_value_mode, + cutoff_rollouts=cfg.cutoff_rollouts, + cutoff_rollout_policy=cfg.cutoff_rollout_policy, + cutoff_rollout_max_steps=cfg.cutoff_rollout_max_steps, + rng=np.random.default_rng(batch.worker_seed), + ) + game_config = LostCitiesConfig(**batch.game_config) + total_stats = TraversalStats() + for seed in batch.seeds: + state = GameState.new_game(game_config, seed=seed) + _, stats = traverser.traverse(state, batch.player, batch.iteration) + total_stats.accumulate(stats) + return TraversalWorkerResult( + player=batch.player, + stats=total_stats, + advantage_samples=advantage_memory.all(), + strategy_samples=strategy_memory.all(), + traversals=len(batch.seeds), + ) diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 609e229..e458854 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -3,6 +3,7 @@ from __future__ import annotations import numpy as np from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig +from coolrl_lost_cities.games.classic.deep_cfr.benchmark import benchmark_traversal from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer @@ -194,3 +195,42 @@ def test_deep_cfr_trainer_saves_loads_and_evaluates_checkpoint(tmp_path) -> None assert (checkpoint_dir / "train.log").exists() assert restored.iteration == 1 assert "eval_random_games" in metrics[0].eval_metrics + + +def test_deep_cfr_trainer_multiprocessing_smoke_run(tmp_path) -> None: + trainer = DeepCFRTrainer( + DeepCFRConfig( + iterations=1, + traversals_per_iteration=2, + max_traversal_depth=2, + max_nodes_per_traversal=32, + batch_size=2, + hidden_size=16, + seed=43, + checkpoint_dir=str(tmp_path / "mp"), + save_every_iteration=False, + num_workers=2, + traversal_worker_chunk_size=1, + ), + LostCitiesConfig(seed=43), + ) + + metrics = trainer.train() + + assert metrics[0].traversal_nodes > 0 + assert metrics[0].advantage_samples > 0 + + +def test_deep_cfr_traversal_benchmark_smoke() -> None: + result = benchmark_traversal( + DeepCFRConfig( + traversals_per_iteration=1, + max_traversal_depth=2, + hidden_size=16, + save_every_iteration=False, + seed=47, + ) + ) + + assert result["traversal_nodes"] > 0 + assert result["nodes_per_second"] > 0.0