From 7ef7e5b927ecf2ec3dd7810baa251462fa38cd38 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:54:16 +0900 Subject: [PATCH] =?UTF-8?q?Deep=20CFR=20endpoint=20depth=20metrics=20?= =?UTF-8?q?=EB=B3=B4=EA=B0=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../games/classic/deep_cfr/trainer.py | 17 +++++++++++++++++ .../games/classic/deep_cfr/traverser.py | 12 ++++++++++-- .../games/classic/deep_cfr/workers.py | 2 ++ tests/games/classic/test_deep_cfr_trainer.py | 2 ++ 4 files changed, 31 insertions(+), 2 deletions(-) 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 9fb086b..b875534 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/trainer.py @@ -47,6 +47,10 @@ class IterationMetrics: traversal_depth_cutoffs: int traversal_node_limit_cutoffs: int traversal_max_depth_reached: int + traversal_endpoint_depth_sum: int + traversal_endpoints: int + traversal_avg_endpoint_depth: float + traversal_endpoint_depth_buckets: dict[str, int] eval_metrics: dict[str, float | int] def to_dict(self) -> dict[str, float | int]: @@ -61,6 +65,13 @@ class IterationMetrics: "traversal_depth_cutoffs": self.traversal_depth_cutoffs, "traversal_node_limit_cutoffs": self.traversal_node_limit_cutoffs, "traversal_max_depth_reached": self.traversal_max_depth_reached, + "traversal_endpoint_depth_sum": self.traversal_endpoint_depth_sum, + "traversal_endpoints": self.traversal_endpoints, + "traversal_avg_endpoint_depth": self.traversal_avg_endpoint_depth, + **{ + f"traversal_endpoint_depth_bucket_{key}": value + for key, value in self.traversal_endpoint_depth_buckets.items() + }, } data.update(self.eval_metrics) return data @@ -174,6 +185,10 @@ class DeepCFRTrainer: 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, + traversal_endpoint_depth_sum=total_stats.endpoint_depth_sum, + traversal_endpoints=total_stats.endpoints, + traversal_avg_endpoint_depth=total_stats.avg_endpoint_depth, + traversal_endpoint_depth_buckets=dict(total_stats.endpoint_depth_buckets), eval_metrics=eval_metrics, ) @@ -206,6 +221,8 @@ class DeepCFRTrainer: 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, + endpoint_depth_bucket_width=self.config.traversal.endpoint_depth_bucket_width, + endpoint_depth_bucket_max=self.config.traversal.endpoint_depth_bucket_max, encoding=self.config.encoding, rng=self.rng, ) 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 4926bdb..07d9a51 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/traverser.py @@ -65,6 +65,8 @@ class TraversalStats: "traversal_cutoff_rollouts": self.cutoff_rollouts, "traversal_cutoff_rollout_steps": self.cutoff_rollout_steps, "traversal_cutoff_rollout_timeouts": self.cutoff_rollout_timeouts, + "traversal_endpoint_depth_sum": self.endpoint_depth_sum, + "traversal_endpoints": self.endpoints, "traversal_avg_endpoint_depth": self.avg_endpoint_depth, **{ f"traversal_endpoint_depth_bucket_{key}": value @@ -103,6 +105,8 @@ class DeepCFRTraverser: self_play_older_weight: float = 0.2, self_play_anchor_weight: float = 0.0, self_play_recent_window: int = 5, + endpoint_depth_bucket_width: int = 10, + endpoint_depth_bucket_max: int = 100, encoding=None, rng: np.random.Generator | None = None, ) -> None: @@ -146,6 +150,8 @@ class DeepCFRTraverser: 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.endpoint_depth_bucket_width = max(1, int(endpoint_depth_bucket_width)) + self.endpoint_depth_bucket_max = max(1, int(endpoint_depth_bucket_max)) self.encoding = encoding self.rng = rng or np.random.default_rng() self._safe_heuristic_rollout_bot = ( @@ -451,6 +457,8 @@ class DeepCFRTraverser: def _record_endpoint(self, stats: TraversalStats, depth: int) -> None: stats.endpoint_depth_sum += depth - start = min(depth // 10 * 10, 100) - key = "100_plus" if start >= 100 else f"{start}_{start + 9}" + width = self.endpoint_depth_bucket_width + max_depth = self.endpoint_depth_bucket_max + start = min(depth // width * width, max_depth) + key = f"{max_depth}_plus" if start >= max_depth else f"{start}_{start + width - 1}" stats.endpoint_depth_buckets[key] = stats.endpoint_depth_buckets.get(key, 0) + 1 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 d3791be..0dce38b 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/workers.py @@ -85,6 +85,8 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe 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, + endpoint_depth_bucket_width=cfg.traversal.endpoint_depth_bucket_width, + endpoint_depth_bucket_max=cfg.traversal.endpoint_depth_bucket_max, encoding=cfg.encoding, rng=np.random.default_rng(batch.worker_seed), ) diff --git a/tests/games/classic/test_deep_cfr_trainer.py b/tests/games/classic/test_deep_cfr_trainer.py index 504755a..0ed24c9 100644 --- a/tests/games/classic/test_deep_cfr_trainer.py +++ b/tests/games/classic/test_deep_cfr_trainer.py @@ -124,6 +124,8 @@ def test_deep_cfr_trainer_smoke_run() -> None: assert metrics[0].strategy_samples > 0 assert metrics[0].traversal_nodes > 0 assert metrics[0].traversal_max_depth_reached <= 3 + assert metrics[0].traversal_endpoints > 0 + assert metrics[0].traversal_avg_endpoint_depth >= 0.0 assert metrics[0].advantage_loss >= 0.0 assert metrics[0].strategy_loss >= 0.0