diff --git a/CLAUDE.md b/CLAUDE.md index 33106f1..ab1abc6 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -73,7 +73,7 @@ GIANT is a conditional generative surrogate for the Geant4 step function. It rep **Validation** (`giant/validate.py`): step-level marginal comparisons. -**Analysis** (`giant/analysis/`, `giant analyze` CLI): a lean, streaming rollout-vs-reference plotting pipeline that compares one autoregressive `giant rollout` (for a given checkpoint) against a held-out miniCaloSim reference steps file, and produces publication-styled PDFs assembled into an HTML gallery. It exploits the fact that rollout output and a raw reference file share a world-frame physical column subset under identical names (`pre_*`/`post_*`/`edep`/`step_length`/`pdg`/`material`/`event_id`), so no ALR/local-frame decode is needed — everything is world-frame mm/MeV. Structure: `sources.py` (canonical LazyFrames + synthetic-termination-row filtering + the secondary view, which is `generation>0 & step_no==0` rollout tracks vs exploded `sec_*_list` reference columns), `reduce.py` (the streaming primitives — a single `hist1d` `group_by([group,bin]).len()` pass, per-event scalars, edep-weighted depth/transverse profiles, species share, leakage), `grouping.py`/`context.py` (fixed bin edges + energy-quantile/pdg/material group sets resolved once by `prep` into `shared.json`, so every compute job is one pass with no range scan), `catalog.py` (the declarative `PlotSpec` registry — marginals × {overall,energy,pdg,material}, per-event totals, shower profiles, species/leakage, secondaries), and `render.py` (the only module importing ETPlot's `plotstyle`/LaTeX; dispatches on `Reduced.kind`, writes PDFs + `metadata.yaml`). **Input is a `giant rollout` YAML sidecar** (`condor.py:load_rollout_yaml`): its `output`/`dataset` keys name the rollout parquet and the seed file (= the reference truth), and the rest of the YAML (checkpoint, geometry oracle, cutoffs) flows into each plot's gallery metadata. `prep` derives its own **run directory** next to the rollout parquet (`<...>/analysis_/`) holding `shared.json`, `run_meta.json`, `reduced_partial/`, `reduced/`, `plots/`. **Compute/merge/render split:** `giant analyze submit rollout.yaml --chunks N` runs `prep` (recording the run's chunk count `N` in `run_meta.json`) then submits one HTCondor job per (plot, chunk) pair (`compute-one --id --chunk --run-dir`, polars/numpy only — no LaTeX on workers), each streaming over an `event_id`-disjoint slice (`event_id % N == chunk`) and writing a small `reduced_partial/__.json`; every `PlotSpec` (`catalog.py`) splits into a `compute_partial`/`finalize` pair so a plot's chunks can be summed/concatenated back together correctly (`chunkable=False` specs — the router diagnostics, already bounded/subsampled — always run as a single chunk regardless of `N`). The local `giant analyze render ` first joins every plot's chunk partials into `reduced/.json` (`merge_all`, a no-op join when `N=1`), then turns those into the styled PDF/gallery tree. See `giant/analysis/__init__.py`. +**Analysis** (`giant/analysis/`, `giant analyze` CLI): a lean, streaming rollout-vs-reference plotting pipeline that compares one or more autoregressive `giant rollout` runs against a single held-out miniCaloSim reference steps file shared by all of them, and produces publication-styled PDFs assembled into an HTML gallery — one distinctly colored series per rollout, one reference line/panel. It exploits the fact that rollout output and a raw reference file share a world-frame physical column subset under identical names (`pre_*`/`post_*`/`edep`/`step_length`/`pdg`/`material`/`event_id`), so no ALR/local-frame decode is needed — everything is world-frame mm/MeV. Structure: `sources.py` (canonical LazyFrames + `RolloutSpec`/`RolloutSide` — a rollout's opened frames + per-checkpoint diagnostic inputs — + synthetic-termination-row filtering + the secondary view, which is `generation>0 & step_no==0` rollout tracks vs exploded `sec_*_list` reference columns), `reduce.py` (the streaming primitives — a single `hist1d` `group_by([group,bin]).len()` pass, per-event scalars, edep-weighted depth/transverse profiles, species share, leakage), `grouping.py`/`context.py` (fixed bin edges + energy-quantile/pdg/material group sets resolved once by `prep` into `shared.json` over the union of the reference and every rollout, so every compute job is one pass with no range scan), `catalog.py` (the declarative `PlotSpec` registry — marginals × {overall,energy,pdg,material}, per-event totals, shower profiles, species/leakage, secondaries; `Bundle.rollouts` is a name-keyed dict of `RolloutSide`, and every `compute_partial`/`finalize` builds a `Reduced.payload["series"]` dict keyed the same way, with `payload["reference"]` as the one distinguished non-rollout entry), and `render.py` (the only module importing ETPlot's `plotstyle`/LaTeX; dispatches on `Reduced.kind`, writes PDFs + `metadata.yaml`; each rollout gets a stable `ps.get_color(i)` slot by its position in `series`, the reference always draws in one fixed dashed-ink style). The two heatmap-shaped specs (`marginal_distance_summary`, `n_sec_confusion`) and the router/type-embedding diagnostics (`router_gating.py`, `type_embedding_distance.py`) are inherently one-matrix/one-checkpoint per rollout, so they render as one panel per rollout instead of one line/bar per rollout. **Input is one or more `giant rollout` YAML sidecars** (`condor.py:load_rollout_yamls`, wrapping the single-YAML `load_rollout_yaml`): each YAML's `output`/`dataset` keys name its rollout parquet and seed file (= the reference truth); every supplied YAML must resolve to the same `dataset`, checked up front with a clear error otherwise (the premise is "N candidates vs one ground truth"). Each rollout's series name comes from a repeated `--label` CLI flag, else the YAML stem (N>1), else `"rollout"` (a single YAML — matching pre-multi-rollout output exactly). `prep` derives its own **run directory** next to the *first* rollout's parquet (`<...>/analysis_/`) holding `shared.json`, `run_meta.json` (`RunMeta.rollouts: list[{name,path,plot_meta}]`, insertion order = CLI order = every plot's series order), `reduced_partial/`, `reduced/`, `plots/`. **Compute/merge/render split:** `giant analyze submit a.yaml [b.yaml ...] --chunks N` runs `prep` (recording the run's chunk count `N` in `run_meta.json`) then submits one HTCondor job per (plot, chunk) pair (`compute-one --id --chunk --run-dir`, polars/numpy only — no LaTeX on workers), each streaming over an `event_id`-disjoint slice (`event_id % N == chunk`) of the reference **and every rollout** and writing a small `reduced_partial/__.json`; every `PlotSpec` (`catalog.py`) splits into a `compute_partial`/`finalize` pair so a plot's chunks can be summed/concatenated back together correctly per rollout (`chunkable=False` specs — the router diagnostics, already bounded/subsampled — always run as a single chunk regardless of `N`). The local `giant analyze render ` first joins every plot's chunk partials into `reduced/.json` (`merge_all`, a no-op join when `N=1`), then turns those into the styled PDF/gallery tree. See `giant/analysis/__init__.py`. **Shower rollout** (`giant/rollout.py`, `giant rollout` CLI): autoregressively steps the two-stage model into a full shower — each primary post-step becomes the next pre-step, secondaries are pushed as new tracks, and per-step `material`/`layer_id` come from a `GeometryOracle` (`giant/geometry.py`, built via `dwarf build-geometry-oracle`) that learns position → (material, layer_id) from data and flags detector escape by nearest-neighbour distance. Tracks terminate on energy cutoff, per-track max steps, escape, or natural end; energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction. diff --git a/README.md b/README.md index 7fbc7da..e771387 100644 --- a/README.md +++ b/README.md @@ -152,10 +152,11 @@ Config-file-only knobs (no CLI flag — use `--config config.toml`): `stage2_mod ```bash giant analyze submit rollout.yaml --accounting-group cms # prep + one HTCondor job per plot (compute only) +giant analyze submit a.yaml b.yaml --accounting-group cms --label flow --label wgan # N rollouts vs one shared reference giant analyze render --gallery # local: styled PDFs + HTML gallery (needs LaTeX) ``` -`` is derived next to the rollout parquet (`analyze prep`/`submit` print it). Compute jobs are polars/numpy only; only `render` needs LaTeX, so it always runs locally. +`` is derived next to the first rollout's parquet (`analyze prep`/`submit` print it). Multiple rollout YAMLs must all name the same reference (`dataset`) file; each renders as its own colored series against one reference line/panel. Compute jobs are polars/numpy only; only `render` needs LaTeX, so it always runs locally. ## Development diff --git a/giant/analysis/__init__.py b/giant/analysis/__init__.py index 48c4343..1056d85 100644 --- a/giant/analysis/__init__.py +++ b/giant/analysis/__init__.py @@ -1,9 +1,10 @@ """Rollout-vs-reference analysis: streaming compute + plotstyle rendering. -Compares one autoregressive ``giant rollout`` against a held-out miniCaloSim -reference file, producing publication-styled comparison plots generated in -parallel on HTCondor (one job per plot x data chunk, compute/merge/render -split). +Compares one or more autoregressive ``giant rollout`` runs against a single +held-out miniCaloSim reference file shared by all of them, producing +publication-styled comparison plots (one colored series per rollout, one +reference line) generated in parallel on HTCondor (one job per plot x data +chunk, compute/merge/render split). Only ``render`` (and the ``render`` CLI path) imports plotstyle/LaTeX; everything re-exported here is plotstyle-free so it runs on a compute worker. Import @@ -12,12 +13,14 @@ re-exported here is plotstyle-free so it runs on a compute worker. Import from giant.analysis.catalog import build_catalog, catalog_ids, get_spec from giant.analysis.condor import ( + LoadedRollout, RunMeta, SubmitConfig, compute_one, compute_reduced, derive_run_dir, load_rollout_yaml, + load_rollout_yamls, merge_all, merge_one, prep, @@ -26,18 +29,20 @@ from giant.analysis.condor import ( from giant.analysis.context import Context, build_context from giant.analysis.reduced import Partial, Reduced from giant.analysis.runtime_estimate import RUNTIME_SAFETY_MARGIN, estimate_runtime_s -from giant.analysis.sources import Side +from giant.analysis.sources import RolloutSpec, Side __all__ = [ "build_catalog", "catalog_ids", "get_spec", + "LoadedRollout", "RunMeta", "SubmitConfig", "compute_one", "compute_reduced", "derive_run_dir", "load_rollout_yaml", + "load_rollout_yamls", "merge_all", "merge_one", "prep", @@ -46,6 +51,7 @@ __all__ = [ "build_context", "Partial", "Reduced", + "RolloutSpec", "Side", "RUNTIME_SAFETY_MARGIN", "estimate_runtime_s", diff --git a/giant/analysis/catalog.py b/giant/analysis/catalog.py index e9b4157..18fbc1b 100644 --- a/giant/analysis/catalog.py +++ b/giant/analysis/catalog.py @@ -4,7 +4,7 @@ Each spec knows its stable ``id`` (used for the reduced-data filename, the PDF stem and the condor queue item), its gallery ``family`` (subdirectory), and a ``compute_partial(bundle) -> dict`` / ``finalize(parts, ctx) -> Reduced`` pair that together run the streaming reduction. ``compute_partial`` runs once per -``(plot, chunk)`` condor job against a ``Bundle`` whose four LazyFrames are +``(plot, chunk)`` condor job against a ``Bundle`` whose LazyFrames are already filtered to that chunk (see ``Bundle.open``'s ``chunk`` argument); it returns a small JSON-safe partial artifact — either a raw sum-mergeable count dict (histograms/species sums against fixed edges) or a raw per-event/ @@ -16,6 +16,19 @@ exactly what a single unchunked pass would produce. Specs marked ``chunkable=False`` (the router ones) always run as a single chunk regardless of the configured chunk count. +Every ``compute_partial`` here returns ``{"r": {rollout_name: }, "t": +}`` — one entry per rollout in ``Bundle.rollouts`` (insertion order, +which is the order rollouts were given on the CLI) plus the single reference. +``finalize`` merges each rollout's chunks independently and assembles a +``Reduced.payload`` keyed the same way: ``"series": {name: ...}`` for the +rollouts, ``"reference": ...`` as one distinguished entry (omitted on +rollout-only plots like ``leakage_fraction``). The two heatmap-shaped specs +(``marginal_distance_summary``, ``n_sec_confusion``) and the router +diagnostics are inherently one-matrix/one-checkpoint per rollout, so their +``"series"`` entries are whole per-rollout artifacts (a matrix, a gating +dict) rather than a single number/array — ``render.py`` draws those as one +panel per rollout instead of one line/bar per rollout. + Rendering lives in ``render.py`` and dispatches on ``Reduced.kind`` — the catalog itself never imports plotstyle, so ``compute-one`` jobs stay LaTeX-free. @@ -59,7 +72,7 @@ from giant.analysis.router_gating import ( compute_router_share_by_process, compute_router_specialization, ) -from giant.analysis.sources import Side, open_side, physical_steps, secondaries +from giant.analysis.sources import RolloutSide, RolloutSpec, Side, open_side, physical_steps, secondaries from giant.analysis.type_embedding_distance import compute_type_embedding_l1_distance from giant.analysis.variables import RANGED_VARS, cos_scatter_expr @@ -69,51 +82,44 @@ class Bundle: """Everything a compute runs against — built once per ``compute-one`` job.""" ctx: Context - r_all: pl.LazyFrame # rollout, all rows (incl. synthetic termination rows) + rollouts: dict[str, RolloutSide] # name -> frames, insertion order = CLI order t_all: pl.LazyFrame # reference, all rows - r_phys: pl.LazyFrame # rollout, physical steps only t_phys: pl.LazyFrame # reference, physical steps only - checkpoint: str | None = None # from the rollout YAML; router_gating only - # Diagnostic pre-aggregated at rollout time (giant.rollout. - # L1DistCollector.summary()) — from the rollout YAML, type_embedding_l1_distance - # only. Unlike checkpoint/router_gating, this needs no live model: it's - # already a finished histogram, just passed through. - type_embedding_l1_dist: dict | None = None @classmethod def open( cls, - rollout, + rollouts: list[RolloutSpec], reference, ctx: Context, - checkpoint=None, chunk: tuple[int, int] | None = None, - type_embedding_l1_dist: dict | None = None, ) -> "Bundle": - """Open both sides, optionally restricted to one event-disjoint chunk. + """Open the reference + every rollout, optionally restricted to one event-disjoint chunk. - ``chunk = (chunk_index, n_chunks)`` filters both sides to + ``chunk = (chunk_index, n_chunks)`` filters every side to ``event_id % n_chunks == chunk_index`` *before* deriving the physical/ secondary views, so every downstream reduction (which is either row-local or a ``group_by("event_id")``) sees a self-contained, event-disjoint slice — no cross-chunk lookups are ever needed. """ - r_all = open_side(rollout, Side.rollout) t_all = open_side(reference, Side.reference) + pred = None if chunk is not None: idx, n = chunk pred = pl.col("event_id") % n == idx - r_all = r_all.filter(pred) t_all = t_all.filter(pred) - return cls( - ctx=ctx, - r_all=r_all, - t_all=t_all, - r_phys=physical_steps(r_all, Side.rollout), - t_phys=physical_steps(t_all, Side.reference), - checkpoint=checkpoint, - type_embedding_l1_dist=type_embedding_l1_dist, - ) + sides: dict[str, RolloutSide] = {} + for rs in rollouts: + r_all = open_side(rs.source, Side.rollout) + if pred is not None: + r_all = r_all.filter(pred) + sides[rs.name] = RolloutSide( + all=r_all, + phys=physical_steps(r_all, Side.rollout), + checkpoint=rs.checkpoint, + type_embedding_l1_dist=rs.type_embedding_l1_dist, + ) + return cls(ctx=ctx, rollouts=sides, t_all=t_all, t_phys=physical_steps(t_all, Side.reference)) @dataclass @@ -146,8 +152,10 @@ def _unchunkable( # small numpy/hist helpers # --------------------------------------------------------------------------- -_ROLL = "rollout" -_REF = "reference" + +def _per_rollout(b: Bundle, fn: Callable[[RolloutSide], object]) -> dict[str, object]: + """``{name: fn(rollout_side)}`` over every rollout, preserving CLI order.""" + return {name: fn(rs) for name, rs in b.rollouts.items()} def _counts(h: dict, key, nbins: int) -> list[int]: @@ -168,14 +176,20 @@ def _finalize_counts(merged: dict[str, list], key, nbins: int) -> list[int]: return list(merged.get(str(key), [0] * nbins)) -def _np_hist_pair(r: np.ndarray, t: np.ndarray, nbins: int) -> tuple[np.ndarray, np.ndarray, np.ndarray]: - """Shared-edge histogram of two small per-event arrays (robust range).""" - both = np.concatenate([r, t]) if (len(r) or len(t)) else np.array([0.0, 1.0]) +def _np_hist_shared_edges(arrays: list[np.ndarray], nbins: int) -> tuple[np.ndarray, list[np.ndarray]]: + """Shared-edge histogram of several small per-event arrays (robust range). + + The edges are sized from the union of every array (reference + all + rollouts), so every series in the resulting overlay is directly + comparable on one axis. + """ + non_empty = [a for a in arrays if len(a)] + both = np.concatenate(non_empty) if non_empty else np.array([0.0, 1.0]) lo, hi = float(np.quantile(both, 0.001)), float(np.quantile(both, 0.999)) if not (hi - lo > 1e-6 * max(abs(hi), 1.0)): lo, hi = lo - 0.5, hi + 0.5 edges = np.linspace(lo, hi, nbins + 1) - return edges, np.histogram(r, edges)[0], np.histogram(t, edges)[0] + return edges, [np.histogram(a, edges)[0] for a in arrays] def _ks_statistic(r_counts, t_counts) -> float: @@ -197,15 +211,22 @@ def _ks_statistic(r_counts, t_counts) -> float: return float(np.max(np.abs(r_cdf - t_cdf))) -def _integer_confusion(t: np.ndarray, r: np.ndarray, max_bins: int = 21) -> tuple[list[str], np.ndarray]: +def _integer_confusion( + t: np.ndarray, r: np.ndarray, max_bins: int = 21, cap: int | None = None +) -> tuple[list[str], np.ndarray]: """Confusion matrix of two paired small-integer arrays (e.g. secondary counts). Bins are consecutive integers ``0..cap``, with the last bin an overflow ``"cap+"`` bucket, so an occasional pathological count doesn't blow up the heatmap. Returns ``(labels, matrix)`` with ``matrix[i, j]`` counting pairs with ``t == i`` and ``r == j`` (both clipped into ``[0, cap]``). + + ``cap``, if given, is used as-is instead of being derived from ``t``/``r`` + — lets a multi-rollout caller fix one shared cap (and so one shared label + set) across every rollout's matrix rather than each panel picking its own. """ - cap = min(max(int(t.max()) if len(t) else 0, int(r.max()) if len(r) else 0, 1), max_bins - 1) + if cap is None: + cap = min(max(int(t.max()) if len(t) else 0, int(r.max()) if len(r) else 0, 1), max_bins - 1) t_c = np.clip(t.astype(np.int64), 0, cap) r_c = np.clip(r.astype(np.int64), 0, cap) n = cap + 1 @@ -274,7 +295,7 @@ def _marginal_overall_partial(b: Bundle, var: str) -> dict: _, expr = _var(var) edges = _marginal_edges(b.ctx, var) return { - "r": _partial_hist(b.r_phys, expr, edges), + "r": _per_rollout(b, lambda rs: _partial_hist(rs.phys, expr, edges)), "t": _partial_hist(b.t_phys, expr, edges), } @@ -283,7 +304,8 @@ def _marginal_overall_finalize(parts: list[dict], ctx: Context, var: str) -> Red label, _ = _var(var) edges = _marginal_edges(ctx, var) nb = len(edges) - 1 - r = sum_merge([p["r"] for p in parts]) + names = list(parts[0]["r"]) + series = {name: _finalize_counts(sum_merge([p["r"][name] for p in parts]), 0, nb) for name in names} t = sum_merge([p["t"] for p in parts]) return Reduced( id=f"marginal_{var}", @@ -293,8 +315,8 @@ def _marginal_overall_finalize(parts: list[dict], ctx: Context, var: str) -> Red xlabel=label, payload={ "edges": edges.tolist(), - _ROLL: _finalize_counts(r, 0, nb), - _REF: _finalize_counts(t, 0, nb), + "series": series, + "reference": _finalize_counts(t, 0, nb), "log_y": True, }, ) @@ -305,23 +327,24 @@ def _energy_group_expr(lf: pl.LazyFrame, edges: np.ndarray) -> pl.Expr: return pl.col("event_id").replace_strict(ids, bins, default=-1, return_dtype=pl.Int64) +def _grouped_hist_dict(lf: pl.LazyFrame, expr: pl.Expr, edges: np.ndarray, axis: str, ctx: Context, nb: int) -> dict: + if axis == "pdg": + h = hist1d(lf, expr, edges, group=pl.col("pdg")) + elif axis == "material": + h = hist1d(lf, expr, edges, group=pl.col("material")) + else: # energy + e_edges = np.asarray(ctx.energy_edges) + h = hist1d(lf, expr, edges, group=_energy_group_expr(lf, e_edges)) + return {str(k): _counts(h, k, nb) for k in h} + + def _marginal_grouped_partial(b: Bundle, var: str, axis: str) -> dict: _, expr = _var(var) edges = _marginal_edges(b.ctx, var) - if axis == "pdg": - r = hist1d(b.r_phys, expr, edges, group=pl.col("pdg")) - t = hist1d(b.t_phys, expr, edges, group=pl.col("pdg")) - elif axis == "material": - r = hist1d(b.r_phys, expr, edges, group=pl.col("material")) - t = hist1d(b.t_phys, expr, edges, group=pl.col("material")) - else: # energy - e_edges = np.asarray(b.ctx.energy_edges) - r = hist1d(b.r_phys, expr, edges, group=_energy_group_expr(b.r_phys, e_edges)) - t = hist1d(b.t_phys, expr, edges, group=_energy_group_expr(b.t_phys, e_edges)) nb = len(edges) - 1 return { - "r": {str(k): _counts(r, k, nb) for k in r}, - "t": {str(k): _counts(t, k, nb) for k in t}, + "r": _per_rollout(b, lambda rs: _grouped_hist_dict(rs.phys, expr, edges, axis, b.ctx, nb)), + "t": _grouped_hist_dict(b.t_phys, expr, edges, axis, b.ctx, nb), } @@ -329,29 +352,24 @@ def _marginal_grouped_finalize(parts: list[dict], ctx: Context, var: str, axis: label, _ = _var(var) edges = _marginal_edges(ctx, var) nb = len(edges) - 1 - r = sum_merge([p["r"] for p in parts]) - t = sum_merge([p["t"] for p in parts]) - groups: dict[str, dict] = {} + names = list(parts[0]["r"]) + r_merged = {name: sum_merge([p["r"][name] for p in parts]) for name in names} + t_merged = sum_merge([p["t"] for p in parts]) if axis == "pdg": - for k in ctx.top_pdgs: - groups[pdg_label(k)] = { - _ROLL: _finalize_counts(r, k, nb), - _REF: _finalize_counts(t, k, nb), - } + keys, labels = ctx.top_pdgs, [pdg_label(k) for k in ctx.top_pdgs] elif axis == "material": - for m in ctx.materials: - groups[material_label(m)] = { - _ROLL: _finalize_counts(r, m, nb), - _REF: _finalize_counts(t, m, nb), - } + keys, labels = ctx.materials, [material_label(m) for m in ctx.materials] else: # energy e_edges = np.asarray(ctx.energy_edges) - for bi, lbl in enumerate(energy_bin_labels(e_edges)): - groups[lbl] = { - _ROLL: _finalize_counts(r, bi, nb), - _REF: _finalize_counts(t, bi, nb), - } + keys, labels = list(range(len(e_edges) - 1)), energy_bin_labels(e_edges) + + groups: dict[str, dict] = {} + for k, lbl in zip(keys, labels): + groups[lbl] = { + "series": {name: _finalize_counts(r_merged[name], k, nb) for name in names}, + "reference": _finalize_counts(t_merged, k, nb), + } return Reduced( id=f"marginal_{var}_by_{axis}", @@ -364,7 +382,7 @@ def _marginal_grouped_finalize(parts: list[dict], ctx: Context, var: str, axis: # --------------------------------------------------------------------------- -# distance summary: a var x group-axis scorecard, reusing the marginal hists +# distance summary: a var x group-axis scorecard per rollout, reusing the marginal hists # --------------------------------------------------------------------------- @@ -379,29 +397,37 @@ def _distance_summary_partial(b: Bundle) -> dict: def _distance_summary_finalize(parts: list[dict], ctx: Context) -> Reduced: col_labels = ["overall", *GROUPING_AXES] - matrix: list[list[float]] = [] + names = list(parts[0][MARGINAL_VARS[0]]["overall"]["r"]) + matrices: dict[str, list[list[float]]] = {name: [] for name in names} + for var in MARGINAL_VARS: edges = _marginal_edges(ctx, var) nb = len(edges) - 1 - row: list[float] = [] - r = sum_merge([p[var]["overall"]["r"] for p in parts]) - t = sum_merge([p[var]["overall"]["t"] for p in parts]) - row.append(_ks_statistic(_finalize_counts(r, 0, nb), _finalize_counts(t, 0, nb))) + t_overall = sum_merge([p[var]["overall"]["t"] for p in parts]) + r_overall = {name: sum_merge([p[var]["overall"]["r"][name] for p in parts]) for name in names} + row: dict[str, list[float]] = {name: [] for name in names} + for name in names: + row[name].append( + _ks_statistic(_finalize_counts(r_overall[name], 0, nb), _finalize_counts(t_overall, 0, nb)) + ) for axis in GROUPING_AXES: - r = sum_merge([p[var][axis]["r"] for p in parts]) - t = sum_merge([p[var][axis]["t"] for p in parts]) - dists, weights = [], [] - for k in _group_keys(ctx, axis): - rc, tc = _finalize_counts(r, k, nb), _finalize_counts(t, k, nb) - w = sum(rc) + sum(tc) - if w == 0: - continue - dists.append(_ks_statistic(rc, tc)) - weights.append(w) - row.append(float(np.average(dists, weights=weights)) if dists else float("nan")) - matrix.append(row) + t_grp = sum_merge([p[var][axis]["t"] for p in parts]) + r_grp = {name: sum_merge([p[var][axis]["r"][name] for p in parts]) for name in names} + for name in names: + dists, weights = [], [] + for k in _group_keys(ctx, axis): + rc, tc = _finalize_counts(r_grp[name], k, nb), _finalize_counts(t_grp, k, nb) + w = sum(rc) + sum(tc) + if w == 0: + continue + dists.append(_ks_statistic(rc, tc)) + weights.append(w) + row[name].append(float(np.average(dists, weights=weights)) if dists else float("nan")) + + for name in names: + matrices[name].append(row[name]) return Reduced( id="marginal_distance_summary", @@ -410,7 +436,7 @@ def _distance_summary_finalize(parts: list[dict], ctx: Context) -> Reduced: title="Marginal distance summary (KS statistic, rollout vs reference)", xlabel="grouping axis", payload={ - "matrix": matrix, + "series": matrices, "row_labels": [_TITLE_NAMES[v] for v in MARGINAL_VARS], "col_labels": col_labels, "ylabel": "marginal variable", @@ -427,16 +453,24 @@ def _distance_summary_finalize(parts: list[dict], ctx: Context) -> Reduced: def _event_scalar_partial(b: Bundle, col: str, use_all: bool) -> dict: - r_lf, t_lf = (b.r_all, b.t_all) if use_all else (b.r_phys, b.t_phys) - r = event_scalars(r_lf)[col].to_numpy() - t = event_scalars(t_lf)[col].to_numpy() - return {"r": r.tolist(), "t": t.tolist()} + t_lf = b.t_all if use_all else b.t_phys + + def _vals(rs: RolloutSide) -> list[float]: + lf = rs.all if use_all else rs.phys + return event_scalars(lf)[col].to_numpy().tolist() + + return { + "r": _per_rollout(b, _vals), + "t": event_scalars(t_lf)[col].to_numpy().tolist(), + } def _event_scalar_finalize(parts: list[dict], ctx: Context, spec_id: str, title: str, xlabel: str) -> Reduced: - r = np.concatenate([np.asarray(p["r"], dtype=float) for p in parts]) + names = list(parts[0]["r"]) + r_arrays = {name: np.concatenate([np.asarray(p["r"][name], dtype=float) for p in parts]) for name in names} t = np.concatenate([np.asarray(p["t"], dtype=float) for p in parts]) - edges, rc, tc = _np_hist_pair(r, t, ctx.n_marginal_bins) + edges, counts = _np_hist_shared_edges([t, *(r_arrays[n] for n in names)], ctx.n_marginal_bins) + t_counts, *r_counts = counts return Reduced( id=spec_id, family="event", @@ -445,40 +479,44 @@ def _event_scalar_finalize(parts: list[dict], ctx: Context, spec_id: str, title: xlabel=xlabel, payload={ "edges": edges.tolist(), - _ROLL: rc.astype(np.int64).tolist(), - _REF: tc.astype(np.int64).tolist(), + "series": {name: c.astype(np.int64).tolist() for name, c in zip(names, r_counts)}, + "reference": t_counts.astype(np.int64).tolist(), "log_y": False, }, ) def _event_total_edep_by_energy_partial(b: Bundle) -> dict: - r = event_scalars(b.r_all) t = event_scalars(b.t_all) + + def _vals(rs: RolloutSide) -> dict: + r = event_scalars(rs.all) + return {"incident": r["incident_E"].to_list(), "edep": r["total_edep"].to_list()} + return { - "r_incident": r["incident_E"].to_list(), - "r_edep": r["total_edep"].to_list(), - "t_incident": t["incident_E"].to_list(), - "t_edep": t["total_edep"].to_list(), + "r": _per_rollout(b, _vals), + "t": {"incident": t["incident_E"].to_list(), "edep": t["total_edep"].to_list()}, } def _event_total_edep_by_energy_finalize(parts: list[dict], ctx: Context) -> Reduced: e_edges = np.asarray(ctx.energy_edges) - r_inc = np.concatenate([np.asarray(p["r_incident"], dtype=float) for p in parts]) - r_val = np.concatenate([np.asarray(p["r_edep"], dtype=float) for p in parts]) - t_inc = np.concatenate([np.asarray(p["t_incident"], dtype=float) for p in parts]) - t_val = np.concatenate([np.asarray(p["t_edep"], dtype=float) for p in parts]) - r_bin = np.clip(np.digitize(r_inc, e_edges[1:-1]), 0, len(e_edges) - 2) + names = list(parts[0]["r"]) + t_inc = np.concatenate([np.asarray(p["t"]["incident"], dtype=float) for p in parts]) + t_val = np.concatenate([np.asarray(p["t"]["edep"], dtype=float) for p in parts]) + r_inc = {n: np.concatenate([np.asarray(p["r"][n]["incident"], dtype=float) for p in parts]) for n in names} + r_val = {n: np.concatenate([np.asarray(p["r"][n]["edep"], dtype=float) for p in parts]) for n in names} + + edges, _ = _np_hist_shared_edges([t_val, *(r_val[n] for n in names)], ctx.n_marginal_bins) t_bin = np.clip(np.digitize(t_inc, e_edges[1:-1]), 0, len(e_edges) - 2) - edges, _, _ = _np_hist_pair(r_val, t_val, ctx.n_marginal_bins) + r_bin = {n: np.clip(np.digitize(r_inc[n], e_edges[1:-1]), 0, len(e_edges) - 2) for n in names} + groups: dict[str, dict] = {} for bi, lbl in enumerate(energy_bin_labels(e_edges)): - rc = np.histogram(r_val[r_bin == bi], edges)[0] tc = np.histogram(t_val[t_bin == bi], edges)[0] groups[lbl] = { - _ROLL: rc.astype(np.int64).tolist(), - _REF: tc.astype(np.int64).tolist(), + "series": {n: np.histogram(r_val[n][r_bin[n] == bi], edges)[0].astype(np.int64).tolist() for n in names}, + "reference": tc.astype(np.int64).tolist(), } return Reduced( id="event_total_edep_by_energy", @@ -497,15 +535,15 @@ def _event_total_edep_by_energy_finalize(parts: list[dict], ctx: Context) -> Red def _profile_partial(b: Bundle, coord_fn, edges_key: str) -> dict: edges = np.asarray(getattr(b.ctx, edges_key)) - r_lf = attach_entry_axis(b.r_all, entry_axis(b.r_all)) - t_lf = attach_entry_axis(b.t_all, entry_axis(b.t_all)) - r_ids, r_mat = profile_partial(r_lf, coord_fn(), edges, pl.col("edep")) - t_ids, t_mat = profile_partial(t_lf, coord_fn(), edges, pl.col("edep")) + + def _mat(lf: pl.LazyFrame) -> dict: + lf2 = attach_entry_axis(lf, entry_axis(lf)) + ids, mat = profile_partial(lf2, coord_fn(), edges, pl.col("edep")) + return {"ids": ids.tolist(), "mat": mat.tolist()} + return { - "r_ids": r_ids.tolist(), - "r_mat": r_mat.tolist(), - "t_ids": t_ids.tolist(), - "t_mat": t_mat.tolist(), + "r": _per_rollout(b, lambda rs: _mat(rs.all)), + "t": _mat(b.t_all), } @@ -537,12 +575,19 @@ def _profile_finalize( ) -> Reduced: edges = np.asarray(getattr(ctx, edges_key)) nb = len(edges) - 1 - _assert_event_disjoint([p["r_ids"] for p in parts], spec_id, "rollout") - _assert_event_disjoint([p["t_ids"] for p in parts], spec_id, "reference") - r_mats = [np.asarray(p["r_mat"], dtype=float).reshape(-1, nb) for p in parts] - t_mats = [np.asarray(p["t_mat"], dtype=float).reshape(-1, nb) for p in parts] - r_mean, r_std = profile_finalize(r_mats) + names = list(parts[0]["r"]) + + _assert_event_disjoint([p["t"]["ids"] for p in parts], spec_id, "reference") + t_mats = [np.asarray(p["t"]["mat"], dtype=float).reshape(-1, nb) for p in parts] t_mean, t_std = profile_finalize(t_mats) + + series: dict[str, dict] = {} + for name in names: + _assert_event_disjoint([p["r"][name]["ids"] for p in parts], spec_id, name) + mats = [np.asarray(p["r"][name]["mat"], dtype=float).reshape(-1, nb) for p in parts] + mean, std = profile_finalize(mats) + series[name] = {"mean": mean.tolist(), "std": std.tolist()} + return Reduced( id=spec_id, family="shower", @@ -551,10 +596,8 @@ def _profile_finalize( xlabel=xlabel, payload={ "edges": edges.tolist(), - "rollout_mean": r_mean.tolist(), - "rollout_std": r_std.tolist(), - "reference_mean": t_mean.tolist(), - "reference_std": t_std.tolist(), + "series": series, + "reference": {"mean": t_mean.tolist(), "std": t_std.tolist()}, "ylabel": "mean deposited energy per event [MeV]", }, ) @@ -573,13 +616,20 @@ _CONTAINMENT_QUANTILES: list[tuple[float, str]] = [ def _containment_finalize(parts: list[dict], ctx: Context, spec_id: str, quantile: float) -> Reduced: edges = np.asarray(ctx.depth_edges) nb = len(edges) - 1 - _assert_event_disjoint([p["r_ids"] for p in parts], spec_id, "rollout") - _assert_event_disjoint([p["t_ids"] for p in parts], spec_id, "reference") - r_full = np.concatenate([np.asarray(p["r_mat"], dtype=float).reshape(-1, nb) for p in parts], axis=0) - t_full = np.concatenate([np.asarray(p["t_mat"], dtype=float).reshape(-1, nb) for p in parts], axis=0) - r_depth = _containment_depths(r_full, edges, quantile) + names = list(parts[0]["r"]) + + _assert_event_disjoint([p["t"]["ids"] for p in parts], spec_id, "reference") + t_full = np.concatenate([np.asarray(p["t"]["mat"], dtype=float).reshape(-1, nb) for p in parts], axis=0) t_depth = _containment_depths(t_full, edges, quantile) - hedges, rc, tc = _np_hist_pair(r_depth, t_depth, ctx.n_marginal_bins) + + r_depths: dict[str, np.ndarray] = {} + for name in names: + _assert_event_disjoint([p["r"][name]["ids"] for p in parts], spec_id, name) + full = np.concatenate([np.asarray(p["r"][name]["mat"], dtype=float).reshape(-1, nb) for p in parts], axis=0) + r_depths[name] = _containment_depths(full, edges, quantile) + + hedges, counts = _np_hist_shared_edges([t_depth, *(r_depths[n] for n in names)], ctx.n_marginal_bins) + t_counts, *r_counts = counts return Reduced( id=spec_id, family="shower", @@ -588,8 +638,8 @@ def _containment_finalize(parts: list[dict], ctx: Context, spec_id: str, quantil xlabel=f"depth containing {quantile:.0%} of deposited energy [mm]", payload={ "edges": hedges.tolist(), - _ROLL: rc.astype(np.int64).tolist(), - _REF: tc.astype(np.int64).tolist(), + "series": {name: c.astype(np.int64).tolist() for name, c in zip(names, r_counts)}, + "reference": t_counts.astype(np.int64).tolist(), "log_y": False, }, ) @@ -601,20 +651,30 @@ def _containment_finalize(parts: list[dict], ctx: Context, spec_id: str, quantil def _species_share_partial(b: Bundle) -> dict: - r = species_share(b.r_all) t = species_share(b.t_all) + + def _map(rs: RolloutSide) -> dict[str, float]: + r = species_share(rs.all) + return {str(k): v for k, v in zip(r["pdg"].to_list(), r["total_edep"].to_list())} + return { - "r": {str(k): v for k, v in zip(r["pdg"].to_list(), r["total_edep"].to_list())}, + "r": _per_rollout(b, _map), "t": {str(k): v for k, v in zip(t["pdg"].to_list(), t["total_edep"].to_list())}, } def _species_share_finalize(parts: list[dict], ctx: Context) -> Reduced: - r_map = sum_merge([p["r"] for p in parts]) + names = list(parts[0]["r"]) + r_maps = {n: sum_merge([p["r"][n] for p in parts]) for n in names} t_map = sum_merge([p["t"] for p in parts]) - r_tot = sum(r_map.values()) or 1.0 t_tot = sum(t_map.values()) or 1.0 labels = [pdg_label(k) for k in ctx.top_pdgs] + + series: dict[str, list[float]] = {} + for n in names: + r_tot = sum(r_maps[n].values()) or 1.0 + series[n] = [r_maps[n].get(str(k), 0.0) / r_tot for k in ctx.top_pdgs] + return Reduced( id="species_edep_share", family="species", @@ -623,22 +683,23 @@ def _species_share_finalize(parts: list[dict], ctx: Context) -> Reduced: xlabel="species", payload={ "labels": labels, - _ROLL: [r_map.get(str(k), 0.0) / r_tot for k in ctx.top_pdgs], - _REF: [t_map.get(str(k), 0.0) / t_tot for k in ctx.top_pdgs], + "series": series, + "reference": [t_map.get(str(k), 0.0) / t_tot for k in ctx.top_pdgs], "ylabel": "fraction of total deposited energy", }, ) def _leakage_partial(b: Bundle) -> dict: - frac = leakage_fraction(b.r_all) - return {"frac": frac.tolist()} + return {"r": _per_rollout(b, lambda rs: leakage_fraction(rs.all).tolist())} def _leakage_finalize(parts: list[dict], ctx: Context) -> Reduced: - frac = np.concatenate([np.asarray(p["frac"], dtype=float) for p in parts]) - edges = np.linspace(0.0, max(float(frac.max()) if len(frac) else 1.0, 1e-3), ctx.n_marginal_bins + 1) - counts = np.histogram(frac, edges)[0] + names = list(parts[0]["r"]) + arrays = {n: np.concatenate([np.asarray(p["r"][n], dtype=float) for p in parts]) for n in names} + max_val = max((float(a.max()) for a in arrays.values() if len(a)), default=1e-3) + edges = np.linspace(0.0, max(max_val, 1e-3), ctx.n_marginal_bins + 1) + series = {n: np.histogram(arrays[n], edges)[0].astype(np.int64).tolist() for n in names} return Reduced( id="leakage_fraction", family="species", @@ -647,7 +708,7 @@ def _leakage_finalize(parts: list[dict], ctx: Context) -> Reduced: xlabel="escaped energy fraction", payload={ "edges": edges.tolist(), - _ROLL: counts.astype(np.int64).tolist(), + "series": series, "log_y": True, "note": "rollout only; the reference has no detector-escape concept", }, @@ -659,24 +720,36 @@ def _leakage_finalize(parts: list[dict], ctx: Context) -> Reduced: # --------------------------------------------------------------------------- -def _sec_frames(b: Bundle): - return ( - secondaries(b.r_phys, Side.rollout), - secondaries(b.t_all, Side.reference), - ) +def _t_sec(b: Bundle) -> pl.LazyFrame: + return secondaries(b.t_all, Side.reference) + + +def _r_sec(rs: RolloutSide) -> pl.LazyFrame: + return secondaries(rs.phys, Side.rollout) def _sec_count_per_event_partial(b: Bundle) -> dict: - r_sec, t_sec = _sec_frames(b) - r = r_sec.group_by("event_id").agg(pl.len().alias("n")).collect(engine="streaming")["n"].to_numpy() - t = t_sec.group_by("event_id").agg(pl.len().alias("n")).collect(engine="streaming")["n"].to_numpy() - return {"r": r.tolist(), "t": t.tolist()} + t = _t_sec(b).group_by("event_id").agg(pl.len().alias("n")).collect(engine="streaming")["n"].to_numpy() + + def _r(rs: RolloutSide) -> list[float]: + return ( + _r_sec(rs) + .group_by("event_id") + .agg(pl.len().alias("n")) + .collect(engine="streaming")["n"] + .to_numpy() + .tolist() + ) + + return {"r": _per_rollout(b, _r), "t": t.tolist()} def _sec_count_per_event_finalize(parts: list[dict], ctx: Context) -> Reduced: - r = np.concatenate([np.asarray(p["r"], dtype=float) for p in parts]) + names = list(parts[0]["r"]) t = np.concatenate([np.asarray(p["t"], dtype=float) for p in parts]) - edges, rc, tc = _np_hist_pair(r, t, min(ctx.n_marginal_bins, 40)) + r = {n: np.concatenate([np.asarray(p["r"][n], dtype=float) for p in parts]) for n in names} + edges, counts = _np_hist_shared_edges([t, *(r[n] for n in names)], min(ctx.n_marginal_bins, 40)) + t_c, *r_cs = counts return Reduced( id="sec_count_per_event", family="secondaries", @@ -685,8 +758,8 @@ def _sec_count_per_event_finalize(parts: list[dict], ctx: Context) -> Reduced: xlabel="secondaries per event", payload={ "edges": edges.tolist(), - _ROLL: rc.astype(np.int64).tolist(), - _REF: tc.astype(np.int64).tolist(), + "series": {n: c.astype(np.int64).tolist() for n, c in zip(names, r_cs)}, + "reference": t_c.astype(np.int64).tolist(), "log_y": False, }, ) @@ -698,14 +771,22 @@ def _counts_by_pdg(sec_lf: pl.LazyFrame) -> dict[str, int]: def _sec_count_per_species_partial(b: Bundle) -> dict: - r_sec, t_sec = _sec_frames(b) - return {"r": _counts_by_pdg(r_sec), "t": _counts_by_pdg(t_sec)} + return {"r": _per_rollout(b, lambda rs: _counts_by_pdg(_r_sec(rs))), "t": _counts_by_pdg(_t_sec(b))} def _sec_count_per_species_finalize(parts: list[dict], ctx: Context) -> Reduced: - r = sum_merge([p["r"] for p in parts]) + names = list(parts[0]["r"]) + r_maps = {n: sum_merge([p["r"][n] for p in parts]) for n in names} t = sum_merge([p["t"] for p in parts]) - keys = sorted(set(r) | set(t), key=lambda k: -(r.get(k, 0) + t.get(k, 0)))[: len(ctx.top_pdgs)] + + all_keys = set(t) + for m in r_maps.values(): + all_keys |= set(m) + + def _total(k: str) -> float: + return t.get(k, 0) + sum(m.get(k, 0) for m in r_maps.values()) + + keys = sorted(all_keys, key=lambda k: -_total(k))[: len(ctx.top_pdgs)] return Reduced( id="sec_count_per_species", family="secondaries", @@ -714,39 +795,34 @@ def _sec_count_per_species_finalize(parts: list[dict], ctx: Context) -> Reduced: xlabel="species", payload={ "labels": [pdg_label(int(k)) for k in keys], - _ROLL: [float(r.get(k, 0)) for k in keys], - _REF: [float(t.get(k, 0)) for k in keys], + "series": {n: [float(r_maps[n].get(k, 0)) for k in keys] for n in names}, + "reference": [float(t.get(k, 0)) for k in keys], "ylabel": "secondary count", }, ) def _sec_energy_partial(b: Bundle) -> dict: - r_sec, t_sec = _sec_frames(b) edges = np.linspace(*b.ctx.sec_energy_range, b.ctx.n_sec_bins + 1) return { - "r": _partial_hist(r_sec, pl.col("energy"), edges), - "t": _partial_hist(t_sec, pl.col("energy"), edges), + "r": _per_rollout(b, lambda rs: _partial_hist(_r_sec(rs), pl.col("energy"), edges)), + "t": _partial_hist(_t_sec(b), pl.col("energy"), edges), } def _sec_energy_finalize(parts: list[dict], ctx: Context) -> Reduced: edges = np.linspace(*ctx.sec_energy_range, ctx.n_sec_bins + 1) nb = len(edges) - 1 - r = sum_merge([p["r"] for p in parts]) + names = list(parts[0]["r"]) t = sum_merge([p["t"] for p in parts]) + series = {name: _finalize_counts(sum_merge([p["r"][name] for p in parts]), 0, nb) for name in names} return Reduced( id="sec_energy", family="secondaries", kind="overlay_hist", title="Secondary birth energy", xlabel="secondary energy [MeV]", - payload={ - "edges": edges.tolist(), - _ROLL: _finalize_counts(r, 0, nb), - _REF: _finalize_counts(t, 0, nb), - "log_y": True, - }, + payload={"edges": edges.tolist(), "series": series, "reference": _finalize_counts(t, 0, nb), "log_y": True}, ) @@ -760,50 +836,67 @@ def _sec_cos_angle_partial(b: Bundle) -> dict: ea = entry_axis(steps_lf) return _partial_hist(attach_entry_axis(sec_lf, ea), cos, edges) - r_sec, t_sec = _sec_frames(b) - return {"r": _side(r_sec, b.r_phys), "t": _side(t_sec, b.t_all)} + return { + "r": _per_rollout(b, lambda rs: _side(_r_sec(rs), rs.phys)), + "t": _side(_t_sec(b), b.t_all), + } def _sec_cos_angle_finalize(parts: list[dict], ctx: Context) -> Reduced: edges = np.linspace(-1.0, 1.0, ctx.n_sec_bins + 1) nb = len(edges) - 1 - r = sum_merge([p["r"] for p in parts]) + names = list(parts[0]["r"]) t = sum_merge([p["t"] for p in parts]) + series = {name: _finalize_counts(sum_merge([p["r"][name] for p in parts]), 0, nb) for name in names} return Reduced( id="sec_cos_angle", family="secondaries", kind="overlay_hist", title="Secondary emission angle relative to the shower axis", xlabel="cos of emission angle", - payload={ - "edges": edges.tolist(), - _ROLL: _finalize_counts(r, 0, nb), - _REF: _finalize_counts(t, 0, nb), - "log_y": False, - }, + payload={"edges": edges.tolist(), "series": series, "reference": _finalize_counts(t, 0, nb), "log_y": False}, ) def _n_sec_confusion_partial(b: Bundle) -> dict: - r_sec, t_sec = _sec_frames(b) - r_ids, r_n = sec_count_by_event(b.r_phys, r_sec) - t_ids, t_n = sec_count_by_event(b.t_all, t_sec) - return {"r_ids": r_ids.tolist(), "r_n": r_n.tolist(), "t_ids": t_ids.tolist(), "t_n": t_n.tolist()} + t_ids, t_n = sec_count_by_event(b.t_all, _t_sec(b)) + + def _r(rs: RolloutSide) -> dict: + ids, n = sec_count_by_event(rs.phys, _r_sec(rs)) + return {"ids": ids.tolist(), "n": n.tolist()} + + return {"r": _per_rollout(b, _r), "t": {"ids": t_ids.tolist(), "n": t_n.tolist()}} def _n_sec_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced: - r_ids = np.concatenate([np.asarray(p["r_ids"], dtype=np.int64) for p in parts]) - r_n = np.concatenate([np.asarray(p["r_n"], dtype=np.int64) for p in parts]) - t_ids = np.concatenate([np.asarray(p["t_ids"], dtype=np.int64) for p in parts]) - t_n = np.concatenate([np.asarray(p["t_n"], dtype=np.int64) for p in parts]) + names = list(parts[0]["r"]) # event-disjoint chunking (see Bundle.open) means each event_id appears in # exactly one part on each side, so a plain dict build is a safe merge. - r_map = dict(zip(r_ids.tolist(), r_n.tolist())) + t_ids = np.concatenate([np.asarray(p["t"]["ids"], dtype=np.int64) for p in parts]) + t_n = np.concatenate([np.asarray(p["t"]["n"], dtype=np.int64) for p in parts]) t_map = dict(zip(t_ids.tolist(), t_n.tolist())) - common = sorted(set(r_map) & set(t_map)) - true_n = np.array([t_map[e] for e in common], dtype=np.int64) - pred_n = np.array([r_map[e] for e in common], dtype=np.int64) - labels, mat = _integer_confusion(true_n, pred_n) + + pairs: dict[str, tuple[np.ndarray, np.ndarray]] = {} + max_val = 0 + for name in names: + r_ids = np.concatenate([np.asarray(p["r"][name]["ids"], dtype=np.int64) for p in parts]) + r_n = np.concatenate([np.asarray(p["r"][name]["n"], dtype=np.int64) for p in parts]) + r_map = dict(zip(r_ids.tolist(), r_n.tolist())) + common = sorted(set(r_map) & set(t_map)) + true_n = np.array([t_map[e] for e in common], dtype=np.int64) + pred_n = np.array([r_map[e] for e in common], dtype=np.int64) + pairs[name] = (true_n, pred_n) + if len(true_n): + max_val = max(max_val, int(true_n.max()), int(pred_n.max())) + + cap = min(max(max_val, 1), 20) + matrices: dict[str, list[list[int]]] = {} + labels: list[str] = [] + for name in names: + true_n, pred_n = pairs[name] + labels, mat = _integer_confusion(true_n, pred_n, cap=cap) + matrices[name] = mat.tolist() + return Reduced( id="n_sec_confusion", family="secondaries", @@ -811,7 +904,7 @@ def _n_sec_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced: title="Predicted vs true secondary count per event", xlabel="predicted secondaries (rollout)", payload={ - "matrix": mat.tolist(), + "series": matrices, "row_labels": labels, "col_labels": labels, "ylabel": "true secondaries (reference)", @@ -825,20 +918,18 @@ def _n_sec_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced: # router diagnostics (not chunked — already bounded/subsampled) # --------------------------------------------------------------------------- -_router_gating_partial, _router_gating_finalize = _unchunkable( - lambda b: compute_router_gating(b.checkpoint, b.r_phys, b.t_phys) -) +_router_gating_partial, _router_gating_finalize = _unchunkable(lambda b: compute_router_gating(b.rollouts, b.t_phys)) _router_share_pdg_partial, _router_share_pdg_finalize = _unchunkable( - lambda b: compute_router_share_by_pdg(b.checkpoint, b.r_phys, b.t_phys, b.ctx.top_pdgs) + lambda b: compute_router_share_by_pdg(b.rollouts, b.t_phys, b.ctx.top_pdgs) ) _router_share_process_partial, _router_share_process_finalize = _unchunkable( - lambda b: compute_router_share_by_process(b.checkpoint, b.t_phys) + lambda b: compute_router_share_by_process(b.rollouts, b.t_phys) ) _router_specialization_partial, _router_specialization_finalize = _unchunkable( - lambda b: compute_router_specialization(b.checkpoint, b.r_phys, b.t_phys) + lambda b: compute_router_specialization(b.rollouts, b.t_phys) ) _type_embedding_l1_distance_partial, _type_embedding_l1_distance_finalize = _unchunkable( - lambda b: compute_type_embedding_l1_distance(b.type_embedding_l1_dist) + lambda b: compute_type_embedding_l1_distance(b.rollouts) ) diff --git a/giant/analysis/condor.py b/giant/analysis/condor.py index f512621..09f2f68 100644 --- a/giant/analysis/condor.py +++ b/giant/analysis/condor.py @@ -1,4 +1,4 @@ -"""HTCondor orchestration driven by a ``giant rollout`` YAML sidecar. +"""HTCondor orchestration driven by one or more ``giant rollout`` YAML sidecars. A rollout writes a YAML sidecar (``giant/cli.py:_write_prediction_ref`` + rollout extras) that already names both files we need and carries the run's @@ -10,8 +10,11 @@ provenance: * ``checkpoint``, ``geometry_oracle``, ``energy_cutoff``, ``steps``, ... — metadata that flows straight into every plot's gallery ``metadata.yaml``. -So the analysis takes that one YAML as input, derives its own **run directory** -next to the rollout parquet, and lays everything out under it: +The analysis takes N such YAMLs — one series per rollout, all required to +share the same ``dataset`` (the premise is "N candidates vs one ground +truth") — resolves each one's series name (``load_rollout_yamls``), derives +its own **run directory** next to the first rollout's parquet, and lays +everything out under it: /shared.json fixed bin edges / group sets (prep) /run_meta.json resolved rollout/reference paths + plot metadata @@ -44,6 +47,7 @@ from __future__ import annotations import json import shutil import sys +from collections.abc import Sequence from dataclasses import dataclass, field from pathlib import Path @@ -54,7 +58,7 @@ from giant.analysis.catalog import Bundle, catalog_ids, get_spec from giant.analysis.context import Context, build_context from giant.analysis.reduced import Partial from giant.analysis.runtime_estimate import estimate_runtime_s -from giant.analysis.sources import Side, open_side +from giant.analysis.sources import RolloutSpec, Side, open_side # Keys copied verbatim from a rollout YAML into each plot's gallery metadata. _PLOT_META_KEYS = ( @@ -111,8 +115,60 @@ def load_rollout_yaml(path: str | Path) -> dict: return d +@dataclass +class LoadedRollout: + """One rollout YAML plus its resolved series ``name`` (see ``load_rollout_yamls``).""" + + name: str + yaml: dict + + +def load_rollout_yamls( + paths: Sequence[str | Path], labels: Sequence[str] | None = None +) -> tuple[list[LoadedRollout], str]: + """Load every rollout YAML, resolve each one's series name, and verify they + all share one reference (``dataset``) file — the premise is "N candidates + vs one ground truth", not N independent comparisons. + + Names: an explicit ``labels[i]`` if given (``labels`` must be empty or + exactly ``len(paths)`` long); otherwise the YAML's stem for N>1, or + ``"rollout"`` for the single-YAML case — matching today's one-series + legend/payload key, so a single-rollout run renders identically to + before this feature existed. Raises ``ValueError`` if two rollouts + resolve to the same name, or if the YAMLs don't all name the same + ``dataset``. + """ + if labels and len(labels) != len(paths): + raise ValueError(f"--label given {len(labels)} time(s) but {len(paths)} rollout YAML(s) were passed") + yamls = [load_rollout_yaml(p) for p in paths] + if labels: + names = list(labels) + elif len(paths) == 1: + names = ["rollout"] + else: + names = [Path(p).stem for p in paths] + if len(set(names)) != len(names): + dupes = sorted({n for n in names if names.count(n) > 1}) + raise ValueError(f"rollout series names collide: {dupes} — pass --label to disambiguate") + + references = {str(y["dataset"]) for y in yamls} + if len(references) > 1: + detail = "\n".join(f" {p}: dataset={y['dataset']!r}" for p, y in zip(paths, yamls)) + raise ValueError( + "all rollout YAMLs must be seeded from the same reference (dataset) " + f"file — got {len(references)} distinct ones:\n{detail}" + ) + + return [LoadedRollout(name=n, yaml=y) for n, y in zip(names, yamls)], yamls[0]["dataset"] + + +def _run_tag(y: dict) -> str: + rollout = Path(y["output"]) + return str(y.get("prediction_id") or rollout.stem)[:8] + + def derive_run_dir( - rollout_yaml: dict, + rollout_yamls: list[dict], run_dir: str | Path | None = None, default_base: str | Path | None = None, ) -> Path: @@ -122,14 +178,24 @@ def derive_run_dir( ``default_base / analysis_`` if ``default_base`` is given (the CLI passes the repo's gitignored ``analysis_runs/``, so run directories don't pile up on ``/ceph`` next to the rollout parquet). Falls back to next to - the rollout parquet — the original convention — for callers that don't - care where the run directory lives. + the *first* rollout's parquet — the original convention — for callers + that don't care where the run directory lives. + + ``tag`` is a single rollout's ``prediction_id``/output stem (matching + today's single-rollout convention exactly) when there's only one; for + N>1 it joins up to three tags with ``-``, then ``-plus`` for any + beyond that, so a many-rollout run still gets a short, stable directory + name. """ if run_dir is not None: return Path(run_dir) - rollout = Path(rollout_yaml["output"]) - tag = str(rollout_yaml.get("prediction_id") or rollout.stem)[:8] - base = Path(default_base) if default_base is not None else rollout.parent + tags = [_run_tag(y) for y in rollout_yamls] + if len(tags) == 1: + tag = tags[0] + else: + shown, rest = tags[:3], tags[3:] + tag = "-".join(shown) + (f"-plus{len(rest)}" if rest else "") + base = Path(default_base) if default_base is not None else Path(rollout_yamls[0]["output"]).parent return base / f"analysis_{tag}" @@ -139,17 +205,21 @@ def _plot_meta(rollout_yaml: dict) -> dict: @dataclass class RunMeta: - """Resolved paths + plot metadata for one analysis run (``run_meta.json``).""" + """Resolved paths + plot metadata for one analysis run (``run_meta.json``). - rollout: str + ``rollouts`` is ``[{"name", "path", "plot_meta"}, ...]``, insertion order + = the order rollouts were given on the CLI (and so the order every + ``Reduced.payload["series"]`` dict is built in — see ``catalog.py``). + """ + + rollouts: list[dict] reference: str run_dir: str title: str - plot_meta: dict n_chunks: int = 1 - # rollout+reference row count of each event_id-disjoint chunk, and the - # dataset total — inputs to `runtime_estimate.estimate_runtime_s`. Empty/0 - # on run directories written before this field existed. + # combined rollout+reference row count of each event_id-disjoint chunk, + # and the dataset total — inputs to `runtime_estimate.estimate_runtime_s`. + # Empty/0 on run directories written before this field existed. rows_per_chunk: list[int] = field(default_factory=list) total_rows: int = 0 @@ -161,8 +231,8 @@ class RunMeta: return cls(**json.loads(Path(path).read_text())) -def _rows_per_chunk(rollout: str | Path, reference: str | Path, n_chunks: int) -> list[int]: - """Rollout+reference row count of each ``event_id % n_chunks`` chunk. +def _rows_per_chunk(rollouts: list[str | Path], reference: str | Path, n_chunks: int) -> list[int]: + """Combined rollout+reference row count of each ``event_id % n_chunks`` chunk. One cheap streaming ``group_by`` per side (just the ``event_id`` column) — the sizing input every job's estimated walltime @@ -178,7 +248,8 @@ def _rows_per_chunk(rollout: str | Path, reference: str | Path, n_chunks: int) - ) out = [0] * n_chunks - for lf in (open_side(rollout, Side.rollout), open_side(reference, Side.reference)): + sides = [open_side(reference, Side.reference)] + [open_side(r, Side.rollout) for r in rollouts] + for lf in sides: df = counts(lf) for c, n in zip(df["_c"].to_list(), df["n"].to_list()): out[c] += n @@ -186,20 +257,22 @@ def _rows_per_chunk(rollout: str | Path, reference: str | Path, n_chunks: int) - def prep( - rollout_yaml: str | Path, + rollout_yamls: Sequence[str | Path], run_dir: str | Path | None = None, n_chunks: int = 1, default_base: str | Path | None = None, + labels: Sequence[str] | None = None, **ctx_kwargs, ) -> Path: - """Read the rollout YAML, build the shared context, and lay out the run dir. + """Read the rollout YAML(s), build the shared context, and lay out the run dir. Writes ``shared.json`` + ``run_meta.json`` and returns the run directory. ``n_chunks`` is the run-level chunk count every ``compute-one``/``merge-one`` job reads back out of ``run_meta.json`` (via ``RunMeta.n_chunks``), so it is resolved once here rather than re-passed (and risking disagreement) at every later step. See ``derive_run_dir`` for how ``run_dir``/``default_base`` - resolve the actual directory. + resolve the actual directory, and ``load_rollout_yamls`` for how + ``labels``/YAML stems resolve each rollout's series name. Clears any existing ``reduced_partial/``/``reduced/`` from a prior prep of this same ``run_dir``: partial files carry no record of what context @@ -208,8 +281,8 @@ def prep( rollout/reference files changed) would otherwise let ``merge_one`` silently merge stale partials against the new ``shared.json``. """ - y = load_rollout_yaml(rollout_yaml) - run_path = derive_run_dir(y, run_dir, default_base=default_base) + loaded, reference = load_rollout_yamls(list(rollout_yamls), labels) + run_path = derive_run_dir([lr.yaml for lr in loaded], run_dir, default_base=default_base) run_path.mkdir(parents=True, exist_ok=True) for stale in ("reduced_partial", "reduced"): @@ -217,19 +290,22 @@ def prep( if stale_dir.exists(): shutil.rmtree(stale_dir) - rollout, reference = y["output"], y["dataset"] - ctx = build_context(rollout, reference, **ctx_kwargs) + rollout_specs = [RolloutSpec(name=lr.name, source=lr.yaml["output"]) for lr in loaded] + ctx = build_context(rollout_specs, reference, **ctx_kwargs) ctx.save(run_path / "shared.json") - rows_per_chunk = _rows_per_chunk(rollout, reference, n_chunks) + rows_per_chunk = _rows_per_chunk([lr.yaml["output"] for lr in loaded], reference, n_chunks) + + rollouts_meta = [ + {"name": lr.name, "path": str(lr.yaml["output"]), "plot_meta": _plot_meta(lr.yaml)} for lr in loaded + ] + ckpts = ", ".join(Path(lr.yaml.get("checkpoint", "")).name or "rollout" for lr in loaded) - ckpt = Path(y.get("checkpoint", "")).name or "rollout" RunMeta( - rollout=str(rollout), + rollouts=rollouts_meta, reference=str(reference), run_dir=str(run_path), - title=f"GIANT rollout analysis — {ckpt}", - plot_meta=_plot_meta(y), + title=f"GIANT rollout analysis — {ckpts}", n_chunks=n_chunks, rows_per_chunk=rows_per_chunk, total_rows=sum(rows_per_chunk), @@ -244,17 +320,19 @@ def prep( def compute_reduced( spec_id: str, - rollout: str | Path, + rollouts: list[dict], reference: str | Path, shared: str | Path, out: str | Path, - checkpoint: str | None = None, chunk_index: int = 0, n_chunks: int = 1, - type_embedding_l1_dist: dict | None = None, ) -> Path: """Core: run one (plot, chunk)'s partial reduction against explicit paths. + ``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?}, + ...]``, one per rollout series (insertion order preserved through to every + plot's ``Reduced.payload["series"]``). + Writes a ``Partial`` JSON — the raw, not-yet-merged output of ``PlotSpec.compute_partial`` — never a finished ``Reduced``; ``merge_one`` is what combines every chunk's ``Partial`` for a plot into the final @@ -268,14 +346,16 @@ def compute_reduced( raise ValueError( f"{spec_id}: chunk_index={chunk_index} out of range for n_chunks={effective_n} (chunkable={spec.chunkable})" ) - bundle = Bundle.open( - rollout, - reference, - ctx, - checkpoint=checkpoint, - chunk=(chunk_index, effective_n), - type_embedding_l1_dist=type_embedding_l1_dist, - ) + rollout_specs = [ + RolloutSpec( + name=r["name"], + source=r["path"], + checkpoint=r.get("checkpoint"), + type_embedding_l1_dist=r.get("type_embedding_l1_dist"), + ) + for r in rollouts + ] + bundle = Bundle.open(rollout_specs, reference, ctx, chunk=(chunk_index, effective_n)) partial = Partial( id=spec_id, family=spec.family, @@ -291,16 +371,23 @@ def compute_one(spec_id: str, run_dir: str | Path, chunk_index: int = 0) -> Path """Run one (plot, chunk)'s partial reduction from a prepped run directory.""" run_path = Path(run_dir) meta = RunMeta.load(run_path / "run_meta.json") + rollouts = [ + { + "name": ro["name"], + "path": ro["path"], + "checkpoint": ro["plot_meta"].get("checkpoint"), + "type_embedding_l1_dist": ro["plot_meta"].get("type_embedding_l1_dist"), + } + for ro in meta.rollouts + ] return compute_reduced( spec_id, - meta.rollout, + rollouts, meta.reference, run_path / "shared.json", run_path / "reduced_partial" / f"{spec_id}__{chunk_index}.json", - checkpoint=meta.plot_meta.get("checkpoint"), chunk_index=chunk_index, n_chunks=meta.n_chunks, - type_embedding_l1_dist=meta.plot_meta.get("type_embedding_l1_dist"), ) diff --git a/giant/analysis/context.py b/giant/analysis/context.py index 090ead5..2b273c2 100644 --- a/giant/analysis/context.py +++ b/giant/analysis/context.py @@ -26,7 +26,7 @@ from giant.analysis.reduce import ( entry_axis, transverse_expr, ) -from giant.analysis.sources import Side, open_side, physical_steps, secondaries +from giant.analysis.sources import RolloutSpec, Side, open_side, physical_steps, secondaries from giant.analysis.variables import RANGED_VARS @@ -74,9 +74,9 @@ def _row_subsample(lf: pl.LazyFrame, sample_rows: int, seed: int) -> pl.LazyFram return lf.filter((pl.col("pre_E").hash(seed=seed) % 2**32) < threshold) -def _combined_quantiles(r_vals: np.ndarray, t_vals: np.ndarray, lo_q: float, hi_q: float) -> tuple[float, float]: - """Robust (lo_q, hi_q) range over the union of two value samples.""" - both = np.concatenate([r_vals, t_vals]) +def _combined_quantiles(vals: list[np.ndarray], lo_q: float, hi_q: float) -> tuple[float, float]: + """Robust (lo_q, hi_q) range over the union of several value samples.""" + both = np.concatenate(vals) lo, hi = float(np.quantile(both, lo_q)), float(np.quantile(both, hi_q)) if not (hi - lo > 1e-6 * max(abs(hi), 1.0)): lo, hi = lo - 0.5, hi + 0.5 @@ -84,7 +84,7 @@ def _combined_quantiles(r_vals: np.ndarray, t_vals: np.ndarray, lo_q: float, hi_ def build_context( - rollout: str | Path | pl.LazyFrame, + rollouts: list[RolloutSpec], reference: str | Path | pl.LazyFrame, *, n_energy_bins: int = 4, @@ -94,41 +94,51 @@ def build_context( sample_rows: int = 1_000_000, seed: int = 0, ) -> Context: - """Resolve the shared context from the two files (the ``prep`` step).""" - r_all = open_side(rollout, Side.rollout) + """Resolve the shared context from the reference + every rollout (the ``prep`` step). + + Every range/quantile below is the union of the reference and *all* + rollouts, so a single set of fixed bin edges/group sets is valid for + every series a compute job streams over. + """ t_all = open_side(reference, Side.reference) - r_lf = physical_steps(r_all, Side.rollout) t_lf = physical_steps(t_all, Side.reference) + r_lfs = {rs.name: physical_steps(open_side(rs.source, Side.rollout), Side.rollout) for rs in rollouts} # Ranged marginal variables: robust ranges over a shared row subsample. exprs = [e.alias(n) for n, (_, e) in RANGED_VARS.items()] - r_s = _row_subsample(r_lf, sample_rows, seed).select(exprs).collect(engine="streaming") t_s = _row_subsample(t_lf, sample_rows, seed).select(exprs).collect(engine="streaming") + r_s = { + name: _row_subsample(lf, sample_rows, seed).select(exprs).collect(engine="streaming") + for name, lf in r_lfs.items() + } var_ranges = { - name: _combined_quantiles(r_s[name].to_numpy(), t_s[name].to_numpy(), _LO_Q, _HI_Q) for name in RANGED_VARS + name: _combined_quantiles([t_s[name].to_numpy(), *(df[name].to_numpy() for df in r_s.values())], _LO_Q, _HI_Q) + for name in RANGED_VARS } # Energy-bin edges from exact per-event incident energies (cheap group_by). def _incident(lf: pl.LazyFrame) -> np.ndarray: return lf.group_by("event_id").agg(pl.col("pre_E").max()).collect(engine="streaming")["pre_E"].to_numpy() - r_inc, t_inc = _incident(r_lf), _incident(t_lf) - energy_edges = energy_bin_edges(np.concatenate([r_inc, t_inc]), n_energy_bins) + t_inc = _incident(t_lf) + r_inc = {name: _incident(lf) for name, lf in r_lfs.items()} + energy_edges = energy_bin_edges(np.concatenate([t_inc, *r_inc.values()]), n_energy_bins) # Top PDG species and material list (cheap single-column group_bys). def _counts(lf: pl.LazyFrame, col: str) -> pl.DataFrame: return lf.group_by(col).agg(pl.len().alias("n")).collect(engine="streaming") pdg_counts = ( - pl.concat([_counts(r_lf, "pdg"), _counts(t_lf, "pdg")]) + pl.concat([_counts(t_lf, "pdg"), *(_counts(lf, "pdg") for lf in r_lfs.values())]) .group_by("pdg") .agg(pl.col("n").sum()) .sort("n", descending=True) ) top_pdgs = [int(x) for x in pdg_counts["pdg"].to_list()[:top_k_pdg]] - materials = sorted( - set(_counts(r_lf, "material")["material"].to_list()) | set(_counts(t_lf, "material")["material"].to_list()) - ) + material_set: set[str] = set(_counts(t_lf, "material")["material"].to_list()) + for lf in r_lfs.values(): + material_set |= set(_counts(lf, "material")["material"].to_list()) + materials = sorted(material_set) # Shower depth / transverse ranges from a subsampled proxy. def _proxy(lf: pl.LazyFrame) -> tuple[np.ndarray, np.ndarray]: @@ -140,19 +150,20 @@ def build_context( ) return sub["d"].to_numpy(), sub["t"].to_numpy() - r_d, r_t = _proxy(r_lf) t_d, t_t = _proxy(t_lf) - d_lo, d_hi = _combined_quantiles(r_d, t_d, _LO_Q, _HI_Q) + r_proxy = {name: _proxy(lf) for name, lf in r_lfs.items()} + d_lo, d_hi = _combined_quantiles([t_d, *(p[0] for p in r_proxy.values())], _LO_Q, _HI_Q) depth_edges = np.linspace(d_lo, d_hi, n_marginal_bins + 1) - t_hi = max(float(np.quantile(np.concatenate([r_t, t_t]), _HI_Q)), 1e-6) + t_hi = max(float(np.quantile(np.concatenate([t_t, *(p[1] for p in r_proxy.values())]), _HI_Q)), 1e-6) transverse_edges = np.linspace(0.0, t_hi, n_marginal_bins + 1) # Secondary energy range. - r_se = secondaries(r_lf, Side.rollout).select("energy") - t_se = secondaries(t_all, Side.reference).select("energy") - r_se = _row_sample_col(r_se, sample_rows, seed) - t_se = _row_sample_col(t_se, sample_rows, seed) - sec_energy_range = _combined_quantiles(r_se, t_se, _LO_Q, _HI_Q) + t_se = _row_sample_col(secondaries(t_all, Side.reference).select("energy"), sample_rows, seed) + r_se = { + name: _row_sample_col(secondaries(lf, Side.rollout).select("energy"), sample_rows, seed) + for name, lf in r_lfs.items() + } + sec_energy_range = _combined_quantiles([t_se, *r_se.values()], _LO_Q, _HI_Q) return Context( n_marginal_bins=n_marginal_bins, @@ -165,8 +176,8 @@ def build_context( sec_energy_range=sec_energy_range, n_sec_bins=n_sec_bins, n_events={ - "rollout": len(r_inc), "reference": len(t_inc), + **{name: len(arr) for name, arr in r_inc.items()}, }, ) diff --git a/giant/analysis/reduced.py b/giant/analysis/reduced.py index d3b2c24..bbf79ac 100644 --- a/giant/analysis/reduced.py +++ b/giant/analysis/reduced.py @@ -11,19 +11,24 @@ import json from dataclasses import asdict, dataclass, field from pathlib import Path -# Reduced.kind values: -# "overlay_hist" rollout vs reference density histogram over shared edges +# Reduced.kind values (payload keys a rollout series by name under +# payload["series"], with the reference — where one exists — kept as one +# distinguished payload["reference"] entry; see catalog.py's module +# docstring for the full per-kind payload shape): +# "overlay_hist" N-rollout-series vs reference density histogram over shared edges # "grouped_hist" one panel per group (energy/pdg/material), each an overlay -# "profile" edep-weighted mean +/- event-RMS vs depth/radius, two series -# "bar" per-category rollout vs reference bars (share / counts) -# "single_hist" one series only (e.g. rollout leakage; reference has none) -# "router_gating" stacked mean MoE gate weight vs energy, rollout + reference -# "router_share" stacked bar of MoE top-1 dispatch share by category -# "router_specialization" max gate weight vs energy, rollout + reference (one -# scalar trend line summarizing "router_gating") -# "heatmap" row x col matrix + colorbar (distance scorecard or a -# predicted-vs-true confusion matrix) -# "unavailable" plot not applicable to this run (e.g. non-MoE checkpoint) +# "profile" edep-weighted mean +/- event-RMS vs depth/radius, N series + reference +# "bar" per-category N-rollout-series vs reference bars (share / counts) +# "single_hist" rollout-only series (e.g. leakage; reference has none) +# "router_gating" stacked mean MoE gate weight vs energy, one rollout+reference +# panel-pair per rollout with an enabled MoE router +# "router_share" stacked bar of MoE top-1 dispatch share by category, one +# panel per rollout with an enabled MoE router +# "router_specialization" max gate weight vs energy (one scalar trend line +# summarizing "router_gating"), per rollout with an enabled router +# "heatmap" row x col matrix + colorbar, one panel per rollout (a +# distance scorecard or a predicted-vs-true confusion matrix) +# "unavailable" plot not applicable to this run (e.g. no MoE checkpoint) @dataclass diff --git a/giant/analysis/render.py b/giant/analysis/render.py index 41ad7dc..afb1ffd 100644 --- a/giant/analysis/render.py +++ b/giant/analysis/render.py @@ -9,6 +9,14 @@ streaming compute. For each reduced artifact it writes ``//.pdf`` plus a sibling ``.yaml`` (per-plot gallery metadata) and a per-family ``metadata.yaml``. Optionally runs ``gallery generate`` to build the static HTML site. + +Every rollout series gets a stable color via ``ps.get_color(i)``, ``i`` being +its position in ``payload["series"]`` — that position is fixed by the run's +YAML/``--label`` order (threaded unchanged from ``condor.RunMeta.rollouts`` +through every ``PlotSpec``), so a given rollout keeps the same color across +every plot in a run. The reference, where a plot has one, always draws in one +fixed, distinct style (dark ink, dashed) instead of taking a slot in that +cycle. """ from __future__ import annotations @@ -22,7 +30,11 @@ import yaml from giant.analysis.reduced import Reduced -_SERIES_LABELS = {"rollout": "rollout", "reference": "reference (Geant4)"} +_REFERENCE_LABEL = "reference (Geant4)" + + +def _ref_color() -> str: + return ps.colors.INK["primary"] def _density(counts: list[int] | np.ndarray, edges: np.ndarray) -> np.ndarray: @@ -33,10 +45,13 @@ def _density(counts: list[int] | np.ndarray, edges: np.ndarray) -> np.ndarray: return counts / (total * (edges[1] - edges[0])) -def _overlay(ax, edges: np.ndarray, series: dict[str, list], log_y: bool) -> None: - for key in ("reference", "rollout"): - if key in series: - ax.stairs(_density(series[key], edges), edges, label=_SERIES_LABELS[key]) +def _overlay(ax, edges: np.ndarray, payload: dict, log_y: bool) -> None: + if "reference" in payload: + ax.stairs( + _density(payload["reference"], edges), edges, label=_REFERENCE_LABEL, color=_ref_color(), linestyle="--" + ) + for i, (name, counts) in enumerate(payload.get("series", {}).items()): + ax.stairs(_density(counts, edges), edges, label=name, color=ps.get_color(i)) if log_y: ax.set_yscale("log") @@ -47,8 +62,8 @@ def _router_summary(router_cfg: dict) -> str: return f"{router_cfg.get('type', '?')}×{router_cfg.get('n_experts', '?')}" -def _figure_params_v2(mc: dict, run_meta: dict) -> dict: - """`_figure_params` for a new-shape (nested) `model_config` — has a +def _figure_params_v2(mc: dict, meta: dict) -> dict: + """`_figure_params_single` for a new-shape (nested) `model_config` — has a `stage1_model` key. Reports stage 1's architecture (the headline generator); stage 2's generator is only added (`mode_s2`) when it differs from stage 1's, since a mixed run (the `stage1=flow` + @@ -71,23 +86,24 @@ def _figure_params_v2(mc: dict, run_meta: dict) -> dict: if particle_type is not None: params["conditioning"] = particle_type params["router"] = _router_summary(s1.get("router") or {}) - 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) + if meta.get("training_epoch") is not None: + params["epoch"] = meta["training_epoch"] + if meta.get("best_val_loss") is not None: + params["best_val_loss"] = round(meta["best_val_loss"], 4) if mode == "wgan": noise_dim = (s1.get("wgan") or {}).get("noise_dim") if noise_dim is not None: params["noise_dim"] = noise_dim - elif run_meta.get("steps") is not None: - params["steps"] = run_meta["steps"] + elif meta.get("steps") is not None: + params["steps"] = meta["steps"] return params -def _figure_params(run_meta: dict) -> dict: - """Curated run identity for the figure subtitle (``new_figure(params=...)``). +def _figure_params_single(meta: dict) -> dict: + """Curated run identity for the figure subtitle (``new_figure(params=...)``), + for exactly one rollout's ``plot_meta``. - ``run_meta``/each plot's own ``.yaml`` (see ``_plot_metadata``) already + ``meta``/each plot's own ``.yaml`` (see ``_plot_metadata``) already carry every threaded model/training/rollout/dataset parameter for after-the-fact lookup — this picks only the handful that matter for telling figures apart at a glance while flipping through a gallery, since @@ -99,9 +115,9 @@ def _figure_params(run_meta: dict) -> dict: Handles both a v0.2 checkpoint's flat ``model_config`` and a v0.3.0 nested one (has a ``stage1_model`` key — see ``_figure_params_v2``). """ - mc = run_meta.get("model_config") or {} + mc = meta.get("model_config") or {} if "stage1_model" in mc: - return _figure_params_v2(mc, run_meta) + return _figure_params_v2(mc, meta) mode = mc.get("mode") params: dict = {} @@ -114,18 +130,35 @@ def _figure_params(run_meta: dict) -> dict: if mc.get("conditioning") is not None: params["conditioning"] = mc["conditioning"] params["router"] = _router_summary(mc.get("router") or {}) - 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) + if meta.get("training_epoch") is not None: + params["epoch"] = meta["training_epoch"] + if meta.get("best_val_loss") is not None: + params["best_val_loss"] = round(meta["best_val_loss"], 4) if mode == "wgan": if mc.get("noise_dim") is not None: params["noise_dim"] = mc["noise_dim"] - elif run_meta.get("steps") is not None: - params["steps"] = run_meta["steps"] + elif meta.get("steps") is not None: + params["steps"] = meta["steps"] return params +def _figure_params(run_meta: dict) -> dict: + """Curated run identity for the figure subtitle. + + A single-rollout run reuses that rollout's ``plot_meta`` (same curated + model/training/rollout subset as always — see ``_figure_params_single``); + a multi-rollout run instead names the series being compared, since no + single ``model_config`` applies to the figure as a whole (each plot's own + gallery YAML still carries every rollout's full ``plot_meta`` for + after-the-fact lookup, via ``_plot_metadata``). + """ + rollouts = run_meta.get("rollouts") or {} + if len(rollouts) == 1: + ((_, meta),) = rollouts.items() + return _figure_params_single(meta) + return {"rollouts": ", ".join(rollouts)} if rollouts else {} + + def _render_overlay(r: Reduced, params: dict): edges = np.asarray(r.payload["edges"]) fig, ax = ps.new_figure("thesis-single", title=r.title, params=params) @@ -139,7 +172,8 @@ def _render_overlay(r: Reduced, params: dict): def _render_single(r: Reduced, params: dict): edges = np.asarray(r.payload["edges"]) fig, ax = ps.new_figure("thesis-single", title=r.title, params=params) - ax.stairs(_density(r.payload["rollout"], edges), edges, label=_SERIES_LABELS["rollout"]) + for i, (name, counts) in enumerate(r.payload.get("series", {}).items()): + ax.stairs(_density(counts, edges), edges, label=name, color=ps.get_color(i)) if r.payload.get("log_y"): ax.set_yscale("log") if r.payload.get("log_x"): @@ -181,11 +215,17 @@ 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, params=params) - for key in ("reference", "rollout"): - mean = np.asarray(r.payload[f"{key}_mean"]) - std = np.asarray(r.payload[f"{key}_std"]) - (line,) = ax.plot(centers, mean, label=_SERIES_LABELS[key]) - ax.fill_between(centers, mean - std, mean + std, alpha=0.2, color=line.get_color()) + if "reference" in r.payload: + ref = r.payload["reference"] + mean, std = np.asarray(ref["mean"]), np.asarray(ref["std"]) + color = _ref_color() + ax.plot(centers, mean, label=_REFERENCE_LABEL, color=color, linestyle="--") + ax.fill_between(centers, mean - std, mean + std, alpha=0.2, color=color) + for i, (name, side) in enumerate(r.payload.get("series", {}).items()): + mean, std = np.asarray(side["mean"]), np.asarray(side["std"]) + color = ps.get_color(i) + ax.plot(centers, mean, label=name, color=color) + ax.fill_between(centers, mean - std, mean + std, alpha=0.2, color=color) ax.set_xlabel(r.xlabel) ax.set_ylabel(r.payload.get("ylabel", "mean deposited energy [MeV]")) ps.style_legend(ax, title="source") @@ -195,10 +235,19 @@ def _render_profile(r: Reduced, params: dict): def _render_bar(r: Reduced, params: dict): labels = r.payload["labels"] x = np.arange(len(labels)) - width = 0.4 + series = r.payload.get("series", {}) + has_ref = "reference" in r.payload + n_bars = len(series) + (1 if has_ref else 0) + width = 0.8 / max(n_bars, 1) + offsets = np.linspace(-0.4 + width / 2, 0.4 - width / 2, n_bars) 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"]) - ax.bar(x + width / 2, r.payload["rollout"], width, label=_SERIES_LABELS["rollout"]) + idx = 0 + if has_ref: + ax.bar(x + offsets[idx], r.payload["reference"], width, label=_REFERENCE_LABEL, color=_ref_color()) + idx += 1 + for i, (name, vals) in enumerate(series.items()): + ax.bar(x + offsets[idx], vals, width, label=name, color=ps.get_color(i)) + idx += 1 ax.set_xticks(x) ax.set_xticklabels(labels, rotation=45, ha="right") ax.set_ylabel(r.payload.get("ylabel", "value")) @@ -207,97 +256,137 @@ def _render_bar(r: Reduced, params: dict): def _render_router_gating(r: Reduced, params: dict): - n_experts = r.payload["n_experts"] + series = r.payload.get("series", {}) + names = list(series) log_x = r.payload.get("log_x", False) - fig, axes = ps.new_figure("slide-16x9", title=r.title, params=params, nrows=1, ncols=2, squeeze=False) - flat = axes.ravel() - for ax, key in zip(flat, ("rollout", "reference")): - side = r.payload.get(key, {}) - centers = np.asarray(side.get("centers", [])) - means = np.asarray(side.get("means", [])) - if len(centers) and means.size: - cum = np.zeros(len(centers)) - for i in range(n_experts): - ax.fill_between(centers, cum, cum + means[:, i], alpha=0.7, label=f"expert {i}") - cum = cum + means[:, i] - if log_x: - ax.set_xscale("log") - ax.set_ylim(0, 1) - ax.set_title(_SERIES_LABELS[key], fontsize=8) - ax.set_xlabel(r.xlabel) - flat[0].set_ylabel("mean gate weight") - ps.style_legend(flat[0], title=f"{r.payload.get('router_type', '')} router") + fig, axes = ps.new_figure("slide-16x9", title=r.title, params=params, nrows=len(names), ncols=2, squeeze=False) + for row, name in enumerate(names): + entry = series[name] + n_experts = entry["n_experts"] + for col, key in enumerate(("rollout", "reference")): + ax = axes[row, col] + side = entry.get(key, {}) + centers = np.asarray(side.get("centers", [])) + means = np.asarray(side.get("means", [])) + if len(centers) and means.size: + cum = np.zeros(len(centers)) + for i in range(n_experts): + ax.fill_between(centers, cum, cum + means[:, i], alpha=0.7, label=f"expert {i}") + cum = cum + means[:, i] + if log_x: + ax.set_xscale("log") + ax.set_ylim(0, 1) + panel_label = _REFERENCE_LABEL if key == "reference" else "rollout" + ax.set_title(f"{name} — {panel_label}", fontsize=8) + if row == len(names) - 1: + ax.set_xlabel(r.xlabel) + axes[row, 0].set_ylabel("mean gate weight") + if names: + ps.style_legend(axes[0, 0], title=f"{series[names[0]]['router_type']} router") return fig def _render_router_share(r: Reduced, params: dict): - categories = r.payload["categories"] - n_experts = r.payload["n_experts"] - x = np.arange(len(categories)) - present = [k for k in ("rollout", "reference") if k in r.payload] - fig, axes = ps.new_figure( - "slide-16x9", - title=r.title, - params=params, - nrows=1, - ncols=len(present), - squeeze=False, - ) - flat = axes.ravel() - for ax, key in zip(flat, present): - side = r.payload[key] - shares = np.array([side[c] for c in categories]) # (n_cat, n_experts) - bottom = np.zeros(len(categories)) - for i in range(n_experts): - ax.bar(x, shares[:, i], bottom=bottom, label=f"expert {i}") - bottom += shares[:, i] - ax.set_xticks(x) - ax.set_xticklabels(categories, rotation=45, ha="right") - ax.set_ylim(0, 1) - ax.set_title(_SERIES_LABELS[key], fontsize=8) - flat[0].set_ylabel("share of rows dispatched to expert") - ps.style_legend(flat[0], title=f"{r.payload.get('router_type', '')} router") + series = r.payload.get("series", {}) + names = list(series) + present: tuple[str, ...] = ("rollout", "reference") + if names: + present = tuple(k for k in ("rollout", "reference") if k in series[names[0]]) + ncols = max(len(present), 1) + fig, axes = ps.new_figure("slide-16x9", title=r.title, params=params, nrows=len(names), ncols=ncols, squeeze=False) + for row, name in enumerate(names): + entry = series[name] + n_experts = entry["n_experts"] + cats = entry["categories"] + x = np.arange(len(cats)) + for col, key in enumerate(present): + ax = axes[row, col] + side = entry.get(key) + if side is not None: + shares = np.array([side[c] for c in cats]) # (n_cat, n_experts) + bottom = np.zeros(len(cats)) + for i in range(n_experts): + ax.bar(x, shares[:, i], bottom=bottom, label=f"expert {i}") + bottom += shares[:, i] + ax.set_xticks(x) + ax.set_xticklabels(cats, rotation=45, ha="right") + ax.set_ylim(0, 1) + panel_label = _REFERENCE_LABEL if key == "reference" else "rollout" + ax.set_title(f"{name} — {panel_label}", fontsize=8) + axes[row, 0].set_ylabel("share of rows dispatched to expert") + if names: + ps.style_legend(axes[0, 0], title=f"{series[names[0]]['router_type']} router") return fig def _render_router_specialization(r: Reduced, params: dict): fig, ax = ps.new_figure("thesis-single", title=r.title, params=params) - for key in ("reference", "rollout"): - side = r.payload.get(key) - if side and side["centers"]: - ax.plot(side["centers"], side["score"], label=_SERIES_LABELS[key], marker="o", markersize=3) - chance = r.payload.get("chance_level") - if chance is not None: - ax.axhline(chance, linestyle="--", color="gray", label="chance level (1/n_experts)") + series = r.payload.get("series", {}) + chance_levels: set[float] = set() + for i, (name, entry) in enumerate(series.items()): + color = ps.get_color(i) + if entry.get("chance_level") is not None: + chance_levels.add(entry["chance_level"]) + for key, linestyle, label in ( + ("rollout", "-", name), + ("reference", "--", f"{name} ({_REFERENCE_LABEL})"), + ): + side = entry.get(key) + if side and side["centers"]: + ax.plot( + side["centers"], + side["score"], + label=label, + color=color, + linestyle=linestyle, + marker="o", + markersize=3, + ) + for lvl in sorted(chance_levels): + ax.axhline(lvl, linestyle=":", color="gray") if r.payload.get("log_x"): ax.set_xscale("log") ax.set_ylim(0, 1) ax.set_xlabel(r.xlabel) ax.set_ylabel("max gate weight") - ps.style_legend(ax, title=f"{r.payload.get('router_type', '')} router") + ps.style_legend(ax, title="router") return fig def _render_heatmap(r: Reduced, params: dict): - mat = np.asarray(r.payload["matrix"], dtype=float) + series = r.payload["series"] row_labels = r.payload["row_labels"] col_labels = r.payload["col_labels"] - fig, ax = ps.new_figure("thesis-single", title=r.title, params=params) - im = ax.imshow( - mat, - origin="upper", - aspect="auto", - cmap=r.payload.get("cmap", "viridis"), - vmin=r.payload.get("vmin"), - vmax=r.payload.get("vmax"), + names = list(series) + fig, axes = ps.new_figure( + "slide-16x9" if len(names) > 1 else "thesis-single", + title=r.title, + params=params, + nrows=1, + ncols=len(names), + squeeze=False, ) - ax.set_xticks(range(len(col_labels))) - ax.set_xticklabels(col_labels, rotation=45, ha="right") - ax.set_yticks(range(len(row_labels))) - ax.set_yticklabels(row_labels) - ax.set_xlabel(r.xlabel) - ax.set_ylabel(r.payload.get("ylabel", "")) - fig.colorbar(im, ax=ax, label=r.payload.get("cbar_label", "value")) + flat = axes.ravel() + im = None + for ax, name in zip(flat, names): + mat = np.asarray(series[name], dtype=float) + im = ax.imshow( + mat, + origin="upper", + aspect="auto", + cmap=r.payload.get("cmap", "viridis"), + vmin=r.payload.get("vmin"), + vmax=r.payload.get("vmax"), + ) + ax.set_xticks(range(len(col_labels))) + ax.set_xticklabels(col_labels, rotation=45, ha="right") + ax.set_yticks(range(len(row_labels))) + ax.set_yticklabels(row_labels) + ax.set_xlabel(r.xlabel) + if len(names) > 1: + ax.set_title(name, fontsize=8) + flat[0].set_ylabel(r.payload.get("ylabel", "")) + fig.colorbar(im, ax=list(flat), label=r.payload.get("cbar_label", "value")) return fig @@ -392,7 +481,7 @@ def render_all( yaml.safe_dump( { "title": run_meta.get("title", "GIANT rollout analysis"), - "description": "Autoregressive rollout compared against held-out Geant4 reference steps.", + "description": "Autoregressive rollout(s) compared against a held-out Geant4 reference steps file.", "experiment": "GIANT", "parameters": {k: v for k, v in run_meta.items() if k != "title"}, }, @@ -425,8 +514,7 @@ def render_run(run_dir: str | Path, *, run_gallery: bool = False) -> list[Path]: meta = RunMeta.load(run_dir / "run_meta.json") run_meta = { "title": meta.title, - "rollout": meta.rollout, "reference": meta.reference, - **meta.plot_meta, + "rollouts": {ro["name"]: ro["plot_meta"] for ro in meta.rollouts}, } return render_all(run_dir / "reduced", run_dir / "plots", run_meta, run_gallery=run_gallery) diff --git a/giant/analysis/router_gating.py b/giant/analysis/router_gating.py index 8df6225..e82de47 100644 --- a/giant/analysis/router_gating.py +++ b/giant/analysis/router_gating.py @@ -35,6 +35,7 @@ from giant.analysis.reduced import Reduced if TYPE_CHECKING: import torch + from giant.analysis.sources import RolloutSide from giant.data.transforms import Normalizer _SAMPLE_ROWS = 200_000 @@ -218,47 +219,46 @@ def _unavailable(spec_id: str) -> Reduced: ) -def compute_router_gating( - checkpoint: str | Path | None, - r_phys: pl.LazyFrame, - t_phys: pl.LazyFrame, - seed: int = 0, -) -> Reduced: - """`Reduced` for the router-gating figure, or an explanatory note if n/a.""" +def _gating_entry(checkpoint: str | Path | None, r_phys: pl.LazyFrame, t_phys: pl.LazyFrame, seed: int) -> dict | None: + """One rollout's ``router_gating`` panel data, or ``None`` if not a MoE checkpoint.""" handle = load_router(checkpoint) if checkpoint else None if handle is None: - return _unavailable("router_gating") - + return None sides: dict[str, dict] = {} for name, lf in (("rollout", r_phys), ("reference", t_phys)): df = _subsample(lf, _SAMPLE_ROWS, seed) df, gate = _gate_for_df(handle, df) x = df["pre_E"].to_numpy() sides[name] = _quantile_bins(x, gate, _N_BINS) if len(x) else {"centers": [], "means": []} + return {"router_type": handle.router_type, "n_experts": handle.router.n_experts, **sides} + +def compute_router_gating(rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, seed: int = 0) -> Reduced: + """`Reduced` for the router-gating figure: one panel-pair per rollout with + an enabled MoE router, or an explanatory note if none of them have one.""" + series = {} + for name, rs in rollouts.items(): + entry = _gating_entry(rs.checkpoint, rs.phys, t_phys, seed) + if entry is not None: + series[name] = entry + if not series: + return _unavailable("router_gating") return Reduced( id="router_gating", family="model", kind="router_gating", title=_TITLES["router_gating"], xlabel="pre-step energy [MeV]", - payload={ - "router_type": handle.router_type, - "n_experts": handle.router.n_experts, - "log_x": True, - **sides, - }, + payload={"series": series, "log_x": True}, ) -def compute_router_specialization( - checkpoint: str | Path | None, - r_phys: pl.LazyFrame, - t_phys: pl.LazyFrame, - seed: int = 0, -) -> Reduced: - """Scalar specialization trend: max gate weight vs energy, per side. +def _specialization_entry( + checkpoint: str | Path | None, r_phys: pl.LazyFrame, t_phys: pl.LazyFrame, seed: int +) -> dict | None: + """One rollout's ``router_specialization`` curve data, or ``None`` if not a MoE checkpoint. + Scalar specialization trend: max gate weight vs energy, per side. Summarizes `router_gating`'s full per-expert stacked area into one curve — the routing plan's own "how sharp is the boundary here" number (1/n_experts = uniform/no specialization, 1.0 = one expert fully owns that energy). Same @@ -268,8 +268,7 @@ def compute_router_specialization( """ handle = load_router(checkpoint) if checkpoint else None if handle is None: - return _unavailable("router_specialization") - + return None sides: dict[str, dict] = {} for name, lf in (("rollout", r_phys), ("reference", t_phys)): df = _subsample(lf, _SAMPLE_ROWS, seed) @@ -282,35 +281,41 @@ def compute_router_specialization( sides[name] = {"centers": binned["centers"], "score": score} else: sides[name] = {"centers": [], "score": []} + return { + "router_type": handle.router_type, + "n_experts": handle.router.n_experts, + "chance_level": 1.0 / handle.router.n_experts, + **sides, + } + +def compute_router_specialization(rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, seed: int = 0) -> Reduced: + """`Reduced` for the router-specialization figure, one curve per rollout with + an enabled MoE router (see `_specialization_entry`).""" + series = {} + for name, rs in rollouts.items(): + entry = _specialization_entry(rs.checkpoint, rs.phys, t_phys, seed) + if entry is not None: + series[name] = entry + if not series: + return _unavailable("router_specialization") return Reduced( id="router_specialization", family="model", kind="router_specialization", title=_TITLES["router_specialization"], xlabel="pre-step energy [MeV]", - payload={ - "router_type": handle.router_type, - "n_experts": handle.router.n_experts, - "log_x": True, - "chance_level": 1.0 / handle.router.n_experts, - **sides, - }, + payload={"series": series, "log_x": True}, ) -def compute_router_share_by_pdg( - checkpoint: str | Path | None, - r_phys: pl.LazyFrame, - t_phys: pl.LazyFrame, - top_pdgs: list[int], - seed: int = 0, -) -> Reduced: - """Stacked-bar share of each particle species dispatched to each expert.""" +def _share_by_pdg_entry( + checkpoint: str | Path | None, r_phys: pl.LazyFrame, t_phys: pl.LazyFrame, top_pdgs: list[int], seed: int +) -> dict | None: + """One rollout's ``router_share_by_pdg`` panel-pair data, or ``None`` if not a MoE checkpoint.""" handle = load_router(checkpoint) if checkpoint else None if handle is None: - return _unavailable("router_share_by_pdg") - + return None labels = [pdg_label(p) for p in top_pdgs] sides: dict[str, dict] = {} for name, lf in (("rollout", r_phys), ("reference", t_phys)): @@ -322,41 +327,36 @@ def compute_router_share_by_pdg( else: shares = {str(p): [0.0] * handle.router.n_experts for p in top_pdgs} sides[name] = {labels[i]: shares[str(p)] for i, p in enumerate(top_pdgs)} + return {"router_type": handle.router_type, "n_experts": handle.router.n_experts, "categories": labels, **sides} + +def compute_router_share_by_pdg( + rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, top_pdgs: list[int], seed: int = 0 +) -> Reduced: + """`Reduced` for the router expert-share-by-species figure, one panel-pair + per rollout with an enabled MoE router.""" + series = {} + for name, rs in rollouts.items(): + entry = _share_by_pdg_entry(rs.checkpoint, rs.phys, t_phys, top_pdgs, seed) + if entry is not None: + series[name] = entry + if not series: + return _unavailable("router_share_by_pdg") return Reduced( id="router_share_by_pdg", family="model", kind="router_share", title=_TITLES["router_share_by_pdg"], xlabel="particle species", - payload={ - "router_type": handle.router_type, - "n_experts": handle.router.n_experts, - "categories": labels, - **sides, - }, + payload={"series": series}, ) -def compute_router_share_by_process( - checkpoint: str | Path | None, - t_phys: pl.LazyFrame, - seed: int = 0, - top_k: int = _TOP_K_PROCESS, -) -> Reduced: - """Stacked-bar share of each physics process dispatched to each expert. - - Reference-only: ``process`` is the true post-step physics process — a - label the rollout side has no equivalent of (see - `giant.model.network.ProcessRouter`, which predicts it from pre-step - conditioning alone, never observes it at eval time). This plot instead - checks *after the fact*, on real data, how well the router's conditioning - -based dispatch lines up with the true process. - """ +def _share_by_process_entry(checkpoint: str | Path | None, t_phys: pl.LazyFrame, seed: int, top_k: int) -> dict | None: + """One rollout checkpoint's ``router_share_by_process`` panel data (reference-only), or ``None`` if not MoE.""" handle = load_router(checkpoint) if checkpoint else None if handle is None: - return _unavailable("router_share_by_process") - + return None df = _subsample(t_phys, _SAMPLE_ROWS, seed, extra_cols=("process",)) df, gate = _gate_for_df(handle, df) if len(df): @@ -366,17 +366,39 @@ def compute_router_share_by_process( shares = _top1_shares(df["process"].to_numpy(), idx, order, handle.router.n_experts) else: order, shares = [], {} + return { + "router_type": handle.router_type, + "n_experts": handle.router.n_experts, + "categories": order, + "reference": {p: shares[p] for p in order}, + } + +def compute_router_share_by_process( + rollouts: dict[str, "RolloutSide"], t_phys: pl.LazyFrame, seed: int = 0, top_k: int = _TOP_K_PROCESS +) -> Reduced: + """Stacked-bar share of each physics process dispatched to each expert, one + panel per rollout checkpoint with an enabled MoE router. + + Reference-only: ``process`` is the true post-step physics process — a + label the rollout side has no equivalent of (see + `giant.model.network.ProcessRouter`, which predicts it from pre-step + conditioning alone, never observes it at eval time). This plot instead + checks *after the fact*, on real data, how well each checkpoint's router + -based dispatch lines up with the true process. + """ + series = {} + for name, rs in rollouts.items(): + entry = _share_by_process_entry(rs.checkpoint, t_phys, seed, top_k) + if entry is not None: + series[name] = entry + if not series: + return _unavailable("router_share_by_process") return Reduced( id="router_share_by_process", family="model", kind="router_share", title=_TITLES["router_share_by_process"], xlabel="physics process", - payload={ - "router_type": handle.router_type, - "n_experts": handle.router.n_experts, - "categories": order, - "reference": {p: shares[p] for p in order}, - }, + payload={"series": series}, ) diff --git a/giant/analysis/sources.py b/giant/analysis/sources.py index 250a8e5..15b496a 100644 --- a/giant/analysis/sources.py +++ b/giant/analysis/sources.py @@ -1,7 +1,12 @@ -"""Canonical world-frame LazyFrame builders for the two sides of a comparison. +"""Canonical world-frame LazyFrame builders for the two kinds of comparison input. -The analysis compares one autoregressive ``giant rollout`` (the *generated* side) -against a raw miniCaloSim steps file (the *reference* / real side). Both carry a +The analysis compares one or more autoregressive ``giant rollout`` runs (the +*generated* side — one named series each, see ``RolloutSpec``) against a single +raw miniCaloSim steps file shared by all of them (the *reference* / real side). +Every rollout is the same *kind* of file regardless of how many there are, so +``Side`` stays binary: it describes a file's schema (rollout column layout + +synthetic-termination rows + per-track secondary view, vs. reference +``sec_*_list`` columns), not series identity. Both kinds carry a **shared world-frame physical column subset** under identical names, so no renaming or coordinate decode is needed — everything is already in world-frame mm / MeV: @@ -26,6 +31,7 @@ HTCondor workers that have no LaTeX toolchain. from __future__ import annotations +from dataclasses import dataclass from enum import Enum from pathlib import Path @@ -82,12 +88,44 @@ SYNTHETIC_TERMINATION_REASONS: frozenset[str] = frozenset( class Side(str, Enum): - """Which of the two comparison inputs a file is.""" + """Which of the two comparison-input *kinds* a file is.""" rollout = "rollout" reference = "reference" +@dataclass +class RolloutSpec: + """One named rollout input, as fed to ``build_context``/``Bundle.open``. + + ``name`` is the series' identity throughout the rest of the pipeline (a + plot's ``payload["series"]`` key, a figure's legend label, its color) — + resolved once in ``condor.load_rollout_yamls`` from ``--label`` or the + YAML stem, then threaded through unchanged. ``checkpoint`` / + ``type_embedding_l1_dist`` are only used by the router/type-embedding + diagnostics (``catalog.py``'s ``chunkable=False`` specs). + """ + + name: str + source: str | Path | pl.LazyFrame + checkpoint: str | None = None + type_embedding_l1_dist: dict | None = None + + +@dataclass +class RolloutSide: + """One rollout's opened frames + per-checkpoint diagnostic inputs (``catalog.Bundle.rollouts`` value).""" + + all: pl.LazyFrame # rollout, all rows (incl. synthetic termination rows) + phys: pl.LazyFrame # rollout, physical steps only + checkpoint: str | None = None # from the rollout YAML; router_gating only + # Diagnostic pre-aggregated at rollout time (giant.rollout. + # L1DistCollector.summary()) — from the rollout YAML, type_embedding_l1_distance + # only. Unlike checkpoint, this needs no live model: it's already a + # finished histogram, just passed through. + type_embedding_l1_dist: dict | None = None + + def _check_rollout_metadata(path: Path) -> None: """Raise if ``path`` carries coord metadata that isn't the rollout tag. diff --git a/giant/analysis/type_embedding_distance.py b/giant/analysis/type_embedding_distance.py index 8304a2a..9ad4383 100644 --- a/giant/analysis/type_embedding_distance.py +++ b/giant/analysis/type_embedding_distance.py @@ -20,26 +20,37 @@ redesign exists to fix. from __future__ import annotations +from typing import TYPE_CHECKING + from giant.analysis.reduced import Reduced +if TYPE_CHECKING: + from giant.analysis.sources import RolloutSide + _NOTE_NOT_APPLICABLE = ( - "not applicable: this rollout's checkpoint doesn't use " + "not applicable: none of these rollouts' checkpoints use " "stage2_model.particle_type.target='embedding' (or generated no " "secondaries), so giant rollout recorded no type_embedding_l1_dist " - "diagnostic in its YAML sidecar" + "diagnostic in their YAML sidecar" ) -def compute_type_embedding_l1_distance(l1_dist: dict | None) -> Reduced: - """`Reduced` for the type-embedding-distance figure, or an explanatory - note if this checkpoint never populated the diagnostic. +def compute_type_embedding_l1_distance(rollouts: dict[str, "RolloutSide"]) -> Reduced: + """`Reduced` for the type-embedding-distance figure: one series per rollout + whose checkpoint populated the diagnostic, or an explanatory note if none did. - `l1_dist`: `giant.rollout.L1DistCollector.summary()`'s dict, as recorded - in the rollout YAML's `type_embedding_l1_dist` key (`Bundle. - type_embedding_l1_dist`) — `{"n", "mean", "std", "min", "max", - "hist_edges", "hist_counts"}`. + Each rollout's `RolloutSide.type_embedding_l1_dist` is + `giant.rollout.L1DistCollector.summary()`'s dict, as recorded in that + rollout's YAML `type_embedding_l1_dist` key — `{"n", "mean", "std", + "min", "max", "hist_edges", "hist_counts"}`. Every collector uses the + same fixed log-spaced edges (`L1DistCollector.__init__`'s defaults, never + overridden — see `giant/cli.py`'s rollout command), so it's safe to plot + every rollout's counts against the first one's edges. """ - if l1_dist is None: + entries = { + name: rs.type_embedding_l1_dist for name, rs in rollouts.items() if rs.type_embedding_l1_dist is not None + } + if not entries: return Reduced( id="type_embedding_l1_distance", family="model", @@ -49,6 +60,11 @@ def compute_type_embedding_l1_distance(l1_dist: dict | None) -> Reduced: payload={"note": _NOTE_NOT_APPLICABLE}, ) + edges = next(iter(entries.values()))["hist_edges"] + notes = [ + f"{name}: n={d['n']:,} mean={d['mean']:.4g} std={d['std']:.4g} min={d['min']:.4g} max={d['max']:.4g}" + for name, d in entries.items() + ] return Reduced( id="type_embedding_l1_distance", family="model", @@ -56,15 +72,10 @@ def compute_type_embedding_l1_distance(l1_dist: dict | None) -> Reduced: title="Secondary-type embedding L1 distance (predicted vector -> nearest PDG row)", xlabel="L1 distance", payload={ - "edges": l1_dist["hist_edges"], - "rollout": l1_dist["hist_counts"], + "edges": edges, + "series": {name: d["hist_counts"] for name, d in entries.items()}, "log_y": True, "log_x": True, - "note": ( - f"n={l1_dist['n']:,} mean={l1_dist['mean']:.4g} " - f"std={l1_dist['std']:.4g} min={l1_dist['min']:.4g} " - f"max={l1_dist['max']:.4g}; rollout only, no reference " - "concept for a raw pre-decode vector" - ), + "note": "; ".join(notes) + "; rollout only, no reference concept for a raw pre-decode vector", }, ) diff --git a/giant/cli.py b/giant/cli.py index c2add9b..9f7433c 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -1556,10 +1556,23 @@ app.add_typer(analyze_app, name="analyze") @analyze_app.command("prep") def analyze_prep( - rollout_yaml: Annotated[ - Path, - typer.Argument(help="giant rollout YAML sidecar (names the rollout + reference files)"), + rollout_yamls: Annotated[ + list[Path], + typer.Argument( + help="giant rollout YAML sidecar(s) (names the rollout + reference files). " + "Multiple compare N rollouts against one shared reference — every YAML must " + "name the same `dataset`." + ), ], + label: Annotated[ + Optional[list[str]], + typer.Option( + "--label", + help="Series name for a rollout YAML, positionally matched to it — give none, " + 'or exactly one per YAML. Defaults to the YAML stem (or "rollout" for a ' + "single YAML).", + ), + ] = None, run_dir: Annotated[ Optional[Path], typer.Option( @@ -1576,14 +1589,15 @@ def analyze_prep( typer.Option("--chunks", help="Split each plot's data into this many event_id chunks"), ] = 1, ) -> None: - """Read the rollout YAML → shared.json + run_meta.json in the run directory.""" + """Read the rollout YAML(s) → shared.json + run_meta.json in the run directory.""" from giant.analysis import prep path = prep( - rollout_yaml, + rollout_yamls, run_dir, n_chunks=chunks, default_base=Path.cwd() / "analysis_runs", + labels=label, n_energy_bins=n_energy_bins, n_marginal_bins=n_marginal_bins, top_k_pdg=top_k_pdg, @@ -1665,8 +1679,23 @@ def analyze_metrics( @analyze_app.command("submit") def analyze_submit( - rollout_yaml: Annotated[Path, typer.Argument(help="giant rollout YAML sidecar")], + rollout_yamls: Annotated[ + list[Path], + typer.Argument( + help="giant rollout YAML sidecar(s). Multiple compare N rollouts against one " + "shared reference — every YAML must name the same `dataset`." + ), + ], accounting_group: Annotated[str, typer.Option("--accounting-group")], + label: Annotated[ + Optional[list[str]], + typer.Option( + "--label", + help="Series name for a rollout YAML, positionally matched to it — give none, " + 'or exactly one per YAML. Defaults to the YAML stem (or "rollout" for a ' + "single YAML).", + ), + ] = None, run_dir: Annotated[ Optional[Path], typer.Option( @@ -1699,10 +1728,11 @@ def analyze_submit( from giant.analysis import SubmitConfig, prep, write_submit path = prep( - rollout_yaml, + rollout_yamls, run_dir, n_chunks=chunks, default_base=Path.cwd() / "analysis_runs", + labels=label, n_energy_bins=n_energy_bins, n_marginal_bins=n_marginal_bins, top_k_pdg=top_k_pdg, diff --git a/giant/tools/profile_analysis_costs.py b/giant/tools/profile_analysis_costs.py index 9a37f1b..28dc1af 100644 --- a/giant/tools/profile_analysis_costs.py +++ b/giant/tools/profile_analysis_costs.py @@ -28,6 +28,7 @@ import polars as pl from giant.analysis.catalog import catalog_ids, get_spec from giant.analysis.condor import compute_reduced from giant.analysis.context import build_context +from giant.analysis.sources import RolloutSpec # Row counts (per side) to benchmark at. Kept in local memory/CPU range so the # whole sweep finishes in about a minute; the fit is linear so it extrapolates @@ -165,11 +166,10 @@ def _time(spec_id: str, rollout: Path, reference: Path, shared: Path, out: Path) t0 = time.perf_counter() compute_reduced( spec_id, - rollout, + [{"name": "rollout", "path": str(rollout)}], reference, shared, out, - checkpoint=None, chunk_index=0, n_chunks=1, ) @@ -191,7 +191,7 @@ def main() -> None: shared = tmp_path / f"shared_{n_side}.json" ctx = build_context( - rollout, + [RolloutSpec(name="rollout", source=rollout)], reference, n_energy_bins=4, n_marginal_bins=50, diff --git a/tests/test_catalog.py b/tests/test_catalog.py index ec5a134..377ac05 100644 --- a/tests/test_catalog.py +++ b/tests/test_catalog.py @@ -14,12 +14,21 @@ from giant.analysis.catalog import ( _ks_statistic, ) from giant.analysis.context import Context, build_context +from giant.analysis.sources import RolloutSpec from tests.test_analysis_reduce import _reference_frame, _rollout_frame def _build_ctx() -> Context: r, t = _rollout_frame(), _reference_frame() - return build_context(r, t, n_energy_bins=2, n_marginal_bins=10, top_k_pdg=3, sample_rows=1000) + return build_context( + [RolloutSpec("rollout", r)], t, n_energy_bins=2, n_marginal_bins=10, top_k_pdg=3, sample_rows=1000 + ) + + +def _two_rollout_specs() -> list[RolloutSpec]: + # Two distinct rollout sources so multi-series merging/finalize code is + # exercised even though the underlying frame is the same fixture. + return [RolloutSpec("flow", _rollout_frame()), RolloutSpec("wgan", _rollout_frame())] @pytest.fixture(scope="module") @@ -27,9 +36,20 @@ def ctx() -> Context: return _build_ctx() +@pytest.fixture(scope="module") +def two_ctx() -> Context: + t = _reference_frame() + return build_context(_two_rollout_specs(), t, n_energy_bins=2, n_marginal_bins=10, top_k_pdg=3, sample_rows=1000) + + @pytest.fixture(scope="module") def bundle(ctx: Context) -> Bundle: - return Bundle.open(_rollout_frame(), _reference_frame(), ctx) + return Bundle.open([RolloutSpec("rollout", _rollout_frame())], _reference_frame(), ctx) + + +@pytest.fixture(scope="module") +def two_bundle(two_ctx: Context) -> Bundle: + return Bundle.open(_two_rollout_specs(), _reference_frame(), two_ctx) def test_catalog_ids_unique_and_nonempty(): @@ -64,46 +84,71 @@ def test_every_spec_computes_valid_reduced(bundle: Bundle): "unavailable", } assert r.title and r.xlabel - _validate_payload(r) + _validate_payload(r, ["rollout"]) -def _validate_payload(r) -> None: +def test_every_spec_computes_valid_reduced_with_two_rollouts(two_bundle: Bundle): + for spec in build_catalog(): + r = spec.finalize([spec.compute_partial(two_bundle)], two_bundle.ctx) + assert r.id == spec.id + _validate_payload(r, ["flow", "wgan"]) + + +def _validate_payload(r, names: list[str]) -> None: p = r.payload if r.kind == "overlay_hist": n = len(p["edges"]) - 1 - assert len(p["rollout"]) == n and len(p["reference"]) == n + assert list(p["series"]) == names + for v in p["series"].values(): + assert len(v) == n + assert len(p["reference"]) == n elif r.kind == "single_hist": - assert len(p["rollout"]) == len(p["edges"]) - 1 + assert list(p["series"]) == names + for v in p["series"].values(): + assert len(v) == len(p["edges"]) - 1 elif r.kind == "grouped_hist": n = len(p["edges"]) - 1 assert p["groups"], "grouped hist must have at least one group" for g in p["groups"].values(): - assert len(g["rollout"]) == n and len(g["reference"]) == n + assert list(g["series"]) == names + for v in g["series"].values(): + assert len(v) == n + assert len(g["reference"]) == n elif r.kind == "profile": n = len(p["edges"]) - 1 - for k in ("rollout_mean", "rollout_std", "reference_mean", "reference_std"): - assert len(p[k]) == n + assert list(p["series"]) == names + for side in p["series"].values(): + assert len(side["mean"]) == n and len(side["std"]) == n + assert len(p["reference"]["mean"]) == n and len(p["reference"]["std"]) == n elif r.kind == "bar": - assert len(p["labels"]) == len(p["rollout"]) == len(p["reference"]) + assert list(p["series"]) == names + for v in p["series"].values(): + assert len(p["labels"]) == len(v) + assert len(p["labels"]) == len(p["reference"]) elif r.kind == "unavailable": assert p["note"] elif r.kind == "router_gating": - for side in ("rollout", "reference"): - if side in p: - assert len(p[side]["centers"]) == len(p[side]["means"]) - elif r.kind == "router_share": - for cat in p["categories"]: + for entry in p["series"].values(): for side in ("rollout", "reference"): - if side in p: - assert cat in p[side] + if side in entry: + assert len(entry[side]["centers"]) == len(entry[side]["means"]) + elif r.kind == "router_share": + for entry in p["series"].values(): + for cat in entry["categories"]: + for side in ("rollout", "reference"): + if side in entry: + assert cat in entry[side] elif r.kind == "router_specialization": - for side in ("rollout", "reference"): - if side in p: - assert len(p[side]["centers"]) == len(p[side]["score"]) + for entry in p["series"].values(): + for side in ("rollout", "reference"): + if side in entry: + assert len(entry[side]["centers"]) == len(entry[side]["score"]) elif r.kind == "heatmap": - assert len(p["matrix"]) == len(p["row_labels"]) - for row in p["matrix"]: - assert len(row) == len(p["col_labels"]) + assert list(p["series"]) == names + for mat in p["series"].values(): + assert len(mat) == len(p["row_labels"]) + for row in mat: + assert len(row) == len(p["col_labels"]) # --------------------------------------------------------------------------- @@ -150,20 +195,21 @@ def _assert_payload_close(a, b, path: str = "payload") -> None: @pytest.mark.parametrize("spec_id", _CHUNK_EQUIVALENCE_IDS) -def test_chunked_matches_unchunked(ctx: Context, spec_id: str): +def test_chunked_matches_unchunked(two_ctx: Context, spec_id: str): """A plot computed over N event-disjoint chunks then merged must equal the same plot computed in one unchunked pass — the core chunking correctness - guarantee (see the analysis-rollout-plots chunking plan).""" + guarantee (see the analysis-rollout-plots chunking plan). Exercised with + two rollout series so the per-rollout merge path is covered too.""" spec: PlotSpec = get_spec(spec_id) - r, t = _rollout_frame(), _reference_frame() + rollouts, t = _two_rollout_specs(), _reference_frame() - unchunked_bundle = Bundle.open(r, t, ctx) - unchunked = spec.finalize([spec.compute_partial(unchunked_bundle)], ctx) + unchunked_bundle = Bundle.open(rollouts, t, two_ctx) + unchunked = spec.finalize([spec.compute_partial(unchunked_bundle)], two_ctx) # 4 chunks over only 2 distinct event_ids also exercises empty chunks. n_chunks = 4 if spec.chunkable else 1 - parts = [spec.compute_partial(Bundle.open(r, t, ctx, chunk=(k, n_chunks))) for k in range(n_chunks)] - chunked = spec.finalize(parts, ctx) + parts = [spec.compute_partial(Bundle.open(rollouts, t, two_ctx, chunk=(k, n_chunks))) for k in range(n_chunks)] + chunked = spec.finalize(parts, two_ctx) assert chunked.id == unchunked.id assert chunked.kind == unchunked.kind @@ -196,6 +242,14 @@ def test_integer_confusion_caps_pathological_outliers(): assert mat.sum() == 2 +def test_integer_confusion_explicit_cap_overrides_local_range(): + # Even though this pair's own max is 1, an explicit shared cap forces a + # wider (and so cross-rollout-consistent) label set. + labels, mat = _integer_confusion(np.array([1, 1]), np.array([0, 1]), cap=3) + assert labels == ["0", "1", "2", "3+"] + assert mat.shape == (4, 4) + + def test_containment_depths_simple_ramp(): # one event, edep concentrated in the first bin -> 90%/95% containment # depth is the first bin's right edge; a zero-energy event is dropped. @@ -209,4 +263,13 @@ def test_n_sec_confusion_spec(bundle): spec = get_spec("n_sec_confusion") r = spec.finalize([spec.compute_partial(bundle)], bundle.ctx) assert r.payload["row_labels"] == r.payload["col_labels"] == ["0", "1+"] - assert r.payload["matrix"] == [[0, 0], [1, 1]] + assert r.payload["series"]["rollout"] == [[0, 0], [1, 1]] + + +def test_n_sec_confusion_shares_one_cap_across_rollouts(two_bundle): + spec = get_spec("n_sec_confusion") + r = spec.finalize([spec.compute_partial(two_bundle)], two_bundle.ctx) + assert list(r.payload["series"]) == ["flow", "wgan"] + # both rollouts share the same fixture data here, so their matrices (and + # the shared label set) must be identical. + assert r.payload["series"]["flow"] == r.payload["series"]["wgan"] diff --git a/tests/test_condor.py b/tests/test_condor.py index 02912ed..e35356f 100644 --- a/tests/test_condor.py +++ b/tests/test_condor.py @@ -1,4 +1,4 @@ -"""Tests for the rollout-YAML → run-directory flow, compute, and submit.""" +"""Tests for the rollout-YAML(s) → run-directory flow, compute, and submit.""" from __future__ import annotations @@ -17,6 +17,7 @@ from giant.analysis import ( compute_reduced, derive_run_dir, load_rollout_yaml, + load_rollout_yamls, merge_one, prep, write_submit, @@ -28,13 +29,17 @@ from giant.constants import PREDICT_COORD_METADATA_KEY, ROLLOUT_COORD_VALUE from tests.test_analysis_reduce import _reference_frame, _rollout_frame +def _write_rollout(path: Path) -> None: + tbl = _rollout_frame().collect().to_arrow() + tbl = tbl.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE}) + pq.write_table(tbl, path) + + def _write_inputs(tmp_path: Path) -> Path: """Materialize rollout+reference parquet and a rollout YAML; return the YAML path.""" rollout = tmp_path / "rollout.parquet" reference = tmp_path / "reference.parquet" - tbl = _rollout_frame().collect().to_arrow() - tbl = tbl.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE}) - pq.write_table(tbl, rollout) + _write_rollout(rollout) _reference_frame().collect().write_parquet(reference) yaml_path = tmp_path / "run.yaml" @@ -54,6 +59,33 @@ def _write_inputs(tmp_path: Path) -> Path: return yaml_path +def _write_two_inputs(tmp_path: Path) -> tuple[Path, Path]: + """Two rollout YAMLs (distinct output files) sharing one reference file.""" + reference = tmp_path / "reference.parquet" + _reference_frame().collect().write_parquet(reference) + + paths = [] + for tag, pred_id in (("a", "aaaa1111ef"), ("b", "bbbb2222ef")): + rollout = tmp_path / f"rollout_{tag}.parquet" + _write_rollout(rollout) + yaml_path = tmp_path / f"run_{tag}.yaml" + yaml_path.write_text( + yaml.safe_dump( + { + "prediction_id": pred_id, + "output": str(rollout), + "dataset": str(reference), + "checkpoint": f"/ckpt/{tag}.pt", + "kind": "rollout", + "energy_cutoff": 0.1, + "steps": 10, + } + ) + ) + paths.append(yaml_path) + return paths[0], paths[1] + + def _fake_venv(repo_dir: Path) -> None: """Stand in for a `uv sync`'d venv: write_submit checks `.venv/bin/giant` exists.""" giant = repo_dir / ".venv" / "bin" / "giant" @@ -62,12 +94,13 @@ def _fake_venv(repo_dir: Path) -> None: giant.chmod(0o755) -def _prep(rollout_yaml: Path, run_dir: str | Path | None = None, chunks: int = 1) -> Path: +def _prep(rollout_yamls, run_dir: str | Path | None = None, chunks: int = 1, labels=None) -> Path: """``prep`` with small test-sized context bins/sampling.""" return prep( - rollout_yaml, + rollout_yamls, run_dir, n_chunks=chunks, + labels=labels, n_energy_bins=2, n_marginal_bins=8, top_k_pdg=3, @@ -82,39 +115,108 @@ def test_load_rollout_yaml_requires_paths(tmp_path: Path): load_rollout_yaml(bad) +def test_load_rollout_yamls_single_defaults_to_rollout_name(tmp_path: Path): + yaml_path = _write_inputs(tmp_path) + loaded, reference = load_rollout_yamls([yaml_path]) + assert [lr.name for lr in loaded] == ["rollout"] + assert reference.endswith("reference.parquet") + + +def test_load_rollout_yamls_multi_defaults_to_stem(tmp_path: Path): + a, b = _write_two_inputs(tmp_path) + loaded, _ = load_rollout_yamls([a, b]) + assert [lr.name for lr in loaded] == ["run_a", "run_b"] + + +def test_load_rollout_yamls_explicit_labels(tmp_path: Path): + a, b = _write_two_inputs(tmp_path) + loaded, _ = load_rollout_yamls([a, b], labels=["flow", "wgan"]) + assert [lr.name for lr in loaded] == ["flow", "wgan"] + + +def test_load_rollout_yamls_label_count_mismatch(tmp_path: Path): + a, b = _write_two_inputs(tmp_path) + with pytest.raises(ValueError, match="--label"): + load_rollout_yamls([a, b], labels=["only-one"]) + + +def test_load_rollout_yamls_rejects_duplicate_names(tmp_path: Path): + a, b = _write_two_inputs(tmp_path) + with pytest.raises(ValueError, match="collide"): + load_rollout_yamls([a, b], labels=["same", "same"]) + + +def test_load_rollout_yamls_rejects_mismatched_reference(tmp_path: Path): + a, _ = _write_two_inputs(tmp_path) + other_ref = tmp_path / "other_reference.parquet" + _reference_frame().collect().write_parquet(other_ref) + c = tmp_path / "run_c.yaml" + c.write_text( + yaml.safe_dump( + {"prediction_id": "cccc3333ef", "output": str(tmp_path / "rollout_c.parquet"), "dataset": str(other_ref)} + ) + ) + _write_rollout(tmp_path / "rollout_c.parquet") + with pytest.raises(ValueError, match="same reference"): + load_rollout_yamls([a, c]) + + def test_derive_run_dir_next_to_rollout(): y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"} - assert derive_run_dir(y) == Path("/data/analysis_abcd1234") - assert derive_run_dir(y, "/somewhere") == Path("/somewhere") + assert derive_run_dir([y]) == Path("/data/analysis_abcd1234") + assert derive_run_dir([y], "/somewhere") == Path("/somewhere") def test_derive_run_dir_default_base(): y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"} - assert derive_run_dir(y, default_base="/work/lbogner/giant2/analysis_runs") == Path( + assert derive_run_dir([y], default_base="/work/lbogner/giant2/analysis_runs") == Path( "/work/lbogner/giant2/analysis_runs/analysis_abcd1234" ) # an explicit run_dir still wins over default_base - assert derive_run_dir(y, "/somewhere", default_base="/other") == Path("/somewhere") + assert derive_run_dir([y], "/somewhere", default_base="/other") == Path("/somewhere") + + +def test_derive_run_dir_multi_rollout_joins_tags(): + ys = [{"output": f"/data/roll_{i}.parquet", "prediction_id": f"tag{i}xxxx", "dataset": "d"} for i in range(2)] + assert derive_run_dir(ys, default_base="/base") == Path("/base/analysis_tag0xxxx-tag1xxxx") + + +def test_derive_run_dir_many_rollouts_truncates_with_plus_count(): + ys = [{"output": f"/data/roll_{i}.parquet", "prediction_id": f"tag{i}xxxx", "dataset": "d"} for i in range(5)] + run_dir = derive_run_dir(ys, default_base="/base") + assert run_dir == Path("/base/analysis_tag0xxxx-tag1xxxx-tag2xxxx-plus2") def test_prep_lays_out_run_dir(tmp_path: Path): yaml_path = _write_inputs(tmp_path) - run_dir = _prep(yaml_path) + run_dir = _prep([yaml_path]) assert run_dir == tmp_path / "analysis_abcd1234" assert (run_dir / "shared.json").exists() ctx = Context.load(run_dir / "shared.json") assert set(ctx.var_ranges) == {"step_length", "edep", "delta_e", "post_E"} meta = RunMeta.load(run_dir / "run_meta.json") assert meta.reference.endswith("reference.parquet") - assert meta.plot_meta["checkpoint"] == "/ckpt/best.pt" + assert [ro["name"] for ro in meta.rollouts] == ["rollout"] + assert meta.rollouts[0]["plot_meta"]["checkpoint"] == "/ckpt/best.pt" assert "best.pt" in meta.title assert meta.n_chunks == 1 assert meta.rows_per_chunk == [meta.total_rows] # single chunk holds everything assert meta.total_rows == 8 # 5 rollout rows + 3 reference rows +def test_prep_multi_rollout_lays_out_run_dir(tmp_path: Path): + a, b = _write_two_inputs(tmp_path) + run_dir = _prep([a, b], labels=["flow", "wgan"]) + meta = RunMeta.load(run_dir / "run_meta.json") + assert [ro["name"] for ro in meta.rollouts] == ["flow", "wgan"] + assert meta.rollouts[0]["plot_meta"]["checkpoint"] == "/ckpt/a.pt" + assert meta.rollouts[1]["plot_meta"]["checkpoint"] == "/ckpt/b.pt" + # 5 rows from each rollout + 3 from the shared reference + assert meta.total_rows == 13 + + def test_prep_splits_rows_per_chunk(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path), chunks=2) + run_dir = _prep([_write_inputs(tmp_path)], chunks=2) meta = RunMeta.load(run_dir / "run_meta.json") assert len(meta.rows_per_chunk) == 2 assert sum(meta.rows_per_chunk) == meta.total_rows == 8 @@ -125,7 +227,7 @@ def test_reprep_clears_stale_partials_from_a_different_chunk_count(tmp_path: Pat partials on disk for merge_one to silently merge against the new context (they'd be keyed/sized for the old n_chunks).""" yaml_path = _write_inputs(tmp_path) - run_dir = _prep(yaml_path, chunks=2) + run_dir = _prep([yaml_path], chunks=2) compute_one("marginal_edep", run_dir, chunk_index=0) compute_one("marginal_edep", run_dir, chunk_index=1) stale = run_dir / "reduced_partial" / "marginal_edep__0.json" @@ -133,7 +235,7 @@ def test_reprep_clears_stale_partials_from_a_different_chunk_count(tmp_path: Pat (run_dir / "reduced").mkdir(exist_ok=True) (run_dir / "reduced" / "marginal_edep.json").write_text("{}") - _prep(yaml_path, run_dir, chunks=1) + _prep([yaml_path], run_dir, chunks=1) assert not stale.exists() assert not (run_dir / "reduced" / "marginal_edep.json").exists() @@ -141,20 +243,22 @@ def test_reprep_clears_stale_partials_from_a_different_chunk_count(tmp_path: Pat def test_compute_one_from_run_dir(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path)) + run_dir = _prep([_write_inputs(tmp_path)]) out = compute_one("marginal_edep", run_dir) assert out == run_dir / "reduced_partial" / "marginal_edep__0.json" partial = Partial.load(out) assert partial.id == "marginal_edep" and partial.chunk == 0 assert "r" in partial.data and "t" in partial.data + assert list(partial.data["r"]) == ["rollout"] def test_compute_reduced_explicit_paths(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path)) + run_dir = _prep([_write_inputs(tmp_path)]) meta = RunMeta.load(run_dir / "run_meta.json") + rollouts = [{"name": ro["name"], "path": ro["path"]} for ro in meta.rollouts] out = compute_reduced( "marginal_step_length", - meta.rollout, + rollouts, meta.reference, run_dir / "shared.json", tmp_path / "r.json", @@ -163,17 +267,17 @@ def test_compute_reduced_explicit_paths(tmp_path: Path): def test_merge_one_produces_reduced(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path)) + run_dir = _prep([_write_inputs(tmp_path)]) compute_one("marginal_edep", run_dir) out = merge_one("marginal_edep", run_dir) assert out == run_dir / "reduced" / "marginal_edep.json" reduced = Reduced.load(out) assert reduced.id == "marginal_edep" - assert len(reduced.payload["rollout"]) == len(reduced.payload["edges"]) - 1 + assert len(reduced.payload["series"]["rollout"]) == len(reduced.payload["edges"]) - 1 def test_merge_one_fails_loudly_on_missing_chunk(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path), chunks=2) + run_dir = _prep([_write_inputs(tmp_path)], chunks=2) compute_one("marginal_edep", run_dir, chunk_index=0) # chunk 1 never computed with pytest.raises(FileNotFoundError, match="missing chunk"): merge_one("marginal_edep", run_dir) @@ -182,11 +286,11 @@ def test_merge_one_fails_loudly_on_missing_chunk(tmp_path: Path): def test_chunked_compute_and_merge_matches_unchunked(tmp_path: Path): (tmp_path / "a").mkdir() (tmp_path / "b").mkdir() - unchunked_dir = _prep(_write_inputs(tmp_path / "a")) + unchunked_dir = _prep([_write_inputs(tmp_path / "a")]) compute_one("marginal_step_length", unchunked_dir) unchunked = Reduced.load(merge_one("marginal_step_length", unchunked_dir)) - chunked_dir = _prep(_write_inputs(tmp_path / "b"), chunks=2) + chunked_dir = _prep([_write_inputs(tmp_path / "b")], chunks=2) for k in range(2): compute_one("marginal_step_length", chunked_dir, chunk_index=k) chunked = Reduced.load(merge_one("marginal_step_length", chunked_dir)) @@ -194,14 +298,23 @@ def test_chunked_compute_and_merge_matches_unchunked(tmp_path: Path): assert chunked.payload == unchunked.payload +def test_two_rollout_compute_and_merge_produces_both_series(tmp_path: Path): + a, b = _write_two_inputs(tmp_path) + run_dir = _prep([a, b], labels=["flow", "wgan"]) + compute_one("marginal_edep", run_dir) + reduced = Reduced.load(merge_one("marginal_edep", run_dir)) + assert list(reduced.payload["series"]) == ["flow", "wgan"] + assert "reference" in reduced.payload + + def test_compute_reduced_rejects_out_of_range_chunk(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path)) # n_chunks=1 (default) + run_dir = _prep([_write_inputs(tmp_path)]) # n_chunks=1 (default) with pytest.raises(ValueError, match="out of range"): compute_one("marginal_edep", run_dir, chunk_index=1) def test_write_submit_description(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path)) + run_dir = _prep([_write_inputs(tmp_path)]) _fake_venv(tmp_path) cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path) txt = write_submit(cfg).read_text() @@ -223,7 +336,7 @@ def test_write_submit_description(tmp_path: Path): def test_write_submit_requires_synced_venv(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): - run_dir = _prep(_write_inputs(tmp_path)) + run_dir = _prep([_write_inputs(tmp_path)]) cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path) # No `giant` next to the (fake) active interpreter, so this falls through # to repo_dir/.venv/bin/giant, which _write_inputs/_prep also didn't create. @@ -233,7 +346,7 @@ def test_write_submit_requires_synced_venv(tmp_path: Path, monkeypatch: pytest.M def test_write_submit_remote_flag(tmp_path: Path): - run_dir = _prep(_write_inputs(tmp_path)) + run_dir = _prep([_write_inputs(tmp_path)]) _fake_venv(tmp_path) cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, remote=True) txt = write_submit(cfg).read_text() @@ -243,7 +356,7 @@ def test_write_submit_remote_flag(tmp_path: Path): def test_write_submit_chunks_respect_chunkable(tmp_path: Path): assert get_spec("router_gating").chunkable is False - run_dir = _prep(_write_inputs(tmp_path), chunks=4) + run_dir = _prep([_write_inputs(tmp_path)], chunks=4) _fake_venv(tmp_path) cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, n_chunks=4) write_submit(cfg) @@ -260,7 +373,7 @@ def test_write_submit_rejects_n_chunks_mismatch_with_run_meta(tmp_path: Path): with — RunMeta.rows_per_chunk is sized to the prepped value, so a mismatch would otherwise surface as a confusing IndexError deep inside _job_walltimes instead of a clear error here.""" - run_dir = _prep(_write_inputs(tmp_path), chunks=2) + run_dir = _prep([_write_inputs(tmp_path)], chunks=2) _fake_venv(tmp_path) cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, n_chunks=4) with pytest.raises(ValueError, match="n_chunks"): @@ -282,7 +395,7 @@ def test_write_submit_walltime_grows_with_chunk_rows(tmp_path: Path): """A chunked run's later job walltimes track that chunk's row count.""" from giant.analysis.runtime_estimate import estimate_runtime_s - run_dir = _prep(_write_inputs(tmp_path), chunks=2) + run_dir = _prep([_write_inputs(tmp_path)], chunks=2) meta = RunMeta.load(run_dir / "run_meta.json") _fake_venv(tmp_path) cfg = SubmitConfig(run_dir=run_dir, accounting_group="cms", repo_dir=tmp_path, n_chunks=2) diff --git a/tests/test_render.py b/tests/test_render.py index 064cb32..f55be37 100644 --- a/tests/test_render.py +++ b/tests/test_render.py @@ -25,41 +25,96 @@ def test_render_router_diagnostics_and_edge_cases(tmp_path: Path): reduced = [ Reduced( "rg", - "router", + "model", "router_gating", "Router gating", "pre-step energy [MeV]", { - "n_experts": 2, "log_x": True, - "router_type": "energy", - "rollout": { - "centers": [1.0, 10.0, 100.0], - "means": [[0.6, 0.4], [0.5, 0.5], [0.4, 0.6]], - }, - "reference": { - "centers": [1.0, 10.0, 100.0], - "means": [[0.55, 0.45], [0.5, 0.5], [0.45, 0.55]], + "series": { + "flow": { + "n_experts": 2, + "router_type": "energy", + "rollout": { + "centers": [1.0, 10.0, 100.0], + "means": [[0.6, 0.4], [0.5, 0.5], [0.4, 0.6]], + }, + "reference": { + "centers": [1.0, 10.0, 100.0], + "means": [[0.55, 0.45], [0.5, 0.5], [0.45, 0.55]], + }, + }, + "wgan": { + "n_experts": 2, + "router_type": "energy", + "rollout": {"centers": [1.0], "means": [[0.5, 0.5]]}, + "reference": {"centers": [1.0], "means": [[0.5, 0.5]]}, + }, }, }, ), Reduced( "rs", - "router", + "model", "router_share", "Router share", "species", { - "categories": ["e-", "gamma"], - "n_experts": 2, - "router_type": "energy", - "rollout": {"e-": [0.7, 0.3], "gamma": [0.2, 0.8]}, - "reference": {"e-": [0.6, 0.4], "gamma": [0.3, 0.7]}, + "series": { + "flow": { + "categories": ["e-", "gamma"], + "n_experts": 2, + "router_type": "energy", + "rollout": {"e-": [0.7, 0.3], "gamma": [0.2, 0.8]}, + "reference": {"e-": [0.6, 0.4], "gamma": [0.3, 0.7]}, + }, + }, + }, + ), + Reduced( + "rp", + "model", + "router_share", + "Router share by process (reference-only)", + "process", + { + "series": { + "flow": { + "categories": ["compt", "phot"], + "n_experts": 2, + "router_type": "energy", + "reference": {"compt": [0.4, 0.6], "phot": [0.9, 0.1]}, + }, + }, + }, + ), + Reduced( + "rz", + "model", + "router_specialization", + "Router specialization", + "pre-step energy [MeV]", + { + "log_x": True, + "series": { + "flow": { + "n_experts": 2, + "chance_level": 0.5, + "rollout": {"centers": [1.0, 10.0], "score": [0.6, 0.7]}, + "reference": {"centers": [1.0, 10.0], "score": [0.55, 0.65]}, + }, + "wgan": { + "n_experts": 4, + "chance_level": 0.25, + "rollout": {"centers": [1.0, 10.0], "score": [0.3, 0.4]}, + "reference": {"centers": [], "score": []}, + }, + }, }, ), Reduced( "ru", - "router", + "model", "unavailable", "Router unavailable", "x", @@ -73,7 +128,10 @@ def test_render_router_diagnostics_and_edge_cases(tmp_path: Path): "x", { "edges": [0, 1, 2], - "groups": {lbl: {"rollout": [1, 2], "reference": [2, 1]} for lbl in ("a", "b", "c", "d")}, + "groups": { + lbl: {"series": {"flow": [1, 2], "wgan": [2, 1]}, "reference": [2, 1]} + for lbl in ("a", "b", "c", "d") + }, "log_y": True, }, ), @@ -83,7 +141,22 @@ def test_render_router_diagnostics_and_edge_cases(tmp_path: Path): "single_hist", "Single (log-x)", "x", - {"edges": [1, 10, 100], "rollout": [5, 1], "log_x": True, "log_y": True}, + {"edges": [1, 10, 100], "series": {"flow": [5, 1], "wgan": [3, 2]}, "log_x": True, "log_y": True}, + ), + Reduced( + "hm", + "quality", + "heatmap", + "Distance summary (2 rollouts)", + "grouping axis", + { + "series": {"flow": [[0.1, 0.2], [0.3, 0.4]], "wgan": [[0.5, 0.6], [0.7, 0.8]]}, + "row_labels": ["step_length", "edep"], + "col_labels": ["overall", "energy"], + "cbar_label": "KS statistic", + "vmin": 0.0, + "vmax": 1.0, + }, ), ] try: @@ -117,7 +190,7 @@ def test_render_all_run_gallery_invokes_subprocess(tmp_path: Path, monkeypatch): "single_hist", "Single", "x", - {"edges": [0, 1, 2], "rollout": [5, 1]}, + {"edges": [0, 1, 2], "series": {"rollout": [5, 1]}}, ) ] for r in reduced: @@ -142,15 +215,14 @@ def test_render_run_glues_condor_run_meta_into_render_all(tmp_path: Path, monkey merge_calls = [] monkeypatch.setattr(condor_mod, "merge_all", lambda rd: merge_calls.append(Path(rd))) meta = condor_mod.RunMeta( - rollout="rollout.parquet", + rollouts=[{"name": "rollout", "path": "rollout.parquet", "plot_meta": {"checkpoint": "ckpt/best.pt"}}], reference="reference.parquet", run_dir=str(run_dir), title="my-run", - plot_meta={"checkpoint": "ckpt/best.pt"}, ) monkeypatch.setattr(condor_mod.RunMeta, "load", classmethod(lambda cls, p: meta)) - Reduced("s", "species", "single_hist", "Single", "x", {"edges": [0, 1], "rollout": [1]}).save( + Reduced("s", "species", "single_hist", "Single", "x", {"edges": [0, 1], "series": {"rollout": [1]}}).save( run_dir / "reduced" / "s.json" ) @@ -177,7 +249,7 @@ def test_render_one_of_each_kind(tmp_path: Path): "x", { "edges": [0, 1, 2, 3], - "rollout": [1, 2, 3], + "series": {"flow": [1, 2, 3], "wgan": [2, 2, 2]}, "reference": [3, 2, 1], "log_y": False, }, @@ -190,7 +262,7 @@ def test_render_one_of_each_kind(tmp_path: Path): "x", { "edges": [0, 1, 2], - "groups": {"a": {"rollout": [1, 2], "reference": [2, 1]}}, + "groups": {"a": {"series": {"flow": [1, 2]}, "reference": [2, 1]}}, "log_y": False, }, ), @@ -202,10 +274,8 @@ def test_render_one_of_each_kind(tmp_path: Path): "depth", { "edges": [0, 1, 2], - "rollout_mean": [1, 2], - "rollout_std": [0.1, 0.2], - "reference_mean": [1.1, 1.9], - "reference_std": [0.1, 0.1], + "series": {"flow": {"mean": [1, 2], "std": [0.1, 0.2]}}, + "reference": {"mean": [1.1, 1.9], "std": [0.1, 0.1]}, "ylabel": "e", }, ), @@ -217,7 +287,7 @@ def test_render_one_of_each_kind(tmp_path: Path): "species", { "labels": ["e-", "gamma"], - "rollout": [0.6, 0.4], + "series": {"flow": [0.6, 0.4], "wgan": [0.55, 0.45]}, "reference": [0.5, 0.5], "ylabel": "frac", }, @@ -228,7 +298,20 @@ def test_render_one_of_each_kind(tmp_path: Path): "single_hist", "Single", "x", - {"edges": [0, 1, 2], "rollout": [5, 1], "log_y": True}, + {"edges": [0, 1, 2], "series": {"flow": [5, 1]}, "log_y": True}, + ), + Reduced( + "hm1", + "secondaries", + "heatmap", + "Confusion (single rollout)", + "predicted", + { + "series": {"flow": [[1, 0], [0, 1]]}, + "row_labels": ["0", "1+"], + "col_labels": ["0", "1+"], + "cbar_label": "count", + }, ), ] try: @@ -276,8 +359,8 @@ def test_figure_params_v2_basics_and_router_and_epoch(): }, "conditioning": {"particle": {"type": "physical"}}, } - run_meta = {"training_epoch": 12, "best_val_loss": 0.123456, "steps": 10} - params = render_mod._figure_params(run_meta | {"model_config": mc}) + meta = {"training_epoch": 12, "best_val_loss": 0.123456, "steps": 10, "model_config": mc} + params = render_mod._figure_params({"rollouts": {"rollout": meta}}) assert params == { "hidden_dim": 256, "n_res_blocks": 4, @@ -297,8 +380,8 @@ def test_figure_params_v2_wgan_reports_noise_dim_not_steps(): "wgan": {"noise_dim": 32}, }, } - run_meta = {"model_config": mc, "steps": 10} - params = render_mod._figure_params(run_meta) + meta = {"model_config": mc, "steps": 10} + params = render_mod._figure_params({"rollouts": {"rollout": meta}}) assert params["mode"] == "wgan" assert params["noise_dim"] == 32 assert "steps" not in params @@ -309,18 +392,18 @@ def test_figure_params_v2_reports_mode_s2_only_when_it_differs(): "stage1_model": {"generator": "flow"}, "stage2_model": {"generator": "flow"}, } - assert "mode_s2" not in render_mod._figure_params({"model_config": same}) + assert "mode_s2" not in render_mod._figure_params({"rollouts": {"rollout": {"model_config": same}}}) mixed = { "stage1_model": {"generator": "flow"}, "stage2_model": {"generator": "wgan"}, } - params = render_mod._figure_params({"model_config": mixed}) + params = render_mod._figure_params({"rollouts": {"rollout": {"model_config": mixed}}}) assert params["mode_s2"] == "wgan" def test_figure_params_old_shape_basics(): - run_meta = { + meta = { "model_config": { "hidden_dim": 128, "n_blocks": 3, @@ -332,7 +415,7 @@ def test_figure_params_old_shape_basics(): "best_val_loss": 0.5, "steps": 20, } - params = render_mod._figure_params(run_meta) + params = render_mod._figure_params({"rollouts": {"rollout": meta}}) assert params == { "hidden_dim": 128, "n_blocks": 3, @@ -346,20 +429,30 @@ def test_figure_params_old_shape_basics(): def test_figure_params_old_shape_wgan_reports_noise_dim_not_steps(): - run_meta = { + meta = { "model_config": {"mode": "wgan", "noise_dim": 16}, "steps": 20, } - params = render_mod._figure_params(run_meta) + params = render_mod._figure_params({"rollouts": {"rollout": meta}}) assert params["noise_dim"] == 16 assert "steps" not in params +def test_figure_params_multi_rollout_names_the_series(): + run_meta = {"rollouts": {"flow": {"model_config": {"mode": "flow"}}, "wgan": {"model_config": {"mode": "wgan"}}}} + assert render_mod._figure_params(run_meta) == {"rollouts": "flow, wgan"} + + +def test_figure_params_empty_rollouts_is_empty(): + assert render_mod._figure_params({}) == {} + assert render_mod._figure_params({"rollouts": {}}) == {} + + def test_plot_metadata_includes_note_and_run_meta_parameters(): r = Reduced("u", "router", "unavailable", "Unavailable", "x", {"note": "no router data"}) - meta = render_mod._plot_metadata(r, {"title": "run-1", "checkpoint": "ckpt.pt"}) + meta = render_mod._plot_metadata(r, {"title": "run-1", "reference": "ref.parquet", "rollouts": {"rollout": {}}}) assert meta["note"] == "no router data" - assert meta["parameters"] == {"checkpoint": "ckpt.pt"} + assert meta["parameters"] == {"reference": "ref.parquet", "rollouts": {"rollout": {}}} assert "title" not in meta["parameters"] diff --git a/tests/test_router_gating.py b/tests/test_router_gating.py index 18be8c8..b16d8ee 100644 --- a/tests/test_router_gating.py +++ b/tests/test_router_gating.py @@ -10,7 +10,9 @@ from giant.analysis.router_gating import ( compute_router_gating, compute_router_share_by_pdg, compute_router_share_by_process, + compute_router_specialization, ) +from giant.analysis.sources import RolloutSide from giant.data.transforms import Normalizer from giant.model.network import build_models @@ -34,7 +36,7 @@ def _model_cfg() -> dict: } -def _write_checkpoint(tmp_path) -> str: +def _write_checkpoint(tmp_path, name: str = "ckpt.pt") -> str: cfg = _model_cfg() stage1 = build_models(cfg)["stage1"] assert stage1 is not None @@ -48,7 +50,7 @@ def _write_checkpoint(tmp_path) -> str: "mat_map": _MAT_MAP, "normalizer": {"cond": norm.to_dict()}, } - path = tmp_path / "ckpt.pt" + path = tmp_path / name torch.save(ckpt, path) return str(path) @@ -86,42 +88,80 @@ def _steps_frame(process: bool = False) -> pl.LazyFrame: return pl.DataFrame(data).lazy() +def _side(checkpoint: str | None, lf: pl.LazyFrame) -> RolloutSide: + return RolloutSide(all=lf, phys=lf, checkpoint=checkpoint) + + def test_compute_router_gating_shapes(tmp_path): checkpoint = _write_checkpoint(tmp_path) lf = _steps_frame() - r = compute_router_gating(checkpoint, lf, lf) + r = compute_router_gating({"rollout": _side(checkpoint, lf)}, lf) assert r.kind == "router_gating" - assert r.payload["n_experts"] == 2 + assert list(r.payload["series"]) == ["rollout"] + entry = r.payload["series"]["rollout"] + assert entry["n_experts"] == 2 for side in ("rollout", "reference"): - means = r.payload[side]["means"] + means = entry[side]["means"] assert means, f"{side} produced no bins" assert all(abs(sum(row) - 1.0) < 1e-5 for row in means) def test_compute_router_gating_missing_checkpoint_is_unavailable(): lf = _steps_frame() - r = compute_router_gating(None, lf, lf) + r = compute_router_gating({"rollout": _side(None, lf)}, lf) assert r.kind == "unavailable" assert "note" in r.payload assert r.title +def test_compute_router_gating_two_rollouts_only_moe_ones_included(tmp_path): + lf = _steps_frame() + ckpt = _write_checkpoint(tmp_path) + rollouts = {"flow": _side(None, lf), "moe": _side(ckpt, lf)} + r = compute_router_gating(rollouts, lf) + assert list(r.payload["series"]) == ["moe"] + + +def test_compute_router_specialization_two_rollouts(tmp_path): + lf = _steps_frame() + ckpt_a = _write_checkpoint(tmp_path, "a.pt") + ckpt_b = _write_checkpoint(tmp_path, "b.pt") + rollouts = {"a": _side(ckpt_a, lf), "b": _side(ckpt_b, lf)} + r = compute_router_specialization(rollouts, lf) + assert r.kind == "router_specialization" + assert list(r.payload["series"]) == ["a", "b"] + for entry in r.payload["series"].values(): + assert entry["chance_level"] == 0.5 + assert len(entry["rollout"]["centers"]) == len(entry["rollout"]["score"]) + + def test_compute_router_share_by_pdg(tmp_path): checkpoint = _write_checkpoint(tmp_path) lf = _steps_frame() - r = compute_router_share_by_pdg(checkpoint, lf, lf, top_pdgs=[11, 22]) + r = compute_router_share_by_pdg({"rollout": _side(checkpoint, lf)}, lf, top_pdgs=[11, 22]) assert r.kind == "router_share" + entry = r.payload["series"]["rollout"] for side in ("rollout", "reference"): - assert set(r.payload[side]) == {"e-", "gamma"} - for shares in r.payload[side].values(): + assert set(entry[side]) == {"e-", "gamma"} + for shares in entry[side].values(): assert abs(sum(shares) - 1.0) < 1e-5 def test_compute_router_share_by_process(tmp_path): checkpoint = _write_checkpoint(tmp_path) lf = _steps_frame(process=True) - r = compute_router_share_by_process(checkpoint, lf) + r = compute_router_share_by_process({"rollout": _side(checkpoint, lf)}, lf) assert r.kind == "router_share" - assert set(r.payload["categories"]) <= {"eIoni", "compt"} - for shares in r.payload["reference"].values(): + entry = r.payload["series"]["rollout"] + assert set(entry["categories"]) <= {"eIoni", "compt"} + for shares in entry["reference"].values(): assert abs(sum(shares) - 1.0) < 1e-5 + + +def test_no_moe_rollouts_are_unavailable(tmp_path): + lf = _steps_frame() + rollouts = {"flow": _side(None, lf), "wgan": _side(None, lf)} + assert compute_router_gating(rollouts, lf).kind == "unavailable" + assert compute_router_share_by_pdg(rollouts, lf, top_pdgs=[11, 22]).kind == "unavailable" + assert compute_router_share_by_process(rollouts, lf).kind == "unavailable" + assert compute_router_specialization(rollouts, lf).kind == "unavailable" diff --git a/tests/test_type_embedding_distance.py b/tests/test_type_embedding_distance.py index 6b3b8dd..c8e21de 100644 --- a/tests/test_type_embedding_distance.py +++ b/tests/test_type_embedding_distance.py @@ -3,6 +3,9 @@ from __future__ import annotations +import polars as pl + +from giant.analysis.sources import RolloutSide from giant.analysis.type_embedding_distance import compute_type_embedding_l1_distance @@ -18,26 +21,47 @@ def _summary(n=100): } +def _side(l1_dist: dict | None) -> RolloutSide: + empty = pl.LazyFrame() + return RolloutSide(all=empty, phys=empty, type_embedding_l1_dist=l1_dist) + + def test_none_is_unavailable(): - r = compute_type_embedding_l1_distance(None) + r = compute_type_embedding_l1_distance({"rollout": _side(None)}) assert r.kind == "unavailable" assert r.id == "type_embedding_l1_distance" assert r.payload["note"] def test_summary_produces_single_hist(): - r = compute_type_embedding_l1_distance(_summary()) + r = compute_type_embedding_l1_distance({"rollout": _side(_summary())}) assert r.kind == "single_hist" assert r.id == "type_embedding_l1_distance" assert r.payload["edges"] == [0.0, 1.0, 2.0, 3.0] - assert r.payload["rollout"] == [30, 40, 30] + assert r.payload["series"]["rollout"] == [30, 40, 30] assert r.payload["log_x"] is True assert r.payload["log_y"] is True assert "n=100" in r.payload["note"] def test_single_hist_payload_shape_matches_render_contract(): - """_render_single (giant.analysis.render) requires len(rollout) == + """_render_single (giant.analysis.render) requires each series' length == len(edges) - 1.""" - r = compute_type_embedding_l1_distance(_summary()) - assert len(r.payload["rollout"]) == len(r.payload["edges"]) - 1 + r = compute_type_embedding_l1_distance({"rollout": _side(_summary())}) + assert len(r.payload["series"]["rollout"]) == len(r.payload["edges"]) - 1 + + +def test_two_rollouts_both_populated(): + r = compute_type_embedding_l1_distance({"flow": _side(_summary(50)), "wgan": _side(_summary(80))}) + assert list(r.payload["series"]) == ["flow", "wgan"] + assert "n=50" in r.payload["note"] and "n=80" in r.payload["note"] + + +def test_one_of_two_rollouts_populated_only_that_one_appears(): + r = compute_type_embedding_l1_distance({"flow": _side(None), "wgan": _side(_summary())}) + assert list(r.payload["series"]) == ["wgan"] + + +def test_none_populated_across_rollouts_is_unavailable(): + r = compute_type_embedding_l1_distance({"flow": _side(None), "wgan": _side(None)}) + assert r.kind == "unavailable"