5 Commits

Author SHA1 Message Date
gitea-actions 8cdeba088e chore: release v0.3.22
CI / Type check (ty) (push) Successful in 2m54s
CI / Format (ruff format) (push) Successful in 2m56s
CI / Lint (ruff check) (push) Successful in 3m0s
CI / Sync project version with tag (hand-pushed tags only) (push) Has been skipped
CI / Publish package to Gitea package registry (push) Has been skipped
CI / Tests (push) Successful in 3m29s
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 14s
2026-09-07 10:02:33 +00:00
lars 48208e6d18 Merge pull request 'feat(analyze): add paired truth/pred plots from giant predict' (#98) from feature/predict-analysis-plots into master
CI / Sync project version with tag (hand-pushed tags only) (push) Has been skipped
CI / Publish package to Gitea package registry (push) Has been skipped
CI / Format (ruff format) (push) Successful in 42s
CI / Type check (ty) (push) Successful in 50s
CI / Lint (ruff check) (push) Successful in 52s
CI / Tests (push) Successful in 3m9s
CI / Release (bump, changelog, badges, tag) on merge to master (push) Successful in 1m47s
Reviewed-on: #98
2026-09-07 11:57:34 +02:00
lars 51790d3e0a feat(predict): enrich YAML sidecar with provenance and timing
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 48s
CI / Lint (ruff check) (pull_request) Successful in 49s
CI / Format (ruff format) (pull_request) Successful in 49s
CI / Tests (pull_request) Successful in 3m14s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
`giant predict`'s sidecar previously stopped at kind/prediction_id/
output/dataset/checkpoint/timestamp, unlike `giant rollout`'s, which
carries full run provenance (model_config, training_epoch,
training_config, timing, ...) that flows into analysis gallery
metadata. `analyze --prediction` consumed the same thin sidecar, so a
prediction series in an analysis run was nearly unlabeled compared to
its rollout counterparts.

- `_write_prediction_ref` takes an `extra: dict | None` merged into
  the sidecar; `giant rollout` now uses it instead of a
  load/update/rewrite round trip (identical output).
- New `_build_predict_timing`, key-compatible with
  `_build_rollout_timing`, from timers now wrapping predict's setup/
  sample/write phases.
- `giant predict` writes coord, has_truth, schema_version, steps,
  weights, device, batch_size(+auto), row/skip/unknown-pdg counts,
  timing, and the checkpoint's model_config/config_overrides/
  training_epoch/best_val_loss/training_config/training_meta.
- `giant/analysis/condor.py`'s `_PLOT_META_KEYS` forwards the new
  predict-only keys (plus rollout's previously-unforwarded
  config_overrides) into each plot's gallery metadata.yaml.
- Fixes a `ty` regression from the prior commit in
  tests/test_cli_predict.py (Command has no static `.commands`).

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01PpxE9nij3ujg9XcuzvQ26q
2026-09-07 11:52:45 +02:00
lars 51f9dad3b0 fix(tests): make predict --truth flag test robust to terminal rendering
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 1m12s
CI / Type check (ty) (pull_request) Failing after 1m12s
CI / Format (ruff format) (pull_request) Successful in 1m12s
CI / Tests (pull_request) Has been skipped
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
The --help-text assertion was brittle to CI's terminal width/color
settings (rich can wrap or re-color the flag name mid-word), causing a
false CI failure even though the flag itself is fine. Inspect the
click command's registered option directly instead of parsing rendered
--help output.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01PpxE9nij3ujg9XcuzvQ26q
2026-09-07 11:35:34 +02:00
lars ac01966a1f feat(analyze): add paired truth/pred plots from giant predict
CI / Sync project version with tag (hand-pushed tags only) (pull_request) Has been skipped
CI / Publish package to Gitea package registry (pull_request) Has been skipped
CI / Lint (ruff check) (pull_request) Successful in 29s
CI / Format (ruff format) (pull_request) Successful in 48s
CI / Type check (ty) (pull_request) Successful in 49s
CI / Tests (pull_request) Failing after 3m5s
CI / Release (bump, changelog, badges, tag) on merge to master (pull_request) Has been skipped
Adds a `prediction` plot family to `giant analyze`, alongside the existing
rollout-vs-reference comparison, and extends `giant predict` to make it
possible:

- `giant predict --coord global` gains schema v3 (`--truth/--no-truth`,
  default on): writes true_* physical columns and true secondary lists
  alongside the predictions, so the output is fully paired.
- New `giant/analysis/prediction.py` builds one canonical true/pred frame
  (`paired_frame`) from either predict coord mode.
- `catalog.py` gains 35 `pred_*` specs: marginals, 2D truth-vs-pred scatter
  (new `heatmap2d` kind), residuals/relative-residuals/calibration profiles,
  KS/bias/RMSE scorecards, n_sec + secondary-species confusion matrices,
  direction-alignment and constraint-violation checks, and a correlation
  delta. Two new Reduced kinds (`paired_hist`, `heatmap2d`) get renderers.
  Every spec degrades to kind="unavailable" with no --prediction given.
- `condor.py`/`cli.py`: `--prediction`/`--prediction-label` on
  `analyze prep`/`submit`, threaded through RunMeta and every compute job.

Full test suite (1162 tests), ruff, and ty all pass.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01PpxE9nij3ujg9XcuzvQ26q
2026-09-07 11:19:19 +02:00
20 changed files with 2440 additions and 72 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
[tool.bumpversion] [tool.bumpversion]
current_version = "0.3.21" current_version = "0.3.22"
parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)" parse = "(?P<major>\\d+)\\.(?P<minor>\\d+)\\.(?P<patch>\\d+)"
serialize = ["{major}.{minor}.{patch}"] serialize = ["{major}.{minor}.{patch}"]
search = "{current_version}" search = "{current_version}"
+10
View File
@@ -1,5 +1,15 @@
# Changelog # 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 ## [0.3.21] - 2026-09-07
### Changed ### Changed
+3 -1
View File
@@ -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 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 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 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 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 giant analyze metrics <train_run_dir> # training-progress plots from metrics.csv
dwarf --help # dataset/tooling CLI: convert, migrate, bump-gen, 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`. **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`. **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 ## Roadmap
+2 -2
View File
@@ -10,8 +10,8 @@ calorimeter showers.
[![python](https://img.shields.io/badge/python-3.12%2B-3776AB?logo=python&logoColor=white)](pyproject.toml) [![python](https://img.shields.io/badge/python-3.12%2B-3776AB?logo=python&logoColor=white)](pyproject.toml)
[![torch](https://img.shields.io/badge/torch-2.3.x-EE4C2C?logo=pytorch&logoColor=white)](pyproject.toml) [![torch](https://img.shields.io/badge/torch-2.3.x-EE4C2C?logo=pytorch&logoColor=white)](pyproject.toml)
[![version](https://img.shields.io/badge/version-0.3.21-informational)](CHANGELOG.md) [![version](https://img.shields.io/badge/version-0.3.22-informational)](CHANGELOG.md)
[![tests](https://img.shields.io/badge/tests-1138%20passing-brightgreen)](tests/) [![tests](https://img.shields.io/badge/tests-1167%20passing-brightgreen)](tests/)
[![CI](https://git.larsbogner.de/lars/giant/actions/workflows/ci.yml/badge.svg?branch=master)](https://git.larsbogner.de/lars/giant/actions) [![CI](https://git.larsbogner.de/lars/giant/actions/workflows/ci.yml/badge.svg?branch=master)](https://git.larsbogner.de/lars/giant/actions)
[![license](https://img.shields.io/badge/license-unlicensed-lightgrey)](#license) [![license](https://img.shields.io/badge/license-unlicensed-lightgrey)](#license)
+8
View File
@@ -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.catalog import build_catalog, catalog_ids, get_spec
from giant.analysis.condor import ( from giant.analysis.condor import (
LoadedPrediction,
LoadedRollout, LoadedRollout,
RunMeta, RunMeta,
SubmitConfig, SubmitConfig,
compute_one, compute_one,
compute_reduced, compute_reduced,
derive_run_dir, derive_run_dir,
load_prediction_yaml,
load_prediction_yamls,
load_rollout_yaml, load_rollout_yaml,
load_rollout_yamls, load_rollout_yamls,
merge_all, merge_all,
@@ -27,6 +30,7 @@ from giant.analysis.condor import (
write_submit, write_submit,
) )
from giant.analysis.context import Context, build_context from giant.analysis.context import Context, build_context
from giant.analysis.prediction import PredictionSpec
from giant.analysis.reduced import Partial, Reduced from giant.analysis.reduced import Partial, Reduced
from giant.analysis.runtime_estimate import RUNTIME_SAFETY_MARGIN, estimate_runtime_s from giant.analysis.runtime_estimate import RUNTIME_SAFETY_MARGIN, estimate_runtime_s
from giant.analysis.sources import RolloutSpec, Side from giant.analysis.sources import RolloutSpec, Side
@@ -34,8 +38,10 @@ from giant.analysis.sources import RolloutSpec, Side
__all__ = [ __all__ = [
"RUNTIME_SAFETY_MARGIN", "RUNTIME_SAFETY_MARGIN",
"Context", "Context",
"LoadedPrediction",
"LoadedRollout", "LoadedRollout",
"Partial", "Partial",
"PredictionSpec",
"Reduced", "Reduced",
"RolloutSpec", "RolloutSpec",
"RunMeta", "RunMeta",
@@ -49,6 +55,8 @@ __all__ = [
"derive_run_dir", "derive_run_dir",
"estimate_runtime_s", "estimate_runtime_s",
"get_spec", "get_spec",
"load_prediction_yaml",
"load_prediction_yamls",
"load_rollout_yaml", "load_rollout_yaml",
"load_rollout_yamls", "load_rollout_yamls",
"merge_all", "merge_all",
+943 -4
View File
@@ -39,7 +39,8 @@ variable x grouping, secondaries, ...) into concrete specs.
from __future__ import annotations from __future__ import annotations
from collections.abc import Callable 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 numpy as np
import polars as pl import polars as pl
@@ -52,12 +53,25 @@ from giant.analysis.grouping import (
material_label, material_label,
pdg_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 ( from giant.analysis.reduce import (
attach_entry_axis, attach_entry_axis,
binned_moments,
depth_expr, depth_expr,
entry_axis, entry_axis,
event_scalars, event_scalars,
hist1d, hist1d,
hist2d,
leakage_fraction, leakage_fraction,
profile_finalize, profile_finalize,
profile_partial, profile_partial,
@@ -93,6 +107,7 @@ class Bundle:
rollouts: dict[str, RolloutSide] # name -> frames, insertion order = CLI order rollouts: dict[str, RolloutSide] # name -> frames, insertion order = CLI order
t_all: pl.LazyFrame # reference, all rows t_all: pl.LazyFrame # reference, all rows
t_phys: pl.LazyFrame # reference, physical steps only t_phys: pl.LazyFrame # reference, physical steps only
predictions: dict[str, PredictionSide] = field(default_factory=dict) # name -> paired frames, CLI order
@classmethod @classmethod
def open( def open(
@@ -101,12 +116,14 @@ class Bundle:
reference, reference,
ctx: Context, ctx: Context,
chunk: tuple[int, int] | None = None, chunk: tuple[int, int] | None = None,
predictions: list[PredictionSpec] | None = None,
) -> Bundle: ) -> 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 ``chunk = (chunk_index, n_chunks)`` filters every side to
``event_id % n_chunks == chunk_index`` *before* deriving the physical/ ``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, row-local or a ``group_by("event_id")``) sees a self-contained,
event-disjoint slice no cross-chunk lookups are ever needed. 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, type_embedding_l1_dist=rs.type_embedding_l1_dist,
timing=rs.timing, 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 @dataclass
@@ -1076,6 +1110,807 @@ MARGINAL_VARS = ["step_length", "edep", "delta_e", "post_E", "cos_scatter"]
GROUPING_AXES = ["energy", "pdg", "material"] 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]: def build_catalog() -> list[PlotSpec]:
"""All concrete plot specs, each with a unique id.""" """All concrete plot specs, each with a unique id."""
specs: list[PlotSpec] = [] specs: list[PlotSpec] = []
@@ -1279,6 +2114,110 @@ def build_catalog() -> list[PlotSpec]:
chunkable=False, 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 return specs
+125 -13
View File
@@ -56,11 +56,14 @@ import yaml
from giant.analysis.catalog import Bundle, catalog_ids, get_spec from giant.analysis.catalog import Bundle, catalog_ids, get_spec
from giant.analysis.context import Context, build_context 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.reduced import Partial
from giant.analysis.runtime_estimate import estimate_runtime_s from giant.analysis.runtime_estimate import estimate_runtime_s
from giant.analysis.sources import RolloutSpec, Side, open_side from giant.analysis.sources import RolloutSpec, Side, open_side
# Keys copied verbatim from a rollout YAML into each plot's gallery metadata. # 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 = ( _PLOT_META_KEYS = (
"prediction_id", "prediction_id",
"checkpoint", "checkpoint",
@@ -84,10 +87,20 @@ _PLOT_META_KEYS = (
"termination_reason_counts", "termination_reason_counts",
"timing", "timing",
"model_config", "model_config",
"config_overrides",
"training_epoch", "training_epoch",
"best_val_loss", "best_val_loss",
"training_config", "training_config",
"training_meta", "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 # Diagnostic — only present when giant rollout ran under
# stage2_model.particle_type.target="embedding" (see giant/cli.py's # stage2_model.particle_type.target="embedding" (see giant/cli.py's
# rollout command and giant.rollout.L1DistCollector); absent otherwise, # 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"] 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: def _run_tag(y: dict) -> str:
rollout = Path(y["output"]) rollout = Path(y["output"])
return str(y.get("prediction_id") or rollout.stem)[:8] 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. # Empty/0 on run directories written before this field existed.
rows_per_chunk: list[int] = field(default_factory=list) rows_per_chunk: list[int] = field(default_factory=list)
total_rows: int = 0 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: def save(self, path: str | Path) -> None:
Path(path).write_text(json.dumps(self.__dict__, indent=2)) Path(path).write_text(json.dumps(self.__dict__, indent=2))
@@ -232,8 +315,13 @@ class RunMeta:
return cls(**json.loads(Path(path).read_text())) return cls(**json.loads(Path(path).read_text()))
def _rows_per_chunk(rollouts: list[str | Path], reference: str | Path, n_chunks: int) -> list[int]: def _rows_per_chunk(
"""Combined rollout+reference row count of each ``event_id % n_chunks`` 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) One cheap streaming ``group_by`` per side (just the ``event_id`` column)
the sizing input every job's estimated walltime 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 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: for lf in sides:
df = counts(lf) df = counts(lf)
for c, n in zip(df["_c"].to_list(), df["n"].to_list()): for c, n in zip(df["_c"].to_list(), df["n"].to_list()):
@@ -263,26 +355,33 @@ def prep(
n_chunks: int = 1, n_chunks: int = 1,
default_base: str | Path | None = None, default_base: str | Path | None = None,
labels: Sequence[str] | None = None, labels: Sequence[str] | None = None,
prediction_yamls: Sequence[str | Path] = (),
prediction_labels: Sequence[str] | None = None,
**ctx_kwargs, **ctx_kwargs,
) -> Path: ) -> 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. Writes ``shared.json`` + ``run_meta.json`` and returns the run directory.
``n_chunks`` is the run-level chunk count every ``compute-one``/``merge-one`` ``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 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 resolved once here rather than re-passed (and risking disagreement) at every
later step. See ``derive_run_dir`` for how ``run_dir``/``default_base`` later step. See ``derive_run_dir`` for how ``run_dir``/``default_base``
resolve the actual directory, and ``load_rollout_yamls`` for how resolve the actual directory, ``load_rollout_yamls`` for how
``labels``/YAML stems resolve each rollout's series name. ``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 Clears any existing ``reduced_partial/``/``reduced/`` from a prior prep of
this same ``run_dir``: partial files carry no record of what context this same ``run_dir``: partial files carry no record of what context
(``n_chunks``, bin edges, group sets) they were computed under, so (``n_chunks``, bin edges, group sets) they were computed under, so
re-prepping with a different ``n_chunks``/``**ctx_kwargs`` (or after the re-prepping with a different ``n_chunks``/``**ctx_kwargs`` (or after the
rollout/reference files changed) would otherwise let ``merge_one`` silently rollout/reference/prediction files changed) would otherwise let
merge stale partials against the new ``shared.json``. ``merge_one`` silently merge stale partials against the new ``shared.json``.
""" """
loaded, reference = load_rollout_yamls(list(rollout_yamls), labels) 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 = derive_run_dir([lr.yaml for lr in loaded], run_dir, default_base=default_base)
run_path.mkdir(parents=True, exist_ok=True) run_path.mkdir(parents=True, exist_ok=True)
@@ -292,14 +391,20 @@ def prep(
shutil.rmtree(stale_dir) shutil.rmtree(stale_dir)
rollout_specs = [RolloutSpec(name=lr.name, source=lr.yaml["output"]) for lr in loaded] 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") 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 = [ rollouts_meta = [
{"name": lr.name, "path": str(lr.yaml["output"]), "plot_meta": _plot_meta(lr.yaml)} for lr in loaded {"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) ckpts = ", ".join(Path(lr.yaml.get("checkpoint", "")).name or "rollout" for lr in loaded)
RunMeta( RunMeta(
@@ -310,6 +415,7 @@ def prep(
n_chunks=n_chunks, n_chunks=n_chunks,
rows_per_chunk=rows_per_chunk, rows_per_chunk=rows_per_chunk,
total_rows=sum(rows_per_chunk), total_rows=sum(rows_per_chunk),
predictions=predictions_meta,
).save(run_path / "run_meta.json") ).save(run_path / "run_meta.json")
return run_path return run_path
@@ -327,12 +433,15 @@ def compute_reduced(
out: str | Path, out: str | Path,
chunk_index: int = 0, chunk_index: int = 0,
n_chunks: int = 1, n_chunks: int = 1,
predictions: Sequence[dict] = (),
) -> Path: ) -> Path:
"""Core: run one (plot, chunk)'s partial reduction against explicit paths. """Core: run one (plot, chunk)'s partial reduction against explicit paths.
``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?, ``rollouts``: ``[{"name", "path", "checkpoint"?, "type_embedding_l1_dist"?,
"timing"?}, ...]``, one per rollout series (insertion order preserved "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 Writes a ``Partial`` JSON the raw, not-yet-merged output of
``PlotSpec.compute_partial`` never a finished ``Reduced``; ``merge_one`` ``PlotSpec.compute_partial`` never a finished ``Reduced``; ``merge_one``
@@ -357,7 +466,8 @@ def compute_reduced(
) )
for r in rollouts 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( partial = Partial(
id=spec_id, id=spec_id,
family=spec.family, 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 for ro in meta.rollouts
] ]
predictions = [{"name": p["name"], "path": p["path"]} for p in meta.predictions]
return compute_reduced( return compute_reduced(
spec_id, spec_id,
rollouts, 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", run_path / "reduced_partial" / f"{spec_id}__{chunk_index}.json",
chunk_index=chunk_index, chunk_index=chunk_index,
n_chunks=meta.n_chunks, n_chunks=meta.n_chunks,
predictions=predictions,
) )
+64
View File
@@ -20,6 +20,7 @@ import numpy as np
import polars as pl import polars as pl
from giant.analysis.grouping import energy_bin_edges 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 ( from giant.analysis.reduce import (
attach_entry_axis, attach_entry_axis,
depth_expr, depth_expr,
@@ -44,6 +45,12 @@ class Context:
sec_energy_range: tuple[float, float] sec_energy_range: tuple[float, float]
n_sec_bins: int n_sec_bins: int
n_events: dict[str, int] = field(default_factory=dict) 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 ------------------------------------------------- # -- (de)serialization -------------------------------------------------
def save(self, path: str | Path) -> None: def save(self, path: str | Path) -> None:
@@ -54,6 +61,10 @@ class Context:
d = json.loads(Path(path).read_text()) d = json.loads(Path(path).read_text())
d["var_ranges"] = {k: tuple(v) for k, v in d["var_ranges"].items()} d["var_ranges"] = {k: tuple(v) for k, v in d["var_ranges"].items()}
d["sec_energy_range"] = tuple(d["sec_energy_range"]) 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) return cls(**d)
# -- convenience ------------------------------------------------------- # -- convenience -------------------------------------------------------
@@ -61,6 +72,14 @@ class Context:
lo, hi = self.var_ranges[var] lo, hi = self.var_ranges[var]
return np.linspace(lo, hi, self.n_marginal_bins + 1) 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 _LO_Q, _HI_Q = 0.001, 0.999
@@ -87,10 +106,13 @@ def build_context(
rollouts: list[RolloutSpec], rollouts: list[RolloutSpec],
reference: str | Path | pl.LazyFrame, reference: str | Path | pl.LazyFrame,
*, *,
predictions: list[PredictionSpec] | None = None,
n_energy_bins: int = 4, n_energy_bins: int = 4,
n_marginal_bins: int = 50, n_marginal_bins: int = 50,
n_sec_bins: int = 40, n_sec_bins: int = 40,
top_k_pdg: int = 6, top_k_pdg: int = 6,
pred_n_sec_cap: int = 10,
top_k_sec_pdg: int = 8,
sample_rows: int = 1_000_000, sample_rows: int = 1_000_000,
seed: int = 0, seed: int = 0,
) -> Context: ) -> Context:
@@ -165,6 +187,44 @@ def build_context(
} }
sec_energy_range = _combined_quantiles([t_se, *r_se.values()], _LO_Q, _HI_Q) 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( return Context(
n_marginal_bins=n_marginal_bins, n_marginal_bins=n_marginal_bins,
var_ranges=var_ranges, var_ranges=var_ranges,
@@ -173,6 +233,10 @@ def build_context(
materials=materials, materials=materials,
depth_edges=[float(x) for x in depth_edges], depth_edges=[float(x) for x in depth_edges],
transverse_edges=[float(x) for x in transverse_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, sec_energy_range=sec_energy_range,
n_sec_bins=n_sec_bins, n_sec_bins=n_sec_bins,
n_events={ n_events={
+315
View File
@@ -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"),
)
)
+58
View File
@@ -70,6 +70,64 @@ def hist1d(
return out 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]: def sum_merge(dicts: list[dict[str, Any]]) -> dict[str, Any]:
"""Elementwise-sum a list of sum-mergeable count/total dicts (JSON-safe keys). """Elementwise-sum a list of sum-mergeable count/total dicts (JSON-safe keys).
+7 -2
View File
@@ -27,8 +27,13 @@ from pathlib import Path
# "router_specialization" max gate weight vs energy (one scalar trend line # "router_specialization" max gate weight vs energy (one scalar trend line
# summarizing "router_gating"), per rollout with an enabled router # summarizing "router_gating"), per rollout with an enabled router
# "heatmap" row x col matrix + colorbar, one panel per rollout (a # "heatmap" row x col matrix + colorbar, one panel per rollout (a
# distance scorecard) # distance scorecard) or per prediction (a confusion matrix)
# "unavailable" plot not applicable to this run (e.g. no MoE checkpoint) # "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 @dataclass
+74
View File
@@ -422,6 +422,78 @@ def _render_heatmap(r: Reduced, params: dict):
return fig 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): def _render_unavailable(r: Reduced, params: dict):
fig, ax = ps.new_figure("thesis-single", title=r.title, params=params) fig, ax = ps.new_figure("thesis-single", title=r.title, params=params)
ax.axis("off") ax.axis("off")
@@ -448,6 +520,8 @@ _RENDERERS = {
"router_share": _render_router_share, "router_share": _render_router_share,
"router_specialization": _render_router_specialization, "router_specialization": _render_router_specialization,
"heatmap": _render_heatmap, "heatmap": _render_heatmap,
"paired_hist": _render_paired_hist,
"heatmap2d": _render_heatmap2d,
"unavailable": _render_unavailable, "unavailable": _render_unavailable,
} }
+244 -44
View File
@@ -20,6 +20,7 @@ from giant.constants import (
PREDICT_COORD_METADATA_KEY, PREDICT_COORD_METADATA_KEY,
PREDICT_SCHEMA_VERSION, PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY, PREDICT_SCHEMA_VERSION_KEY,
PREDICT_TRUTH_METADATA_KEY,
ROLLOUT_COORD_VALUE, ROLLOUT_COORD_VALUE,
) )
@@ -186,11 +187,18 @@ def _write_prediction_ref(
out: Path, out: Path,
dataset_path: Path, dataset_path: Path,
comment: str | None = None, comment: str | None = None,
extra: dict | None = None,
) -> Path: ) -> 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 import yaml
ref = { ref = {
"kind": "prediction",
"prediction_id": pred_uuid, "prediction_id": pred_uuid,
"output": str(out), "output": str(out),
"dataset": str(dataset_path), "dataset": str(dataset_path),
@@ -199,6 +207,8 @@ def _write_prediction_ref(
} }
if comment is not None: if comment is not None:
ref["comment"] = comment ref["comment"] = comment
if extra:
ref.update(extra)
ref_path = checkpoint.parent / f"{pred_uuid}.yaml" ref_path = checkpoint.parent / f"{pred_uuid}.yaml"
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False)) ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
return ref_path 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() @app.callback()
def _main() -> None: def _main() -> None:
"""GIANT — Geant4 step-function surrogate.""" """GIANT — Geant4 step-function surrogate."""
@@ -1020,9 +1065,21 @@ def predict(
typer.Option( typer.Option(
"--out", "--out",
"-o", "-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, ] = 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[ batch_size: Annotated[
str, str,
typer.Option( typer.Option(
@@ -1071,6 +1128,8 @@ def predict(
] = None, ] = None,
) -> None: ) -> None:
"""Run trained model on a parquet file and save predictions.""" """Run trained model on a parquet file and save predictions."""
import time
import numpy as np import numpy as np
import pyarrow as pa import pyarrow as pa
import pyarrow.parquet as pq import pyarrow.parquet as pq
@@ -1090,6 +1149,8 @@ def predict(
from giant.rollout import decode_secondary_identity from giant.rollout import decode_secondary_identity
from giant.sample import resolve_n_sec, sample_stage1, sample_stage2 from giant.sample import resolve_n_sec, sample_stage1, sample_stage2
_t_setup_start = time.perf_counter()
batch_size_auto = False batch_size_auto = False
batch_size_value: int | None = None batch_size_value: int | None = None
if batch_size.strip().lower() == "auto": if batch_size.strip().lower() == "auto":
@@ -1161,8 +1222,17 @@ def predict(
unknown_pdg_counts: Counter[int] = Counter() unknown_pdg_counts: Counter[int] = Counter()
total_rows = sum(pq.ParquetFile(path).metadata.num_rows for path in files) 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): 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_file_chunks(path, offset=offset, k_max=stage2_k_max)
return iter_cond_chunks(path, offset=offset) 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} return {k: np.concatenate([a[k], b[k]], axis=0) for k in a}
def _process(piece: dict[str, np.ndarray]) -> None: def _process(piece: dict[str, np.ndarray]) -> None:
nonlocal writer, total nonlocal writer, total, _write_s
if coord == Coord.local: if coord == Coord.local:
feats = build_features( 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_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)] sec_dz_list = [sec_dir_world[i, :n, 2].tolist() for i, n in enumerate(n_sec_pred_np)]
table = pa.table( columns = {
{ "event_id": piece["event_id"],
"event_id": piece["event_id"], "pdg": piece["pdg"],
"pdg": piece["pdg"], "pre_x": piece["pre_pos"][:, 0],
"pre_x": piece["pre_pos"][:, 0], "pre_y": piece["pre_pos"][:, 1],
"pre_y": piece["pre_pos"][:, 1], "pre_z": piece["pre_pos"][:, 2],
"pre_z": piece["pre_pos"][:, 2], "pre_E": piece["pre_E"],
"pre_E": piece["pre_E"], "pre_dx": piece["pre_dir"][:, 0],
"pre_dx": piece["pre_dir"][:, 0], "pre_dy": piece["pre_dir"][:, 1],
"pre_dy": piece["pre_dir"][:, 1], "pre_dz": piece["pre_dir"][:, 2],
"pre_dz": piece["pre_dir"][:, 2], "material": piece["material"],
"material": piece["material"], "layer_id": piece["layer_id"],
"layer_id": piece["layer_id"], "n_sec": piece["n_sec"],
"n_sec": piece["n_sec"], "n_sec_pred": n_sec_pred_np,
"n_sec_pred": n_sec_pred_np, "step_length": step_length,
"step_length": step_length, "delta_e": delta_e,
"delta_e": delta_e, "edep": edep,
"edep": edep, "post_dx": post_dir_world[:, 0],
"post_dx": post_dir_world[:, 0], "post_dy": post_dir_world[:, 1],
"post_dy": post_dir_world[:, 1], "post_dz": post_dir_world[:, 2],
"post_dz": post_dir_world[:, 2], "post_x": post_pos_world[:, 0],
"post_x": post_pos_world[:, 0], "post_y": post_pos_world[:, 1],
"post_y": post_pos_world[:, 1], "post_z": post_pos_world[:, 2],
"post_z": post_pos_world[:, 2], "sec_pdg_list": sec_pdg_list,
"sec_pdg_list": sec_pdg_list, "sec_E_list": sec_E_list,
"sec_E_list": sec_E_list, "sec_dx_list": sec_dx_list,
"sec_dx_list": sec_dx_list, "sec_dy_list": sec_dy_list,
"sec_dy_list": sec_dy_list, "sec_dz_list": sec_dz_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( table = table.replace_schema_metadata(
{ {
PREDICT_COORD_METADATA_KEY: coord.value, PREDICT_COORD_METADATA_KEY: coord.value,
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION, 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: if writer is None:
writer = pq.ParquetWriter(out, table.schema) writer = pq.ParquetWriter(out, table.schema)
writer.write_table(table) writer.write_table(table)
_write_s += time.perf_counter() - _t0
total += len(piece["event_id"]) total += len(piece["event_id"])
# Buffer rows across row-group boundaries so the inference batch size # Buffer rows across row-group boundaries so the inference batch size
# isn't capped by however the source file happens to be chunked. # isn't capped by however the source file happens to be chunked.
buffer: dict[str, np.ndarray] | None = None 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) bar = tqdm(total=total_rows, desc="predict", unit="row", dynamic_ncols=True)
for i, path in enumerate(files): for i, path in enumerate(files):
for chunk in chunk_iter(path, offset=event_id_offset(i)): for chunk in chunk_iter(path, offset=event_id_offset(i)):
@@ -1365,9 +1476,56 @@ def predict(
bar.close() bar.close()
if writer is not None: if writer is not None:
writer.close() 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}") 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: if skipped:
codes = ", ".join(f"{pdg} ({count})" for pdg, count in sorted(unknown_pdg_counts.items())) 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 as pa
import pyarrow.parquet as pq import pyarrow.parquet as pq
import torch import torch
import yaml
from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference from giant.checkpoint_io import CheckpointCompatibilityError, load_for_inference
from giant.data.loader import find_parquet_files from giant.data.loader import find_parquet_files
@@ -1628,10 +1785,12 @@ def rollout(
l1_summary = l1_dist_collector.summary() l1_summary = l1_dist_collector.summary()
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path) ref_path = _write_prediction_ref(
ref = yaml.safe_load(ref_path.read_text()) checkpoint,
ref.update( pred_uuid,
{ out,
dataset_path,
extra={
"kind": "rollout", "kind": "rollout",
"geometry_oracle": str(geometry.resolve()), "geometry_oracle": str(geometry.resolve()),
"energy_cutoff": energy_cutoff, "energy_cutoff": energy_cutoff,
@@ -1668,9 +1827,8 @@ def rollout(
# — empty dicts if the checkpoint has no config.toml next to it. # — empty dicts if the checkpoint has no config.toml next to it.
"training_config": dict(training_cfg.get("train", {})), "training_config": dict(training_cfg.get("train", {})),
"training_meta": dict(training_cfg.get("meta", {})), "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"wrote {summary['n_rows']:,} step rows → {out}")
typer.echo(f"terminations: {summary['termination_reason_counts']}") typer.echo(f"terminations: {summary['termination_reason_counts']}")
@@ -1708,6 +1866,25 @@ def analyze_prep(
"single YAML).", "single YAML).",
), ),
] = None, ] = 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[ run_dir: Annotated[
Path | None, Path | None,
typer.Option( typer.Option(
@@ -1724,7 +1901,7 @@ def analyze_prep(
typer.Option("--chunks", help="Split each plot's data into this many event_id chunks"), typer.Option("--chunks", help="Split each plot's data into this many event_id chunks"),
] = 1, ] = 1,
) -> None: ) -> 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 from giant.analysis import prep
path = prep( path = prep(
@@ -1733,6 +1910,8 @@ def analyze_prep(
n_chunks=chunks, n_chunks=chunks,
default_base=Path.cwd() / "analysis_runs", default_base=Path.cwd() / "analysis_runs",
labels=label, labels=label,
prediction_yamls=prediction or (),
prediction_labels=prediction_label,
n_energy_bins=n_energy_bins, n_energy_bins=n_energy_bins,
n_marginal_bins=n_marginal_bins, n_marginal_bins=n_marginal_bins,
top_k_pdg=top_k_pdg, top_k_pdg=top_k_pdg,
@@ -1831,6 +2010,25 @@ def analyze_submit(
"single YAML).", "single YAML).",
), ),
] = None, ] = 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[ run_dir: Annotated[
Path | None, Path | None,
typer.Option( typer.Option(
@@ -1868,6 +2066,8 @@ def analyze_submit(
n_chunks=chunks, n_chunks=chunks,
default_base=Path.cwd() / "analysis_runs", default_base=Path.cwd() / "analysis_runs",
labels=label, labels=label,
prediction_yamls=prediction or (),
prediction_labels=prediction_label,
n_energy_bins=n_energy_bins, n_energy_bins=n_energy_bins,
n_marginal_bins=n_marginal_bins, n_marginal_bins=n_marginal_bins,
top_k_pdg=top_k_pdg, top_k_pdg=top_k_pdg,
+6 -1
View File
@@ -64,7 +64,12 @@ LOCAL_TARGET_NAMES = [
# guessing from its column names. # guessing from its column names.
PREDICT_COORD_METADATA_KEY = "giant.predict.coord" PREDICT_COORD_METADATA_KEY = "giant.predict.coord"
PREDICT_SCHEMA_VERSION_KEY = "giant.predict.schema_version" 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, # Coord-metadata value tagging a `giant rollout` steps parquet (world frame,
# autoregressive shower output). Distinct from predict's "global"/"local". # autoregressive shower output). Distinct from predict's "global"/"local".
+1 -1
View File
@@ -1,6 +1,6 @@
[project] [project]
name = "giant" name = "giant"
version = "0.3.21" version = "0.3.22"
description = "Geant4 step-function surrogate via conditional flow matching" description = "Geant4 step-function surrogate via conditional flow matching"
readme = "README.md" readme = "README.md"
requires-python = ">=3.12" requires-python = ">=3.12"
+278
View File
@@ -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"}
+111
View File
@@ -6,6 +6,7 @@ from typer.testing import CliRunner
from giant.cli import ( from giant.cli import (
_CEPH_PREDICTIONS, _CEPH_PREDICTIONS,
_build_predict_timing,
_resolve_prediction_output, _resolve_prediction_output,
_write_prediction_ref, _write_prediction_ref,
app, app,
@@ -103,6 +104,7 @@ def test_ref_yaml_contains_expected_fields(tmp_path):
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset) ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset)
data = yaml.safe_load(ref_path.read_text()) data = yaml.safe_load(ref_path.read_text())
assert data["kind"] == "prediction"
assert data["prediction_id"] == pred_uuid assert data["prediction_id"] == pred_uuid
assert data["output"] == str(out) assert data["output"] == str(out)
assert data["dataset"] == str(dataset) assert data["dataset"] == str(dataset)
@@ -144,6 +146,46 @@ def test_ref_timestamp_is_iso_format(tmp_path):
assert ts.tzinfo is not None 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): def test_ref_checkpoint_path_is_absolute(tmp_path):
ckpt_dir = tmp_path / "checkpoints" ckpt_dir = tmp_path / "checkpoints"
ckpt_dir.mkdir() 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 result.exit_code == 1
assert "not an inference-safe override" in result.output 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
View File
@@ -16,6 +16,8 @@ from giant.analysis import (
compute_one, compute_one,
compute_reduced, compute_reduced,
derive_run_dir, derive_run_dir,
load_prediction_yaml,
load_prediction_yamls,
load_rollout_yaml, load_rollout_yaml,
load_rollout_yamls, load_rollout_yamls,
merge_one, merge_one,
@@ -25,7 +27,8 @@ from giant.analysis import (
from giant.analysis.catalog import get_spec from giant.analysis.catalog import get_spec
from giant.analysis.condor import Context from giant.analysis.condor import Context
from giant.analysis.reduced import Partial, Reduced 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 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] 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: def _fake_venv(repo_dir: Path) -> None:
"""Stand in for a `uv sync`'d venv: write_submit checks `.venv/bin/giant` exists.""" """Stand in for a `uv sync`'d venv: write_submit checks `.venv/bin/giant` exists."""
giant = repo_dir / ".venv" / "bin" / "giant" giant = repo_dir / ".venv" / "bin" / "giant"
@@ -94,13 +121,14 @@ def _fake_venv(repo_dir: Path) -> None:
giant.chmod(0o755) 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.""" """``prep`` with small test-sized context bins/sampling."""
return prep( return prep(
rollout_yamls, rollout_yamls,
run_dir, run_dir,
n_chunks=chunks, n_chunks=chunks,
labels=labels, labels=labels,
prediction_yamls=prediction_yamls,
n_energy_bins=2, n_energy_bins=2,
n_marginal_bins=8, n_marginal_bins=8,
top_k_pdg=3, top_k_pdg=3,
@@ -161,6 +189,113 @@ def test_load_rollout_yamls_rejects_mismatched_reference(tmp_path: Path):
load_rollout_yamls([a, c]) 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(): def test_derive_run_dir_next_to_rollout():
y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"} y = {"output": "/data/roll.parquet", "prediction_id": "abcd1234ef", "dataset": "d"}
assert derive_run_dir([y]) == Path("/data/analysis_abcd1234") assert derive_run_dir([y]) == Path("/data/analysis_abcd1234")
+52
View File
@@ -329,6 +329,58 @@ def test_render_one_of_each_kind(tmp_path: Path):
"log_color": True, "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: try:
pdfs = _try_render(reduced, tmp_path) pdfs = _try_render(reduced, tmp_path)
Generated
+1 -1
View File
@@ -825,7 +825,7 @@ wheels = [
[[package]] [[package]]
name = "giant" name = "giant"
version = "0.3.21" version = "0.3.22"
source = { editable = "." } source = { editable = "." }
dependencies = [ dependencies = [
{ name = "numpy" }, { name = "numpy" },