Refactor the analysis plot creation with focus on rollout #16
@@ -57,6 +57,9 @@ _PLOT_META_KEYS = (
|
||||
"n_seed_events",
|
||||
"timestamp",
|
||||
"comment",
|
||||
"model_config",
|
||||
"training_epoch",
|
||||
"best_val_loss",
|
||||
)
|
||||
|
||||
|
||||
|
||||
+30
-13
@@ -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)
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user