From 5555e781af4e3d17d2d1346bc69a0e9774c125c9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Thu, 7 May 2026 00:09:36 +0900 Subject: [PATCH] =?UTF-8?q?Deep=20CFR=20weighted=20self-play=20league=20?= =?UTF-8?q?=EC=B6=94=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../games/classic/deep_cfr/config.py | 5 ++ .../games/classic/deep_cfr/trainer.py | 5 ++ .../games/classic/deep_cfr/traverser.py | 64 ++++++++++++++++--- .../games/classic/deep_cfr/workers.py | 5 ++ tests/games/classic/test_deep_cfr_trainer.py | 31 +++++++++ 5 files changed, 102 insertions(+), 8 deletions(-) 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 a6809a4..f419a4c 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/config.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/config.py @@ -23,6 +23,11 @@ class DeepCFRConfig: self_play_snapshot_every: int = 1 self_play_max_snapshots: int = 20 self_play_anchor_probability: float = 0.0 + self_play_current_weight: float = 0.5 + self_play_recent_weight: float = 0.3 + self_play_older_weight: float = 0.2 + self_play_anchor_weight: float = 0.0 + self_play_recent_window: int = 5 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 9ee021f..75bafc0 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -184,6 +184,11 @@ class DeepCFRTrainer: opponent_policy=self.config.opponent_policy, league_advantage_networks=self._materialize_league_networks(), self_play_anchor_probability=self.config.self_play_anchor_probability, + self_play_current_weight=self.config.self_play_current_weight, + self_play_recent_weight=self.config.self_play_recent_weight, + self_play_older_weight=self.config.self_play_older_weight, + self_play_anchor_weight=self.config.self_play_anchor_weight, + self_play_recent_window=self.config.self_play_recent_window, rng=self.rng, ) for network in self.advantage_networks: 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 20b27c7..461de46 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py @@ -98,6 +98,11 @@ class DeepCFRTraverser: opponent_policy: str = "network", league_advantage_networks: list[list[torch.nn.Module]] | None = None, self_play_anchor_probability: float = 0.0, + self_play_current_weight: float = 0.5, + self_play_recent_weight: float = 0.3, + self_play_older_weight: float = 0.2, + self_play_anchor_weight: float = 0.0, + self_play_recent_window: int = 5, rng: np.random.Generator | None = None, ) -> None: self.advantage_networks = advantage_networks @@ -135,6 +140,11 @@ class DeepCFRTraverser: ) 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.self_play_current_weight = max(0.0, float(self_play_current_weight)) + self.self_play_recent_weight = max(0.0, float(self_play_recent_weight)) + self.self_play_older_weight = max(0.0, float(self_play_older_weight)) + self.self_play_anchor_weight = max(0.0, float(self_play_anchor_weight)) + self.self_play_recent_window = max(0, int(self_play_recent_window)) self.rng = rng or np.random.default_rng() self._safe_heuristic_rollout_bot = ( SafeHeuristicBot() if self.cutoff_rollout_policy == "safe_heuristic" else None @@ -285,18 +295,16 @@ class DeepCFRTraverser: 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 - ): + bucket = self._self_play_bucket() + if bucket == "current": + return None + if bucket == "anchor": 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: + networks = self._self_play_snapshot_networks(bucket) + if networks is None: 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: @@ -304,6 +312,46 @@ class DeepCFRTraverser: unified_action = self._sample_action(policy, legal_actions) return state.from_unified_action(unified_action) + def _self_play_bucket(self) -> str: + if ( + self.self_play_anchor_probability > 0.0 + and self.rng.random() < self.self_play_anchor_probability + ): + return "anchor" + recent_count = min(len(self.league_advantage_networks), self.self_play_recent_window) + older_count = max(0, len(self.league_advantage_networks) - recent_count) + labels = ["current", "recent", "older", "anchor"] + weights = np.asarray( + [ + self.self_play_current_weight, + self.self_play_recent_weight if recent_count > 0 else 0.0, + self.self_play_older_weight if older_count > 0 else 0.0, + self.self_play_anchor_weight, + ], + dtype=np.float64, + ) + total = float(weights.sum()) + if total <= 0.0: + return "current" + weights /= total + return str(self.rng.choice(labels, p=weights)) + + def _self_play_snapshot_networks(self, bucket: str) -> list[torch.nn.Module] | None: + if not self.league_advantage_networks: + return None + recent_count = min(len(self.league_advantage_networks), self.self_play_recent_window) + if bucket == "recent" and recent_count > 0: + candidates = self.league_advantage_networks[-recent_count:] + elif bucket == "older": + candidates = self.league_advantage_networks[ + : max(0, len(self.league_advantage_networks) - recent_count) + ] + else: + candidates = self.league_advantage_networks + if not candidates: + return None + return candidates[int(self.rng.integers(0, len(candidates)))] + 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 19ad11f..ca0ea77 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -79,6 +79,11 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe opponent_policy=cfg.opponent_policy, league_advantage_networks=league_networks, self_play_anchor_probability=cfg.self_play_anchor_probability, + self_play_current_weight=cfg.self_play_current_weight, + self_play_recent_weight=cfg.self_play_recent_weight, + self_play_older_weight=cfg.self_play_older_weight, + self_play_anchor_weight=cfg.self_play_anchor_weight, + self_play_recent_window=cfg.self_play_recent_window, 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 3cdba37..5e52a46 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -260,3 +260,34 @@ def test_deep_cfr_self_play_league_records_snapshots(tmp_path) -> None: assert len(metrics) == 2 assert len(trainer.self_play_league_snapshots) == 1 + + +def test_deep_cfr_weighted_self_play_league_uses_snapshot_bucket(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=59, + checkpoint_dir=str(tmp_path / "weighted-league"), + save_every_iteration=False, + opponent_policy="self_play_league", + self_play_snapshot_every=1, + self_play_max_snapshots=2, + self_play_current_weight=0.0, + self_play_recent_weight=1.0, + self_play_older_weight=0.0, + self_play_anchor_weight=0.0, + self_play_recent_window=1, + ), + LostCitiesConfig(seed=59), + ) + + metrics = trainer.train() + + assert len(metrics) == 2 + assert len(trainer.self_play_league_snapshots) == 2 + assert metrics[1].traversal_nodes > 0