Refactor the analysis plot creation with focus on rollout #16

Merged
lars merged 30 commits from analysis-rollout-plots into master 2026-07-27 14:58:51 +02:00
3 changed files with 43 additions and 13 deletions
Showing only changes of commit 1e7d7d7efd - Show all commits
+3
View File
@@ -57,6 +57,9 @@ _PLOT_META_KEYS = (
"n_seed_events",
"timestamp",
"comment",
"model_config",
"training_epoch",
"best_val_loss",
)
+30 -13
View File
@@ -41,9 +41,21 @@ def _overlay(ax, edges: np.ndarray, series: dict[str, list], log_y: bool) -> Non
ax.set_yscale("log")
def _render_overlay(r: Reduced):
def _nn_params(run_meta: dict) -> dict:
"""Flatten the rollout's model/training provenance for the figure subtitle."""
params = {
k: v for k, v in (run_meta.get("model_config") or {}).items() if v is not None
}
if run_meta.get("training_epoch") is not None:
params["epoch"] = run_meta["training_epoch"]
if run_meta.get("best_val_loss") is not None:
params["best_val_loss"] = round(run_meta["best_val_loss"], 4)
return params
def _render_overlay(r: Reduced, params: dict):
edges = np.asarray(r.payload["edges"])
fig, ax = ps.new_figure("thesis-single", title=r.title)
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
_overlay(ax, edges, r.payload, r.payload.get("log_y", False))
ax.set_xlabel(r.xlabel)
ax.set_ylabel("density")
@@ -51,9 +63,9 @@ def _render_overlay(r: Reduced):
return fig
def _render_single(r: Reduced):
def _render_single(r: Reduced, params: dict):
edges = np.asarray(r.payload["edges"])
fig, ax = ps.new_figure("thesis-single", title=r.title)
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
ax.stairs(
_density(r.payload["rollout"], edges), edges, label=_SERIES_LABELS["rollout"]
)
@@ -65,7 +77,7 @@ def _render_single(r: Reduced):
return fig
def _render_grouped(r: Reduced):
def _render_grouped(r: Reduced, params: dict):
edges = np.asarray(r.payload["edges"])
groups = r.payload["groups"]
labels = list(groups)
@@ -73,7 +85,12 @@ def _render_grouped(r: Reduced):
ncols = min(3, n) or 1
nrows = (n + ncols - 1) // ncols
fig, axes = ps.new_figure(
"slide-16x9", title=r.title, nrows=nrows, ncols=ncols, squeeze=False
"slide-16x9",
title=r.title,
params=params,
nrows=nrows,
ncols=ncols,
squeeze=False,
)
flat = axes.ravel()
for i, lbl in enumerate(labels):
@@ -87,10 +104,10 @@ def _render_grouped(r: Reduced):
return fig
def _render_profile(r: Reduced):
def _render_profile(r: Reduced, params: dict):
edges = np.asarray(r.payload["edges"])
centers = 0.5 * (edges[:-1] + edges[1:])
fig, ax = ps.new_figure("thesis-single", title=r.title)
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
for key in ("reference", "rollout"):
mean = np.asarray(r.payload[f"{key}_mean"])
std = np.asarray(r.payload[f"{key}_std"])
@@ -104,11 +121,11 @@ def _render_profile(r: Reduced):
return fig
def _render_bar(r: Reduced):
def _render_bar(r: Reduced, params: dict):
labels = r.payload["labels"]
x = np.arange(len(labels))
width = 0.4
fig, ax = ps.new_figure("thesis-single", title=r.title)
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
ax.bar(
x - width / 2, r.payload["reference"], width, label=_SERIES_LABELS["reference"]
)
@@ -129,9 +146,9 @@ _RENDERERS = {
}
def render(r: Reduced):
def render(r: Reduced, run_meta: dict | None = None):
"""Build the matplotlib figure for one reduced artifact (dispatch on kind)."""
return _RENDERERS[r.kind](r)
return _RENDERERS[r.kind](r, _nn_params(run_meta or {}))
def _plot_metadata(r: Reduced, run_meta: dict) -> dict:
@@ -171,7 +188,7 @@ def render_all(
family_dir = out_dir / r.family
family_dir.mkdir(parents=True, exist_ok=True)
families.add(r.family)
fig = render(r)
fig = render(r, run_meta)
ps.savefig(fig, str(family_dir / r.id), formats=("pdf",))
(family_dir / f"{r.id}.yaml").write_text(
yaml.safe_dump(_plot_metadata(r, run_meta), sort_keys=False)
+10
View File
@@ -1105,6 +1105,16 @@ def rollout(
"steps": steps,
"max_tracks_per_event": max_tracks_per_event,
"n_seed_events": int(len(seeds["event_id"])),
"model_config": {
"mode": model_cfg.get("mode", "flow"),
"hidden_dim": model_cfg.get("hidden_dim"),
"n_blocks": model_cfg.get("n_blocks"),
"emb_dim": model_cfg.get("emb_dim"),
"dropout": model_cfg.get("dropout"),
"conditioning": conditioning,
},
"training_epoch": ckpt.get("epoch"),
"best_val_loss": ckpt.get("best_val_loss"),
}
)
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))