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 7ae77e8..a6809a4 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -19,6 +19,10 @@ class DeepCFRConfig: cutoff_rollouts: int = 0 cutoff_rollout_policy: str = "random" cutoff_rollout_max_steps: int = 10_000 + opponent_policy: str = "network" + self_play_snapshot_every: int = 1 + self_play_max_snapshots: int = 20 + self_play_anchor_probability: float = 0.0 strategy_sample_interval: int = 1 store_strategy_on_traverser_nodes: bool = True store_strategy_on_opponent_nodes: bool = True 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 bcde4f0..9ee021f 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -98,6 +98,7 @@ class DeepCFRTrainer: self.metrics_path = self.run_dir / "metrics.jsonl" self.progress_path = self.run_dir / "runtime_progress.json" self.log_path = self.run_dir / "train.log" + self.self_play_league_snapshots: list[list[dict]] = [] def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict: return { @@ -108,6 +109,7 @@ class DeepCFRTrainer: "action_size": self.action_size, "advantage_networks": [network.state_dict() for network in self.advantage_networks], "strategy_network": self.strategy_network.state_dict(), + "self_play_league_snapshots": self.self_play_league_snapshots, "advantage_optimizers": [ optimizer.state_dict() for optimizer in self.advantage_optimizers ], @@ -132,6 +134,7 @@ class DeepCFRTrainer: optimizer.load_state_dict(state_dict) if "strategy_optimizer" in payload: self.strategy_optimizer.load_state_dict(payload["strategy_optimizer"]) + self.self_play_league_snapshots = payload.get("self_play_league_snapshots", []) def run_iteration(self, iteration: int) -> IterationMetrics: self.iteration = iteration @@ -178,6 +181,9 @@ class DeepCFRTrainer: cutoff_rollouts=self.config.cutoff_rollouts, cutoff_rollout_policy=self.config.cutoff_rollout_policy, cutoff_rollout_max_steps=self.config.cutoff_rollout_max_steps, + opponent_policy=self.config.opponent_policy, + league_advantage_networks=self._materialize_league_networks(), + self_play_anchor_probability=self.config.self_play_anchor_probability, rng=self.rng, ) for network in self.advantage_networks: @@ -233,12 +239,49 @@ class DeepCFRTrainer: input_dim=self.input_dim, action_size=self.action_size, advantage_networks=network_payloads, + league_advantage_networks=self._league_payloads(), worker_seed=self.config.seed + iteration * 1_000_003 + batch_index, ) ) batch_index += 1 return batches + def _frozen_advantage_state_dicts(self) -> list[dict]: + return [ + {name: value.detach().cpu().clone() for name, value in network.state_dict().items()} + for network in self.advantage_networks + ] + + def _maybe_record_self_play_snapshot(self, iteration: int) -> None: + if self.config.opponent_policy != "self_play_league": + return + if self.config.self_play_max_snapshots <= 0: + return + if iteration % max(1, self.config.self_play_snapshot_every) != 0: + return + self.self_play_league_snapshots.append(self._frozen_advantage_state_dicts()) + overflow = len(self.self_play_league_snapshots) - self.config.self_play_max_snapshots + if overflow > 0: + del self.self_play_league_snapshots[:overflow] + + def _materialize_league_networks(self) -> list[list[nn.Module]]: + league: list[list[nn.Module]] = [] + for snapshot in self.self_play_league_snapshots: + networks = [ + DeepCFRMLP(self.input_dim, self.action_size, self.config.hidden_size).to( + self.device + ) + for _ in range(2) + ] + for network, state_dict in zip(networks, snapshot, strict=True): + network.load_state_dict(state_dict) + network.eval() + league.append(networks) + return league + + def _league_payloads(self) -> list[list[dict]]: + return self.self_play_league_snapshots + def train(self) -> list[IterationMetrics]: self._start_run_logging() metrics: list[IterationMetrics] = [] @@ -250,6 +293,7 @@ class DeepCFRTrainer: elapsed = time.perf_counter() - started metrics.append(item) self._append_metrics(item, elapsed) + self._maybe_record_self_play_snapshot(iteration) if self.config.save_every_iteration: checkpoint_dir = self.run_dir self.save_checkpoint(checkpoint_dir / f"iteration_{iteration:05d}.pt", item) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py b/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py index 76f8bc6..20b27c7 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py @@ -95,6 +95,9 @@ class DeepCFRTraverser: cutoff_rollouts: int = 0, cutoff_rollout_policy: str = "random", cutoff_rollout_max_steps: int = 10_000, + opponent_policy: str = "network", + league_advantage_networks: list[list[torch.nn.Module]] | None = None, + self_play_anchor_probability: float = 0.0, rng: np.random.Generator | None = None, ) -> None: self.advantage_networks = advantage_networks @@ -125,10 +128,22 @@ class DeepCFRTraverser: if self.cutoff_rollout_policy not in {"random", "safe_heuristic"}: raise ValueError("cutoff_rollout_policy must be 'random' or 'safe_heuristic'") self.cutoff_rollout_max_steps = max(1, int(cutoff_rollout_max_steps)) + self.opponent_policy = opponent_policy + if self.opponent_policy not in {"network", "safe_heuristic", "self_play_league"}: + raise ValueError( + "opponent_policy must be 'network', 'safe_heuristic', or 'self_play_league'" + ) + self.league_advantage_networks = league_advantage_networks or [] + self.self_play_anchor_probability = min(1.0, max(0.0, float(self_play_anchor_probability))) self.rng = rng or np.random.default_rng() self._safe_heuristic_rollout_bot = ( SafeHeuristicBot() if self.cutoff_rollout_policy == "safe_heuristic" else None ) + self._safe_heuristic_opponent_bot = ( + SafeHeuristicBot() + if self.opponent_policy == "safe_heuristic" or self.self_play_anchor_probability > 0.0 + else None + ) def traverse( self, state: GameState, traverser: int, iteration: int @@ -163,6 +178,24 @@ class DeepCFRTraverser: return self._cutoff_value(state, traverser, stats) player = state.current_player + fixed_action = self._fixed_opponent_action(state, player, traverser) + if fixed_action is not None: + unified_action = state.to_unified_action(fixed_action) + swapped_deck_index = self._sample_deck_draw_chance(state, unified_action) + state.push_action(fixed_action) + try: + return self._traverse( + state, + traverser, + iteration, + depth=depth + 1, + stats=stats, + ) + finally: + state.pop_action() + if swapped_deck_index is not None: + state.swap_deck_cards(swapped_deck_index, len(state.deck) - 1) + info_state, legal, policy = self._policy(state, player) self._record_strategy(info_state, legal, policy, player, traverser, iteration, depth, stats) @@ -224,21 +257,53 @@ class DeepCFRTraverser: return node_value def _policy(self, state: GameState, player: int) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + return self._policy_from_networks(self.advantage_networks, state, player) + + def _policy_from_networks( + self, + networks: list[torch.nn.Module], + state: GameState, + player: int, + ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: info_state = encode_info_state(state, player) legal = np.asarray(state.unified_legal_mask(), dtype=bool) with torch.inference_mode(): x = torch.as_tensor(info_state, dtype=torch.float32, device=self.device).unsqueeze(0) - advantages = ( - self.advantage_networks[player](x) - .squeeze(0) - .detach() - .cpu() - .numpy() - .astype(np.float32) - ) + advantages = networks[player](x).squeeze(0).detach().cpu().numpy().astype(np.float32) policy = regret_matching(advantages, legal, self.epsilon).astype(np.float32) return info_state, legal, policy + def _fixed_opponent_action( + self, + state: GameState, + player: int, + traverser: int, + ) -> int | None: + if player == traverser or self.opponent_policy == "network": + return None + if self.opponent_policy == "safe_heuristic": + if self._safe_heuristic_opponent_bot is None: + self._safe_heuristic_opponent_bot = SafeHeuristicBot() + return self._safe_heuristic_opponent_bot.act(state) + if ( + self.self_play_anchor_probability > 0.0 + and self.rng.random() < self.self_play_anchor_probability + ): + if self._safe_heuristic_opponent_bot is None: + self._safe_heuristic_opponent_bot = SafeHeuristicBot() + return self._safe_heuristic_opponent_bot.act(state) + if not self.league_advantage_networks: + return None + networks = self.league_advantage_networks[ + int(self.rng.integers(0, len(self.league_advantage_networks))) + ] + _, legal, policy = self._policy_from_networks(networks, state, player) + legal_actions = np.flatnonzero(legal) + if len(legal_actions) == 0: + return None + unified_action = self._sample_action(policy, legal_actions) + return state.from_unified_action(unified_action) + def _sample_action(self, policy: np.ndarray, legal_actions: np.ndarray) -> int: probs = policy[legal_actions].astype(np.float64) total = float(probs.sum()) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py index f750488..19ad11f 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -23,6 +23,7 @@ class TraversalWorkerBatch: input_dim: int action_size: int advantage_networks: list[dict[str, Any]] + league_advantage_networks: list[list[dict[str, Any]]] worker_seed: int @@ -44,6 +45,16 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe for network, state_dict in zip(networks, batch.advantage_networks, strict=True): network.load_state_dict(state_dict) network.eval() + league_networks: list[list[torch.nn.Module]] = [] + for snapshot in batch.league_advantage_networks: + snapshot_networks = [ + DeepCFRMLP(batch.input_dim, batch.action_size, cfg.hidden_size).to(device) + for _ in range(2) + ] + for network, state_dict in zip(snapshot_networks, snapshot, strict=True): + network.load_state_dict(state_dict) + network.eval() + league_networks.append(snapshot_networks) advantage_memory = ReservoirMemory() strategy_memory = ReservoirMemory() traverser = DeepCFRTraverser( @@ -65,6 +76,9 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe cutoff_rollouts=cfg.cutoff_rollouts, cutoff_rollout_policy=cfg.cutoff_rollout_policy, cutoff_rollout_max_steps=cfg.cutoff_rollout_max_steps, + opponent_policy=cfg.opponent_policy, + league_advantage_networks=league_networks, + self_play_anchor_probability=cfg.self_play_anchor_probability, rng=np.random.default_rng(batch.worker_seed), ) game_config = LostCitiesConfig(**batch.game_config) diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index e458854..3cdba37 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -234,3 +234,29 @@ def test_deep_cfr_traversal_benchmark_smoke() -> None: assert result["traversal_nodes"] > 0 assert result["nodes_per_second"] > 0.0 + + +def test_deep_cfr_self_play_league_records_snapshots(tmp_path) -> None: + trainer = DeepCFRTrainer( + DeepCFRConfig( + iterations=2, + traversals_per_iteration=1, + max_traversal_depth=2, + max_nodes_per_traversal=32, + batch_size=2, + hidden_size=16, + seed=53, + checkpoint_dir=str(tmp_path / "league"), + save_every_iteration=False, + opponent_policy="self_play_league", + self_play_snapshot_every=1, + self_play_max_snapshots=1, + self_play_anchor_probability=1.0, + ), + LostCitiesConfig(seed=53), + ) + + metrics = trainer.train() + + assert len(metrics) == 2 + assert len(trainer.self_play_league_snapshots) == 1