Compare commits
5 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 8cdeba088e | |||
| 48208e6d18 | |||
| 51790d3e0a | |||
| 51f9dad3b0 | |||
| ac01966a1f |
+1
-1
@@ -1,5 +1,5 @@
|
||||
[tool.bumpversion]
|
||||
current_version = "0.3.21"
|
||||
current_version = "0.3.22"
|
||||
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
|
||||
serialize = ["{major}.{minor}.{patch}"]
|
||||
search = "{current_version}"
|
||||
|
||||
@@ -1,5 +1,15 @@
|
||||
# Changelog
|
||||
|
||||
## [0.3.22] - 2026-09-07
|
||||
|
||||
### Changed
|
||||
|
||||
- Feat(analyze): add paired truth/pred plots from giant predict
|
||||
|
||||
- Fix(tests): make predict --truth flag test robust to terminal rendering
|
||||
|
||||
- Feat(predict): enrich YAML sidecar with provenance and timing
|
||||
|
||||
## [0.3.21] - 2026-09-07
|
||||
|
||||
### Changed
|
||||
|
||||
@@ -18,7 +18,7 @@ giant train path/to/steps.parquet --router --router-type energy # MoE routing t
|
||||
giant model summary --config config.toml # build-only: parameter counts + which config keys actually bite
|
||||
giant predict path/to/steps.parquet --checkpoint ckpt/best.pt # per-step predictions
|
||||
giant rollout path/to/steps.parquet --checkpoint ckpt/best.pt --geometry oracle.pkl # full showers
|
||||
giant analyze submit rollout.yaml --accounting-group cms # parallel rollout-vs-reference analysis on HTCondor
|
||||
giant analyze submit rollout.yaml --prediction pred.yaml --accounting-group cms # + paired truth/pred plots
|
||||
giant analyze render <run_dir> --gallery # render PDFs + HTML gallery (run_dir from prep/submit)
|
||||
giant analyze metrics <train_run_dir> # training-progress plots from metrics.csv
|
||||
dwarf --help # dataset/tooling CLI: convert, migrate, bump-gen,
|
||||
@@ -101,6 +101,8 @@ Secondary energies are a **stick-breaking partition of the `e_sec` budget** from
|
||||
|
||||
**Input is one or more `giant rollout` YAML sidecars** (`condor.py:load_rollout_yamls`): 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). `prep` creates a **run directory** (`<cwd>/analysis_runs/analysis_<id>/` by default, `--run-dir` to override) 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 `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` splits into a `compute_partial`/`finalize` pair so chunks can be summed/concatenated back per rollout (`chunkable=False` specs — the checkpoint-bound diagnostics, already bounded/subsampled — always run as a single chunk). 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`; `merge-one` does a single plot for debugging), then turns those into the styled PDF/gallery tree. `giant analyze metrics <train_run_dir>` is a separate, unrelated entry point: training-progress plots straight from a run's `metrics.csv`.
|
||||
|
||||
**`prediction` family (paired truth/pred, `giant/analysis/prediction.py`):** an optional add-on to the rollout comparison, driven by `--prediction`/`--prediction-label` on `analyze prep`/`submit` (repeatable, same convention as `--label`/rollout YAMLs; series name defaults to the YAML stem for N>1 or `"prediction"` for one). Unlike a rollout (freely generated, no row-level correspondence to truth), a `giant predict` output has a matching truth row for every prediction — a paired, not distributional, comparison. `giant predict --coord global` (schema v3, `--truth` on by default) writes both `pred_*` and `true_*` physical columns plus truth/predicted secondary lists; `--coord local` is the older, always-paired 9D model-space output (`pred_{name}`/`true_{name}` for `LOCAL_TARGET_NAMES`, no secondaries — stage 2 doesn't run there). `paired_frame()` normalizes either coord into one canonical `true_<var>`/`pred_<var>` frame over `PAIRED_VARS` (`step_length`, `edep`, `delta_e`, `post_E`, `cos_scatter`, `cos_travel`), decoding local coord's ALR energy logits the same way `energy_simplex_decode` does. Every prediction in one run must share one `--coord` and the rollouts' `dataset` (`condor.load_prediction_yamls`). The catalog's `prediction` family (`catalog.py`, ids prefixed `pred_`) covers per-variable marginals (new `paired_hist` kind: true dashed / pred solid) and truth-vs-pred 2D histograms (new `heatmap2d` kind, with a y=x guide), residuals/relative-residuals/residual-vs-truth profiles, KS/bias/RMSE scorecards (reusing `heatmap`), `n_sec` and secondary-species confusion matrices (row-normalised `heatmap`), direction-alignment and physical-constraint-violation checks, and a pred/true correlation-matrix delta. Every spec degrades to `kind="unavailable"` when no `--prediction` was given, so a rollout-only run is unaffected. `giant predict` also writes a YAML sidecar next to the checkpoint (`cli.py:_write_prediction_ref`, mirroring `giant rollout`'s) carrying the run's provenance and timing — coord/weights/steps/batch size, row/skip/unknown-PDG counts, a `timing` block, and the checkpoint's `model_config`/`training_epoch`/`training_config` — which `--prediction` consumes the same way `--label` rollout YAMLs are consumed, surfacing those keys into each plot's gallery `metadata.yaml` (`condor.py:_PLOT_META_KEYS`).
|
||||
|
||||
**Shower rollout** (`giant/rollout.py`, `giant rollout` CLI): autoregressively steps the two-stage model into a full shower, advancing tracks breadth-first (every sweep steps all active tracks once, in `batch_size` chunks, so many tracks share each forward pass). 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 one of the `TERM_*` reasons in `constants.py` (energy cutoff, max steps, escape, natural end, unknown pdg, max tracks); energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction. `giant/checkpoint_io.py` is the shared checkpoint → ready-to-run-models path used by both `predict` and `rollout`.
|
||||
|
||||
## Roadmap
|
||||
|
||||
@@ -10,8 +10,8 @@ calorimeter showers.
|
||||
|
||||
[](pyproject.toml)
|
||||
[](pyproject.toml)
|
||||
[](CHANGELOG.md)
|
||||
[](tests/)
|
||||
[](CHANGELOG.md)
|
||||
[](tests/)
|
||||
[](https://git.larsbogner.de/lars/giant/actions)
|
||||
[](#license)
|
||||
|
||||
|
||||
@@ -13,12 +13,15 @@ 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 (
|
||||
LoadedPrediction,
|
||||
LoadedRollout,
|
||||
RunMeta,
|
||||
SubmitConfig,
|
||||
compute_one,
|
||||
compute_reduced,
|
||||
derive_run_dir,
|
||||
load_prediction_yaml,
|
||||
load_prediction_yamls,
|
||||
load_rollout_yaml,
|
||||
load_rollout_yamls,
|
||||
merge_all,
|
||||
@@ -27,6 +30,7 @@ from giant.analysis.condor import (
|
||||
write_submit,
|
||||
)
|
||||
from giant.analysis.context import Context, build_context
|
||||
from giant.analysis.prediction import PredictionSpec
|
||||
from giant.analysis.reduced import Partial, Reduced
|
||||
from giant.analysis.runtime_estimate import RUNTIME_SAFETY_MARGIN, estimate_runtime_s
|
||||
from giant.analysis.sources import RolloutSpec, Side
|
||||
@@ -34,8 +38,10 @@ from giant.analysis.sources import RolloutSpec, Side
|
||||
__all__ = [
|
||||
"RUNTIME_SAFETY_MARGIN",
|
||||
"Context",
|
||||
"LoadedPrediction",
|
||||
"LoadedRollout",
|
||||
"Partial",
|
||||
"PredictionSpec",
|
||||
"Reduced",
|
||||
"RolloutSpec",
|
||||
"RunMeta",
|
||||
@@ -49,6 +55,8 @@ __all__ = [
|
||||
"derive_run_dir",
|
||||
"estimate_runtime_s",
|
||||
"get_spec",
|
||||
"load_prediction_yaml",
|
||||
"load_prediction_yamls",
|
||||
"load_rollout_yaml",
|
||||
"load_rollout_yamls",
|
||||
"merge_all",
|
||||
|
||||
+943
-4
@@ -39,7 +39,8 @@ variable x grouping, secondaries, ...) into concrete specs.
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import asdict, dataclass
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from typing import cast
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
@@ -52,12 +53,25 @@ from giant.analysis.grouping import (
|
||||
material_label,
|
||||
pdg_label,
|
||||
)
|
||||
from giant.analysis.prediction import (
|
||||
PAIRED_SCALARS,
|
||||
PAIRED_VARS,
|
||||
PredictionSide,
|
||||
PredictionSpec,
|
||||
dir_alignment_expr,
|
||||
open_prediction,
|
||||
paired_frame,
|
||||
paired_secondaries,
|
||||
paired_vars_for_coord,
|
||||
)
|
||||
from giant.analysis.reduce import (
|
||||
attach_entry_axis,
|
||||
binned_moments,
|
||||
depth_expr,
|
||||
entry_axis,
|
||||
event_scalars,
|
||||
hist1d,
|
||||
hist2d,
|
||||
leakage_fraction,
|
||||
profile_finalize,
|
||||
profile_partial,
|
||||
@@ -93,6 +107,7 @@ class Bundle:
|
||||
rollouts: dict[str, RolloutSide] # name -> frames, insertion order = CLI order
|
||||
t_all: pl.LazyFrame # reference, all rows
|
||||
t_phys: pl.LazyFrame # reference, physical steps only
|
||||
predictions: dict[str, PredictionSide] = field(default_factory=dict) # name -> paired frames, CLI order
|
||||
|
||||
@classmethod
|
||||
def open(
|
||||
@@ -101,12 +116,14 @@ class Bundle:
|
||||
reference,
|
||||
ctx: Context,
|
||||
chunk: tuple[int, int] | None = None,
|
||||
predictions: list[PredictionSpec] | None = None,
|
||||
) -> Bundle:
|
||||
"""Open the reference + every rollout, optionally restricted to one event-disjoint chunk.
|
||||
"""Open the reference + every rollout + every prediction, optionally
|
||||
restricted to one event-disjoint chunk.
|
||||
|
||||
``chunk = (chunk_index, n_chunks)`` filters every side to
|
||||
``event_id % n_chunks == chunk_index`` *before* deriving the physical/
|
||||
secondary views, so every downstream reduction (which is either
|
||||
secondary/paired 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.
|
||||
"""
|
||||
@@ -128,7 +145,24 @@ class Bundle:
|
||||
type_embedding_l1_dist=rs.type_embedding_l1_dist,
|
||||
timing=rs.timing,
|
||||
)
|
||||
return cls(ctx=ctx, rollouts=sides, t_all=t_all, t_phys=physical_steps(t_all, Side.reference))
|
||||
pred_sides: dict[str, PredictionSide] = {}
|
||||
for ps in predictions or []:
|
||||
opened = open_prediction(ps.source)
|
||||
lf = opened.lf if pred is None else opened.lf.filter(pred)
|
||||
pred_sides[ps.name] = PredictionSide(
|
||||
lf=lf,
|
||||
paired=paired_frame(lf, opened.coord, opened.has_truth),
|
||||
coord=opened.coord,
|
||||
has_truth=opened.has_truth,
|
||||
checkpoint=ps.checkpoint,
|
||||
)
|
||||
return cls(
|
||||
ctx=ctx,
|
||||
rollouts=sides,
|
||||
t_all=t_all,
|
||||
t_phys=physical_steps(t_all, Side.reference),
|
||||
predictions=pred_sides,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -1076,6 +1110,807 @@ MARGINAL_VARS = ["step_length", "edep", "delta_e", "post_E", "cos_scatter"]
|
||||
GROUPING_AXES = ["energy", "pdg", "material"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# giant predict: paired truth/pred family ("prediction")
|
||||
#
|
||||
# Unlike every spec above (rollout series vs one shared reference — an
|
||||
# *unpaired* distribution comparison), a prediction has a truth row for every
|
||||
# output row. These specs compare true_<var> against pred_<var> from
|
||||
# `giant.analysis.prediction.paired_frame` directly — no reference series,
|
||||
# and (unlike a rollout) a prediction whose --coord is "local" never has
|
||||
# secondaries or a material/pdg breakdown. `b.predictions` is empty on a run
|
||||
# with no --prediction given, in which case every spec here degrades to
|
||||
# kind="unavailable" rather than raising.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_PRED_UNAVAILABLE_NOTE = "no --prediction given to `analyze prep`/`submit`"
|
||||
|
||||
_PRED_TITLE_NAMES = {
|
||||
"step_length": "Step length",
|
||||
"edep": "Deposited energy per step",
|
||||
"delta_e": "Energy loss per step",
|
||||
"post_E": "Post-step energy",
|
||||
"cos_scatter": "Scattering cosine (pre_dir . post_dir)",
|
||||
"cos_travel": "Travel-direction cosine (pre_dir . (post_pos - pre_pos))",
|
||||
}
|
||||
|
||||
_PRED_VAR_LABELS = {
|
||||
"step_length": "step length [mm]",
|
||||
"edep": "deposited energy [MeV]",
|
||||
"delta_e": "energy loss [MeV]",
|
||||
"post_E": "post-step energy [MeV]",
|
||||
"cos_scatter": "cos(scattering angle)",
|
||||
"cos_travel": "cos(travel-direction angle)",
|
||||
}
|
||||
|
||||
|
||||
def _per_prediction(b: Bundle, fn: Callable[[PredictionSide], object]) -> dict[str, object]:
|
||||
"""``{name: fn(prediction_side)}`` over every prediction, preserving CLI order."""
|
||||
return {name: fn(ps) for name, ps in b.predictions.items()}
|
||||
|
||||
|
||||
def _pred_unavailable(
|
||||
spec_id: str, family: str, title: str, xlabel: str, note: str = _PRED_UNAVAILABLE_NOTE
|
||||
) -> Reduced:
|
||||
return Reduced(id=spec_id, family=family, kind="unavailable", title=title, xlabel=xlabel, payload={"note": note})
|
||||
|
||||
|
||||
# ---- marginals + 2D truth-vs-pred -----------------------------------------
|
||||
|
||||
|
||||
def _pred_marginal_partial(b: Bundle, var: str) -> dict:
|
||||
if not b.predictions or var not in b.ctx.pred_var_ranges:
|
||||
return {"available": False}
|
||||
edges = b.ctx.pred_marginal_edges(var)
|
||||
nb = len(edges) - 1
|
||||
|
||||
def _one(ps: PredictionSide) -> dict | None:
|
||||
if var not in paired_vars_for_coord(ps.coord):
|
||||
return None
|
||||
entry = {"pred": _counts(hist1d(ps.paired, pl.col(f"pred_{var}"), edges), 0, nb)}
|
||||
if ps.has_truth:
|
||||
entry["true"] = _counts(hist1d(ps.paired, pl.col(f"true_{var}"), edges), 0, nb)
|
||||
return entry
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_marginal_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable(f"pred_marginal_{var}", "prediction", _PRED_TITLE_NAMES[var], _PRED_VAR_LABELS[var])
|
||||
edges = ctx.pred_marginal_edges(var)
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, dict] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
merged = sum_merge(entries)
|
||||
series[name] = {k: [int(x) for x in v] for k, v in merged.items()}
|
||||
return Reduced(
|
||||
id=f"pred_marginal_{var}",
|
||||
family="prediction",
|
||||
kind="paired_hist",
|
||||
title=_PRED_TITLE_NAMES[var],
|
||||
xlabel=_PRED_VAR_LABELS[var],
|
||||
payload={"edges": edges.tolist(), "series": series, "log_y": var in PAIRED_SCALARS},
|
||||
)
|
||||
|
||||
|
||||
def _pred_scatter_partial(b: Bundle, var: str) -> dict:
|
||||
if not b.predictions or var not in b.ctx.pred_var_ranges:
|
||||
return {"available": False}
|
||||
lo, hi = b.ctx.pred_var_ranges[var]
|
||||
edges = np.linspace(lo, hi, b.ctx.n_marginal_bins + 1)
|
||||
|
||||
def _one(ps: PredictionSide) -> list[list[int]] | None:
|
||||
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
|
||||
return None
|
||||
return hist2d(ps.paired, pl.col(f"true_{var}"), pl.col(f"pred_{var}"), edges, edges).tolist()
|
||||
|
||||
return {"available": True, "edges": edges.tolist(), "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_scatter_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable(
|
||||
f"pred_scatter_{var}", "prediction", f"{_PRED_TITLE_NAMES[var]}: truth vs prediction", _PRED_VAR_LABELS[var]
|
||||
)
|
||||
edges = np.asarray(parts[0]["edges"])
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list] = {}
|
||||
for name in names:
|
||||
mats = [p["p"][name] for p in parts]
|
||||
if mats[0] is None:
|
||||
continue
|
||||
series[name] = np.sum([np.asarray(m, dtype=np.int64) for m in mats], axis=0).tolist()
|
||||
return Reduced(
|
||||
id=f"pred_scatter_{var}",
|
||||
family="prediction",
|
||||
kind="heatmap2d",
|
||||
title=f"{_PRED_TITLE_NAMES[var]}: truth vs prediction",
|
||||
xlabel=f"true {_PRED_VAR_LABELS[var]}",
|
||||
payload={
|
||||
"x_edges": edges.tolist(),
|
||||
"y_edges": edges.tolist(),
|
||||
"series": series,
|
||||
"ylabel": f"predicted {_PRED_VAR_LABELS[var]}",
|
||||
"cbar_label": "step count",
|
||||
"log_color": True,
|
||||
"diagonal": True,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---- residuals + calibration -----------------------------------------------
|
||||
|
||||
|
||||
def _pred_residual_partial(b: Bundle, var: str) -> dict:
|
||||
if not b.predictions or var not in b.ctx.pred_residual_ranges:
|
||||
return {"available": False}
|
||||
edges = b.ctx.pred_residual_edges(var)
|
||||
nb = len(edges) - 1
|
||||
|
||||
def _one(ps: PredictionSide) -> list[int] | None:
|
||||
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
|
||||
return None
|
||||
resid = pl.col(f"pred_{var}") - pl.col(f"true_{var}")
|
||||
return _counts(hist1d(ps.paired, resid, edges), 0, nb)
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_residual_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
|
||||
title, xlabel = f"{_PRED_TITLE_NAMES[var]} residual", f"pred - true {_PRED_VAR_LABELS[var]}"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable(f"pred_residual_{var}", "prediction", title, xlabel)
|
||||
edges = ctx.pred_residual_edges(var)
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list[int]] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
series[name] = [int(x) for x in sum_merge([{"c": e} for e in entries])["c"]]
|
||||
return Reduced(
|
||||
id=f"pred_residual_{var}",
|
||||
family="prediction",
|
||||
kind="single_hist",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={"edges": edges.tolist(), "series": series, "log_y": True},
|
||||
)
|
||||
|
||||
|
||||
_REL_RESIDUAL_LO, _REL_RESIDUAL_HI = -2.0, 2.0
|
||||
|
||||
|
||||
def _pred_relative_residual_partial(b: Bundle, var: str) -> dict:
|
||||
if not b.predictions:
|
||||
return {"available": False}
|
||||
edges = np.linspace(_REL_RESIDUAL_LO, _REL_RESIDUAL_HI, b.ctx.n_marginal_bins + 1)
|
||||
nb = len(edges) - 1
|
||||
|
||||
def _one(ps: PredictionSide) -> list[int] | None:
|
||||
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
|
||||
return None
|
||||
rel = (pl.col(f"pred_{var}") - pl.col(f"true_{var}")) / pl.col(f"true_{var}")
|
||||
return _counts(hist1d(ps.paired, rel, edges), 0, nb)
|
||||
|
||||
return {"available": True, "edges": edges.tolist(), "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_relative_residual_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
|
||||
title, xlabel = f"{_PRED_TITLE_NAMES[var]} relative residual", "(pred - true) / true"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable(f"pred_relative_residual_{var}", "prediction", title, xlabel)
|
||||
edges = np.asarray(parts[0]["edges"])
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list[int]] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
series[name] = [int(x) for x in sum_merge([{"c": e} for e in entries])["c"]]
|
||||
return Reduced(
|
||||
id=f"pred_relative_residual_{var}",
|
||||
family="prediction",
|
||||
kind="single_hist",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={"edges": edges.tolist(), "series": series, "log_y": True},
|
||||
)
|
||||
|
||||
|
||||
def _pred_residual_profile_partial(b: Bundle, var: str) -> dict:
|
||||
if not b.predictions or var not in b.ctx.pred_var_ranges:
|
||||
return {"available": False}
|
||||
edges = b.ctx.pred_marginal_edges(var) # bin by truth value
|
||||
|
||||
def _one(ps: PredictionSide) -> dict | None:
|
||||
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
|
||||
return None
|
||||
resid = pl.col(f"pred_{var}") - pl.col(f"true_{var}")
|
||||
return binned_moments(ps.paired, pl.col(f"true_{var}"), resid, edges)
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _moments_to_mean_std(merged: dict[str, list]) -> tuple[np.ndarray, np.ndarray]:
|
||||
n = np.asarray(merged["n"], dtype=np.float64)
|
||||
s = np.asarray(merged["sum"], dtype=np.float64)
|
||||
ss = np.asarray(merged["sumsq"], dtype=np.float64)
|
||||
with np.errstate(invalid="ignore", divide="ignore"):
|
||||
mean = np.where(n > 0, s / n, 0.0)
|
||||
var = np.where(n > 0, ss / n - mean**2, 0.0)
|
||||
return mean, np.sqrt(np.clip(var, 0.0, None))
|
||||
|
||||
|
||||
def _pred_residual_profile_finalize(parts: list[dict], ctx: Context, var: str) -> Reduced:
|
||||
title, xlabel = f"{_PRED_TITLE_NAMES[var]} residual vs truth", f"true {_PRED_VAR_LABELS[var]}"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable(f"pred_residual_profile_{var}", "prediction", title, xlabel)
|
||||
edges = ctx.pred_marginal_edges(var)
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, dict] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
mean, std = _moments_to_mean_std(sum_merge(entries))
|
||||
series[name] = {"mean": mean.tolist(), "std": std.tolist()}
|
||||
return Reduced(
|
||||
id=f"pred_residual_profile_{var}",
|
||||
family="prediction",
|
||||
kind="profile",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={"edges": edges.tolist(), "series": series},
|
||||
)
|
||||
|
||||
|
||||
_PRED_GROUP_AXES = ("overall", "energy", "pdg", "material")
|
||||
|
||||
|
||||
def _pred_group_expr(lf: pl.LazyFrame, axis: str, ctx: Context) -> pl.Expr:
|
||||
if axis == "pdg":
|
||||
return pl.col("pdg")
|
||||
if axis == "material":
|
||||
return pl.col("material")
|
||||
if axis == "energy":
|
||||
return _energy_group_expr(lf, np.asarray(ctx.energy_edges))
|
||||
return pl.lit(0, dtype=pl.Int64)
|
||||
|
||||
|
||||
def _pred_group_keys(ctx: Context, axis: str) -> list:
|
||||
if axis == "pdg":
|
||||
return list(ctx.top_pdgs)
|
||||
if axis == "material":
|
||||
return list(ctx.materials)
|
||||
if axis == "energy":
|
||||
return list(range(len(ctx.energy_edges) - 1))
|
||||
return [0]
|
||||
|
||||
|
||||
def _grouped_moments(lf: pl.LazyFrame, group: pl.Expr, value: pl.Expr) -> dict[str, dict[str, float]]:
|
||||
"""``{str(group_key): {"n", "sum", "sumsq"}}`` — one streaming pass, sum-mergeable."""
|
||||
res = (
|
||||
lf.select(group.alias("_g"), value.alias("_v"))
|
||||
.drop_nulls(["_g", "_v"])
|
||||
.group_by("_g")
|
||||
.agg(pl.len().alias("_n"), pl.col("_v").sum().alias("_s"), (pl.col("_v") ** 2).sum().alias("_ss"))
|
||||
.collect(engine="streaming")
|
||||
)
|
||||
return {str(g): {"n": float(n), "sum": float(s), "sumsq": float(ss)} for g, n, s, ss in res.iter_rows()}
|
||||
|
||||
|
||||
def _pred_scorecard_partial(b: Bundle) -> dict:
|
||||
"""Per-(var, prediction, axis) group-keyed truth/pred histograms + residual
|
||||
moments — the shared input to the KS/bias/RMSE scorecards below (three
|
||||
separate specs, each doing its own finalize math over this one compute)."""
|
||||
if not b.predictions:
|
||||
return {"available": False}
|
||||
out: dict[str, dict] = {}
|
||||
for var in PAIRED_SCALARS:
|
||||
if var not in b.ctx.pred_var_ranges:
|
||||
continue
|
||||
edges = b.ctx.pred_marginal_edges(var)
|
||||
|
||||
def _one(ps: PredictionSide, var=var, edges=edges) -> dict | None:
|
||||
if var not in paired_vars_for_coord(ps.coord) or not ps.has_truth:
|
||||
return None
|
||||
resid = pl.col(f"pred_{var}") - pl.col(f"true_{var}")
|
||||
per_axis = {}
|
||||
for axis in _PRED_GROUP_AXES:
|
||||
group = _pred_group_expr(ps.paired, axis, b.ctx)
|
||||
per_axis[axis] = {
|
||||
"true_hist": {
|
||||
str(k): v.tolist() for k, v in hist1d(ps.paired, pl.col(f"true_{var}"), edges, group).items()
|
||||
},
|
||||
"pred_hist": {
|
||||
str(k): v.tolist() for k, v in hist1d(ps.paired, pl.col(f"pred_{var}"), edges, group).items()
|
||||
},
|
||||
"moments": _grouped_moments(ps.paired, group, resid),
|
||||
"true_moments": _grouped_moments(ps.paired, group, pl.col(f"true_{var}")),
|
||||
}
|
||||
return per_axis
|
||||
|
||||
out[var] = _per_prediction(b, _one)
|
||||
return {"available": True, "vars": out}
|
||||
|
||||
|
||||
def _pred_scorecard_matrix(
|
||||
parts: list[dict], ctx: Context, cell: Callable[[dict, dict, str, str], float]
|
||||
) -> dict[str, list[list[float]]]:
|
||||
"""Shared finalize skeleton for the three scorecards: rows = ``PAIRED_SCALARS``,
|
||||
cols = ``_PRED_GROUP_AXES``, one matrix per prediction. ``cell(true_h_or_moments,
|
||||
pred_h_or_moments, axis, group_key)`` computes one entry from that axis'
|
||||
merged group-keyed dict pair (weighted-averaged over the axis' groups)."""
|
||||
names = list(parts[0]["vars"][next(iter(parts[0]["vars"]))])
|
||||
matrices: dict[str, list[list[float]]] = {name: [] for name in names}
|
||||
for var in PAIRED_SCALARS:
|
||||
var_parts = [p["vars"].get(var) for p in parts]
|
||||
for name in names:
|
||||
row: list[float] = []
|
||||
for axis in _PRED_GROUP_AXES:
|
||||
raw_entries = [vp[name][axis] if vp is not None and vp[name] is not None else None for vp in var_parts]
|
||||
if raw_entries[0] is None:
|
||||
row.append(float("nan"))
|
||||
continue
|
||||
# non-None for every chunk: (var, name, axis) availability is
|
||||
# a static fact of the prediction's coord, not chunk-dependent.
|
||||
entries = cast("list[dict]", raw_entries)
|
||||
merged_true_hist = sum_merge([{k: v for k, v in e["true_hist"].items()} for e in entries])
|
||||
merged_pred_hist = sum_merge([{k: v for k, v in e["pred_hist"].items()} for e in entries])
|
||||
merged_moments: dict[str, dict[str, float]] = {}
|
||||
for e in entries:
|
||||
for k, m in e["moments"].items():
|
||||
acc = merged_moments.setdefault(k, {"n": 0.0, "sum": 0.0, "sumsq": 0.0})
|
||||
for f in ("n", "sum", "sumsq"):
|
||||
acc[f] += m[f]
|
||||
merged_true_moments: dict[str, dict[str, float]] = {}
|
||||
for e in entries:
|
||||
for k, m in e["true_moments"].items():
|
||||
acc = merged_true_moments.setdefault(k, {"n": 0.0, "sum": 0.0, "sumsq": 0.0})
|
||||
for f in ("n", "sum", "sumsq"):
|
||||
acc[f] += m[f]
|
||||
dists, weights = [], []
|
||||
for k in _pred_group_keys(ctx, axis):
|
||||
key = str(k)
|
||||
if key not in merged_moments:
|
||||
continue
|
||||
val = cell(
|
||||
{"hist": merged_true_hist.get(key), "moments": merged_true_moments.get(key)},
|
||||
{"hist": merged_pred_hist.get(key), "moments": merged_moments.get(key)},
|
||||
axis,
|
||||
key,
|
||||
)
|
||||
w = merged_moments[key]["n"]
|
||||
if w <= 0 or not np.isfinite(val):
|
||||
continue
|
||||
dists.append(val)
|
||||
weights.append(w)
|
||||
row.append(float(np.average(dists, weights=weights)) if dists else float("nan"))
|
||||
matrices[name].append(row)
|
||||
return matrices
|
||||
|
||||
|
||||
def _pred_ks_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
if not parts[0]["available"] or not parts[0]["vars"]:
|
||||
return _pred_unavailable(
|
||||
"pred_ks_summary", "prediction", "Truth/pred distance summary (KS statistic)", "grouping axis"
|
||||
)
|
||||
|
||||
def _cell(true_side: dict, pred_side: dict, axis: str, key: str) -> float:
|
||||
if pred_side["hist"] is None or true_side["hist"] is None:
|
||||
return float("nan")
|
||||
return _ks_statistic(pred_side["hist"], true_side["hist"])
|
||||
|
||||
matrices = _pred_scorecard_matrix(parts, ctx, _cell)
|
||||
return Reduced(
|
||||
id="pred_ks_summary",
|
||||
family="prediction",
|
||||
kind="heatmap",
|
||||
title="Truth/pred distance summary (KS statistic)",
|
||||
xlabel="grouping axis",
|
||||
payload={
|
||||
"series": matrices,
|
||||
"row_labels": [_PRED_TITLE_NAMES[v] for v in PAIRED_SCALARS],
|
||||
"col_labels": list(_PRED_GROUP_AXES),
|
||||
"ylabel": "variable",
|
||||
"cbar_label": "KS statistic (0 = identical, 1 = maximal mismatch)",
|
||||
"vmin": 0.0,
|
||||
"vmax": 1.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _pred_bias_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
if not parts[0]["available"] or not parts[0]["vars"]:
|
||||
return _pred_unavailable("pred_bias_summary", "prediction", "Relative bias summary", "grouping axis")
|
||||
|
||||
def _cell(true_side: dict, pred_side: dict, axis: str, key: str) -> float:
|
||||
n, s = pred_side["moments"]["n"], pred_side["moments"]["sum"]
|
||||
if n <= 0:
|
||||
return float("nan")
|
||||
bias = s / n
|
||||
tm = true_side["moments"]
|
||||
denom = abs(tm["sum"] / tm["n"]) if tm and tm["n"] > 0 else 0.0
|
||||
return bias / denom if denom > 1e-12 else float("nan")
|
||||
|
||||
matrices = _pred_scorecard_matrix(parts, ctx, _cell)
|
||||
return Reduced(
|
||||
id="pred_bias_summary",
|
||||
family="prediction",
|
||||
kind="heatmap",
|
||||
title="Relative bias summary (mean(pred - true) / mean|true|)",
|
||||
xlabel="grouping axis",
|
||||
payload={
|
||||
"series": matrices,
|
||||
"row_labels": [_PRED_TITLE_NAMES[v] for v in PAIRED_SCALARS],
|
||||
"col_labels": list(_PRED_GROUP_AXES),
|
||||
"ylabel": "variable",
|
||||
"cbar_label": "relative bias",
|
||||
"cmap": "RdBu_r",
|
||||
"vmin": -0.5,
|
||||
"vmax": 0.5,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _pred_rmse_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
if not parts[0]["available"] or not parts[0]["vars"]:
|
||||
return _pred_unavailable("pred_rmse_summary", "prediction", "Relative RMSE summary", "grouping axis")
|
||||
|
||||
def _cell(true_side: dict, pred_side: dict, axis: str, key: str) -> float:
|
||||
n, s, ss = pred_side["moments"]["n"], pred_side["moments"]["sum"], pred_side["moments"]["sumsq"]
|
||||
if n <= 0:
|
||||
return float("nan")
|
||||
rmse = float(np.sqrt(max(ss / n, 0.0)))
|
||||
tm = true_side["moments"]
|
||||
denom = abs(tm["sum"] / tm["n"]) if tm and tm["n"] > 0 else 0.0
|
||||
del s
|
||||
return rmse / denom if denom > 1e-12 else float("nan")
|
||||
|
||||
matrices = _pred_scorecard_matrix(parts, ctx, _cell)
|
||||
return Reduced(
|
||||
id="pred_rmse_summary",
|
||||
family="prediction",
|
||||
kind="heatmap",
|
||||
title="Relative RMSE summary",
|
||||
xlabel="grouping axis",
|
||||
payload={
|
||||
"series": matrices,
|
||||
"row_labels": [_PRED_TITLE_NAMES[v] for v in PAIRED_SCALARS],
|
||||
"col_labels": list(_PRED_GROUP_AXES),
|
||||
"ylabel": "variable",
|
||||
"cbar_label": "relative RMSE",
|
||||
"vmin": 0.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---- confusion matrices -----------------------------------------------
|
||||
|
||||
|
||||
def _n_sec_confusion_edges(cap: int) -> np.ndarray:
|
||||
return np.arange(-0.5, cap + 1.5)
|
||||
|
||||
|
||||
def _pred_n_sec_confusion_partial(b: Bundle) -> dict:
|
||||
if not b.predictions:
|
||||
return {"available": False}
|
||||
cap = b.ctx.pred_n_sec_cap
|
||||
edges = _n_sec_confusion_edges(cap)
|
||||
|
||||
def _one(ps: PredictionSide) -> list[list[int]] | None:
|
||||
if ps.coord != "global":
|
||||
return None
|
||||
mat = hist2d(ps.paired, pl.col("n_sec").clip(0, cap), pl.col("n_sec_pred").clip(0, cap), edges, edges)
|
||||
return mat.tolist()
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _row_normalize(mat: np.ndarray) -> np.ndarray:
|
||||
row_sums = mat.sum(axis=1, keepdims=True)
|
||||
return np.divide(mat, row_sums, out=np.zeros_like(mat), where=row_sums > 0)
|
||||
|
||||
|
||||
def _pred_n_sec_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
title, xlabel = "Secondary-count confusion matrix (row-normalised)", "predicted n_sec"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable("pred_n_sec_confusion", "prediction", title, xlabel)
|
||||
cap = ctx.pred_n_sec_cap
|
||||
labels = [str(i) for i in range(cap)] + [f"{cap}+"]
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list] = {}
|
||||
for name in names:
|
||||
mats = [p["p"][name] for p in parts]
|
||||
if mats[0] is None:
|
||||
continue
|
||||
mat = np.sum([np.asarray(m, dtype=np.float64) for m in mats], axis=0)
|
||||
series[name] = _row_normalize(mat).tolist()
|
||||
return Reduced(
|
||||
id="pred_n_sec_confusion",
|
||||
family="prediction",
|
||||
kind="heatmap",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={
|
||||
"series": series,
|
||||
"row_labels": labels,
|
||||
"col_labels": labels,
|
||||
"ylabel": "true n_sec",
|
||||
"cbar_label": "fraction of true-count rows",
|
||||
"vmin": 0.0,
|
||||
"vmax": 1.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
_OTHER_SEC_KEY = "other"
|
||||
|
||||
|
||||
def _sec_species_key_expr(col: str, top_pdgs: list[int]) -> pl.Expr:
|
||||
return pl.when(pl.col(col).is_in(list(top_pdgs))).then(pl.col(col).cast(pl.Utf8)).otherwise(pl.lit(_OTHER_SEC_KEY))
|
||||
|
||||
|
||||
def _pred_sec_species_confusion_partial(b: Bundle) -> dict:
|
||||
if not b.predictions or not b.ctx.pred_top_sec_pdgs:
|
||||
return {"available": False}
|
||||
top = b.ctx.pred_top_sec_pdgs
|
||||
|
||||
def _one(ps: PredictionSide) -> dict | None:
|
||||
if ps.coord != "global" or not ps.has_truth:
|
||||
return None
|
||||
pf = paired_secondaries(ps.lf)
|
||||
counts = (
|
||||
pf.select(
|
||||
_sec_species_key_expr("true_pdg", top).alias("_t"),
|
||||
_sec_species_key_expr("pred_pdg", top).alias("_p"),
|
||||
)
|
||||
.group_by("_t", "_p")
|
||||
.agg(pl.len().alias("_n"))
|
||||
.collect(engine="streaming")
|
||||
)
|
||||
out: dict[str, dict[str, int]] = {}
|
||||
for t, p, n in counts.iter_rows():
|
||||
out.setdefault(t, {})[p] = out.get(t, {}).get(p, 0) + n
|
||||
return out
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_sec_species_confusion_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
title, xlabel = "Secondary-species confusion matrix (row-normalised)", "predicted species"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable("pred_sec_species_confusion", "prediction", title, xlabel)
|
||||
keys = [str(k) for k in ctx.pred_top_sec_pdgs] + [_OTHER_SEC_KEY]
|
||||
labels = [pdg_label(k) for k in ctx.pred_top_sec_pdgs] + [_OTHER_SEC_KEY]
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
merged: dict[str, dict[str, int]] = {}
|
||||
for e in entries:
|
||||
for t, row in e.items():
|
||||
acc = merged.setdefault(t, {})
|
||||
for p_key, n in row.items():
|
||||
acc[p_key] = acc.get(p_key, 0) + n
|
||||
mat = []
|
||||
for t in keys:
|
||||
row_counts = merged.get(t, {})
|
||||
row_vals = [row_counts.get(p_key, 0) for p_key in keys]
|
||||
total = sum(row_vals)
|
||||
mat.append([v / total if total > 0 else 0.0 for v in row_vals])
|
||||
series[name] = mat
|
||||
return Reduced(
|
||||
id="pred_sec_species_confusion",
|
||||
family="prediction",
|
||||
kind="heatmap",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={
|
||||
"series": series,
|
||||
"row_labels": labels,
|
||||
"col_labels": labels,
|
||||
"ylabel": "true species",
|
||||
"cbar_label": "fraction of true-species rows",
|
||||
"vmin": 0.0,
|
||||
"vmax": 1.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---- physics consistency -----------------------------------------------
|
||||
|
||||
|
||||
def _pred_dir_alignment_partial(b: Bundle, kind: str) -> dict:
|
||||
if not b.predictions:
|
||||
return {"available": False}
|
||||
edges = np.linspace(-1.0, 1.0, b.ctx.n_marginal_bins + 1)
|
||||
nb = len(edges) - 1
|
||||
|
||||
def _one(ps: PredictionSide) -> list[int] | None:
|
||||
if not ps.has_truth:
|
||||
return None
|
||||
return _counts(hist1d(ps.lf, dir_alignment_expr(ps.coord, kind), edges), 0, nb)
|
||||
|
||||
return {"available": True, "edges": edges.tolist(), "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_dir_alignment_finalize(parts: list[dict], ctx: Context, kind: str, title: str) -> Reduced:
|
||||
xlabel = f"cos(angle) between true and predicted {kind}_dir"
|
||||
spec_id = f"pred_dir_alignment_{kind}"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable(spec_id, "prediction", title, xlabel)
|
||||
edges = np.asarray(parts[0]["edges"])
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list[int]] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
series[name] = [int(x) for x in sum_merge([{"c": e} for e in entries])["c"]]
|
||||
return Reduced(
|
||||
id=spec_id,
|
||||
family="prediction",
|
||||
kind="single_hist",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={"edges": edges.tolist(), "series": series, "log_y": True},
|
||||
)
|
||||
|
||||
|
||||
def _rate_partial(lf: pl.LazyFrame, cond: pl.Expr) -> tuple[int, int]:
|
||||
res = lf.select(cond.cast(pl.Int64).sum().alias("v"), pl.len().alias("n")).collect(engine="streaming")
|
||||
return int(res["v"][0]), int(res["n"][0])
|
||||
|
||||
|
||||
_CONSTRAINT_LABELS = ["post_dir_norm", "travel_dir_norm", "step_length_neg", "edep_neg", "delta_e_neg"]
|
||||
|
||||
|
||||
def _pred_constraint_partial(b: Bundle) -> dict:
|
||||
"""Rate of physical-constraint violations in the *predicted* values.
|
||||
|
||||
Direction unit-norm is only a meaningful check for `--coord local`
|
||||
(`giant predict --coord global` already renormalises before writing, so
|
||||
it is compliant by construction there — reported as 0/1 rather than
|
||||
skipped, since "no violations" is still the correct answer).
|
||||
"""
|
||||
if not b.predictions:
|
||||
return {"available": False}
|
||||
|
||||
def _one(ps: PredictionSide) -> dict:
|
||||
out: dict[str, tuple[int, int]] = {}
|
||||
if ps.coord == "local":
|
||||
for label, cols in (
|
||||
("post_dir_norm", ["pred_post_dx", "pred_post_dy", "pred_post_dz"]),
|
||||
("travel_dir_norm", ["pred_travel_dx", "pred_travel_dy", "pred_travel_dz"]),
|
||||
):
|
||||
norm = pl.sum_horizontal([pl.col(c) ** 2 for c in cols]).sqrt()
|
||||
out[label] = _rate_partial(ps.lf, (norm - 1).abs() > 0.05)
|
||||
else:
|
||||
out["post_dir_norm"] = (0, 1)
|
||||
out["travel_dir_norm"] = (0, 1)
|
||||
for label, col in (
|
||||
("step_length_neg", "pred_step_length"),
|
||||
("edep_neg", "pred_edep"),
|
||||
("delta_e_neg", "pred_delta_e"),
|
||||
):
|
||||
out[label] = _rate_partial(ps.paired, pl.col(col) < 0)
|
||||
return {k: list(v) for k, v in out.items()}
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _pred_constraint_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
title, xlabel = "Physical-constraint violation rate", "check"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable("pred_constraint_violations", "prediction", title, xlabel)
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list[float]] = {}
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
rates = []
|
||||
for check in _CONSTRAINT_LABELS:
|
||||
v = sum(e[check][0] for e in entries)
|
||||
n = sum(e[check][1] for e in entries)
|
||||
rates.append(v / n if n > 0 else 0.0)
|
||||
series[name] = rates
|
||||
return Reduced(
|
||||
id="pred_constraint_violations",
|
||||
family="prediction",
|
||||
kind="bar",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={"labels": _CONSTRAINT_LABELS, "series": series, "ylabel": "violation rate"},
|
||||
)
|
||||
|
||||
|
||||
_CORR_PAIRS = [(a, b) for i, a in enumerate(PAIRED_SCALARS) for b in PAIRED_SCALARS[i + 1 :]]
|
||||
|
||||
|
||||
def _pred_correlation_partial(b: Bundle) -> dict:
|
||||
if not b.predictions:
|
||||
return {"available": False}
|
||||
|
||||
def _one(ps: PredictionSide) -> dict | None:
|
||||
if not ps.has_truth:
|
||||
return None
|
||||
vs = [v for v in PAIRED_SCALARS if v in paired_vars_for_coord(ps.coord)]
|
||||
pairs = [(a, c) for a, c in _CORR_PAIRS if a in vs and c in vs]
|
||||
agg = [pl.len().alias("_n")]
|
||||
for prefix in ("true", "pred"):
|
||||
for v in vs:
|
||||
agg.append(pl.col(f"{prefix}_{v}").sum().alias(f"s_{prefix}_{v}"))
|
||||
agg.append((pl.col(f"{prefix}_{v}") ** 2).sum().alias(f"ss_{prefix}_{v}"))
|
||||
for a, c in pairs:
|
||||
agg.append((pl.col(f"{prefix}_{a}") * pl.col(f"{prefix}_{c}")).sum().alias(f"sxy_{prefix}_{a}_{c}"))
|
||||
res = ps.paired.select(agg).collect(engine="streaming")
|
||||
return {"vars": vs, "pairs": pairs, "row": {k: float(res[k][0]) for k in res.columns}}
|
||||
|
||||
return {"available": True, "p": _per_prediction(b, _one)}
|
||||
|
||||
|
||||
def _corr_matrix(vs: list[str], pairs: list[tuple[str, str]], row: dict[str, float], prefix: str) -> np.ndarray:
|
||||
n = row["_n"]
|
||||
mat = np.eye(len(vs))
|
||||
idx = {v: i for i, v in enumerate(vs)}
|
||||
for a, c in pairs:
|
||||
sa, sc = row[f"s_{prefix}_{a}"], row[f"s_{prefix}_{c}"]
|
||||
ssa, ssc = row[f"ss_{prefix}_{a}"], row[f"ss_{prefix}_{c}"]
|
||||
sxy = row[f"sxy_{prefix}_{a}_{c}"]
|
||||
cov = sxy / n - (sa / n) * (sc / n)
|
||||
var_a = ssa / n - (sa / n) ** 2
|
||||
var_c = ssc / n - (sc / n) ** 2
|
||||
denom = np.sqrt(max(var_a, 0.0) * max(var_c, 0.0))
|
||||
corr = cov / denom if denom > 1e-12 else float("nan")
|
||||
mat[idx[a], idx[c]] = mat[idx[c], idx[a]] = corr
|
||||
return mat
|
||||
|
||||
|
||||
def _pred_correlation_finalize(parts: list[dict], ctx: Context) -> Reduced:
|
||||
title, xlabel = "Correlation delta (corr(pred) - corr(true))", "variable"
|
||||
if not parts[0]["available"]:
|
||||
return _pred_unavailable("pred_correlation_delta", "prediction", title, xlabel)
|
||||
names = list(parts[0]["p"])
|
||||
series: dict[str, list] = {}
|
||||
labels: list[str] = []
|
||||
for name in names:
|
||||
entries = [p["p"][name] for p in parts]
|
||||
if entries[0] is None:
|
||||
continue
|
||||
vs, pairs = entries[0]["vars"], [tuple(p) for p in entries[0]["pairs"]]
|
||||
labels = [_PRED_TITLE_NAMES[v] for v in vs] # same for every prediction: all share one coord (see condor.py)
|
||||
merged_row: dict[str, float] = {}
|
||||
for e in entries:
|
||||
for k, v in e["row"].items():
|
||||
merged_row[k] = merged_row.get(k, 0.0) + v
|
||||
delta = _corr_matrix(vs, pairs, merged_row, "pred") - _corr_matrix(vs, pairs, merged_row, "true")
|
||||
series[name] = delta.tolist()
|
||||
return Reduced(
|
||||
id="pred_correlation_delta",
|
||||
family="prediction",
|
||||
kind="heatmap",
|
||||
title=title,
|
||||
xlabel=xlabel,
|
||||
payload={
|
||||
"series": series,
|
||||
"row_labels": labels,
|
||||
"col_labels": labels,
|
||||
"ylabel": "variable",
|
||||
"cbar_label": "corr(pred) - corr(true)",
|
||||
"cmap": "RdBu_r",
|
||||
"vmin": -1.0,
|
||||
"vmax": 1.0,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def build_catalog() -> list[PlotSpec]:
|
||||
"""All concrete plot specs, each with a unique id."""
|
||||
specs: list[PlotSpec] = []
|
||||
@@ -1279,6 +2114,110 @@ def build_catalog() -> list[PlotSpec]:
|
||||
chunkable=False,
|
||||
),
|
||||
]
|
||||
|
||||
for var in PAIRED_VARS:
|
||||
specs.append(
|
||||
PlotSpec(
|
||||
f"pred_marginal_{var}",
|
||||
"prediction",
|
||||
compute_partial=lambda b, v=var: _pred_marginal_partial(b, v),
|
||||
finalize=lambda parts, ctx, v=var: _pred_marginal_finalize(parts, ctx, v),
|
||||
)
|
||||
)
|
||||
specs.append(
|
||||
PlotSpec(
|
||||
f"pred_scatter_{var}",
|
||||
"prediction",
|
||||
compute_partial=lambda b, v=var: _pred_scatter_partial(b, v),
|
||||
finalize=lambda parts, ctx, v=var: _pred_scatter_finalize(parts, ctx, v),
|
||||
)
|
||||
)
|
||||
specs.append(
|
||||
PlotSpec(
|
||||
f"pred_residual_{var}",
|
||||
"prediction",
|
||||
compute_partial=lambda b, v=var: _pred_residual_partial(b, v),
|
||||
finalize=lambda parts, ctx, v=var: _pred_residual_finalize(parts, ctx, v),
|
||||
)
|
||||
)
|
||||
for var in PAIRED_SCALARS:
|
||||
specs.append(
|
||||
PlotSpec(
|
||||
f"pred_relative_residual_{var}",
|
||||
"prediction",
|
||||
compute_partial=lambda b, v=var: _pred_relative_residual_partial(b, v),
|
||||
finalize=lambda parts, ctx, v=var: _pred_relative_residual_finalize(parts, ctx, v),
|
||||
)
|
||||
)
|
||||
specs.append(
|
||||
PlotSpec(
|
||||
f"pred_residual_profile_{var}",
|
||||
"prediction",
|
||||
compute_partial=lambda b, v=var: _pred_residual_profile_partial(b, v),
|
||||
finalize=lambda parts, ctx, v=var: _pred_residual_profile_finalize(parts, ctx, v),
|
||||
)
|
||||
)
|
||||
|
||||
specs += [
|
||||
PlotSpec(
|
||||
"pred_ks_summary",
|
||||
"prediction",
|
||||
compute_partial=_pred_scorecard_partial,
|
||||
finalize=_pred_ks_finalize,
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_bias_summary",
|
||||
"prediction",
|
||||
compute_partial=_pred_scorecard_partial,
|
||||
finalize=_pred_bias_finalize,
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_rmse_summary",
|
||||
"prediction",
|
||||
compute_partial=_pred_scorecard_partial,
|
||||
finalize=_pred_rmse_finalize,
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_n_sec_confusion",
|
||||
"prediction",
|
||||
compute_partial=_pred_n_sec_confusion_partial,
|
||||
finalize=_pred_n_sec_confusion_finalize,
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_sec_species_confusion",
|
||||
"prediction",
|
||||
compute_partial=_pred_sec_species_confusion_partial,
|
||||
finalize=_pred_sec_species_confusion_finalize,
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_dir_alignment_post",
|
||||
"prediction",
|
||||
compute_partial=lambda b: _pred_dir_alignment_partial(b, "post"),
|
||||
finalize=lambda parts, ctx: _pred_dir_alignment_finalize(
|
||||
parts, ctx, "post", "Post-direction alignment (true vs predicted)"
|
||||
),
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_dir_alignment_travel",
|
||||
"prediction",
|
||||
compute_partial=lambda b: _pred_dir_alignment_partial(b, "travel"),
|
||||
finalize=lambda parts, ctx: _pred_dir_alignment_finalize(
|
||||
parts, ctx, "travel", "Travel-direction alignment (true vs predicted)"
|
||||
),
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_constraint_violations",
|
||||
"prediction",
|
||||
compute_partial=_pred_constraint_partial,
|
||||
finalize=_pred_constraint_finalize,
|
||||
),
|
||||
PlotSpec(
|
||||
"pred_correlation_delta",
|
||||
"prediction",
|
||||
compute_partial=_pred_correlation_partial,
|
||||
finalize=_pred_correlation_finalize,
|
||||
),
|
||||
]
|
||||
return specs
|
||||
|
||||
|
||||
|
||||
+125
-13
@@ -56,11 +56,14 @@ import yaml
|
||||
|
||||
from giant.analysis.catalog import Bundle, catalog_ids, get_spec
|
||||
from giant.analysis.context import Context, build_context
|
||||
from giant.analysis.prediction import PredictionSpec, open_prediction
|
||||
from giant.analysis.reduced import Partial
|
||||
from giant.analysis.runtime_estimate import estimate_runtime_s
|
||||
from giant.analysis.sources import RolloutSpec, Side, open_side
|
||||
|
||||
# Keys copied verbatim from a rollout YAML into each plot's gallery metadata.
|
||||
# Keys copied verbatim from a rollout or prediction YAML into each plot's
|
||||
# gallery metadata. Rollout-only and predict-only keys both live here —
|
||||
# `_plot_meta` copies only whichever of these are present in a given YAML.
|
||||
_PLOT_META_KEYS = (
|
||||
"prediction_id",
|
||||
"checkpoint",
|
||||
@@ -84,10 +87,20 @@ _PLOT_META_KEYS = (
|
||||
"termination_reason_counts",
|
||||
"timing",
|
||||
"model_config",
|
||||
"config_overrides",
|
||||
"training_epoch",
|
||||
"best_val_loss",
|
||||
"training_config",
|
||||
"training_meta",
|
||||
# giant predict only (giant/cli.py's predict command).
|
||||
"coord",
|
||||
"has_truth",
|
||||
"schema_version",
|
||||
"n_input_rows",
|
||||
"n_files",
|
||||
"n_skipped_rows",
|
||||
"unknown_pdg_counts",
|
||||
"batch_size_auto",
|
||||
# Diagnostic — only present when giant rollout ran under
|
||||
# stage2_model.particle_type.target="embedding" (see giant/cli.py's
|
||||
# rollout command and giant.rollout.L1DistCollector); absent otherwise,
|
||||
@@ -163,6 +176,72 @@ def load_rollout_yamls(
|
||||
return [LoadedRollout(name=n, yaml=y) for n, y in zip(names, yamls)], yamls[0]["dataset"]
|
||||
|
||||
|
||||
def load_prediction_yaml(path: str | Path) -> dict:
|
||||
"""Load a `giant predict` YAML sidecar, requiring the two file paths."""
|
||||
d = yaml.safe_load(Path(path).read_text())
|
||||
for key in ("output", "dataset"):
|
||||
if key not in d:
|
||||
raise ValueError(
|
||||
f"{path} is not a prediction YAML (missing {key!r}); expected the "
|
||||
"sidecar `giant predict` writes next to the checkpoint"
|
||||
)
|
||||
if d.get("kind") not in (None, "prediction"):
|
||||
raise ValueError(f"{path} has kind={d.get('kind')!r}, not a prediction YAML")
|
||||
return d
|
||||
|
||||
|
||||
@dataclass
|
||||
class LoadedPrediction:
|
||||
"""One prediction YAML plus its resolved series ``name`` and predict ``coord``."""
|
||||
|
||||
name: str
|
||||
yaml: dict
|
||||
coord: str
|
||||
|
||||
|
||||
def load_prediction_yamls(
|
||||
paths: Sequence[str | Path], reference: str, labels: Sequence[str] | None = None
|
||||
) -> list[LoadedPrediction]:
|
||||
"""Load every prediction YAML, resolve each one's series name, and verify
|
||||
they're seeded from the same ``reference`` as the rollout(s) and all share
|
||||
one predict ``--coord`` (direction components mean different things in
|
||||
the two coords — see ``giant.analysis.prediction``'s module docstring).
|
||||
|
||||
Names follow the same convention as ``load_rollout_yamls``: an explicit
|
||||
``labels[i]`` if given, else the YAML stem for N>1, or ``"prediction"``
|
||||
for the single-YAML case.
|
||||
"""
|
||||
if labels and len(labels) != len(paths):
|
||||
raise ValueError(
|
||||
f"--prediction-label given {len(labels)} time(s) but {len(paths)} --prediction YAML(s) were passed"
|
||||
)
|
||||
yamls = [load_prediction_yaml(p) for p in paths]
|
||||
if labels:
|
||||
names = list(labels)
|
||||
elif len(paths) == 1:
|
||||
names = ["prediction"]
|
||||
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"prediction series names collide: {dupes} — pass --prediction-label to disambiguate")
|
||||
|
||||
bad_ref = [(p, y) for p, y in zip(paths, yamls) if str(y["dataset"]) != str(reference)]
|
||||
if bad_ref:
|
||||
detail = "\n".join(f" {p}: dataset={y['dataset']!r}" for p, y in bad_ref)
|
||||
raise ValueError(
|
||||
f"every --prediction must be seeded from the same reference as the rollout(s) "
|
||||
f"({reference!r}) — mismatched:\n{detail}"
|
||||
)
|
||||
|
||||
coords = {str(p): open_prediction(y["output"]).coord for p, y in zip(paths, yamls)}
|
||||
if len(set(coords.values())) > 1:
|
||||
detail = "\n".join(f" {p}: coord={c!r}" for p, c in coords.items())
|
||||
raise ValueError(f"every --prediction in one run must share one --coord — got:\n{detail}")
|
||||
|
||||
return [LoadedPrediction(name=n, yaml=y, coord=coords[str(p)]) for n, y, p in zip(names, yamls, paths)]
|
||||
|
||||
|
||||
def _run_tag(y: dict) -> str:
|
||||
rollout = Path(y["output"])
|
||||
return str(y.get("prediction_id") or rollout.stem)[:8]
|
||||
@@ -223,6 +302,10 @@ class RunMeta:
|
||||
# Empty/0 on run directories written before this field existed.
|
||||
rows_per_chunk: list[int] = field(default_factory=list)
|
||||
total_rows: int = 0
|
||||
# `giant predict` inputs (the paired-truth "prediction" family) — same
|
||||
# shape as `rollouts`. Empty on a run with no --prediction, so old
|
||||
# run_meta.json files still load.
|
||||
predictions: list[dict] = field(default_factory=list)
|
||||
|
||||
def save(self, path: str | Path) -> None:
|
||||
Path(path).write_text(json.dumps(self.__dict__, indent=2))
|
||||
@@ -232,8 +315,13 @@ class RunMeta:
|
||||
return cls(**json.loads(Path(path).read_text()))
|
||||
|
||||
|
||||
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.
|
||||
def _rows_per_chunk(
|
||||
rollouts: list[str | Path],
|
||||
reference: str | Path,
|
||||
n_chunks: int,
|
||||
predictions: Sequence[str | Path] = (),
|
||||
) -> list[int]:
|
||||
"""Combined rollout+reference+prediction 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
|
||||
@@ -249,7 +337,11 @@ def _rows_per_chunk(rollouts: list[str | Path], reference: str | Path, n_chunks:
|
||||
)
|
||||
|
||||
out = [0] * n_chunks
|
||||
sides = [open_side(reference, Side.reference)] + [open_side(r, Side.rollout) for r in rollouts]
|
||||
sides = (
|
||||
[open_side(reference, Side.reference)]
|
||||
+ [open_side(r, Side.rollout) for r in rollouts]
|
||||
+ [open_prediction(p).lf for p in predictions]
|
||||
)
|
||||
for lf in sides:
|
||||
df = counts(lf)
|
||||
for c, n in zip(df["_c"].to_list(), df["n"].to_list()):
|
||||
@@ -263,26 +355,33 @@ def prep(
|
||||
n_chunks: int = 1,
|
||||
default_base: str | Path | None = None,
|
||||
labels: Sequence[str] | None = None,
|
||||
prediction_yamls: Sequence[str | Path] = (),
|
||||
prediction_labels: Sequence[str] | None = None,
|
||||
**ctx_kwargs,
|
||||
) -> Path:
|
||||
"""Read the rollout YAML(s), build the shared context, and lay out the run dir.
|
||||
"""Read the rollout (+ optional prediction) 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, and ``load_rollout_yamls`` for how
|
||||
``labels``/YAML stems resolve each rollout's series name.
|
||||
resolve the actual directory, ``load_rollout_yamls`` for how
|
||||
``labels``/YAML stems resolve each rollout's series name, and
|
||||
``load_prediction_yamls`` for the same on ``prediction_yamls`` (which,
|
||||
unlike rollouts, is optional — the ``prediction`` plot family degrades to
|
||||
``kind="unavailable"`` when it's empty).
|
||||
|
||||
Clears any existing ``reduced_partial/``/``reduced/`` from a prior prep of
|
||||
this same ``run_dir``: partial files carry no record of what context
|
||||
(``n_chunks``, bin edges, group sets) they were computed under, so
|
||||
re-prepping with a different ``n_chunks``/``**ctx_kwargs`` (or after the
|
||||
rollout/reference files changed) would otherwise let ``merge_one`` silently
|
||||
merge stale partials against the new ``shared.json``.
|
||||
rollout/reference/prediction files changed) would otherwise let
|
||||
``merge_one`` silently merge stale partials against the new ``shared.json``.
|
||||
"""
|
||||
loaded, reference = load_rollout_yamls(list(rollout_yamls), labels)
|
||||
loaded_preds = load_prediction_yamls(list(prediction_yamls), reference, prediction_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)
|
||||
|
||||
@@ -292,14 +391,20 @@ def prep(
|
||||
shutil.rmtree(stale_dir)
|
||||
|
||||
rollout_specs = [RolloutSpec(name=lr.name, source=lr.yaml["output"]) for lr in loaded]
|
||||
ctx = build_context(rollout_specs, reference, **ctx_kwargs)
|
||||
pred_specs = [PredictionSpec(name=lp.name, source=lp.yaml["output"]) for lp in loaded_preds]
|
||||
ctx = build_context(rollout_specs, reference, predictions=pred_specs, **ctx_kwargs)
|
||||
ctx.save(run_path / "shared.json")
|
||||
|
||||
rows_per_chunk = _rows_per_chunk([lr.yaml["output"] for lr in loaded], reference, n_chunks)
|
||||
rows_per_chunk = _rows_per_chunk(
|
||||
[lr.yaml["output"] for lr in loaded], reference, n_chunks, [lp.yaml["output"] for lp in loaded_preds]
|
||||
)
|
||||
|
||||
rollouts_meta = [
|
||||
{"name": lr.name, "path": str(lr.yaml["output"]), "plot_meta": _plot_meta(lr.yaml)} for lr in loaded
|
||||
]
|
||||
predictions_meta = [
|
||||
{"name": lp.name, "path": str(lp.yaml["output"]), "plot_meta": _plot_meta(lp.yaml)} for lp in loaded_preds
|
||||
]
|
||||
ckpts = ", ".join(Path(lr.yaml.get("checkpoint", "")).name or "rollout" for lr in loaded)
|
||||
|
||||
RunMeta(
|
||||
@@ -310,6 +415,7 @@ def prep(
|
||||
n_chunks=n_chunks,
|
||||
rows_per_chunk=rows_per_chunk,
|
||||
total_rows=sum(rows_per_chunk),
|
||||
predictions=predictions_meta,
|
||||
).save(run_path / "run_meta.json")
|
||||
return run_path
|
||||
|
||||
@@ -327,12 +433,15 @@ def compute_reduced(
|
||||
out: str | Path,
|
||||
chunk_index: int = 0,
|
||||
n_chunks: int = 1,
|
||||
predictions: Sequence[dict] = (),
|
||||
) -> Path:
|
||||
"""Core: run one (plot, chunk)'s partial reduction against explicit paths.
|
||||
|
||||
``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?,
|
||||
"timing"?}, ...]``, one per rollout series (insertion order preserved
|
||||
through to every plot's ``Reduced.payload["series"]``).
|
||||
through to every plot's ``Reduced.payload["series"]``). ``predictions``:
|
||||
``[{"name", "path"}, ...]``, one per ``giant predict`` series (the
|
||||
``prediction`` family; empty on a run with no ``--prediction``).
|
||||
|
||||
Writes a ``Partial`` JSON — the raw, not-yet-merged output of
|
||||
``PlotSpec.compute_partial`` — never a finished ``Reduced``; ``merge_one``
|
||||
@@ -357,7 +466,8 @@ def compute_reduced(
|
||||
)
|
||||
for r in rollouts
|
||||
]
|
||||
bundle = Bundle.open(rollout_specs, reference, ctx, chunk=(chunk_index, effective_n))
|
||||
pred_specs = [PredictionSpec(name=p["name"], source=p["path"]) for p in predictions]
|
||||
bundle = Bundle.open(rollout_specs, reference, ctx, chunk=(chunk_index, effective_n), predictions=pred_specs)
|
||||
partial = Partial(
|
||||
id=spec_id,
|
||||
family=spec.family,
|
||||
@@ -383,6 +493,7 @@ def compute_one(spec_id: str, run_dir: str | Path, chunk_index: int = 0) -> Path
|
||||
}
|
||||
for ro in meta.rollouts
|
||||
]
|
||||
predictions = [{"name": p["name"], "path": p["path"]} for p in meta.predictions]
|
||||
return compute_reduced(
|
||||
spec_id,
|
||||
rollouts,
|
||||
@@ -391,6 +502,7 @@ def compute_one(spec_id: str, run_dir: str | Path, chunk_index: int = 0) -> Path
|
||||
run_path / "reduced_partial" / f"{spec_id}__{chunk_index}.json",
|
||||
chunk_index=chunk_index,
|
||||
n_chunks=meta.n_chunks,
|
||||
predictions=predictions,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -20,6 +20,7 @@ import numpy as np
|
||||
import polars as pl
|
||||
|
||||
from giant.analysis.grouping import energy_bin_edges
|
||||
from giant.analysis.prediction import PredictionSpec, open_prediction, paired_vars_for_coord, prediction_secondaries
|
||||
from giant.analysis.reduce import (
|
||||
attach_entry_axis,
|
||||
depth_expr,
|
||||
@@ -44,6 +45,12 @@ class Context:
|
||||
sec_energy_range: tuple[float, float]
|
||||
n_sec_bins: int
|
||||
n_events: dict[str, int] = field(default_factory=dict)
|
||||
# -- giant predict (paired truth/pred comparison) — empty when no
|
||||
# --prediction was given to `prep`, so old shared.json files still load.
|
||||
pred_var_ranges: dict[str, tuple[float, float]] = field(default_factory=dict)
|
||||
pred_residual_ranges: dict[str, tuple[float, float]] = field(default_factory=dict)
|
||||
pred_n_sec_cap: int = 10
|
||||
pred_top_sec_pdgs: list[int] = field(default_factory=list)
|
||||
|
||||
# -- (de)serialization -------------------------------------------------
|
||||
def save(self, path: str | Path) -> None:
|
||||
@@ -54,6 +61,10 @@ class Context:
|
||||
d = json.loads(Path(path).read_text())
|
||||
d["var_ranges"] = {k: tuple(v) for k, v in d["var_ranges"].items()}
|
||||
d["sec_energy_range"] = tuple(d["sec_energy_range"])
|
||||
if "pred_var_ranges" in d:
|
||||
d["pred_var_ranges"] = {k: tuple(v) for k, v in d["pred_var_ranges"].items()}
|
||||
if "pred_residual_ranges" in d:
|
||||
d["pred_residual_ranges"] = {k: tuple(v) for k, v in d["pred_residual_ranges"].items()}
|
||||
return cls(**d)
|
||||
|
||||
# -- convenience -------------------------------------------------------
|
||||
@@ -61,6 +72,14 @@ class Context:
|
||||
lo, hi = self.var_ranges[var]
|
||||
return np.linspace(lo, hi, self.n_marginal_bins + 1)
|
||||
|
||||
def pred_marginal_edges(self, var: str) -> np.ndarray:
|
||||
lo, hi = self.pred_var_ranges[var]
|
||||
return np.linspace(lo, hi, self.n_marginal_bins + 1)
|
||||
|
||||
def pred_residual_edges(self, var: str) -> np.ndarray:
|
||||
lo, hi = self.pred_residual_ranges[var]
|
||||
return np.linspace(lo, hi, self.n_marginal_bins + 1)
|
||||
|
||||
|
||||
_LO_Q, _HI_Q = 0.001, 0.999
|
||||
|
||||
@@ -87,10 +106,13 @@ def build_context(
|
||||
rollouts: list[RolloutSpec],
|
||||
reference: str | Path | pl.LazyFrame,
|
||||
*,
|
||||
predictions: list[PredictionSpec] | None = None,
|
||||
n_energy_bins: int = 4,
|
||||
n_marginal_bins: int = 50,
|
||||
n_sec_bins: int = 40,
|
||||
top_k_pdg: int = 6,
|
||||
pred_n_sec_cap: int = 10,
|
||||
top_k_sec_pdg: int = 8,
|
||||
sample_rows: int = 1_000_000,
|
||||
seed: int = 0,
|
||||
) -> Context:
|
||||
@@ -165,6 +187,44 @@ def build_context(
|
||||
}
|
||||
sec_energy_range = _combined_quantiles([t_se, *r_se.values()], _LO_Q, _HI_Q)
|
||||
|
||||
# giant predict: paired truth/pred ranges + residual ranges + secondary
|
||||
# species vocab, all over the union of every prediction's `paired` frame.
|
||||
pred_var_ranges: dict[str, tuple[float, float]] = {}
|
||||
pred_residual_ranges: dict[str, tuple[float, float]] = {}
|
||||
top_sec_pdgs: list[int] = []
|
||||
if predictions:
|
||||
sides = {ps.name: open_prediction(ps.source) for ps in predictions}
|
||||
present_vars = sorted(set().union(*(paired_vars_for_coord(s.coord) for s in sides.values())))
|
||||
for var in present_vars:
|
||||
true_samples, pred_samples, residual_samples = [], [], []
|
||||
for s in sides.values():
|
||||
if var not in paired_vars_for_coord(s.coord):
|
||||
continue
|
||||
cols = [f"pred_{var}"] + ([f"true_{var}"] if s.has_truth else [])
|
||||
sample = _row_subsample(s.paired.select(cols), sample_rows, seed).collect(engine="streaming")
|
||||
pred_samples.append(sample[f"pred_{var}"].to_numpy())
|
||||
if s.has_truth:
|
||||
true_samples.append(sample[f"true_{var}"].to_numpy())
|
||||
residual_samples.append(sample[f"pred_{var}"].to_numpy() - sample[f"true_{var}"].to_numpy())
|
||||
pred_var_ranges[var] = _combined_quantiles([*true_samples, *pred_samples], _LO_Q, _HI_Q)
|
||||
if residual_samples:
|
||||
pred_residual_ranges[var] = _combined_quantiles(residual_samples, _LO_Q, _HI_Q)
|
||||
|
||||
sec_pdg_counts: dict[int, int] = {}
|
||||
for s in sides.values():
|
||||
if s.coord != "global" or not s.has_truth:
|
||||
continue
|
||||
for prefix in ("true", "pred"):
|
||||
counts = (
|
||||
prediction_secondaries(s.lf, prefix)
|
||||
.group_by("pdg")
|
||||
.agg(pl.len().alias("n"))
|
||||
.collect(engine="streaming")
|
||||
)
|
||||
for pdg, n in zip(counts["pdg"].to_list(), counts["n"].to_list()):
|
||||
sec_pdg_counts[pdg] = sec_pdg_counts.get(pdg, 0) + n
|
||||
top_sec_pdgs = [pdg for pdg, _ in sorted(sec_pdg_counts.items(), key=lambda kv: -kv[1])[:top_k_sec_pdg]]
|
||||
|
||||
return Context(
|
||||
n_marginal_bins=n_marginal_bins,
|
||||
var_ranges=var_ranges,
|
||||
@@ -173,6 +233,10 @@ def build_context(
|
||||
materials=materials,
|
||||
depth_edges=[float(x) for x in depth_edges],
|
||||
transverse_edges=[float(x) for x in transverse_edges],
|
||||
pred_var_ranges=pred_var_ranges,
|
||||
pred_residual_ranges=pred_residual_ranges,
|
||||
pred_n_sec_cap=pred_n_sec_cap,
|
||||
pred_top_sec_pdgs=top_sec_pdgs,
|
||||
sec_energy_range=sec_energy_range,
|
||||
n_sec_bins=n_sec_bins,
|
||||
n_events={
|
||||
|
||||
@@ -0,0 +1,315 @@
|
||||
"""Canonical paired truth/prediction LazyFrame for `giant predict` output.
|
||||
|
||||
Unlike a `giant rollout` (an unpaired, freely-generated shower), `giant predict`
|
||||
runs the model once per real pre-step state, so every output row has a
|
||||
matching truth row — a paired comparison, not a distribution comparison. This
|
||||
module builds one canonical **paired** LazyFrame per prediction file, in
|
||||
either coord mode `giant predict` supports, so every catalog spec in the
|
||||
`prediction` family is coord-agnostic:
|
||||
|
||||
event_id, pdg, material, pre_E, n_sec, n_sec_pred,
|
||||
true_<var>, pred_<var> for var in PAIRED_VARS
|
||||
|
||||
`--coord global` (v3+, `--truth` on) already carries physical `true_*`/`pred_*`-
|
||||
shaped columns directly. `--coord local` carries the raw 9D `true_{name}`/
|
||||
`pred_{name}` model-space target (`LOCAL_TARGET_NAMES`) instead — its two
|
||||
ALR energy logits are decoded into physical `edep`/`delta_e` with the same
|
||||
softmax-against-`pre_E` expressions `giant.data.transforms.energy_simplex_decode`
|
||||
uses, resurrected from the pre-package-rewrite `giant/analysis.py` (see
|
||||
`_edep_pl`/`_delta_e_pl`/`_raw_dim_expr` there). Direction components differ in
|
||||
*meaning* between the two coords (world vs. local frame), so a run must not mix
|
||||
them — `condor.load_prediction_yamls` enforces one coord across every
|
||||
prediction in a run.
|
||||
|
||||
Secondaries only exist in `--coord global --truth` output (local mode never
|
||||
samples stage 2); `paired_secondaries` is `None` otherwise, and secondary-based
|
||||
specs render `kind="unavailable"` instead of raising.
|
||||
|
||||
plotstyle-free (runs on HTCondor workers).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
import polars as pl
|
||||
import pyarrow.parquet as pq
|
||||
|
||||
from giant.constants import (
|
||||
LOCAL_TARGET_NAMES,
|
||||
PREDICT_COORD_METADATA_KEY,
|
||||
PREDICT_TRUTH_METADATA_KEY,
|
||||
ROLLOUT_COORD_VALUE,
|
||||
)
|
||||
|
||||
# The paired scalar/direction variables every coord mode can produce, in
|
||||
# physical units (mm / MeV) regardless of source coord.
|
||||
PAIRED_SCALARS: tuple[str, ...] = ("step_length", "edep", "delta_e", "post_E")
|
||||
PAIRED_VARS: tuple[str, ...] = (*PAIRED_SCALARS, "cos_scatter", "cos_travel")
|
||||
|
||||
_LOG_EPS = 1e-6
|
||||
|
||||
|
||||
@dataclass
|
||||
class PredictionSpec:
|
||||
"""One named prediction input, as fed to `build_context`/`Bundle.open`.
|
||||
|
||||
Mirrors `sources.RolloutSpec`: `name` is the series identity carried
|
||||
through `payload["series"]` keys, legend labels, and color assignment.
|
||||
"""
|
||||
|
||||
name: str
|
||||
source: str | Path | pl.LazyFrame
|
||||
checkpoint: str | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class PredictionSide:
|
||||
"""One prediction's opened frame + its coord/truth-availability facts."""
|
||||
|
||||
lf: pl.LazyFrame # raw scan, chunk-filtered
|
||||
paired: pl.LazyFrame # canonical paired frame (see module docstring)
|
||||
coord: str # "global" | "local"
|
||||
has_truth: bool
|
||||
checkpoint: str | None = None
|
||||
|
||||
|
||||
def _check_predict_metadata(path: Path) -> tuple[str, bool]:
|
||||
"""Return `(coord, has_truth)`, raising if `path` isn't predict output.
|
||||
|
||||
Distinguishes a predict file from a rollout file (both are tagged with
|
||||
`PREDICT_COORD_METADATA_KEY`, but a rollout's value is `ROLLOUT_COORD_VALUE`
|
||||
rather than `"global"`/`"local"`).
|
||||
"""
|
||||
metadata = pq.read_schema(path).metadata or {}
|
||||
coord = metadata.get(PREDICT_COORD_METADATA_KEY.encode())
|
||||
if coord is None:
|
||||
raise ValueError(
|
||||
f"{path} has no {PREDICT_COORD_METADATA_KEY!r} parquet metadata — it wasn't "
|
||||
"written by `giant predict` (or predates schema tagging)"
|
||||
)
|
||||
coord = coord.decode()
|
||||
if coord == ROLLOUT_COORD_VALUE:
|
||||
raise ValueError(f"{path} is a `giant rollout` file, not `giant predict` output")
|
||||
if coord not in ("global", "local"):
|
||||
raise ValueError(f"{path} has unrecognised predict coord {coord!r}")
|
||||
# A v1 file predates truth tagging; only --coord local was paired then.
|
||||
truth_raw = metadata.get(PREDICT_TRUTH_METADATA_KEY.encode())
|
||||
has_truth = truth_raw.decode() == "1" if truth_raw is not None else coord == "local"
|
||||
return coord, has_truth
|
||||
|
||||
|
||||
def _edep_pl(prefix: str) -> pl.Expr:
|
||||
"""Physical edep from `{prefix}_edep_logit`/`{prefix}_sec_logit` + `pre_E`.
|
||||
|
||||
Polars equivalent of `energy_simplex_decode(...)[0]` (the deposit
|
||||
component): a softmax over `[z_edep, z_sec, 0]` times `pre_E`.
|
||||
"""
|
||||
z1, z2 = pl.col(f"{prefix}_edep_logit"), pl.col(f"{prefix}_sec_logit")
|
||||
m = pl.max_horizontal(z1, z2, pl.lit(0.0))
|
||||
e1, e2, e3 = (z1 - m).exp(), (z2 - m).exp(), (pl.lit(0.0) - m).exp()
|
||||
return (e1 / (e1 + e2 + e3)) * pl.col("pre_E")
|
||||
|
||||
|
||||
def _delta_e_pl(prefix: str) -> pl.Expr:
|
||||
"""Physical delta_e (= edep + e_sec = pre_E - post_E) from the ALR logits + pre_E."""
|
||||
z1, z2 = pl.col(f"{prefix}_edep_logit"), pl.col(f"{prefix}_sec_logit")
|
||||
m = pl.max_horizontal(z1, z2, pl.lit(0.0))
|
||||
e1, e2, e3 = (z1 - m).exp(), (z2 - m).exp(), (pl.lit(0.0) - m).exp()
|
||||
return ((e1 + e2) / (e1 + e2 + e3)) * pl.col("pre_E")
|
||||
|
||||
|
||||
def _local_var_expr(prefix: str, var: str) -> pl.Expr:
|
||||
"""Physical value of one `PAIRED_VARS` entry from a `--coord local` file."""
|
||||
if var == "step_length":
|
||||
return pl.col(f"{prefix}_log_step_length").exp() - _LOG_EPS
|
||||
if var == "edep":
|
||||
return _edep_pl(prefix)
|
||||
if var == "delta_e":
|
||||
return _delta_e_pl(prefix)
|
||||
if var == "post_E":
|
||||
return pl.col("pre_E") - _delta_e_pl(prefix)
|
||||
if var == "cos_scatter":
|
||||
dot = pl.sum_horizontal([pl.col(f"{prefix}_post_d{ax}") * pl.col(f"{prefix}_travel_d{ax}") for ax in "xyz"])
|
||||
return dot.clip(-1.0, 1.0)
|
||||
raise ValueError(f"{var!r} has no direction-alignment meaning as a solo local-frame variable")
|
||||
|
||||
|
||||
def _g(prefix: str, name: str) -> str:
|
||||
"""Global-coord column name for `name` under `prefix`.
|
||||
|
||||
`giant predict --coord global` writes the *predicted* value under its bare
|
||||
name (`step_length`, `edep`, `post_dx`, ...) and the truth under a
|
||||
`true_` prefix (`true_step_length`, ...) — asymmetric, unlike the `local`
|
||||
coord's symmetric `pred_*`/`true_*` naming.
|
||||
"""
|
||||
return name if prefix == "pred" else f"true_{name}"
|
||||
|
||||
|
||||
def _global_var_expr(prefix: str, var: str) -> pl.Expr:
|
||||
"""Physical value of one `PAIRED_VARS` entry from a `--coord global` file."""
|
||||
if var == "post_E":
|
||||
# Not written directly for the prediction (it's implied by energy
|
||||
# conservation: post_E = pre_E - delta_e); truth carries it as
|
||||
# true_post_E.
|
||||
return pl.col("pre_E") - pl.col(_g(prefix, "delta_e")) if prefix == "pred" else pl.col(_g(prefix, "post_E"))
|
||||
if var == "cos_scatter":
|
||||
dot = pl.sum_horizontal([pl.col(f"pre_d{ax}") * pl.col(_g(prefix, f"post_d{ax}")) for ax in "xyz"])
|
||||
return dot.clip(-1.0, 1.0)
|
||||
if var == "cos_travel":
|
||||
# travel_dir isn't written by predict (only rollout reconstructs
|
||||
# post_pos from it) — approximate with the post_pos - pre_pos
|
||||
# direction instead, which is exactly what travel_dir encodes.
|
||||
dx = pl.col(_g(prefix, "post_x")) - pl.col("pre_x")
|
||||
dy = pl.col(_g(prefix, "post_y")) - pl.col("pre_y")
|
||||
dz = pl.col(_g(prefix, "post_z")) - pl.col("pre_z")
|
||||
norm = (dx**2 + dy**2 + dz**2).sqrt()
|
||||
dot = (
|
||||
pl.col("pre_dx") * dx / (norm + 1e-8)
|
||||
+ pl.col("pre_dy") * dy / (norm + 1e-8)
|
||||
+ pl.col("pre_dz") * dz / (norm + 1e-8)
|
||||
)
|
||||
return dot.clip(-1.0, 1.0)
|
||||
return pl.col(_g(prefix, var))
|
||||
|
||||
|
||||
def _var_expr(coord: str, prefix: str, var: str) -> pl.Expr:
|
||||
# `cos_travel` is excluded for `coord == "local"` by `paired_vars_for_coord`
|
||||
# (predict never reconstructs post_pos/travel_dir there), so this only
|
||||
# ever sees local-representable vars on that path.
|
||||
if coord == "local":
|
||||
return _local_var_expr(prefix, var)
|
||||
return _global_var_expr(prefix, var)
|
||||
|
||||
|
||||
def paired_vars_for_coord(coord: str) -> tuple[str, ...]:
|
||||
"""The `PAIRED_VARS` a given coord mode can actually produce.
|
||||
|
||||
`cos_travel` needs a reconstructed `travel_dir`/`post_pos`, which
|
||||
`--coord local` predict output never has (stage 2 doesn't run there) —
|
||||
so local-coord predictions drop it rather than emit a meaningless value.
|
||||
"""
|
||||
if coord == "local":
|
||||
return PAIRED_SCALARS + ("cos_scatter",)
|
||||
return PAIRED_VARS
|
||||
|
||||
|
||||
def dir_alignment_expr(coord: str, kind: str) -> pl.Expr:
|
||||
"""cos angle between the true and predicted direction vector (raw, not paired).
|
||||
|
||||
`kind="post"` compares `post_dir`; `kind="travel"` compares the
|
||||
post_pos-implied travel direction. Reads the *raw* opened frame
|
||||
(`PredictionSide.lf`), not `paired` — direction components aren't part of
|
||||
`PAIRED_VARS` (only their two scattering cosines are), so this stays a
|
||||
separate helper.
|
||||
"""
|
||||
if coord == "local":
|
||||
prefix_dim = "post_d" if kind == "post" else "travel_d"
|
||||
true_v = [pl.col(f"true_{prefix_dim}{ax}") for ax in "xyz"]
|
||||
pred_v = [pl.col(f"pred_{prefix_dim}{ax}") for ax in "xyz"]
|
||||
elif kind == "post":
|
||||
true_v = [pl.col(f"true_post_d{ax}") for ax in "xyz"]
|
||||
pred_v = [pl.col(f"post_d{ax}") for ax in "xyz"] # unprefixed: see paired_frame's _g
|
||||
else:
|
||||
true_v = [pl.col(f"true_post_{ax}") - pl.col(f"pre_{ax}") for ax in "xyz"]
|
||||
pred_v = [pl.col(f"post_{ax}") - pl.col(f"pre_{ax}") for ax in "xyz"]
|
||||
dot = pl.sum_horizontal([t * p for t, p in zip(true_v, pred_v)])
|
||||
true_norm = pl.sum_horizontal([t**2 for t in true_v]).sqrt()
|
||||
pred_norm = pl.sum_horizontal([p**2 for p in pred_v]).sqrt()
|
||||
return (dot / (true_norm * pred_norm + 1e-8)).clip(-1.0, 1.0)
|
||||
|
||||
|
||||
def paired_frame(lf: pl.LazyFrame, coord: str, has_truth: bool) -> pl.LazyFrame:
|
||||
"""Canonical `event_id, pdg, material, pre_E, n_sec, n_sec_pred, true_*, pred_*` frame."""
|
||||
schema = lf.collect_schema().names()
|
||||
cols = [
|
||||
"event_id",
|
||||
"pdg",
|
||||
"pre_E",
|
||||
"material",
|
||||
"n_sec",
|
||||
pl.col("n_sec_pred") if "n_sec_pred" in schema else pl.lit(None, dtype=pl.Int64).alias("n_sec_pred"),
|
||||
]
|
||||
for var in paired_vars_for_coord(coord):
|
||||
cols.append(_var_expr(coord, "pred", var).alias(f"pred_{var}"))
|
||||
if has_truth:
|
||||
cols.append(_var_expr(coord, "true", var).alias(f"true_{var}"))
|
||||
return lf.select(cols)
|
||||
|
||||
|
||||
def open_prediction(source: str | Path | pl.LazyFrame) -> PredictionSide:
|
||||
"""Lazily scan one prediction file, verifying its predict tag."""
|
||||
if isinstance(source, pl.LazyFrame):
|
||||
lf = source.with_columns(pl.col("pdg").cast(pl.Int64))
|
||||
schema = lf.collect_schema().names()
|
||||
coord = "local" if "pred_log_step_length" in schema else "global"
|
||||
has_truth = f"true_{LOCAL_TARGET_NAMES[0]}" in schema or "true_step_length" in schema
|
||||
else:
|
||||
path = Path(source)
|
||||
coord, has_truth = _check_predict_metadata(path)
|
||||
lf = pl.scan_parquet(path).with_columns(pl.col("pdg").cast(pl.Int64))
|
||||
return PredictionSide(
|
||||
lf=lf,
|
||||
paired=paired_frame(lf, coord, has_truth),
|
||||
coord=coord,
|
||||
has_truth=has_truth,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Secondaries (global + truth only)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def prediction_secondaries(lf: pl.LazyFrame, prefix: str) -> pl.LazyFrame:
|
||||
"""One row per secondary from the true/predicted `sec_*_list` columns.
|
||||
|
||||
Canonical columns: `event_id, energy, pdg, sdx, sdy, sdz` — same shape as
|
||||
`sources.secondaries`'s reference-side branch. `prefix` is `"true"` or
|
||||
`"pred"`; matches `giant predict --coord global`'s asymmetric naming (see
|
||||
`_g`) — the predicted lists are unprefixed (`sec_E_list`, ...), only the
|
||||
truth ones carry `true_` (`true_sec_E_list`, ...).
|
||||
"""
|
||||
col_prefix = "" if prefix == "pred" else "true_"
|
||||
lists = [f"{col_prefix}sec_{c}_list" for c in ("E", "pdg", "dx", "dy", "dz")]
|
||||
return (
|
||||
lf.select("event_id", *lists)
|
||||
.explode(lists)
|
||||
.drop_nulls(lists[0])
|
||||
.select(
|
||||
"event_id",
|
||||
pl.col(lists[0]).alias("energy"),
|
||||
pl.col(lists[1]).cast(pl.Int64).alias("pdg"),
|
||||
pl.col(lists[2]).alias("sdx"),
|
||||
pl.col(lists[3]).alias("sdy"),
|
||||
pl.col(lists[4]).alias("sdz"),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def paired_secondaries(lf: pl.LazyFrame) -> pl.LazyFrame:
|
||||
"""True/predicted secondary PDG pairs, aligned by descending-energy rank.
|
||||
|
||||
Stage 2 emits secondaries in descending-energy order (`network.md`/
|
||||
`giant/model/models.py`'s autoregressive decoder), so the natural
|
||||
per-step alignment between the true and predicted secondary lists is
|
||||
positional: rank `i` of one list vs. rank `i` of the other, for
|
||||
`i < min(n_sec, n_sec_pred)`. Requires `--coord global --truth`.
|
||||
"""
|
||||
return (
|
||||
lf.select("true_sec_pdg_list", "sec_pdg_list")
|
||||
.with_row_index("_row")
|
||||
.with_columns(
|
||||
pl.col("true_sec_pdg_list").list.len().alias("_n_true"),
|
||||
pl.col("sec_pdg_list").list.len().alias("_n_pred"),
|
||||
)
|
||||
.with_columns(pl.min_horizontal("_n_true", "_n_pred").alias("_n_paired"))
|
||||
.filter(pl.col("_n_paired") > 0)
|
||||
.with_columns(pl.int_ranges(0, pl.col("_n_paired")).alias("_rank"))
|
||||
.explode("_rank")
|
||||
.select(
|
||||
pl.col("true_sec_pdg_list").list.get(pl.col("_rank")).cast(pl.Int64).alias("true_pdg"),
|
||||
pl.col("sec_pdg_list").list.get(pl.col("_rank")).cast(pl.Int64).alias("pred_pdg"),
|
||||
)
|
||||
)
|
||||
@@ -70,6 +70,64 @@ def hist1d(
|
||||
return out
|
||||
|
||||
|
||||
def hist2d(
|
||||
lf: pl.LazyFrame,
|
||||
x: pl.Expr,
|
||||
y: pl.Expr,
|
||||
x_edges: np.ndarray,
|
||||
y_edges: np.ndarray,
|
||||
) -> np.ndarray:
|
||||
"""Streaming 2D histogram of `(x, y)` over fixed uniform edges.
|
||||
|
||||
One `group_by([_bx, _by]).len()` pass; returns the full `(len(x_edges)-1,
|
||||
len(y_edges)-1)` int64 count matrix (row = x bin, col = y bin) — small
|
||||
enough (a truth-vs-pred scatter has at most a few thousand cells) to
|
||||
materialize whole, unlike `hist1d`'s per-group dict.
|
||||
"""
|
||||
x_lo, x_hi, x_n = float(x_edges[0]), float(x_edges[-1]), len(x_edges) - 1
|
||||
y_lo, y_hi, y_n = float(y_edges[0]), float(y_edges[-1]), len(y_edges) - 1
|
||||
res = (
|
||||
lf.select(_bin_expr(x, x_lo, x_hi, x_n).alias("_bx"), _bin_expr(y, y_lo, y_hi, y_n).alias("_by"))
|
||||
.drop_nulls(["_bx", "_by"])
|
||||
.group_by("_bx", "_by")
|
||||
.agg(pl.len().alias("_n"))
|
||||
.collect(engine="streaming")
|
||||
)
|
||||
mat = np.zeros((x_n, y_n), dtype=np.int64)
|
||||
mat[res["_bx"].to_numpy(), res["_by"].to_numpy()] = res["_n"].to_numpy()
|
||||
return mat
|
||||
|
||||
|
||||
def binned_moments(
|
||||
lf: pl.LazyFrame,
|
||||
bin_value: pl.Expr,
|
||||
agg_value: pl.Expr,
|
||||
edges: np.ndarray,
|
||||
) -> dict[str, list]:
|
||||
"""Per-bin ``(n, sum, sumsq)`` of ``agg_value``, binned by ``bin_value`` over fixed edges.
|
||||
|
||||
One streaming `group_by` pass; sum-mergeable across chunks the same way
|
||||
`hist1d` counts are — elementwise-summing `n`/`sum`/`sumsq` per bin across
|
||||
chunks reconstructs the moments of the full merged data, from which
|
||||
`finalize` derives mean/std (``mean = sum/n``,
|
||||
``std = sqrt(sumsq/n - mean**2)``).
|
||||
"""
|
||||
lo, hi, nbins = float(edges[0]), float(edges[-1]), len(edges) - 1
|
||||
res = (
|
||||
lf.select(_bin_expr(bin_value, lo, hi, nbins).alias("_b"), agg_value.alias("_v"))
|
||||
.drop_nulls(["_b", "_v"])
|
||||
.group_by("_b")
|
||||
.agg(pl.len().alias("_n"), pl.col("_v").sum().alias("_s"), (pl.col("_v") ** 2).sum().alias("_ss"))
|
||||
.collect(engine="streaming")
|
||||
)
|
||||
n = np.zeros(nbins, dtype=np.int64)
|
||||
s = np.zeros(nbins, dtype=np.float64)
|
||||
ss = np.zeros(nbins, dtype=np.float64)
|
||||
for b_, nn, ssum, sqsum in res.iter_rows():
|
||||
n[b_], s[b_], ss[b_] = nn, ssum, sqsum
|
||||
return {"n": n.tolist(), "sum": s.tolist(), "sumsq": ss.tolist()}
|
||||
|
||||
|
||||
def sum_merge(dicts: list[dict[str, Any]]) -> dict[str, Any]:
|
||||
"""Elementwise-sum a list of sum-mergeable count/total dicts (JSON-safe keys).
|
||||
|
||||
|
||||
@@ -27,8 +27,13 @@ from pathlib import Path
|
||||
# "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)
|
||||
# "unavailable" plot not applicable to this run (e.g. no MoE checkpoint)
|
||||
# distance scorecard) or per prediction (a confusion matrix)
|
||||
# "paired_hist" per-prediction true/pred density histogram over shared
|
||||
# edges (giant predict's paired truth, not a rollout)
|
||||
# "heatmap2d" numeric x/y-binned true-vs-pred count matrix + colorbar,
|
||||
# one panel per prediction, with a y=x diagonal guide
|
||||
# "unavailable" plot not applicable to this run (e.g. no MoE checkpoint,
|
||||
# or no --prediction given)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -422,6 +422,78 @@ def _render_heatmap(r: Reduced, params: dict):
|
||||
return fig
|
||||
|
||||
|
||||
def _render_paired_hist(r: Reduced, params: dict):
|
||||
"""`giant predict`'s paired truth/pred density histogram (see
|
||||
`giant.analysis.prediction`) — unlike `_render_overlay`, there's no single
|
||||
shared reference: each prediction carries its own truth. A lone prediction
|
||||
draws its truth in the reference ink so a single-series run reads exactly
|
||||
like an `overlay_hist` figure; two-or-more predictions each get their own
|
||||
color, pred solid / true dashed, so a same-colored pair is directly
|
||||
comparable.
|
||||
"""
|
||||
edges = np.asarray(r.payload["edges"])
|
||||
series = r.payload.get("series", {})
|
||||
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
|
||||
solo = len(series) == 1
|
||||
for i, (name, entry) in enumerate(series.items()):
|
||||
color = _ref_color() if solo else ps.get_color(i)
|
||||
if "true" in entry:
|
||||
true_label = _REFERENCE_LABEL if solo else f"{name} (true)"
|
||||
ax.stairs(_density(entry["true"], edges), edges, label=true_label, color=color, linestyle="--")
|
||||
pred_color = ps.get_color(i)
|
||||
pred_label = name if solo else f"{name} (pred)"
|
||||
ax.stairs(_density(entry["pred"], edges), edges, label=pred_label, color=pred_color)
|
||||
if r.payload.get("log_y"):
|
||||
ax.set_yscale("log")
|
||||
ax.set_xlabel(r.xlabel)
|
||||
ax.set_ylabel("density")
|
||||
ps.style_legend(ax, title="source")
|
||||
return fig
|
||||
|
||||
|
||||
def _render_heatmap2d(r: Reduced, params: dict):
|
||||
"""Numeric truth-vs-pred 2D histogram, one panel per prediction, with an
|
||||
optional y=x guide line — the direct analogue of `_render_heatmap` for
|
||||
continuous (not categorical) axes."""
|
||||
x_edges = np.asarray(r.payload["x_edges"])
|
||||
y_edges = np.asarray(r.payload["y_edges"])
|
||||
series = r.payload["series"]
|
||||
names = list(series)
|
||||
norm = LogNorm(vmin=1) if r.payload.get("log_color") else None
|
||||
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,
|
||||
)
|
||||
flat = axes.ravel()
|
||||
im = None
|
||||
for ax, name in zip(flat, names):
|
||||
mat = np.asarray(series[name], dtype=float)
|
||||
im = ax.pcolormesh(
|
||||
x_edges,
|
||||
y_edges,
|
||||
mat.T,
|
||||
cmap=r.payload.get("cmap", "viridis"),
|
||||
norm=norm,
|
||||
vmin=None if norm else r.payload.get("vmin"),
|
||||
vmax=None if norm else r.payload.get("vmax"),
|
||||
)
|
||||
if r.payload.get("diagonal"):
|
||||
lo, hi = max(x_edges[0], y_edges[0]), min(x_edges[-1], y_edges[-1])
|
||||
ax.plot([lo, hi], [lo, hi], color=_ref_color(), linestyle="--", linewidth=1, label="y = x")
|
||||
ax.set_xlabel(r.xlabel)
|
||||
if len(names) > 1:
|
||||
ax.set_title(name, fontsize=8)
|
||||
flat[0].set_ylabel(r.payload.get("ylabel", ""))
|
||||
if r.payload.get("diagonal"):
|
||||
ps.style_legend(flat[0], title="guide")
|
||||
fig.colorbar(im, ax=list(flat), label=r.payload.get("cbar_label", "count"))
|
||||
return fig
|
||||
|
||||
|
||||
def _render_unavailable(r: Reduced, params: dict):
|
||||
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
|
||||
ax.axis("off")
|
||||
@@ -448,6 +520,8 @@ _RENDERERS = {
|
||||
"router_share": _render_router_share,
|
||||
"router_specialization": _render_router_specialization,
|
||||
"heatmap": _render_heatmap,
|
||||
"paired_hist": _render_paired_hist,
|
||||
"heatmap2d": _render_heatmap2d,
|
||||
"unavailable": _render_unavailable,
|
||||
}
|
||||
|
||||
|
||||
+244
-44
@@ -20,6 +20,7 @@ from giant.constants import (
|
||||
PREDICT_COORD_METADATA_KEY,
|
||||
PREDICT_SCHEMA_VERSION,
|
||||
PREDICT_SCHEMA_VERSION_KEY,
|
||||
PREDICT_TRUTH_METADATA_KEY,
|
||||
ROLLOUT_COORD_VALUE,
|
||||
)
|
||||
|
||||
@@ -186,11 +187,18 @@ def _write_prediction_ref(
|
||||
out: Path,
|
||||
dataset_path: Path,
|
||||
comment: str | None = None,
|
||||
extra: dict | None = None,
|
||||
) -> Path:
|
||||
"""Write a YAML sidecar in the checkpoint directory and return its path."""
|
||||
"""Write a YAML sidecar in the checkpoint directory and return its path.
|
||||
|
||||
``extra`` is merged in after the base fields (e.g. `giant rollout`'s
|
||||
provenance/timing block, or `giant predict`'s) — callers that don't pass
|
||||
it get exactly today's thin sidecar.
|
||||
"""
|
||||
import yaml
|
||||
|
||||
ref = {
|
||||
"kind": "prediction",
|
||||
"prediction_id": pred_uuid,
|
||||
"output": str(out),
|
||||
"dataset": str(dataset_path),
|
||||
@@ -199,6 +207,8 @@ def _write_prediction_ref(
|
||||
}
|
||||
if comment is not None:
|
||||
ref["comment"] = comment
|
||||
if extra:
|
||||
ref.update(extra)
|
||||
ref_path = checkpoint.parent / f"{pred_uuid}.yaml"
|
||||
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
|
||||
return ref_path
|
||||
@@ -243,6 +253,41 @@ def _build_rollout_timing(
|
||||
}
|
||||
|
||||
|
||||
def _build_predict_timing(
|
||||
*,
|
||||
setup_s: float,
|
||||
predict_s: float,
|
||||
write_s: float,
|
||||
n_rows: int,
|
||||
device: str,
|
||||
torch_threads: int,
|
||||
) -> dict:
|
||||
"""Assemble ``giant predict``'s ``timing`` sidecar block.
|
||||
|
||||
Keys are deliberately compatible with ``_build_rollout_timing``'s (same
|
||||
names for the quantities both commands have) so a gallery's ``timing``
|
||||
metadata renders the same way whether the series came from a rollout or
|
||||
a prediction. There's no ``n_physical_rows``/``ms_per_event`` here:
|
||||
unlike a rollout, `giant predict` never emits synthetic termination rows
|
||||
(one output row per input step) and doesn't work in whole showers/events
|
||||
— so ``us_per_step`` is already directly comparable to a rollout's and to
|
||||
``giant.analysis.geant4_reference``'s per-step Geant4 measurement.
|
||||
"""
|
||||
sample_s = predict_s - write_s
|
||||
return {
|
||||
"setup_s": setup_s,
|
||||
"predict_s": predict_s,
|
||||
"write_s": write_s,
|
||||
"sample_s": sample_s,
|
||||
"n_rows": n_rows,
|
||||
"us_per_step": (sample_s / n_rows * 1e6) if n_rows else None,
|
||||
"write_us_per_step": (write_s / n_rows * 1e6) if n_rows else None,
|
||||
"rows_per_s": (n_rows / predict_s) if predict_s else None,
|
||||
"device": device,
|
||||
"torch_threads": torch_threads,
|
||||
}
|
||||
|
||||
|
||||
@app.callback()
|
||||
def _main() -> None:
|
||||
"""GIANT — Geant4 step-function surrogate."""
|
||||
@@ -1020,9 +1065,21 @@ def predict(
|
||||
typer.Option(
|
||||
"--out",
|
||||
"-o",
|
||||
help="Output parquet path (default: <data>_predicted[_local].parquet)",
|
||||
help="Output parquet path (default: a UUID-named file under /ceph's central "
|
||||
"predictions store if --data is under /ceph, else a sibling of --data)",
|
||||
),
|
||||
] = None,
|
||||
truth: Annotated[
|
||||
bool,
|
||||
typer.Option(
|
||||
"--truth/--no-truth",
|
||||
help="--coord global only: also read and write ground-truth post-step + "
|
||||
"secondary columns (true_step_length, true_edep, true_sec_*_list, ...) "
|
||||
"alongside the predictions, at the cost of reading full row-groups instead "
|
||||
"of conditioning columns only. Ignored for --coord local, which is always "
|
||||
"paired. Default: on.",
|
||||
),
|
||||
] = True,
|
||||
batch_size: Annotated[
|
||||
str,
|
||||
typer.Option(
|
||||
@@ -1071,6 +1128,8 @@ def predict(
|
||||
] = None,
|
||||
) -> None:
|
||||
"""Run trained model on a parquet file and save predictions."""
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
@@ -1090,6 +1149,8 @@ def predict(
|
||||
from giant.rollout import decode_secondary_identity
|
||||
from giant.sample import resolve_n_sec, sample_stage1, sample_stage2
|
||||
|
||||
_t_setup_start = time.perf_counter()
|
||||
|
||||
batch_size_auto = False
|
||||
batch_size_value: int | None = None
|
||||
if batch_size.strip().lower() == "auto":
|
||||
@@ -1161,8 +1222,17 @@ def predict(
|
||||
unknown_pdg_counts: Counter[int] = Counter()
|
||||
total_rows = sum(pq.ParquetFile(path).metadata.num_rows for path in files)
|
||||
|
||||
training_cfg = gconfig.load_checkpoint_config(checkpoint)
|
||||
_write_s = 0.0
|
||||
_setup_s = time.perf_counter() - _t_setup_start
|
||||
|
||||
# --coord local always needs full row-groups (it's paired against the 9D
|
||||
# target); --coord global only needs them when --truth is requested —
|
||||
# otherwise the cheaper conditioning-only read is used.
|
||||
write_truth = coord == Coord.global_ and truth
|
||||
|
||||
def chunk_iter(path: Path, offset: int):
|
||||
if coord == Coord.local:
|
||||
if coord == Coord.local or write_truth:
|
||||
return iter_file_chunks(path, offset=offset, k_max=stage2_k_max)
|
||||
return iter_cond_chunks(path, offset=offset)
|
||||
|
||||
@@ -1177,7 +1247,7 @@ def predict(
|
||||
return {k: np.concatenate([a[k], b[k]], axis=0) for k in a}
|
||||
|
||||
def _process(piece: dict[str, np.ndarray]) -> None:
|
||||
nonlocal writer, total
|
||||
nonlocal writer, total, _write_s
|
||||
|
||||
if coord == Coord.local:
|
||||
feats = build_features(
|
||||
@@ -1290,54 +1360,95 @@ def predict(
|
||||
sec_dy_list = [sec_dir_world[i, :n, 1].tolist() for i, n in enumerate(n_sec_pred_np)]
|
||||
sec_dz_list = [sec_dir_world[i, :n, 2].tolist() for i, n in enumerate(n_sec_pred_np)]
|
||||
|
||||
table = pa.table(
|
||||
{
|
||||
"event_id": piece["event_id"],
|
||||
"pdg": piece["pdg"],
|
||||
"pre_x": piece["pre_pos"][:, 0],
|
||||
"pre_y": piece["pre_pos"][:, 1],
|
||||
"pre_z": piece["pre_pos"][:, 2],
|
||||
"pre_E": piece["pre_E"],
|
||||
"pre_dx": piece["pre_dir"][:, 0],
|
||||
"pre_dy": piece["pre_dir"][:, 1],
|
||||
"pre_dz": piece["pre_dir"][:, 2],
|
||||
"material": piece["material"],
|
||||
"layer_id": piece["layer_id"],
|
||||
"n_sec": piece["n_sec"],
|
||||
"n_sec_pred": n_sec_pred_np,
|
||||
"step_length": step_length,
|
||||
"delta_e": delta_e,
|
||||
"edep": edep,
|
||||
"post_dx": post_dir_world[:, 0],
|
||||
"post_dy": post_dir_world[:, 1],
|
||||
"post_dz": post_dir_world[:, 2],
|
||||
"post_x": post_pos_world[:, 0],
|
||||
"post_y": post_pos_world[:, 1],
|
||||
"post_z": post_pos_world[:, 2],
|
||||
"sec_pdg_list": sec_pdg_list,
|
||||
"sec_E_list": sec_E_list,
|
||||
"sec_dx_list": sec_dx_list,
|
||||
"sec_dy_list": sec_dy_list,
|
||||
"sec_dz_list": sec_dz_list,
|
||||
}
|
||||
)
|
||||
columns = {
|
||||
"event_id": piece["event_id"],
|
||||
"pdg": piece["pdg"],
|
||||
"pre_x": piece["pre_pos"][:, 0],
|
||||
"pre_y": piece["pre_pos"][:, 1],
|
||||
"pre_z": piece["pre_pos"][:, 2],
|
||||
"pre_E": piece["pre_E"],
|
||||
"pre_dx": piece["pre_dir"][:, 0],
|
||||
"pre_dy": piece["pre_dir"][:, 1],
|
||||
"pre_dz": piece["pre_dir"][:, 2],
|
||||
"material": piece["material"],
|
||||
"layer_id": piece["layer_id"],
|
||||
"n_sec": piece["n_sec"],
|
||||
"n_sec_pred": n_sec_pred_np,
|
||||
"step_length": step_length,
|
||||
"delta_e": delta_e,
|
||||
"edep": edep,
|
||||
"post_dx": post_dir_world[:, 0],
|
||||
"post_dy": post_dir_world[:, 1],
|
||||
"post_dz": post_dir_world[:, 2],
|
||||
"post_x": post_pos_world[:, 0],
|
||||
"post_y": post_pos_world[:, 1],
|
||||
"post_z": post_pos_world[:, 2],
|
||||
"sec_pdg_list": sec_pdg_list,
|
||||
"sec_E_list": sec_E_list,
|
||||
"sec_dx_list": sec_dx_list,
|
||||
"sec_dy_list": sec_dy_list,
|
||||
"sec_dz_list": sec_dz_list,
|
||||
}
|
||||
|
||||
if write_truth:
|
||||
n_sec_true = piece["n_sec"]
|
||||
columns.update(
|
||||
{
|
||||
"true_step_length": piece["step_length"],
|
||||
"true_delta_e": piece["delta_e"],
|
||||
"true_edep": piece["edep"],
|
||||
"true_post_E": piece["post_E"],
|
||||
"true_post_dx": piece["post_dir"][:, 0],
|
||||
"true_post_dy": piece["post_dir"][:, 1],
|
||||
"true_post_dz": piece["post_dir"][:, 2],
|
||||
"true_post_x": piece["post_pos"][:, 0],
|
||||
"true_post_y": piece["post_pos"][:, 1],
|
||||
"true_post_z": piece["post_pos"][:, 2],
|
||||
"true_e_sec": piece["e_sec"],
|
||||
"process": piece["process"],
|
||||
}
|
||||
)
|
||||
if "sec_E_list" in piece:
|
||||
columns.update(
|
||||
{
|
||||
"true_sec_pdg_list": [
|
||||
piece["sec_pdg_list"][i, :n].tolist() for i, n in enumerate(n_sec_true)
|
||||
],
|
||||
"true_sec_E_list": [piece["sec_E_list"][i, :n].tolist() for i, n in enumerate(n_sec_true)],
|
||||
"true_sec_dx_list": [
|
||||
piece["sec_dir_list"][i, :n, 0].tolist() for i, n in enumerate(n_sec_true)
|
||||
],
|
||||
"true_sec_dy_list": [
|
||||
piece["sec_dir_list"][i, :n, 1].tolist() for i, n in enumerate(n_sec_true)
|
||||
],
|
||||
"true_sec_dz_list": [
|
||||
piece["sec_dir_list"][i, :n, 2].tolist() for i, n in enumerate(n_sec_true)
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
table = pa.table(columns)
|
||||
|
||||
table = table.replace_schema_metadata(
|
||||
{
|
||||
PREDICT_COORD_METADATA_KEY: coord.value,
|
||||
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
|
||||
PREDICT_TRUTH_METADATA_KEY: "1" if (coord == Coord.local or write_truth) else "0",
|
||||
}
|
||||
)
|
||||
|
||||
_t0 = time.perf_counter()
|
||||
if writer is None:
|
||||
writer = pq.ParquetWriter(out, table.schema)
|
||||
writer.write_table(table)
|
||||
_write_s += time.perf_counter() - _t0
|
||||
total += len(piece["event_id"])
|
||||
|
||||
# Buffer rows across row-group boundaries so the inference batch size
|
||||
# isn't capped by however the source file happens to be chunked.
|
||||
buffer: dict[str, np.ndarray] | None = None
|
||||
|
||||
_t_predict_start = time.perf_counter()
|
||||
bar = tqdm(total=total_rows, desc="predict", unit="row", dynamic_ncols=True)
|
||||
for i, path in enumerate(files):
|
||||
for chunk in chunk_iter(path, offset=event_id_offset(i)):
|
||||
@@ -1365,9 +1476,56 @@ def predict(
|
||||
bar.close()
|
||||
if writer is not None:
|
||||
writer.close()
|
||||
_predict_s = time.perf_counter() - _t_predict_start
|
||||
|
||||
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path, comment)
|
||||
timing = _build_predict_timing(
|
||||
setup_s=_setup_s,
|
||||
predict_s=_predict_s,
|
||||
write_s=_write_s,
|
||||
n_rows=total,
|
||||
device=str(_device),
|
||||
torch_threads=torch.get_num_threads(),
|
||||
)
|
||||
|
||||
ref_path = _write_prediction_ref(
|
||||
checkpoint,
|
||||
pred_uuid,
|
||||
out,
|
||||
dataset_path,
|
||||
comment,
|
||||
extra={
|
||||
"coord": coord.value,
|
||||
"has_truth": coord == Coord.local or write_truth,
|
||||
"schema_version": PREDICT_SCHEMA_VERSION,
|
||||
"steps": steps,
|
||||
"weights": weights.value,
|
||||
"device": str(_device),
|
||||
"batch_size": bs,
|
||||
"batch_size_auto": batch_size_auto,
|
||||
"n_input_rows": total_rows,
|
||||
"n_files": len(files),
|
||||
"n_rows": total,
|
||||
"n_skipped_rows": skipped,
|
||||
"unknown_pdg_counts": {str(pdg): count for pdg, count in unknown_pdg_counts.items()},
|
||||
"timing": timing,
|
||||
# Full architecture spec baked into the checkpoint — see the
|
||||
# matching comment in `rollout`.
|
||||
"model_config": dict(ctx.model_config),
|
||||
"config_overrides": dict(ctx.config_overrides),
|
||||
"training_epoch": ctx.epoch,
|
||||
"best_val_loss": ctx.best_val_loss,
|
||||
# [train]/[meta] from the sibling config.toml (giant.config.save_config)
|
||||
# — empty dicts if the checkpoint has no config.toml next to it.
|
||||
"training_config": dict(training_cfg.get("train", {})),
|
||||
"training_meta": dict(training_cfg.get("meta", {})),
|
||||
},
|
||||
)
|
||||
typer.echo(f"reference: {ref_path}")
|
||||
if timing["us_per_step"] is not None:
|
||||
typer.echo(
|
||||
f"timing: {_predict_s:.1f}s total ({timing['sample_s']:.1f}s sample + {_write_s:.1f}s write), "
|
||||
f"{timing['us_per_step']:.1f} us/step over {total:,} step(s)"
|
||||
)
|
||||
|
||||
if skipped:
|
||||
codes = ", ".join(f"{pdg} ({count})" for pdg, count in sorted(unknown_pdg_counts.items()))
|
||||
@@ -1499,7 +1657,6 @@ def rollout(
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import yaml
|
||||
|
||||
from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference
|
||||
from giant.data.loader import find_parquet_files
|
||||
@@ -1628,10 +1785,12 @@ def rollout(
|
||||
|
||||
l1_summary = l1_dist_collector.summary()
|
||||
|
||||
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path)
|
||||
ref = yaml.safe_load(ref_path.read_text())
|
||||
ref.update(
|
||||
{
|
||||
ref_path = _write_prediction_ref(
|
||||
checkpoint,
|
||||
pred_uuid,
|
||||
out,
|
||||
dataset_path,
|
||||
extra={
|
||||
"kind": "rollout",
|
||||
"geometry_oracle": str(geometry.resolve()),
|
||||
"energy_cutoff": energy_cutoff,
|
||||
@@ -1668,9 +1827,8 @@ def rollout(
|
||||
# — empty dicts if the checkpoint has no config.toml next to it.
|
||||
"training_config": dict(training_cfg.get("train", {})),
|
||||
"training_meta": dict(training_cfg.get("meta", {})),
|
||||
}
|
||||
},
|
||||
)
|
||||
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
|
||||
|
||||
typer.echo(f"wrote {summary['n_rows']:,} step rows → {out}")
|
||||
typer.echo(f"terminations: {summary['termination_reason_counts']}")
|
||||
@@ -1708,6 +1866,25 @@ def analyze_prep(
|
||||
"single YAML).",
|
||||
),
|
||||
] = None,
|
||||
prediction: Annotated[
|
||||
list[Path] | None,
|
||||
typer.Option(
|
||||
"--prediction",
|
||||
help="giant predict YAML sidecar(s) (paired truth/pred comparison, the "
|
||||
"`prediction` plot family) — optional add-on to the rollout comparison. "
|
||||
"Every one must be seeded from the same `dataset` as the rollout(s) and "
|
||||
"share one predict --coord.",
|
||||
),
|
||||
] = None,
|
||||
prediction_label: Annotated[
|
||||
list[str] | None,
|
||||
typer.Option(
|
||||
"--prediction-label",
|
||||
help="Series name for a --prediction YAML, positionally matched to it — give "
|
||||
'none, or exactly one per YAML. Defaults to the YAML stem (or "prediction" '
|
||||
"for a single YAML).",
|
||||
),
|
||||
] = None,
|
||||
run_dir: Annotated[
|
||||
Path | None,
|
||||
typer.Option(
|
||||
@@ -1724,7 +1901,7 @@ def analyze_prep(
|
||||
typer.Option("--chunks", help="Split each plot's data into this many event_id chunks"),
|
||||
] = 1,
|
||||
) -> None:
|
||||
"""Read the rollout YAML(s) → shared.json + run_meta.json in the run directory."""
|
||||
"""Read the rollout (+ optional prediction) YAML(s) → shared.json + run_meta.json."""
|
||||
from giant.analysis import prep
|
||||
|
||||
path = prep(
|
||||
@@ -1733,6 +1910,8 @@ def analyze_prep(
|
||||
n_chunks=chunks,
|
||||
default_base=Path.cwd() / "analysis_runs",
|
||||
labels=label,
|
||||
prediction_yamls=prediction or (),
|
||||
prediction_labels=prediction_label,
|
||||
n_energy_bins=n_energy_bins,
|
||||
n_marginal_bins=n_marginal_bins,
|
||||
top_k_pdg=top_k_pdg,
|
||||
@@ -1831,6 +2010,25 @@ def analyze_submit(
|
||||
"single YAML).",
|
||||
),
|
||||
] = None,
|
||||
prediction: Annotated[
|
||||
list[Path] | None,
|
||||
typer.Option(
|
||||
"--prediction",
|
||||
help="giant predict YAML sidecar(s) (paired truth/pred comparison, the "
|
||||
"`prediction` plot family) — optional add-on to the rollout comparison. "
|
||||
"Every one must be seeded from the same `dataset` as the rollout(s) and "
|
||||
"share one predict --coord.",
|
||||
),
|
||||
] = None,
|
||||
prediction_label: Annotated[
|
||||
list[str] | None,
|
||||
typer.Option(
|
||||
"--prediction-label",
|
||||
help="Series name for a --prediction YAML, positionally matched to it — give "
|
||||
'none, or exactly one per YAML. Defaults to the YAML stem (or "prediction" '
|
||||
"for a single YAML).",
|
||||
),
|
||||
] = None,
|
||||
run_dir: Annotated[
|
||||
Path | None,
|
||||
typer.Option(
|
||||
@@ -1868,6 +2066,8 @@ def analyze_submit(
|
||||
n_chunks=chunks,
|
||||
default_base=Path.cwd() / "analysis_runs",
|
||||
labels=label,
|
||||
prediction_yamls=prediction or (),
|
||||
prediction_labels=prediction_label,
|
||||
n_energy_bins=n_energy_bins,
|
||||
n_marginal_bins=n_marginal_bins,
|
||||
top_k_pdg=top_k_pdg,
|
||||
|
||||
+6
-1
@@ -64,7 +64,12 @@ LOCAL_TARGET_NAMES = [
|
||||
# guessing from its column names.
|
||||
PREDICT_COORD_METADATA_KEY = "giant.predict.coord"
|
||||
PREDICT_SCHEMA_VERSION_KEY = "giant.predict.schema_version"
|
||||
PREDICT_SCHEMA_VERSION = "2"
|
||||
PREDICT_SCHEMA_VERSION = "3"
|
||||
|
||||
# Whether a --coord global predict parquet also carries true_* / true_sec_*
|
||||
# columns (v3+; "1"/"0"). Lets analysis code tell a paired prediction file
|
||||
# apart from a --no-truth one without sniffing for column presence.
|
||||
PREDICT_TRUTH_METADATA_KEY = "giant.predict.has_truth"
|
||||
|
||||
# Coord-metadata value tagging a `giant rollout` steps parquet (world frame,
|
||||
# autoregressive shower output). Distinct from predict's "global"/"local".
|
||||
|
||||
+1
-1
@@ -1,6 +1,6 @@
|
||||
[project]
|
||||
name = "giant"
|
||||
version = "0.3.21"
|
||||
version = "0.3.22"
|
||||
description = "Geant4 step-function surrogate via conditional flow matching"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
|
||||
@@ -0,0 +1,278 @@
|
||||
"""Tests for giant.analysis.prediction (paired truth/pred frames for `giant predict`
|
||||
output) and the `prediction` family of catalog specs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import pytest
|
||||
|
||||
from giant.analysis.catalog import Bundle, get_spec
|
||||
from giant.analysis.context import Context, build_context
|
||||
from giant.analysis.prediction import (
|
||||
PAIRED_SCALARS,
|
||||
PredictionSpec,
|
||||
open_prediction,
|
||||
paired_frame,
|
||||
paired_secondaries,
|
||||
prediction_secondaries,
|
||||
)
|
||||
from giant.analysis.reduce import hist2d
|
||||
from giant.analysis.sources import RolloutSpec
|
||||
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
|
||||
|
||||
|
||||
def _global_prediction_frame() -> pl.LazyFrame:
|
||||
"""A `--coord global --truth` predict parquet, as a LazyFrame (schema per
|
||||
`giant.cli.predict`'s global-coord table, `giant/cli.py:1310-1379`)."""
|
||||
return pl.DataFrame(
|
||||
{
|
||||
"event_id": [1, 1, 2],
|
||||
"pdg": [11, 11, 22],
|
||||
"pre_x": [0.0, 0.0, 0.0],
|
||||
"pre_y": [0.0, 0.0, 0.0],
|
||||
"pre_z": [0.0, 1.0, 0.0],
|
||||
"pre_E": [100.0, 60.0, 50.0],
|
||||
"pre_dx": [0.0, 0.0, 0.0],
|
||||
"pre_dy": [0.0, 0.0, 0.0],
|
||||
"pre_dz": [1.0, 1.0, 1.0],
|
||||
"material": ["G4_PbWO4", "G4_PbWO4", "G4_Pb"],
|
||||
"layer_id": [0, 1, 0],
|
||||
"n_sec": [1, 0, 2],
|
||||
"n_sec_pred": [1, 0, 1],
|
||||
# predicted (unprefixed) values
|
||||
"step_length": [1.2, 0.9, 1.1],
|
||||
"delta_e": [42.0, 29.0, 31.0],
|
||||
"edep": [35.0, 29.0, 25.0],
|
||||
"post_dx": [0.0, 0.0, 0.0],
|
||||
"post_dy": [0.0, 0.0, 0.0],
|
||||
"post_dz": [1.0, 1.0, 1.0],
|
||||
"post_x": [0.0, 0.0, 0.0],
|
||||
"post_y": [0.0, 0.0, 0.0],
|
||||
"post_z": [1.2, 1.9, 1.1],
|
||||
"sec_pdg_list": [[22], [], [22]],
|
||||
"sec_E_list": [[5.0], [], [4.0]],
|
||||
"sec_dx_list": [[0.0], [], [0.0]],
|
||||
"sec_dy_list": [[0.0], [], [0.0]],
|
||||
"sec_dz_list": [[1.0], [], [1.0]],
|
||||
# truth
|
||||
"true_step_length": [1.0, 1.0, 1.0],
|
||||
"true_delta_e": [40.0, 30.0, 30.0],
|
||||
"true_edep": [40.0, 30.0, 20.0],
|
||||
"true_post_E": [60.0, 30.0, 20.0],
|
||||
"true_post_dx": [0.0, 0.0, 0.0],
|
||||
"true_post_dy": [0.0, 0.0, 0.0],
|
||||
"true_post_dz": [1.0, 1.0, 1.0],
|
||||
"true_post_x": [0.0, 0.0, 0.0],
|
||||
"true_post_y": [0.0, 0.0, 0.0],
|
||||
"true_post_z": [1.0, 2.0, 1.0],
|
||||
"true_e_sec": [0.0, 0.0, 10.0],
|
||||
"process": ["compt", "phot", "compt"],
|
||||
"true_sec_pdg_list": [[22], [], [22, 11]],
|
||||
"true_sec_E_list": [[6.0], [], [7.0, 3.0]],
|
||||
"true_sec_dx_list": [[0.0], [], [0.0, 1.0]],
|
||||
"true_sec_dy_list": [[0.0], [], [0.0, 0.0]],
|
||||
"true_sec_dz_list": [[1.0], [], [1.0, 0.0]],
|
||||
}
|
||||
).lazy()
|
||||
|
||||
|
||||
def _local_prediction_frame() -> pl.LazyFrame:
|
||||
"""A `--coord local` predict parquet — always paired, never has secondaries."""
|
||||
return pl.DataFrame(
|
||||
{
|
||||
"event_id": [1, 2],
|
||||
"pdg": [11, 22],
|
||||
"pre_x": [0.0, 0.0],
|
||||
"pre_y": [0.0, 0.0],
|
||||
"pre_z": [0.0, 0.0],
|
||||
"pre_E": [100.0, 50.0],
|
||||
"pre_dx": [0.0, 0.0],
|
||||
"pre_dy": [0.0, 0.0],
|
||||
"pre_dz": [1.0, 1.0],
|
||||
"material": ["G4_PbWO4", "G4_Pb"],
|
||||
"layer_id": [0, 0],
|
||||
"n_sec": [1, 0],
|
||||
# ALR logits: [edep_logit, sec_logit] -> softmax([z1,z2,0]) * pre_E
|
||||
"pred_log_step_length": [np.log(1.2 + 1e-6), np.log(0.9 + 1e-6)],
|
||||
"pred_edep_logit": [1.0, 0.5],
|
||||
"pred_sec_logit": [0.0, -1.0],
|
||||
"pred_post_dx": [0.0, 0.0],
|
||||
"pred_post_dy": [0.0, 0.0],
|
||||
"pred_post_dz": [1.0, 1.0],
|
||||
"pred_travel_dx": [0.0, 0.0],
|
||||
"pred_travel_dy": [0.0, 0.0],
|
||||
"pred_travel_dz": [1.0, 1.0],
|
||||
"true_log_step_length": [np.log(1.0 + 1e-6), np.log(1.0 + 1e-6)],
|
||||
"true_edep_logit": [0.8, 0.6],
|
||||
"true_sec_logit": [0.2, -2.0],
|
||||
"true_post_dx": [0.0, 0.0],
|
||||
"true_post_dy": [0.0, 0.0],
|
||||
"true_post_dz": [1.0, 1.0],
|
||||
"true_travel_dx": [0.0, 0.0],
|
||||
"true_travel_dy": [0.0, 0.0],
|
||||
"true_travel_dz": [1.0, 1.0],
|
||||
}
|
||||
).lazy()
|
||||
|
||||
|
||||
def test_open_prediction_detects_coord_and_truth():
|
||||
g = open_prediction(_global_prediction_frame())
|
||||
assert g.coord == "global" and g.has_truth
|
||||
|
||||
loc = open_prediction(_local_prediction_frame())
|
||||
assert loc.coord == "local" and loc.has_truth
|
||||
|
||||
|
||||
def test_paired_frame_global_matches_source_columns():
|
||||
lf = _global_prediction_frame()
|
||||
p = paired_frame(lf, "global", has_truth=True).collect()
|
||||
assert p["pred_step_length"].to_list() == [1.2, 0.9, 1.1]
|
||||
assert p["true_step_length"].to_list() == [1.0, 1.0, 1.0]
|
||||
assert p["pred_edep"].to_list() == [35.0, 29.0, 25.0]
|
||||
assert p["true_edep"].to_list() == [40.0, 30.0, 20.0]
|
||||
# post_E isn't written directly for the prediction (energy conservation:
|
||||
# pre_E - delta_e); truth carries it verbatim.
|
||||
assert p["pred_post_E"].to_list() == pytest.approx([100.0 - 42.0, 60.0 - 29.0, 50.0 - 31.0])
|
||||
assert p["true_post_E"].to_list() == [60.0, 30.0, 20.0]
|
||||
# cos_scatter: pre_dir . post_dir, both (0,0,1) here -> 1.0
|
||||
assert p["pred_cos_scatter"].to_list() == pytest.approx([1.0, 1.0, 1.0])
|
||||
assert p["true_cos_scatter"].to_list() == pytest.approx([1.0, 1.0, 1.0])
|
||||
|
||||
|
||||
def test_paired_frame_local_decodes_energy_simplex():
|
||||
lf = _local_prediction_frame()
|
||||
p = paired_frame(lf, "local", has_truth=True).collect()
|
||||
# softmax([1.0, 0.0, 0.0]) * 100 for row 0's pred edep
|
||||
z = np.exp([1.0, 0.0, 0.0])
|
||||
expected_edep_0 = (z[0] / z.sum()) * 100.0
|
||||
assert p["pred_edep"][0] == pytest.approx(expected_edep_0)
|
||||
assert p["pred_step_length"][0] == pytest.approx(1.2, abs=1e-4)
|
||||
# local coord never has a meaningful cos_travel (no reconstructed post_pos)
|
||||
assert "cos_travel" not in [c.rsplit("_", 1)[-1] for c in ["pred_cos_travel"] if c in p.columns] or True
|
||||
assert "pred_cos_travel" not in p.columns
|
||||
|
||||
|
||||
def test_prediction_secondaries_and_pairing():
|
||||
lf = _global_prediction_frame()
|
||||
true_sec = prediction_secondaries(lf, "true").collect()
|
||||
pred_sec = prediction_secondaries(lf, "pred").collect()
|
||||
assert true_sec["pdg"].to_list() == [22, 22, 11]
|
||||
assert pred_sec["pdg"].to_list() == [22, 22]
|
||||
|
||||
pairs = paired_secondaries(lf).collect()
|
||||
# event 1: 1 true, 1 pred -> paired (22, 22); event 2: 2 true, 1 pred -> paired rank0 only (22, 22)
|
||||
assert pairs["true_pdg"].to_list() == [22, 22]
|
||||
assert pairs["pred_pdg"].to_list() == [22, 22]
|
||||
|
||||
|
||||
def test_hist2d_basic():
|
||||
lf = pl.DataFrame({"x": [0.1, 0.5, 0.9, 0.5], "y": [0.1, 0.9, 0.9, 0.1]}).lazy()
|
||||
edges = np.linspace(0.0, 1.0, 3) # 2 bins: [0,0.5), [0.5,1]
|
||||
mat = hist2d(lf, pl.col("x"), pl.col("y"), edges, edges)
|
||||
assert mat.sum() == 4
|
||||
assert mat.shape == (2, 2)
|
||||
|
||||
|
||||
def _ctx_with_predictions(n_marginal_bins: int = 10) -> Context:
|
||||
return build_context(
|
||||
[RolloutSpec("rollout", _rollout_frame())],
|
||||
_reference_frame(),
|
||||
predictions=[PredictionSpec("pred", _global_prediction_frame())],
|
||||
n_energy_bins=2,
|
||||
n_marginal_bins=n_marginal_bins,
|
||||
top_k_pdg=3,
|
||||
sample_rows=1000,
|
||||
)
|
||||
|
||||
|
||||
def test_build_context_resolves_prediction_ranges():
|
||||
ctx = _ctx_with_predictions()
|
||||
assert "edep" in ctx.pred_var_ranges
|
||||
assert "edep" in ctx.pred_residual_ranges
|
||||
assert ctx.pred_top_sec_pdgs # secondaries present in the fixture
|
||||
|
||||
|
||||
def test_prediction_specs_compute_valid_reduced():
|
||||
ctx = _ctx_with_predictions()
|
||||
bundle = Bundle.open(
|
||||
[RolloutSpec("rollout", _rollout_frame())],
|
||||
_reference_frame(),
|
||||
ctx,
|
||||
predictions=[PredictionSpec("pred", _global_prediction_frame())],
|
||||
)
|
||||
for spec_id in (
|
||||
"pred_marginal_edep",
|
||||
"pred_scatter_edep",
|
||||
"pred_residual_edep",
|
||||
"pred_relative_residual_edep",
|
||||
"pred_residual_profile_edep",
|
||||
"pred_ks_summary",
|
||||
"pred_bias_summary",
|
||||
"pred_rmse_summary",
|
||||
"pred_n_sec_confusion",
|
||||
"pred_sec_species_confusion",
|
||||
"pred_dir_alignment_post",
|
||||
"pred_dir_alignment_travel",
|
||||
"pred_constraint_violations",
|
||||
"pred_correlation_delta",
|
||||
):
|
||||
spec = get_spec(spec_id)
|
||||
r = spec.finalize([spec.compute_partial(bundle)], ctx)
|
||||
assert r.id == spec_id
|
||||
assert r.kind != "unavailable", f"{spec_id} unexpectedly unavailable"
|
||||
assert "pred" in r.payload["series"]
|
||||
|
||||
|
||||
def test_prediction_specs_unavailable_without_predictions():
|
||||
ctx = _ctx_with_predictions()
|
||||
bundle = Bundle.open([RolloutSpec("rollout", _rollout_frame())], _reference_frame(), ctx)
|
||||
for spec_id in ("pred_marginal_edep", "pred_scatter_edep", "pred_n_sec_confusion", "pred_ks_summary"):
|
||||
spec = get_spec(spec_id)
|
||||
r = spec.finalize([spec.compute_partial(bundle)], ctx)
|
||||
assert r.kind == "unavailable"
|
||||
assert r.payload["note"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"spec_id",
|
||||
["pred_marginal_edep", "pred_scatter_edep", "pred_n_sec_confusion", "pred_ks_summary", "pred_correlation_delta"],
|
||||
)
|
||||
def test_prediction_chunked_matches_unchunked(spec_id: str):
|
||||
ctx = _ctx_with_predictions()
|
||||
specs = [RolloutSpec("rollout", _rollout_frame())]
|
||||
preds = [PredictionSpec("pred", _global_prediction_frame())]
|
||||
spec = get_spec(spec_id)
|
||||
|
||||
unchunked_bundle = Bundle.open(specs, _reference_frame(), ctx, predictions=preds)
|
||||
unchunked = spec.finalize([spec.compute_partial(unchunked_bundle)], ctx)
|
||||
|
||||
n_chunks = 2
|
||||
parts = [
|
||||
spec.compute_partial(Bundle.open(specs, _reference_frame(), ctx, chunk=(k, n_chunks), predictions=preds))
|
||||
for k in range(n_chunks)
|
||||
]
|
||||
chunked = spec.finalize(parts, ctx)
|
||||
|
||||
assert chunked.kind == unchunked.kind
|
||||
_assert_close(unchunked.payload, chunked.payload)
|
||||
|
||||
|
||||
def _assert_close(a, b) -> None:
|
||||
"""Recursively compare two JSON-shaped payloads (float-tolerant)."""
|
||||
if isinstance(a, dict):
|
||||
assert set(a) == set(b)
|
||||
for k in a:
|
||||
_assert_close(a[k], b[k])
|
||||
elif isinstance(a, list):
|
||||
assert len(a) == len(b)
|
||||
for x, y in zip(a, b):
|
||||
_assert_close(x, y)
|
||||
elif isinstance(a, float):
|
||||
assert np.isclose(a, b, atol=1e-9) or (np.isnan(a) and np.isnan(b))
|
||||
else:
|
||||
assert a == b
|
||||
|
||||
|
||||
def test_paired_scalars_are_subset_of_all_vars():
|
||||
assert set(PAIRED_SCALARS) <= {"step_length", "edep", "delta_e", "post_E"}
|
||||
@@ -6,6 +6,7 @@ from typer.testing import CliRunner
|
||||
|
||||
from giant.cli import (
|
||||
_CEPH_PREDICTIONS,
|
||||
_build_predict_timing,
|
||||
_resolve_prediction_output,
|
||||
_write_prediction_ref,
|
||||
app,
|
||||
@@ -103,6 +104,7 @@ def test_ref_yaml_contains_expected_fields(tmp_path):
|
||||
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
|
||||
data = yaml.safe_load(ref_path.read_text())
|
||||
|
||||
assert data["kind"] == "prediction"
|
||||
assert data["prediction_id"] == pred_uuid
|
||||
assert data["output"] == str(out)
|
||||
assert data["dataset"] == str(dataset)
|
||||
@@ -144,6 +146,46 @@ def test_ref_timestamp_is_iso_format(tmp_path):
|
||||
assert ts.tzinfo is not None
|
||||
|
||||
|
||||
def test_ref_yaml_merges_extra_after_base_fields(tmp_path):
|
||||
ckpt_dir = tmp_path / "checkpoints"
|
||||
ckpt_dir.mkdir()
|
||||
checkpoint = ckpt_dir / "best.pt"
|
||||
checkpoint.touch()
|
||||
|
||||
out = tmp_path / "pred.parquet"
|
||||
dataset = tmp_path / "full.manifest"
|
||||
pred_uuid = str(uuid.uuid4())
|
||||
|
||||
ref_path = _write_prediction_ref(
|
||||
checkpoint,
|
||||
pred_uuid,
|
||||
out,
|
||||
dataset,
|
||||
extra={"coord": "global", "n_rows": 42, "timing": {"setup_s": 1.0}},
|
||||
)
|
||||
data = yaml.safe_load(ref_path.read_text())
|
||||
|
||||
# Base fields untouched, extras layered on top.
|
||||
assert data["kind"] == "prediction"
|
||||
assert data["prediction_id"] == pred_uuid
|
||||
assert data["coord"] == "global"
|
||||
assert data["n_rows"] == 42
|
||||
assert data["timing"] == {"setup_s": 1.0}
|
||||
|
||||
|
||||
def test_ref_yaml_without_extra_matches_today(tmp_path):
|
||||
ckpt_dir = tmp_path / "checkpoints"
|
||||
ckpt_dir.mkdir()
|
||||
checkpoint = ckpt_dir / "best.pt"
|
||||
checkpoint.touch()
|
||||
|
||||
pred_uuid = str(uuid.uuid4())
|
||||
ref_path = _write_prediction_ref(checkpoint, pred_uuid, tmp_path / "p.parquet", tmp_path / "d")
|
||||
data = yaml.safe_load(ref_path.read_text())
|
||||
|
||||
assert set(data) == {"kind", "prediction_id", "output", "dataset", "checkpoint", "timestamp"}
|
||||
|
||||
|
||||
def test_ref_checkpoint_path_is_absolute(tmp_path):
|
||||
ckpt_dir = tmp_path / "checkpoints"
|
||||
ckpt_dir.mkdir()
|
||||
@@ -212,3 +254,72 @@ def test_predict_set_flag_disallowed_path_surfaces_compat_error(tmp_path):
|
||||
|
||||
assert result.exit_code == 1
|
||||
assert "not an inference-safe override" in result.output
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# schema v3 constants (truth-column tagging)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_predict_schema_version_is_v3():
|
||||
from giant.constants import PREDICT_SCHEMA_VERSION
|
||||
|
||||
assert PREDICT_SCHEMA_VERSION == "3"
|
||||
|
||||
|
||||
def test_predict_truth_metadata_key_exists():
|
||||
from giant.constants import PREDICT_TRUTH_METADATA_KEY
|
||||
|
||||
assert PREDICT_TRUTH_METADATA_KEY == "giant.predict.has_truth"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _build_predict_timing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_build_predict_timing_computes_per_step_cost():
|
||||
timing = _build_predict_timing(
|
||||
setup_s=1.0,
|
||||
predict_s=10.0,
|
||||
write_s=2.0,
|
||||
n_rows=100,
|
||||
device="cpu",
|
||||
torch_threads=4,
|
||||
)
|
||||
assert timing["n_rows"] == 100
|
||||
assert timing["sample_s"] == 8.0 # predict_s - write_s
|
||||
assert timing["us_per_step"] == 8.0 / 100 * 1e6
|
||||
assert timing["write_us_per_step"] == 2.0 / 100 * 1e6
|
||||
assert timing["rows_per_s"] == 10.0
|
||||
assert timing["device"] == "cpu" and timing["torch_threads"] == 4
|
||||
|
||||
|
||||
def test_build_predict_timing_handles_zero_rows():
|
||||
timing = _build_predict_timing(
|
||||
setup_s=1.0,
|
||||
predict_s=0.0,
|
||||
write_s=0.0,
|
||||
n_rows=0,
|
||||
device="cpu",
|
||||
torch_threads=1,
|
||||
)
|
||||
assert timing["us_per_step"] is None
|
||||
assert timing["write_us_per_step"] is None
|
||||
assert timing["rows_per_s"] is None
|
||||
|
||||
|
||||
def test_predict_has_truth_flag_default_on():
|
||||
# Inspecting rendered --help text is brittle across terminal
|
||||
# widths/color settings (wraps or re-colors mid-flag); go straight to
|
||||
# the underlying click command's registered option instead.
|
||||
from typing import cast
|
||||
|
||||
import typer
|
||||
from click import Group
|
||||
|
||||
predict_cmd = cast(Group, typer.main.get_command(app)).commands["predict"]
|
||||
truth_param = next(p for p in predict_cmd.params if p.name == "truth")
|
||||
assert truth_param.opts == ["--truth"]
|
||||
assert truth_param.secondary_opts == ["--no-truth"]
|
||||
assert truth_param.default is True
|
||||
|
||||
+137
-2
@@ -16,6 +16,8 @@ from giant.analysis import (
|
||||
compute_one,
|
||||
compute_reduced,
|
||||
derive_run_dir,
|
||||
load_prediction_yaml,
|
||||
load_prediction_yamls,
|
||||
load_rollout_yaml,
|
||||
load_rollout_yamls,
|
||||
merge_one,
|
||||
@@ -25,7 +27,8 @@ from giant.analysis import (
|
||||
from giant.analysis.catalog import get_spec
|
||||
from giant.analysis.condor import Context
|
||||
from giant.analysis.reduced import Partial, Reduced
|
||||
from giant.constants import PREDICT_COORD_METADATA_KEY, ROLLOUT_COORD_VALUE
|
||||
from giant.constants import PREDICT_COORD_METADATA_KEY, PREDICT_TRUTH_METADATA_KEY, ROLLOUT_COORD_VALUE
|
||||
from tests.test_analysis_prediction import _global_prediction_frame
|
||||
from tests.test_analysis_reduce import _reference_frame, _rollout_frame
|
||||
|
||||
|
||||
@@ -86,6 +89,30 @@ def _write_two_inputs(tmp_path: Path) -> tuple[Path, Path]:
|
||||
return paths[0], paths[1]
|
||||
|
||||
|
||||
def _write_prediction(path: Path, coord: str = "global") -> None:
|
||||
tbl = _global_prediction_frame().collect().to_arrow()
|
||||
tbl = tbl.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: coord, PREDICT_TRUTH_METADATA_KEY: "1"})
|
||||
pq.write_table(tbl, path)
|
||||
|
||||
|
||||
def _write_prediction_yaml(tmp_path: Path, reference: Path, tag: str = "p", coord: str = "global") -> Path:
|
||||
pred = tmp_path / f"pred_{tag}.parquet"
|
||||
_write_prediction(pred, coord=coord)
|
||||
yaml_path = tmp_path / f"pred_{tag}.yaml"
|
||||
yaml_path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"kind": "prediction",
|
||||
"prediction_id": f"{tag}pred1234",
|
||||
"output": str(pred),
|
||||
"dataset": str(reference),
|
||||
"checkpoint": f"/ckpt/{tag}.pt",
|
||||
}
|
||||
)
|
||||
)
|
||||
return yaml_path
|
||||
|
||||
|
||||
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"
|
||||
@@ -94,13 +121,14 @@ def _fake_venv(repo_dir: Path) -> None:
|
||||
giant.chmod(0o755)
|
||||
|
||||
|
||||
def _prep(rollout_yamls, run_dir: str | Path | None = None, chunks: int = 1, labels=None) -> Path:
|
||||
def _prep(rollout_yamls, run_dir: str | Path | None = None, chunks: int = 1, labels=None, prediction_yamls=()) -> Path:
|
||||
"""``prep`` with small test-sized context bins/sampling."""
|
||||
return prep(
|
||||
rollout_yamls,
|
||||
run_dir,
|
||||
n_chunks=chunks,
|
||||
labels=labels,
|
||||
prediction_yamls=prediction_yamls,
|
||||
n_energy_bins=2,
|
||||
n_marginal_bins=8,
|
||||
top_k_pdg=3,
|
||||
@@ -161,6 +189,113 @@ def test_load_rollout_yamls_rejects_mismatched_reference(tmp_path: Path):
|
||||
load_rollout_yamls([a, c])
|
||||
|
||||
|
||||
def test_load_prediction_yaml_requires_paths(tmp_path: Path):
|
||||
bad = tmp_path / "bad.yaml"
|
||||
bad.write_text(yaml.safe_dump({"output": "x.parquet"})) # no dataset
|
||||
with pytest.raises(ValueError):
|
||||
load_prediction_yaml(bad)
|
||||
|
||||
|
||||
def test_load_prediction_yaml_rejects_rollout_kind(tmp_path: Path):
|
||||
y = tmp_path / "r.yaml"
|
||||
y.write_text(yaml.safe_dump({"output": "x.parquet", "dataset": "d.parquet", "kind": "rollout"}))
|
||||
with pytest.raises(ValueError, match="kind"):
|
||||
load_prediction_yaml(y)
|
||||
|
||||
|
||||
def test_load_prediction_yamls_single_defaults_to_prediction_name(tmp_path: Path):
|
||||
reference = tmp_path / "reference.parquet"
|
||||
_reference_frame().collect().write_parquet(reference)
|
||||
y = _write_prediction_yaml(tmp_path, reference)
|
||||
loaded = load_prediction_yamls([y], str(reference))
|
||||
assert [lp.name for lp in loaded] == ["prediction"]
|
||||
assert loaded[0].coord == "global"
|
||||
|
||||
|
||||
def test_load_prediction_yamls_multi_defaults_to_stem_and_labels(tmp_path: Path):
|
||||
reference = tmp_path / "reference.parquet"
|
||||
_reference_frame().collect().write_parquet(reference)
|
||||
a = _write_prediction_yaml(tmp_path, reference, tag="a")
|
||||
b = _write_prediction_yaml(tmp_path, reference, tag="b")
|
||||
loaded = load_prediction_yamls([a, b], str(reference))
|
||||
assert [lp.name for lp in loaded] == ["pred_a", "pred_b"]
|
||||
loaded = load_prediction_yamls([a, b], str(reference), labels=["ep20", "ep50"])
|
||||
assert [lp.name for lp in loaded] == ["ep20", "ep50"]
|
||||
|
||||
|
||||
def test_load_prediction_yamls_rejects_mismatched_reference(tmp_path: Path):
|
||||
reference = tmp_path / "reference.parquet"
|
||||
_reference_frame().collect().write_parquet(reference)
|
||||
other_ref = tmp_path / "other_reference.parquet"
|
||||
_reference_frame().collect().write_parquet(other_ref)
|
||||
y = _write_prediction_yaml(tmp_path, other_ref)
|
||||
with pytest.raises(ValueError, match="same reference"):
|
||||
load_prediction_yamls([y], str(reference))
|
||||
|
||||
|
||||
def test_load_prediction_yamls_rejects_mixed_coord(tmp_path: Path):
|
||||
reference = tmp_path / "reference.parquet"
|
||||
_reference_frame().collect().write_parquet(reference)
|
||||
a = _write_prediction_yaml(tmp_path, reference, tag="a", coord="global")
|
||||
b = _write_prediction_yaml(tmp_path, reference, tag="b", coord="local")
|
||||
with pytest.raises(ValueError, match="coord"):
|
||||
load_prediction_yamls([a, b], str(reference))
|
||||
|
||||
|
||||
def test_prep_with_prediction_writes_run_meta(tmp_path: Path):
|
||||
rollout_yaml = _write_inputs(tmp_path)
|
||||
reference = load_rollout_yaml(rollout_yaml)["dataset"]
|
||||
pred_yaml = _write_prediction_yaml(tmp_path, Path(reference))
|
||||
run_dir = _prep([rollout_yaml], prediction_yamls=[pred_yaml])
|
||||
meta = RunMeta.load(run_dir / "run_meta.json")
|
||||
assert [p["name"] for p in meta.predictions] == ["prediction"]
|
||||
assert meta.predictions[0]["plot_meta"]["checkpoint"] == "/ckpt/p.pt"
|
||||
|
||||
computed = compute_one("pred_marginal_edep", run_dir, chunk_index=0)
|
||||
partial = Partial.load(computed)
|
||||
assert partial.data["available"]
|
||||
|
||||
|
||||
def test_prep_forwards_predict_only_metadata_keys(tmp_path: Path):
|
||||
"""A rich `giant predict` sidecar's provenance/timing keys reach
|
||||
run_meta.json's plot_meta, same as a rollout's do — a thin legacy
|
||||
sidecar (no such keys) still loads fine (see _write_prediction_yaml)."""
|
||||
rollout_yaml = _write_inputs(tmp_path)
|
||||
reference = load_rollout_yaml(rollout_yaml)["dataset"]
|
||||
pred = tmp_path / "pred_rich.parquet"
|
||||
_write_prediction(pred, coord="global")
|
||||
yaml_path = tmp_path / "pred_rich.yaml"
|
||||
yaml_path.write_text(
|
||||
yaml.safe_dump(
|
||||
{
|
||||
"kind": "prediction",
|
||||
"prediction_id": "richpred12",
|
||||
"output": str(pred),
|
||||
"dataset": str(reference),
|
||||
"checkpoint": "/ckpt/rich.pt",
|
||||
"coord": "global",
|
||||
"has_truth": True,
|
||||
"schema_version": "3",
|
||||
"n_input_rows": 1000,
|
||||
"n_files": 1,
|
||||
"n_skipped_rows": 3,
|
||||
"unknown_pdg_counts": {"999999": 3},
|
||||
"batch_size_auto": False,
|
||||
"timing": {"us_per_step": 12.5},
|
||||
}
|
||||
)
|
||||
)
|
||||
run_dir = _prep([rollout_yaml], prediction_yamls=[yaml_path])
|
||||
meta = RunMeta.load(run_dir / "run_meta.json")
|
||||
plot_meta = meta.predictions[0]["plot_meta"]
|
||||
assert plot_meta["coord"] == "global"
|
||||
assert plot_meta["has_truth"] is True
|
||||
assert plot_meta["n_input_rows"] == 1000
|
||||
assert plot_meta["n_skipped_rows"] == 3
|
||||
assert plot_meta["unknown_pdg_counts"] == {"999999": 3}
|
||||
assert plot_meta["timing"] == {"us_per_step": 12.5}
|
||||
|
||||
|
||||
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")
|
||||
|
||||
@@ -329,6 +329,58 @@ def test_render_one_of_each_kind(tmp_path: Path):
|
||||
"log_color": True,
|
||||
},
|
||||
),
|
||||
Reduced(
|
||||
"ph1",
|
||||
"prediction",
|
||||
"paired_hist",
|
||||
"Paired hist (single prediction)",
|
||||
"x",
|
||||
{"edges": [0, 1, 2, 3], "series": {"pred": {"pred": [1, 2, 3], "true": [2, 2, 2]}}, "log_y": False},
|
||||
),
|
||||
Reduced(
|
||||
"ph2",
|
||||
"prediction",
|
||||
"paired_hist",
|
||||
"Paired hist (two predictions)",
|
||||
"x",
|
||||
{
|
||||
"edges": [0, 1, 2, 3],
|
||||
"series": {"a": {"pred": [1, 2, 3], "true": [2, 2, 2]}, "b": {"pred": [3, 2, 1]}},
|
||||
"log_y": False,
|
||||
},
|
||||
),
|
||||
Reduced(
|
||||
"hm2d",
|
||||
"prediction",
|
||||
"heatmap2d",
|
||||
"Scatter (truth vs pred)",
|
||||
"true x",
|
||||
{
|
||||
"x_edges": [0, 1, 2],
|
||||
"y_edges": [0, 1, 2],
|
||||
"series": {"pred": [[2, 0], [1, 3]]},
|
||||
"ylabel": "predicted x",
|
||||
"cbar_label": "count",
|
||||
"log_color": True,
|
||||
"diagonal": True,
|
||||
},
|
||||
),
|
||||
Reduced(
|
||||
"profile_noref",
|
||||
"prediction",
|
||||
"profile",
|
||||
"Residual profile (no reference)",
|
||||
"true x",
|
||||
{"edges": [0, 1, 2], "series": {"pred": {"mean": [0.1, -0.1], "std": [0.2, 0.2]}}},
|
||||
),
|
||||
Reduced(
|
||||
"bar_noref",
|
||||
"prediction",
|
||||
"bar",
|
||||
"Constraint violations (no reference)",
|
||||
"check",
|
||||
{"labels": ["a", "b"], "series": {"pred": [0.01, 0.0]}, "ylabel": "rate"},
|
||||
),
|
||||
]
|
||||
try:
|
||||
pdfs = _try_render(reduced, tmp_path)
|
||||
|
||||
Reference in New Issue
Block a user