diff --git a/giant/analysis/condor.py b/giant/analysis/condor.py index 28cbcd8..e2ef763 100644 --- a/giant/analysis/condor.py +++ b/giant/analysis/condor.py @@ -57,6 +57,9 @@ _PLOT_META_KEYS = ( "n_seed_events", "timestamp", "comment", + "model_config", + "training_epoch", + "best_val_loss", ) diff --git a/giant/analysis/render.py b/giant/analysis/render.py index a7d9014..aef6c63 100644 --- a/giant/analysis/render.py +++ b/giant/analysis/render.py @@ -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) diff --git a/giant/cli.py b/giant/cli.py index d989d08..e371120 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -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))