Files
giant/giant/analysis/catalog.py
T
lars ebd3e0dc71
CI / Lint (ruff check) (push) Successful in 32s
CI / Format (ruff format) (push) Successful in 30s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 35s
CI / Lint (ruff check) (pull_request) Successful in 33s
CI / Format (ruff format) (pull_request) Successful in 30s
CI / Type check (ty) (pull_request) Successful in 34s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Tests (push) Successful in 5m59s
CI / Bump version, tag, and update changelog on merge to master (push) Has been skipped
CI / Tests (pull_request) Successful in 4m22s
CI / Bump version, tag, and update changelog on merge to master (pull_request) Has been skipped
Add multi-rollout support to giant analyze (gitea #77)
giant analyze compares N rollout YAMLs against one shared reference file
(all must name the same dataset, checked up front) instead of exactly one
rollout vs one reference, rendering each rollout as its own colored series
against a single reference line/panel. Series names come from a repeated
--label flag, else the YAML stem, else "rollout" for a single YAML — a
single-rollout run keeps rendering identically to before this change.

Bundle now holds a name-keyed dict of rollout sides instead of one fixed
pair, every catalog compute_partial/finalize builds a Reduced.payload
keyed the same way ("series": {name: ...}, "reference": ... as the one
distinguished non-rollout entry), and every renderer draws N series (or
N panels, for the two heatmap-shaped specs and the router/type-embedding
diagnostics, which are inherently one-matrix/one-checkpoint per rollout)
against the reference's fixed dashed-ink style.
2026-08-24 13:23:50 +02:00

1146 lines
42 KiB
Python

"""The declarative plot catalog: one ``PlotSpec`` per figure.
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 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/
per-secondary array to be concatenated (anything that derives its own edges or
a mean/std from the full dataset). ``finalize`` merges the per-chunk partials
(in chunk order) and does the actual histogramming/edge-selection/mean-std
collapse, once, over the merged data — for ``n_chunks=1`` this reproduces
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: <shape>}, "t":
<shape>}`` — 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.
The registry is built by expanding parametric families (marginals over
variable x grouping, secondaries, ...) into concrete specs.
"""
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Callable
import numpy as np
import polars as pl
from giant.analysis.context import Context
from giant.analysis.grouping import (
energy_bin_labels,
event_energy_bins,
material_label,
pdg_label,
)
from giant.analysis.reduce import (
attach_entry_axis,
depth_expr,
entry_axis,
event_scalars,
hist1d,
leakage_fraction,
profile_finalize,
profile_partial,
sec_count_by_event,
species_share,
sum_merge,
transverse_expr,
)
from giant.analysis.reduced import Reduced
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, 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
@dataclass
class Bundle:
"""Everything a compute runs against — built once per ``compute-one`` job."""
ctx: Context
rollouts: dict[str, RolloutSide] # name -> frames, insertion order = CLI order
t_all: pl.LazyFrame # reference, all rows
t_phys: pl.LazyFrame # reference, physical steps only
@classmethod
def open(
cls,
rollouts: list[RolloutSpec],
reference,
ctx: Context,
chunk: tuple[int, int] | None = None,
) -> "Bundle":
"""Open the reference + every rollout, optionally restricted to one event-disjoint chunk.
``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.
"""
t_all = open_side(reference, Side.reference)
pred = None
if chunk is not None:
idx, n = chunk
pred = pl.col("event_id") % n == idx
t_all = t_all.filter(pred)
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
class PlotSpec:
id: str
family: str
compute_partial: Callable[[Bundle], dict]
finalize: Callable[[list[dict], Context], Reduced]
chunkable: bool = True
def _unchunkable(
compute: Callable[[Bundle], Reduced],
) -> tuple[Callable[[Bundle], dict], Callable[[list[dict], Context], Reduced]]:
"""Wrap a whole-dataset ``compute(bundle) -> Reduced`` as a trivial
``(compute_partial, finalize)`` pair, for specs marked ``chunkable=False``
(which always run as a single chunk, so ``parts`` is always one element).
"""
def partial(b: Bundle) -> dict:
return {"reduced": asdict(compute(b))}
def finalize(parts: list[dict], ctx: Context) -> Reduced:
return Reduced(**parts[0]["reduced"])
return partial, finalize
# ---------------------------------------------------------------------------
# small numpy/hist helpers
# ---------------------------------------------------------------------------
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]:
return h.get(key, np.zeros(nbins, dtype=np.int64)).astype(np.int64).tolist()
def _partial_hist(
lf: pl.LazyFrame, value: pl.Expr, edges: np.ndarray, group: pl.Expr | None = None
) -> dict[str, list[int]]:
"""One chunk's raw ``hist1d`` result as a JSON-safe, sum-mergeable dict."""
nb = len(edges) - 1
h = hist1d(lf, value, edges, group=group)
return {str(k): _counts(h, k, nb) for k in h}
def _finalize_counts(merged: dict[str, list], key, nbins: int) -> list[int]:
"""One group's merged counts (zero-filled if the group never appeared)."""
return list(merged.get(str(key), [0] * nbins))
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(a, edges)[0] for a in arrays]
def _ks_statistic(r_counts, t_counts) -> float:
"""KS statistic (max |CDF diff|) between two same-edge binned histograms.
``nan`` when neither side has any mass (nothing to compare); 1.0 (maximal
mismatch) when exactly one side is entirely empty and the other isn't —
correctly the worst score rather than an undefined one.
"""
r_counts = np.asarray(r_counts, dtype=np.float64)
t_counts = np.asarray(t_counts, dtype=np.float64)
r_tot, t_tot = r_counts.sum(), t_counts.sum()
if r_tot == 0 and t_tot == 0:
return float("nan")
if r_tot == 0 or t_tot == 0:
return 1.0
r_cdf = np.cumsum(r_counts) / r_tot
t_cdf = np.cumsum(t_counts) / t_tot
return float(np.max(np.abs(r_cdf - t_cdf)))
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.
"""
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
mat = np.zeros((n, n), dtype=np.int64)
np.add.at(mat, (t_c, r_c), 1)
labels = [str(i) for i in range(cap)] + [f"{cap}+"]
return labels, mat
def _containment_depths(mat: np.ndarray, edges: np.ndarray, quantile: float) -> np.ndarray:
"""Per-event depth containing ``quantile`` of that event's deposited energy.
``mat`` is a ``(n_events, n_bins)`` edep-per-depth-bin sum matrix (see
``reduce.profile_partial``); bins are ordered by increasing depth (matching
``edges``, monotonic). Zero-energy events are dropped — containment depth is
undefined for them.
"""
totals = mat.sum(axis=1)
valid = totals > 0
mat, totals = mat[valid], totals[valid]
cum = np.cumsum(mat, axis=1) / totals[:, None]
idx = (cum >= quantile).argmax(axis=1) # first bin whose cumulative fraction reaches quantile
return edges[1:][idx]
def _group_keys(ctx: Context, axis: str) -> list:
"""The group keys ``_marginal_grouped_finalize`` iterates for ``axis``."""
if axis == "pdg":
return list(ctx.top_pdgs)
if axis == "material":
return list(ctx.materials)
return list(range(len(ctx.energy_edges) - 1)) # energy
# Human-readable figure titles per marginal variable (the axis labels carry units;
# these read cleanly as a title without them).
_TITLE_NAMES = {
"step_length": "Step length",
"edep": "Deposited energy per step",
"delta_e": "Energy loss per step",
"post_E": "Post-step energy",
"cos_scatter": "Scattering cosine",
}
def _var(var: str):
"""(axis label, value expr) for a marginal variable name."""
if var == "cos_scatter":
return ("cos of scattering angle", cos_scatter_expr())
label, expr = RANGED_VARS[var]
return (label, expr)
def _marginal_edges(ctx: Context, var: str) -> np.ndarray:
if var == "cos_scatter":
return np.linspace(-1.0, 1.0, ctx.n_marginal_bins + 1)
return ctx.marginal_edges(var)
# ---------------------------------------------------------------------------
# marginals: variable x {overall, energy, pdg, material}
# ---------------------------------------------------------------------------
def _marginal_overall_partial(b: Bundle, var: str) -> dict:
_, expr = _var(var)
edges = _marginal_edges(b.ctx, var)
return {
"r": _per_rollout(b, lambda rs: _partial_hist(rs.phys, expr, edges)),
"t": _partial_hist(b.t_phys, expr, edges),
}
def _marginal_overall_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
label, _ = _var(var)
edges = _marginal_edges(ctx, var)
nb = len(edges) - 1
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}",
family="marginals",
kind="overlay_hist",
title=_TITLE_NAMES[var],
xlabel=label,
payload={
"edges": edges.tolist(),
"series": series,
"reference": _finalize_counts(t, 0, nb),
"log_y": True,
},
)
def _energy_group_expr(lf: pl.LazyFrame, edges: np.ndarray) -> pl.Expr:
ids, bins = event_energy_bins(lf, edges)
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)
nb = len(edges) - 1
return {
"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),
}
def _marginal_grouped_finalize(parts: list[dict], ctx: Context, var: str, axis: str) -> Reduced:
label, _ = _var(var)
edges = _marginal_edges(ctx, var)
nb = len(edges) - 1
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":
keys, labels = ctx.top_pdgs, [pdg_label(k) for k in ctx.top_pdgs]
elif axis == "material":
keys, labels = ctx.materials, [material_label(m) for m in ctx.materials]
else: # energy
e_edges = np.asarray(ctx.energy_edges)
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}",
family="marginals",
kind="grouped_hist",
title=f"{_TITLE_NAMES[var]} by {axis}",
xlabel=label,
payload={"edges": edges.tolist(), "groups": groups, "log_y": True},
)
# ---------------------------------------------------------------------------
# distance summary: a var x group-axis scorecard per rollout, reusing the marginal hists
# ---------------------------------------------------------------------------
def _distance_summary_partial(b: Bundle) -> dict:
out: dict[str, dict] = {}
for var in MARGINAL_VARS:
out[var] = {"overall": _marginal_overall_partial(b, var)}
for axis in GROUPING_AXES:
out[var][axis] = _marginal_grouped_partial(b, var, axis)
return out
def _distance_summary_finalize(parts: list[dict], ctx: Context) -> Reduced:
col_labels = ["overall", *GROUPING_AXES]
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
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:
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",
family="quality",
kind="heatmap",
title="Marginal distance summary (KS statistic, rollout vs reference)",
xlabel="grouping axis",
payload={
"series": matrices,
"row_labels": [_TITLE_NAMES[v] for v in MARGINAL_VARS],
"col_labels": col_labels,
"ylabel": "marginal variable",
"cbar_label": "KS statistic (0 = identical, 1 = maximal mismatch)",
"vmin": 0.0,
"vmax": 1.0,
},
)
# ---------------------------------------------------------------------------
# per-event scalar observables
# ---------------------------------------------------------------------------
def _event_scalar_partial(b: Bundle, col: str, use_all: bool) -> dict:
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:
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, 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",
kind="overlay_hist",
title=title,
xlabel=xlabel,
payload={
"edges": edges.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:
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": _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)
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)
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)):
tc = np.histogram(t_val[t_bin == bi], edges)[0]
groups[lbl] = {
"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",
family="event",
kind="grouped_hist",
title="Total deposited energy per event by incident energy",
xlabel="total deposited energy [MeV]",
payload={"edges": edges.tolist(), "groups": groups, "log_y": False},
)
# ---------------------------------------------------------------------------
# shower shape profiles
# ---------------------------------------------------------------------------
def _profile_partial(b: Bundle, coord_fn, edges_key: str) -> dict:
edges = np.asarray(getattr(b.ctx, edges_key))
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": _per_rollout(b, lambda rs: _mat(rs.all)),
"t": _mat(b.t_all),
}
def _assert_event_disjoint(id_lists: list[list[int]], spec_id: str, side: str) -> None:
"""Guard the chunking invariant profiles depend on: no event in two chunks.
A violation would silently double-count that event in the merged mean/RMS
with no other symptom, so this is worth a loud failure rather than a
quietly-wrong plot.
"""
seen: set[int] = set()
for ids in id_lists:
overlap = seen & set(ids)
if overlap:
raise ValueError(
f"{spec_id} ({side}): event_id(s) {sorted(overlap)[:5]} appear "
"in more than one chunk — chunking must be event-disjoint"
)
seen.update(ids)
def _profile_finalize(
parts: list[dict],
ctx: Context,
spec_id: str,
title: str,
xlabel: str,
edges_key: str,
) -> Reduced:
edges = np.asarray(getattr(ctx, edges_key))
nb = len(edges) - 1
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",
kind="profile",
title=title,
xlabel=xlabel,
payload={
"edges": edges.tolist(),
"series": series,
"reference": {"mean": t_mean.tolist(), "std": t_std.tolist()},
"ylabel": "mean deposited energy per event [MeV]",
},
)
# ---------------------------------------------------------------------------
# shower containment depth (reuses the longitudinal profile's per-event matrix)
# ---------------------------------------------------------------------------
_CONTAINMENT_QUANTILES: list[tuple[float, str]] = [
(0.90, "shower_containment_depth_90"),
(0.95, "shower_containment_depth_95"),
]
def _containment_finalize(parts: list[dict], ctx: Context, spec_id: str, quantile: float) -> Reduced:
edges = np.asarray(ctx.depth_edges)
nb = len(edges) - 1
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)
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",
kind="overlay_hist",
title=f"Shower containment depth ({quantile:.0%} of deposited energy)",
xlabel=f"depth containing {quantile:.0%} of deposited energy [mm]",
payload={
"edges": hedges.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,
},
)
# ---------------------------------------------------------------------------
# species share + leakage
# ---------------------------------------------------------------------------
def _species_share_partial(b: Bundle) -> dict:
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": _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:
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])
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",
kind="bar",
title="Deposited-energy share by species",
xlabel="species",
payload={
"labels": labels,
"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:
return {"r": _per_rollout(b, lambda rs: leakage_fraction(rs.all).tolist())}
def _leakage_finalize(parts: list[dict], ctx: Context) -> Reduced:
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",
kind="single_hist",
title="Escaped (leakage) energy fraction per shower",
xlabel="escaped energy fraction",
payload={
"edges": edges.tolist(),
"series": series,
"log_y": True,
"note": "rollout only; the reference has no detector-escape concept",
},
)
# ---------------------------------------------------------------------------
# secondaries
# ---------------------------------------------------------------------------
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:
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:
names = list(parts[0]["r"])
t = np.concatenate([np.asarray(p["t"], dtype=float) for p in parts])
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",
kind="overlay_hist",
title="Number of secondaries per event",
xlabel="secondaries per event",
payload={
"edges": edges.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,
},
)
def _counts_by_pdg(sec_lf: pl.LazyFrame) -> dict[str, int]:
df = sec_lf.group_by("pdg").agg(pl.len().alias("n")).collect(engine="streaming")
return {str(k): v for k, v in zip(df["pdg"].to_list(), df["n"].to_list())}
def _sec_count_per_species_partial(b: Bundle) -> dict:
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:
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])
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",
kind="bar",
title="Secondary count by species",
xlabel="species",
payload={
"labels": [pdg_label(int(k)) 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:
edges = np.linspace(*b.ctx.sec_energy_range, b.ctx.n_sec_bins + 1)
return {
"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
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(), "series": series, "reference": _finalize_counts(t, 0, nb), "log_y": True},
)
def _sec_cos_angle_partial(b: Bundle) -> dict:
edges = np.linspace(-1.0, 1.0, b.ctx.n_sec_bins + 1)
cos = (pl.col("sdx") * pl.col("axis_x") + pl.col("sdy") * pl.col("axis_y") + pl.col("sdz") * pl.col("axis_z")).clip(
-1.0, 1.0
)
def _side(sec_lf: pl.LazyFrame, steps_lf: pl.LazyFrame) -> dict[str, list[int]]:
ea = entry_axis(steps_lf)
return _partial_hist(attach_entry_axis(sec_lf, ea), cos, edges)
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
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(), "series": series, "reference": _finalize_counts(t, 0, nb), "log_y": False},
)
def _n_sec_confusion_partial(b: Bundle) -> dict:
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:
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.
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()))
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",
kind="heatmap",
title="Predicted vs true secondary count per event",
xlabel="predicted secondaries (rollout)",
payload={
"series": matrices,
"row_labels": labels,
"col_labels": labels,
"ylabel": "true secondaries (reference)",
"cbar_label": "event count",
"vmin": 0.0,
},
)
# ---------------------------------------------------------------------------
# router diagnostics (not chunked — already bounded/subsampled)
# ---------------------------------------------------------------------------
_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.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.rollouts, b.t_phys)
)
_router_specialization_partial, _router_specialization_finalize = _unchunkable(
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.rollouts)
)
# ---------------------------------------------------------------------------
# registry assembly
# ---------------------------------------------------------------------------
MARGINAL_VARS = ["step_length", "edep", "delta_e", "post_E", "cos_scatter"]
GROUPING_AXES = ["energy", "pdg", "material"]
def build_catalog() -> list[PlotSpec]:
"""All concrete plot specs, each with a unique id."""
specs: list[PlotSpec] = []
for var in MARGINAL_VARS:
specs.append(
PlotSpec(
f"marginal_{var}",
"marginals",
compute_partial=lambda b, v=var: _marginal_overall_partial(b, v),
finalize=lambda parts, ctx, v=var: _marginal_overall_finalize(parts, ctx, v),
)
)
for axis in GROUPING_AXES:
specs.append(
PlotSpec(
f"marginal_{var}_by_{axis}",
"marginals",
compute_partial=lambda b, v=var, a=axis: _marginal_grouped_partial(b, v, a),
finalize=lambda parts, ctx, v=var, a=axis: _marginal_grouped_finalize(parts, ctx, v, a),
)
)
specs.append(
PlotSpec(
"marginal_distance_summary",
"quality",
compute_partial=_distance_summary_partial,
finalize=_distance_summary_finalize,
)
)
specs += [
PlotSpec(
"event_total_edep",
"event",
compute_partial=lambda b: _event_scalar_partial(b, "total_edep", use_all=True),
finalize=lambda parts, ctx: _event_scalar_finalize(
parts,
ctx,
"event_total_edep",
"Total deposited energy per event",
"total deposited energy [MeV]",
),
),
PlotSpec(
"event_total_edep_by_energy",
"event",
compute_partial=_event_total_edep_by_energy_partial,
finalize=_event_total_edep_by_energy_finalize,
),
PlotSpec(
"event_mean_length",
"event",
compute_partial=lambda b: _event_scalar_partial(b, "mean_length", use_all=False),
finalize=lambda parts, ctx: _event_scalar_finalize(
parts,
ctx,
"event_mean_length",
"Mean step length per event",
"mean step length [mm]",
),
),
PlotSpec(
"event_n_steps",
"event",
compute_partial=lambda b: _event_scalar_partial(b, "n_steps", use_all=False),
finalize=lambda parts, ctx: _event_scalar_finalize(
parts,
ctx,
"event_n_steps",
"Number of steps per event",
"steps per event",
),
),
PlotSpec(
"shower_longitudinal",
"shower",
compute_partial=lambda b: _profile_partial(b, depth_expr, "depth_edges"),
finalize=lambda parts, ctx: _profile_finalize(
parts,
ctx,
"shower_longitudinal",
"Longitudinal shower profile",
"depth along shower axis [mm]",
"depth_edges",
),
),
PlotSpec(
"shower_transverse",
"shower",
compute_partial=lambda b: _profile_partial(b, transverse_expr, "transverse_edges"),
finalize=lambda parts, ctx: _profile_finalize(
parts,
ctx,
"shower_transverse",
"Transverse shower profile",
"radius from shower axis [mm]",
"transverse_edges",
),
),
]
for quantile, spec_id in _CONTAINMENT_QUANTILES:
specs.append(
PlotSpec(
spec_id,
"shower",
compute_partial=lambda b: _profile_partial(b, depth_expr, "depth_edges"),
finalize=lambda parts, ctx, q=quantile, sid=spec_id: _containment_finalize(parts, ctx, sid, q),
)
)
specs += [
PlotSpec(
"species_edep_share",
"species",
compute_partial=_species_share_partial,
finalize=_species_share_finalize,
),
PlotSpec(
"leakage_fraction",
"species",
compute_partial=_leakage_partial,
finalize=_leakage_finalize,
),
PlotSpec(
"sec_count_per_event",
"secondaries",
compute_partial=_sec_count_per_event_partial,
finalize=_sec_count_per_event_finalize,
),
PlotSpec(
"sec_count_per_species",
"secondaries",
compute_partial=_sec_count_per_species_partial,
finalize=_sec_count_per_species_finalize,
),
PlotSpec(
"sec_energy",
"secondaries",
compute_partial=_sec_energy_partial,
finalize=_sec_energy_finalize,
),
PlotSpec(
"sec_cos_angle",
"secondaries",
compute_partial=_sec_cos_angle_partial,
finalize=_sec_cos_angle_finalize,
),
PlotSpec(
"n_sec_confusion",
"secondaries",
compute_partial=_n_sec_confusion_partial,
finalize=_n_sec_confusion_finalize,
),
PlotSpec(
"router_gating",
"model",
compute_partial=_router_gating_partial,
finalize=_router_gating_finalize,
chunkable=False,
),
PlotSpec(
"router_share_by_pdg",
"model",
compute_partial=_router_share_pdg_partial,
finalize=_router_share_pdg_finalize,
chunkable=False,
),
PlotSpec(
"router_share_by_process",
"model",
compute_partial=_router_share_process_partial,
finalize=_router_share_process_finalize,
chunkable=False,
),
PlotSpec(
"router_specialization",
"model",
compute_partial=_router_specialization_partial,
finalize=_router_specialization_finalize,
chunkable=False,
),
PlotSpec(
"type_embedding_l1_distance",
"model",
compute_partial=_type_embedding_l1_distance_partial,
finalize=_type_embedding_l1_distance_finalize,
chunkable=False,
),
]
return specs
def catalog_ids() -> list[str]:
return [s.id for s in build_catalog()]
def get_spec(spec_id: str) -> PlotSpec:
for s in build_catalog():
if s.id == spec_id:
return s
raise KeyError(f"unknown plot id: {spec_id!r}")