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