Deep CFR endpoint depth metrics 보강
This commit is contained in:
@@ -47,6 +47,10 @@ class IterationMetrics:
|
|||||||
traversal_depth_cutoffs: int
|
traversal_depth_cutoffs: int
|
||||||
traversal_node_limit_cutoffs: int
|
traversal_node_limit_cutoffs: int
|
||||||
traversal_max_depth_reached: 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]
|
eval_metrics: dict[str, float | int]
|
||||||
|
|
||||||
def to_dict(self) -> 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_depth_cutoffs": self.traversal_depth_cutoffs,
|
||||||
"traversal_node_limit_cutoffs": self.traversal_node_limit_cutoffs,
|
"traversal_node_limit_cutoffs": self.traversal_node_limit_cutoffs,
|
||||||
"traversal_max_depth_reached": self.traversal_max_depth_reached,
|
"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)
|
data.update(self.eval_metrics)
|
||||||
return data
|
return data
|
||||||
@@ -174,6 +185,10 @@ class DeepCFRTrainer:
|
|||||||
traversal_depth_cutoffs=total_stats.depth_cutoffs,
|
traversal_depth_cutoffs=total_stats.depth_cutoffs,
|
||||||
traversal_node_limit_cutoffs=total_stats.node_limit_cutoffs,
|
traversal_node_limit_cutoffs=total_stats.node_limit_cutoffs,
|
||||||
traversal_max_depth_reached=total_stats.max_depth_reached,
|
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,
|
eval_metrics=eval_metrics,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -206,6 +221,8 @@ class DeepCFRTrainer:
|
|||||||
self_play_older_weight=self.config.self_play.older_weight,
|
self_play_older_weight=self.config.self_play.older_weight,
|
||||||
self_play_anchor_weight=self.config.self_play.anchor_weight,
|
self_play_anchor_weight=self.config.self_play.anchor_weight,
|
||||||
self_play_recent_window=self.config.self_play.recent_window,
|
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,
|
encoding=self.config.encoding,
|
||||||
rng=self.rng,
|
rng=self.rng,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -65,6 +65,8 @@ class TraversalStats:
|
|||||||
"traversal_cutoff_rollouts": self.cutoff_rollouts,
|
"traversal_cutoff_rollouts": self.cutoff_rollouts,
|
||||||
"traversal_cutoff_rollout_steps": self.cutoff_rollout_steps,
|
"traversal_cutoff_rollout_steps": self.cutoff_rollout_steps,
|
||||||
"traversal_cutoff_rollout_timeouts": self.cutoff_rollout_timeouts,
|
"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,
|
"traversal_avg_endpoint_depth": self.avg_endpoint_depth,
|
||||||
**{
|
**{
|
||||||
f"traversal_endpoint_depth_bucket_{key}": value
|
f"traversal_endpoint_depth_bucket_{key}": value
|
||||||
@@ -103,6 +105,8 @@ class DeepCFRTraverser:
|
|||||||
self_play_older_weight: float = 0.2,
|
self_play_older_weight: float = 0.2,
|
||||||
self_play_anchor_weight: float = 0.0,
|
self_play_anchor_weight: float = 0.0,
|
||||||
self_play_recent_window: int = 5,
|
self_play_recent_window: int = 5,
|
||||||
|
endpoint_depth_bucket_width: int = 10,
|
||||||
|
endpoint_depth_bucket_max: int = 100,
|
||||||
encoding=None,
|
encoding=None,
|
||||||
rng: np.random.Generator | None = None,
|
rng: np.random.Generator | None = 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_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_anchor_weight = max(0.0, float(self_play_anchor_weight))
|
||||||
self.self_play_recent_window = max(0, int(self_play_recent_window))
|
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.encoding = encoding
|
||||||
self.rng = rng or np.random.default_rng()
|
self.rng = rng or np.random.default_rng()
|
||||||
self._safe_heuristic_rollout_bot = (
|
self._safe_heuristic_rollout_bot = (
|
||||||
@@ -451,6 +457,8 @@ class DeepCFRTraverser:
|
|||||||
|
|
||||||
def _record_endpoint(self, stats: TraversalStats, depth: int) -> None:
|
def _record_endpoint(self, stats: TraversalStats, depth: int) -> None:
|
||||||
stats.endpoint_depth_sum += depth
|
stats.endpoint_depth_sum += depth
|
||||||
start = min(depth // 10 * 10, 100)
|
width = self.endpoint_depth_bucket_width
|
||||||
key = "100_plus" if start >= 100 else f"{start}_{start + 9}"
|
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
|
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_older_weight=cfg.self_play.older_weight,
|
||||||
self_play_anchor_weight=cfg.self_play.anchor_weight,
|
self_play_anchor_weight=cfg.self_play.anchor_weight,
|
||||||
self_play_recent_window=cfg.self_play.recent_window,
|
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,
|
encoding=cfg.encoding,
|
||||||
rng=np.random.default_rng(batch.worker_seed),
|
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].strategy_samples > 0
|
||||||
assert metrics[0].traversal_nodes > 0
|
assert metrics[0].traversal_nodes > 0
|
||||||
assert metrics[0].traversal_max_depth_reached <= 3
|
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].advantage_loss >= 0.0
|
||||||
assert metrics[0].strategy_loss >= 0.0
|
assert metrics[0].strategy_loss >= 0.0
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user