Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| a2d55e745f | |||
| f62f12e49e | |||
| d07bac8d32 | |||
| e90eead2af | |||
| ebd3e0dc71 |
+1
-1
@@ -1,5 +1,5 @@
|
||||
[tool.bumpversion]
|
||||
current_version = "0.3.8"
|
||||
current_version = "0.3.9"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
|
||||
@@ -1,5 +1,16 @@
|
||||
# Changelog
|
||||
|
||||
## [0.3.9] - 2026-08-24
|
||||
|
||||
### Added
|
||||
|
||||
- Add multi-rollout support to giant analyze [gitea #77](https://git.larsbogner.de/lars/giant/issues/77)
|
||||
|
||||
|
||||
### Changed
|
||||
|
||||
- Escape LaTeX-special characters in plot titles/xlabels [gitea #81](https://git.larsbogner.de/lars/giant/issues/81)
|
||||
|
||||
## [0.3.8] - 2026-08-24
|
||||
|
||||
### Added
|
||||
|
||||
@@ -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_<id>/`) 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/<id>__<chunk>.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 <run_dir>` first joins every plot's chunk partials into `reduced/<id>.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_<tag(s)>/`) 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/<id>__<chunk>.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 <run_dir>` first joins every plot's chunk partials into `reduced/<id>.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.
|
||||
|
||||
|
||||
@@ -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 <run_dir> --gallery # local: styled PDFs + HTML gallery (needs LaTeX)
|
||||
```
|
||||
|
||||
`<run_dir>` 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.
|
||||
`<run_dir>` 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
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
+296
-205
@@ -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: <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.
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
|
||||
|
||||
+132
-45
@@ -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:
|
||||
|
||||
<run_dir>/shared.json fixed bin edges / group sets (prep)
|
||||
<run_dir>/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_<tag>`` 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<K>`` 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"),
|
||||
)
|
||||
|
||||
|
||||
|
||||
+36
-25
@@ -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()},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
+17
-12
@@ -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
|
||||
|
||||
+223
-107
@@ -9,10 +9,19 @@ streaming compute.
|
||||
For each reduced artifact it writes ``<out>/<family>/<id>.pdf`` plus a sibling
|
||||
``<id>.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
|
||||
|
||||
import dataclasses
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
|
||||
@@ -22,7 +31,32 @@ import yaml
|
||||
|
||||
from giant.analysis.reduced import Reduced
|
||||
|
||||
_SERIES_LABELS = {"rollout": "rollout", "reference": "reference (Geant4)"}
|
||||
_REFERENCE_LABEL = "reference (Geant4)"
|
||||
|
||||
_TEX_ESCAPE_MAP = {
|
||||
"\\": r"\textbackslash{}",
|
||||
"%": r"\%",
|
||||
"&": r"\&",
|
||||
"#": r"\#",
|
||||
"$": r"\$",
|
||||
"_": r"\_",
|
||||
"{": r"\{",
|
||||
"}": r"\}",
|
||||
}
|
||||
|
||||
|
||||
def _tex_escape(text: str) -> str:
|
||||
"""Escape characters LaTeX treats specially in catalog-authored title/xlabel
|
||||
text (e.g. a literal ``%`` in a "90% of deposited energy" title, which
|
||||
``usetex`` otherwise reads as a comment marker and aborts the whole figure —
|
||||
see gitea #81). A single pass over the *original* characters, so the
|
||||
backslashes an escape itself introduces (e.g. ``\textbackslash{}``) are
|
||||
never re-escaped."""
|
||||
return "".join(_TEX_ESCAPE_MAP.get(c, c) for c in text)
|
||||
|
||||
|
||||
def _ref_color() -> str:
|
||||
return ps.colors.INK["primary"]
|
||||
|
||||
|
||||
def _density(counts: list[int] | np.ndarray, edges: np.ndarray) -> np.ndarray:
|
||||
@@ -33,10 +67,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 +84,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 +108,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 ``<id>.yaml`` (see ``_plot_metadata``) already
|
||||
``meta``/each plot's own ``<id>.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 +137,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 +152,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 +194,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 +237,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 +257,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 +278,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
|
||||
|
||||
|
||||
@@ -332,8 +443,14 @@ _RENDERERS = {
|
||||
|
||||
|
||||
def render(r: Reduced, run_meta: dict | None = None):
|
||||
"""Build the matplotlib figure for one reduced artifact (dispatch on kind)."""
|
||||
return _RENDERERS[r.kind](r, _figure_params(run_meta or {}))
|
||||
"""Build the matplotlib figure for one reduced artifact (dispatch on kind).
|
||||
|
||||
``title``/``xlabel`` are LaTeX-escaped here, at the one point every kind's
|
||||
renderer draws them from — ``_plot_metadata`` deliberately keeps using the
|
||||
unescaped ``r`` for the gallery YAML, which isn't LaTeX.
|
||||
"""
|
||||
escaped = dataclasses.replace(r, title=_tex_escape(r.title), xlabel=_tex_escape(r.xlabel))
|
||||
return _RENDERERS[r.kind](escaped, _figure_params(run_meta or {}))
|
||||
|
||||
|
||||
def _plot_metadata(r: Reduced, run_meta: dict) -> dict:
|
||||
@@ -392,7 +509,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 +542,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)
|
||||
|
||||
@@ -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},
|
||||
)
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
)
|
||||
|
||||
+37
-7
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "giant"
|
||||
version = "0.3.8"
|
||||
version = "0.3.9"
|
||||
description = "Geant4 step-function surrogate via conditional flow matching"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
|
||||
+94
-31
@@ -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"]
|
||||
|
||||
+143
-30
@@ -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)
|
||||
|
||||
+163
-43
@@ -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"]
|
||||
|
||||
|
||||
@@ -368,3 +461,30 @@ def test_plot_metadata_omits_parameters_when_run_meta_empty():
|
||||
meta = render_mod._plot_metadata(r, {})
|
||||
assert "parameters" not in meta
|
||||
assert "note" not in meta
|
||||
|
||||
|
||||
def test_tex_escape_handles_percent_and_other_special_chars():
|
||||
assert render_mod._tex_escape("90% of deposited energy") == r"90\% of deposited energy"
|
||||
assert render_mod._tex_escape(r"a_b & c#d $e {f} \bar") == r"a\_b \& c\#d \$e \{f\} \textbackslash{}bar"
|
||||
|
||||
|
||||
def test_render_survives_title_and_xlabel_with_literal_percent(tmp_path: Path):
|
||||
# Regression test for gitea #81: a literal "%" in a catalog title (e.g.
|
||||
# "Shower containment depth (90% of deposited energy)") crashed the whole
|
||||
# LaTeX render, since usetex treats an unescaped "%" as a comment marker.
|
||||
reduced = [
|
||||
Reduced(
|
||||
"shower_containment_depth_90",
|
||||
"shower",
|
||||
"single_hist",
|
||||
"Shower containment depth (90% of deposited energy)",
|
||||
"depth containing 90% of deposited energy [mm]",
|
||||
{"edges": [0, 1, 2], "series": {"flow": [5, 1]}},
|
||||
),
|
||||
]
|
||||
try:
|
||||
pdfs = _try_render(reduced, tmp_path)
|
||||
except RuntimeError as e: # LaTeX missing at render time
|
||||
pytest.skip(f"LaTeX rendering unavailable: {e}")
|
||||
assert len(pdfs) == 1
|
||||
assert pdfs[0].exists()
|
||||
|
||||
+52
-12
@@ -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"
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user