Deep CFR endpoint depth metrics 보강

This commit is contained in:
2026-05-07 00:54:16 +09:00
parent 95d0660b0b
commit 7ef7e5b927
4 changed files with 31 additions and 2 deletions
@@ -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,
)
@@ -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
@@ -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),
)
@@ -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