Slash-namespace logged metric keys for W&B grouping

Adopt a 5-namespace scheme so wandb groups related metrics in the
sidebar and capture-group regex (eval/(random|safe_heuristic)/win_rate0
vs eval/(?:random|safe_heuristic)/win_rate0) controls panel splitting:

- loss/{advantage,strategy}
- samples/{advantage,strategy,advantage_player_N}
- memory/{advantage,strategy,advantage_player_N}
- time/{iteration_seconds,traversal_seconds,advantage_train_seconds,
  strategy_train_seconds,evaluation_seconds,memory_add_seconds,
  checkpoint_seconds,batch_tensor_seconds,nodes_per_second,
  advantage_player_N_sample_seconds,strategy_sample_seconds}
- traversal/{nodes,terminals,depth_cutoffs,node_limit_cutoffs,
  max_depth_reached,endpoints,avg_endpoint_depth,
  endpoint_depth_bucket_*,regret_fallback_*,sampled_actions}
- eval/<opponent>/<metric> (3-level so opponent can be the capture group)

`iteration` keeps no namespace (it's the wandb step axis). Internal
TraversalStats.to_dict() and benchmark.py's standalone result dict
keep their flat names — only the trainer's emitted metrics are
remapped, with the traversal_*→traversal/* translation done at
insertion into runtime_metrics.

analyze.py updated to read the new keys (PlotSpec metrics, color map,
opponent_names parser, _first_existing_eval lookup). Tests updated for
the new eval_metrics dict keys.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-05-07 17:37:41 +09:00
co-authored by Claude Opus 4.7
parent 5aab84c15f
commit 7c7c582c46
3 changed files with 159 additions and 153 deletions
@@ -32,17 +32,17 @@ SECTIONS: tuple[SectionSpec, ...] = (
"Loss",
"analysis_01_loss.png",
(
PlotSpec("Advantage Loss", ("advantage_loss",), "loss", kind="train"),
PlotSpec("Strategy Loss", ("strategy_loss",), "loss", kind="train"),
PlotSpec("Advantage Loss", ("loss/advantage",), "loss", kind="train"),
PlotSpec("Strategy Loss", ("loss/strategy",), "loss", kind="train"),
PlotSpec(
"Samples",
("advantage_samples", "strategy_samples"),
("samples/advantage", "samples/strategy"),
"samples",
kind="train",
),
PlotSpec(
"Memory Size",
("advantage_memory_size", "strategy_memory_size"),
("memory/advantage", "memory/strategy"),
"samples",
kind="train",
),
@@ -209,36 +209,36 @@ SECTIONS: tuple[SectionSpec, ...] = (
"Traversal",
"analysis_08_traversal.png",
(
PlotSpec("Iteration Time", ("iteration_seconds",), "seconds", kind="train"),
PlotSpec("Throughput", ("nodes_per_second",), "nodes / second", kind="train"),
PlotSpec("Iteration Time", ("time/iteration_seconds",), "seconds", kind="train"),
PlotSpec("Throughput", ("time/nodes_per_second",), "nodes / second", kind="train"),
PlotSpec(
"Traversal Depth",
("traversal_avg_endpoint_depth", "traversal_max_depth_reached"),
("traversal/avg_endpoint_depth", "traversal/max_depth_reached"),
"depth",
kind="train",
),
PlotSpec(
"Traversal Endpoints",
("traversal_terminals", "traversal_node_limit_cutoffs", "traversal_depth_cutoffs"),
("traversal/terminals", "traversal/node_limit_cutoffs", "traversal/depth_cutoffs"),
"count",
kind="train",
),
PlotSpec(
"Traversal Endpoint Rates",
(
"traversal_terminal_rate",
"traversal_node_limit_cutoff_rate",
"traversal_depth_cutoff_rate",
"traversal/terminal_rate",
"traversal/node_limit_cutoff_rate",
"traversal/depth_cutoff_rate",
),
"rate (%)",
scale=100.0,
kind="train",
fixed_ylim=(0, 100),
),
PlotSpec("Traversal Nodes", ("traversal_nodes",), "nodes", kind="train"),
PlotSpec("Traversal Nodes", ("traversal/nodes",), "nodes", kind="train"),
PlotSpec(
"Regret Fallback Rate",
("traversal_regret_fallback_rate",),
("traversal/regret_fallback_rate",),
"rate (%)",
scale=100.0,
kind="train",
@@ -246,18 +246,18 @@ SECTIONS: tuple[SectionSpec, ...] = (
),
PlotSpec(
"Regret Fallback Count",
("traversal_regret_fallback_count",),
("traversal/regret_fallback_count",),
"count",
kind="train",
),
PlotSpec(
"Fallback Selected Actions",
(
"traversal_regret_fallback_action_play_existing",
"traversal_regret_fallback_action_open_new",
"traversal_regret_fallback_action_discard",
"traversal_regret_fallback_action_draw_deck",
"traversal_regret_fallback_action_draw_pile",
"traversal/regret_fallback_action_play_existing",
"traversal/regret_fallback_action_open_new",
"traversal/regret_fallback_action_discard",
"traversal/regret_fallback_action_draw_deck",
"traversal/regret_fallback_action_draw_pile",
),
"count",
kind="train",
@@ -265,8 +265,8 @@ SECTIONS: tuple[SectionSpec, ...] = (
PlotSpec(
"Fallback Open-New Rates",
(
"traversal_regret_fallback_open_new_available_rate",
"traversal_regret_fallback_open_new_selected_rate",
"traversal/regret_fallback_open_new_available_rate",
"traversal/regret_fallback_open_new_selected_rate",
),
"rate (%)",
scale=100.0,
@@ -275,34 +275,34 @@ SECTIONS: tuple[SectionSpec, ...] = (
),
PlotSpec(
"Fallback Open-New Bias",
("traversal_regret_fallback_open_new_selection_over_availability",),
("traversal/regret_fallback_open_new_selection_over_availability",),
"selected / available",
kind="train",
),
PlotSpec(
"Fallback Avg Depth",
("traversal_regret_fallback_avg_depth",),
("traversal/regret_fallback_avg_depth",),
"depth",
kind="train",
),
PlotSpec(
"Fallback Opened Colors Before Action",
("traversal_regret_fallback_avg_opened_colors_before_action",),
("traversal/regret_fallback_avg_opened_colors_before_action",),
"colors",
kind="train",
),
PlotSpec(
"Fallback Legal Actions Mean",
("traversal_regret_fallback_legal_actions_mean",),
("traversal/regret_fallback_legal_actions_mean",),
"value",
kind="train",
),
PlotSpec(
"Argmax Tie Diagnostics",
(
"traversal_regret_fallback_argmax_tie_rate",
"traversal_regret_fallback_argmax_full_tie_rate",
"traversal_regret_fallback_argmax_tie_size_mean",
"traversal/regret_fallback_argmax_tie_rate",
"traversal/regret_fallback_argmax_full_tie_rate",
"traversal/regret_fallback_argmax_tie_size_mean",
),
"rate / size",
kind="train",
@@ -349,33 +349,33 @@ OPPONENT_COLORS: dict[str, str] = {
}
TRAVERSAL_COLORS: dict[str, str] = {
"iteration_seconds": "#4c78a8",
"nodes_per_second": "#4c78a8",
"traversal_avg_endpoint_depth": "#72b7b2",
"traversal_max_depth_reached": "#f58518",
"traversal_terminals": "#54a24b",
"traversal_node_limit_cutoffs": "#e45756",
"traversal_depth_cutoffs": "#b279a2",
"traversal_terminal_rate": "#54a24b",
"traversal_node_limit_cutoff_rate": "#e45756",
"traversal_depth_cutoff_rate": "#b279a2",
"traversal_nodes": "#4c78a8",
"traversal_regret_fallback_rate": "#e45756",
"traversal_regret_fallback_count": "#e45756",
"traversal_regret_fallback_action_play_existing": "#4c78a8",
"traversal_regret_fallback_action_open_new": "#f58518",
"traversal_regret_fallback_action_discard": "#54a24b",
"traversal_regret_fallback_action_draw_deck": "#b279a2",
"traversal_regret_fallback_action_draw_pile": "#72b7b2",
"traversal_regret_fallback_open_new_available_rate": "#9d755d",
"traversal_regret_fallback_open_new_selected_rate": "#f58518",
"traversal_regret_fallback_open_new_selection_over_availability": "#f58518",
"traversal_regret_fallback_avg_depth": "#72b7b2",
"traversal_regret_fallback_avg_opened_colors_before_action": "#f58518",
"traversal_regret_fallback_legal_actions_mean": "#54a24b",
"traversal_regret_fallback_argmax_tie_rate": "#e45756",
"traversal_regret_fallback_argmax_full_tie_rate": "#b279a2",
"traversal_regret_fallback_argmax_tie_size_mean": "#4c78a8",
"time/iteration_seconds": "#4c78a8",
"time/nodes_per_second": "#4c78a8",
"traversal/avg_endpoint_depth": "#72b7b2",
"traversal/max_depth_reached": "#f58518",
"traversal/terminals": "#54a24b",
"traversal/node_limit_cutoffs": "#e45756",
"traversal/depth_cutoffs": "#b279a2",
"traversal/terminal_rate": "#54a24b",
"traversal/node_limit_cutoff_rate": "#e45756",
"traversal/depth_cutoff_rate": "#b279a2",
"traversal/nodes": "#4c78a8",
"traversal/regret_fallback_rate": "#e45756",
"traversal/regret_fallback_count": "#e45756",
"traversal/regret_fallback_action_play_existing": "#4c78a8",
"traversal/regret_fallback_action_open_new": "#f58518",
"traversal/regret_fallback_action_discard": "#54a24b",
"traversal/regret_fallback_action_draw_deck": "#b279a2",
"traversal/regret_fallback_action_draw_pile": "#72b7b2",
"traversal/regret_fallback_open_new_available_rate": "#9d755d",
"traversal/regret_fallback_open_new_selected_rate": "#f58518",
"traversal/regret_fallback_open_new_selection_over_availability": "#f58518",
"traversal/regret_fallback_avg_depth": "#72b7b2",
"traversal/regret_fallback_avg_opened_colors_before_action": "#f58518",
"traversal/regret_fallback_legal_actions_mean": "#54a24b",
"traversal/regret_fallback_argmax_tie_rate": "#e45756",
"traversal/regret_fallback_argmax_full_tie_rate": "#b279a2",
"traversal/regret_fallback_argmax_tie_size_mean": "#4c78a8",
}
ACTION_RATE_METRICS = {
@@ -400,13 +400,11 @@ def opponent_names(rows: list[dict[str, Any]]) -> list[str]:
names: set[str] = set()
for row in rows:
for key in row:
if not key.startswith("eval_"):
if not key.startswith("eval/"):
continue
rest = key[len("eval_") :]
for metric in _all_eval_metrics():
suffix = f"_{metric}"
if rest.endswith(suffix):
names.add(rest[: -len(suffix)])
parts = key.split("/", 2)
if len(parts) == 3:
names.add(parts[1])
return sorted(names)
@@ -653,12 +651,12 @@ def _plot_train_spec(
def _train_value(row: dict[str, Any], metric: str) -> float | None:
if metric == "traversal_terminal_rate":
return _ratio(row, "traversal_terminals", "traversal_endpoints")
if metric == "traversal_node_limit_cutoff_rate":
return _ratio(row, "traversal_node_limit_cutoffs", "traversal_endpoints")
if metric == "traversal_depth_cutoff_rate":
return _ratio(row, "traversal_depth_cutoffs", "traversal_endpoints")
if metric == "traversal/terminal_rate":
return _ratio(row, "traversal/terminals", "traversal/endpoints")
if metric == "traversal/node_limit_cutoff_rate":
return _ratio(row, "traversal/node_limit_cutoffs", "traversal/endpoints")
if metric == "traversal/depth_cutoff_rate":
return _ratio(row, "traversal/depth_cutoffs", "traversal/endpoints")
value = row.get(metric)
if value is None:
return None
@@ -789,7 +787,7 @@ def _first_existing_eval(
row: dict[str, Any], opponent: str, metrics: tuple[str, ...]
) -> float | None:
for metric in metrics:
value = row.get(f"eval_{opponent}_{metric}")
value = row.get(f"eval/{opponent}/{metric}")
if value is not None:
return float(value)
return None
@@ -921,37 +919,37 @@ def _label(metric: str) -> str:
def _train_metric_label(metric: str) -> str:
labels = {
"advantage_memory_size": "advantage",
"advantage_samples": "advantage",
"iteration_seconds": "iteration",
"nodes_per_second": "nodes/sec",
"strategy_memory_size": "strategy",
"strategy_samples": "strategy",
"traversal_avg_endpoint_depth": "avg endpoint depth",
"traversal_depth_cutoff_rate": "depth cutoff",
"traversal_depth_cutoffs": "depth cutoff",
"traversal_max_depth_reached": "max depth",
"traversal_node_limit_cutoff_rate": "node limit cutoff",
"traversal_node_limit_cutoffs": "node limit cutoff",
"traversal_nodes": "nodes",
"traversal_regret_fallback_action_discard": "discard",
"traversal_regret_fallback_action_draw_deck": "draw deck",
"traversal_regret_fallback_action_draw_pile": "draw pile",
"traversal_regret_fallback_action_open_new": "open new",
"traversal_regret_fallback_action_play_existing": "play existing",
"traversal_regret_fallback_argmax_full_tie_rate": "full tie rate",
"traversal_regret_fallback_argmax_tie_rate": "tie rate",
"traversal_regret_fallback_argmax_tie_size_mean": "tie size",
"traversal_regret_fallback_avg_depth": "avg depth",
"traversal_regret_fallback_avg_opened_colors_before_action": "opened colors",
"traversal_regret_fallback_count": "fallbacks",
"traversal_regret_fallback_legal_actions_mean": "legal actions",
"traversal_regret_fallback_open_new_available_rate": "available",
"traversal_regret_fallback_open_new_selected_rate": "selected",
"traversal_regret_fallback_open_new_selection_over_availability": "selected / available",
"traversal_regret_fallback_rate": "fallback rate",
"traversal_terminal_rate": "terminal",
"traversal_terminals": "terminal",
"memory/advantage": "advantage",
"samples/advantage": "advantage",
"time/iteration_seconds": "iteration",
"time/nodes_per_second": "nodes/sec",
"memory/strategy": "strategy",
"samples/strategy": "strategy",
"traversal/avg_endpoint_depth": "avg endpoint depth",
"traversal/depth_cutoff_rate": "depth cutoff",
"traversal/depth_cutoffs": "depth cutoff",
"traversal/max_depth_reached": "max depth",
"traversal/node_limit_cutoff_rate": "node limit cutoff",
"traversal/node_limit_cutoffs": "node limit cutoff",
"traversal/nodes": "nodes",
"traversal/regret_fallback_action_discard": "discard",
"traversal/regret_fallback_action_draw_deck": "draw deck",
"traversal/regret_fallback_action_draw_pile": "draw pile",
"traversal/regret_fallback_action_open_new": "open new",
"traversal/regret_fallback_action_play_existing": "play existing",
"traversal/regret_fallback_argmax_full_tie_rate": "full tie rate",
"traversal/regret_fallback_argmax_tie_rate": "tie rate",
"traversal/regret_fallback_argmax_tie_size_mean": "tie size",
"traversal/regret_fallback_avg_depth": "avg depth",
"traversal/regret_fallback_avg_opened_colors_before_action": "opened colors",
"traversal/regret_fallback_count": "fallbacks",
"traversal/regret_fallback_legal_actions_mean": "legal actions",
"traversal/regret_fallback_open_new_available_rate": "available",
"traversal/regret_fallback_open_new_selected_rate": "selected",
"traversal/regret_fallback_open_new_selection_over_availability": "selected / available",
"traversal/regret_fallback_rate": "fallback rate",
"traversal/terminal_rate": "terminal",
"traversal/terminals": "terminal",
}
return labels.get(metric, _label(metric))
@@ -107,20 +107,20 @@ class IterationMetrics:
def to_dict(self) -> dict[str, float | int]:
data = {
"iteration": self.iteration,
"advantage_samples": self.advantage_samples,
"strategy_samples": self.strategy_samples,
"advantage_loss": self.advantage_loss,
"strategy_loss": self.strategy_loss,
"traversal_nodes": self.traversal_nodes,
"traversal_terminals": self.traversal_terminals,
"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,
"samples/advantage": self.advantage_samples,
"samples/strategy": self.strategy_samples,
"loss/advantage": self.advantage_loss,
"loss/strategy": self.strategy_loss,
"traversal/nodes": self.traversal_nodes,
"traversal/terminals": self.traversal_terminals,
"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
f"traversal/endpoint_depth_bucket_{key}": value
for key, value in self.traversal_endpoint_depth_buckets.items()
},
}
@@ -130,24 +130,24 @@ class IterationMetrics:
def _format_iteration_summary(metrics: IterationMetrics, data: dict[str, float | int]) -> str:
iteration_seconds = float(data.get("time/iteration_seconds", 0.0) or 0.0)
parts = [
f"[i={metrics.iteration}] Iteration complete",
f"traversal_nodes={metrics.traversal_nodes}",
f"nodes_per_second={_format_summary_value(data['nodes_per_second'])}",
f"nodes_per_second={_format_summary_value(data['time/nodes_per_second'])}",
f"advantage_loss={_format_summary_value(metrics.advantage_loss)}",
f"strategy_loss={_format_summary_value(metrics.strategy_loss)}",
f"iteration_seconds={_format_summary_value(data['iteration_seconds'])}",
f"iters_per_hour={3600.0 / data['iteration_seconds']:.1f}"
if data.get("iteration_seconds", 0.0)
f"iteration_seconds={_format_summary_value(iteration_seconds)}",
f"iters_per_hour={3600.0 / iteration_seconds:.1f}"
if iteration_seconds
else "iters_per_hour=n/a",
]
if metrics.eval_metrics:
eval_seconds = float(data.get("evaluation_seconds", 0.0) or 0.0)
iter_seconds = float(data.get("iteration_seconds", 0.0) or 0.0)
fraction = eval_seconds / iter_seconds if iter_seconds > 0.0 else 0.0
eval_seconds = float(data.get("time/evaluation_seconds", 0.0) or 0.0)
fraction = eval_seconds / iteration_seconds if iteration_seconds > 0.0 else 0.0
parts.append(f"eval_seconds={_format_summary_value(eval_seconds)}({fraction * 100:.0f}%)")
for key in sorted(metrics.eval_metrics):
if key.endswith("_win_rate0") or key.endswith("_avg_score_diff0"):
if key.endswith("/win_rate0") or key.endswith("/avg_score_diff0"):
parts.append(f"{key}={_format_summary_value(metrics.eval_metrics[key])}")
return " ".join(parts)
@@ -313,27 +313,33 @@ class DeepCFRTrainer:
total_stats = self._run_traversals_parallel(iteration)
else:
total_stats = self._run_traversals_single_process(iteration)
self._runtime_metrics["traversal_seconds"] = time.perf_counter() - traversal_started
self._runtime_metrics["time/traversal_seconds"] = time.perf_counter() - traversal_started
advantage_started = time.perf_counter()
advantage_loss = self._train_advantage_networks()
self._runtime_metrics["advantage_train_seconds"] = time.perf_counter() - advantage_started
self._runtime_metrics["time/advantage_train_seconds"] = (
time.perf_counter() - advantage_started
)
strategy_started = time.perf_counter()
strategy_loss = self._train_strategy_network()
self._runtime_metrics["strategy_train_seconds"] = time.perf_counter() - strategy_started
self._runtime_metrics["time/strategy_train_seconds"] = (
time.perf_counter() - strategy_started
)
eval_started = time.perf_counter()
eval_metrics = self._evaluate(iteration)
if eval_metrics:
self._runtime_metrics["evaluation_seconds"] = time.perf_counter() - eval_started
self._runtime_metrics["advantage_memory_size"] = self._advantage_memory_size()
self._runtime_metrics["time/evaluation_seconds"] = time.perf_counter() - eval_started
self._runtime_metrics["memory/advantage"] = self._advantage_memory_size()
for player, memory in enumerate(self.advantage_memories):
self._runtime_metrics[f"advantage_player_{player}_memory_size"] = len(memory)
self._runtime_metrics["strategy_memory_size"] = len(self.strategy_memory)
self._runtime_metrics[f"memory/advantage_player_{player}"] = len(memory)
self._runtime_metrics["memory/strategy"] = len(self.strategy_memory)
for key, value in total_stats.to_dict().items():
if key.startswith("traversal_regret_") or key == "traversal_sampled_actions":
self._runtime_metrics[key] = value
if key.startswith("traversal_regret_"):
self._runtime_metrics["traversal/" + key[len("traversal_") :]] = value
elif key == "traversal_sampled_actions":
self._runtime_metrics["traversal/sampled_actions"] = value
return IterationMetrics(
iteration=iteration,
advantage_samples=self._advantage_memory_size(),
@@ -416,8 +422,8 @@ class DeepCFRTrainer:
memory_add_started = time.perf_counter()
self._add_advantage_samples(advantage_samples)
self.strategy_memory.add_many(strategy_samples, self.rng)
self._runtime_metrics["memory_add_seconds"] = (
float(self._runtime_metrics.get("memory_add_seconds", 0.0))
self._runtime_metrics["time/memory_add_seconds"] = (
float(self._runtime_metrics.get("time/memory_add_seconds", 0.0))
+ time.perf_counter()
- memory_add_started
)
@@ -475,8 +481,8 @@ class DeepCFRTrainer:
memory_add_started = time.perf_counter()
self._add_advantage_samples(result.advantage_samples)
self.strategy_memory.add_many(result.strategy_samples, self.rng)
self._runtime_metrics["memory_add_seconds"] = (
float(self._runtime_metrics.get("memory_add_seconds", 0.0))
self._runtime_metrics["time/memory_add_seconds"] = (
float(self._runtime_metrics.get("time/memory_add_seconds", 0.0))
+ time.perf_counter()
- memory_add_started
)
@@ -591,7 +597,7 @@ class DeepCFRTrainer:
self._maybe_record_self_play_snapshot(iteration)
checkpoint_started = time.perf_counter()
self._save_iteration_checkpoints(iteration, item)
item.runtime_metrics["checkpoint_seconds"] = (
item.runtime_metrics["time/checkpoint_seconds"] = (
time.perf_counter() - checkpoint_started
)
elapsed = time.perf_counter() - started
@@ -648,8 +654,8 @@ class DeepCFRTrainer:
def _append_metrics(self, metrics: IterationMetrics, iteration_seconds: float) -> None:
data = metrics.to_dict()
data["iteration_seconds"] = iteration_seconds
data["nodes_per_second"] = metrics.traversal_nodes / max(iteration_seconds, 1.0e-12)
data["time/iteration_seconds"] = iteration_seconds
data["time/nodes_per_second"] = metrics.traversal_nodes / max(iteration_seconds, 1.0e-12)
self.tracker.log_metrics(data, step=metrics.iteration)
self.tracker.log_event(_format_iteration_summary(metrics, data))
@@ -677,7 +683,7 @@ class DeepCFRTrainer:
batch_size=self.config.evaluation.batch_size,
)
for key, value in result.items():
results[f"eval_{opponent}_{key}"] = value
results[f"eval/{opponent}/{key}"] = value
return results
def _evaluate_parallel(
@@ -719,7 +725,7 @@ class DeepCFRTrainer:
for future in as_completed(futures):
opponent, result = future.result()
for key, value in result.items():
results[f"eval_{opponent}_{key}"] = value
results[f"eval/{opponent}/{key}"] = value
return results
def _evaluation_device(self) -> torch.device:
@@ -747,14 +753,14 @@ class DeepCFRTrainer:
for player, (network, memory) in enumerate(
zip(self.advantage_networks, self.advantage_memories, strict=True)
):
self._runtime_metrics[f"advantage_player_{player}_sample_count"] = len(memory)
self._runtime_metrics[f"samples/advantage_player_{player}"] = len(memory)
if len(memory) == 0:
continue
losses.append(self._train_advantage(player, network, self.advantage_optimizers[player]))
return float(np.mean(losses)) if losses else 0.0
def _train_strategy_network(self) -> float:
self._runtime_metrics["strategy_sample_count"] = len(self.strategy_memory)
self._runtime_metrics["samples/strategy"] = len(self.strategy_memory)
if len(self.strategy_memory) == 0:
return 0.0
return self._train_strategy(self.strategy_network, self.strategy_optimizer)
@@ -784,8 +790,8 @@ class DeepCFRTrainer:
dtype=torch.float32,
device=self.device,
)
self._runtime_metrics["batch_tensor_seconds"] = (
float(self._runtime_metrics.get("batch_tensor_seconds", 0.0))
self._runtime_metrics["time/batch_tensor_seconds"] = (
float(self._runtime_metrics.get("time/batch_tensor_seconds", 0.0))
+ time.perf_counter()
- started
)
@@ -812,8 +818,10 @@ class DeepCFRTrainer:
self.config.optimization.advantage_batch_size,
self.rng,
)
self._runtime_metrics[f"advantage_player_{player}_sample_seconds"] = (
float(self._runtime_metrics.get(f"advantage_player_{player}_sample_seconds", 0.0))
self._runtime_metrics[f"time/advantage_player_{player}_sample_seconds"] = (
float(
self._runtime_metrics.get(f"time/advantage_player_{player}_sample_seconds", 0.0)
)
+ time.perf_counter()
- sample_started
)
@@ -866,8 +874,8 @@ class DeepCFRTrainer:
batch = self.strategy_memory.sample(
self.config.optimization.strategy_batch_size, self.rng
)
self._runtime_metrics["strategy_sample_seconds"] = (
float(self._runtime_metrics.get("strategy_sample_seconds", 0.0))
self._runtime_metrics["time/strategy_sample_seconds"] = (
float(self._runtime_metrics.get("time/strategy_sample_seconds", 0.0))
+ time.perf_counter()
- sample_started
)
+7 -7
View File
@@ -670,14 +670,14 @@ def test_deep_cfr_trainer_saves_loads_and_evaluates_checkpoint(tmp_path) -> None
assert "[i=1]" in train_log
assert "reservoir memories and RNG state are not restored" in train_log
assert restored.iteration == 1
assert "eval_random_games" in metrics[0].eval_metrics
assert "eval_random_play_action_rate" in metrics[0].eval_metrics
assert "eval_random_policy_entropy" in metrics[0].eval_metrics
assert "eval_random_avg_opened_colors" in metrics[0].eval_metrics
assert "eval_random_bad_open_actions" in metrics[0].eval_metrics
assert "eval_random_positive_expedition_rate" in metrics[0].eval_metrics
assert "eval/random/games" in metrics[0].eval_metrics
assert "eval/random/play_action_rate" in metrics[0].eval_metrics
assert "eval/random/policy_entropy" in metrics[0].eval_metrics
assert "eval/random/avg_opened_colors" in metrics[0].eval_metrics
assert "eval/random/bad_open_actions" in metrics[0].eval_metrics
assert "eval/random/positive_expedition_rate" in metrics[0].eval_metrics
assert (
"eval_random_first_open_recoverable_score_mean_for_positive_final"
"eval/random/first_open_recoverable_score_mean_for_positive_final"
in metrics[0].eval_metrics
)