Files
giant/analysis/export_rollout_observables.py
T
lars 26a176aeaa Add autoregressive shower rollout driver
Closes the loop from single-step prediction into full showers:

- giant/geometry.py + `dwarf build-geometry-oracle`: learn position ->
  (material, layer_id) from data (KNN/SVM) to supply the conditioning the
  surrogate does not predict; flag detector escape by NN distance.
- giant/rollout.py: breadth-first batched frontier that steps all active
  tracks, spawns secondaries as new tracks, and terminates on energy cutoff,
  per-track max steps, escape, or natural end. Energy is deposited locally on
  every stop except escape (leakage), so showers conserve energy exactly.
- `giant rollout` CLI: seed from real events (argmax pre_E), load checkpoint,
  write a world-frame steps parquet + YAML sidecar.
- giant/analysis.py: compute_rollout_observables + plot_rollout_* for
  single-sided longitudinal/transverse/total-energy shower profiles;
  analysis/export_rollout_observables.py driver.
- scikit-learn added as an optional `geometry` extra (lazy-imported).
- Tests: tests/test_geometry.py, tests/test_rollout.py.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-08 10:10:36 +02:00

47 lines
1.6 KiB
Python

"""Export shower-level plots from a `giant rollout` steps parquet.
Not part of the package; run manually. Optionally overlays the real showers
seeded from the same events (a `giant predict --coord local` file) by passing a
reference path. Usage::
python analysis/export_rollout_observables.py ROLLOUT.parquet [REFERENCE_local.parquet]
"""
import sys
from pathlib import Path
import giant.analysis as a
rollout_file = sys.argv[1] if len(sys.argv) > 1 else "rollout.parquet"
reference_file = sys.argv[2] if len(sys.argv) > 2 and sys.argv[2] not in ("", "-") else None
OUT = Path(sys.argv[3]) if len(sys.argv) > 3 else Path(".")
print(f"=== computing rollout observables: {rollout_file} ===")
obs = a.compute_rollout_observables(rollout_file)
tbl = obs.event_table
print(f"n events: {len(tbl)}")
print(
f"total_edep/event: mean={tbl['total_edep'].mean():.4g} MeV "
f"leaked_E/event: mean={tbl['leaked_E'].mean():.4g} MeV "
f"n_tracks/event: mean={tbl['n_tracks'].mean():.1f} "
f"n_steps/event: mean={tbl['n_steps'].mean():.1f}"
)
reference = None
if reference_file is not None:
print(f"=== computing real reference: {reference_file} ===")
reference = a.compute_event_observables_pl(reference_file)
print("=== plots ===")
for name, fn in [
("longitudinal", a.plot_rollout_longitudinal),
("transverse", a.plot_rollout_transverse),
("total-energy", a.plot_rollout_total_energy),
]:
fig = fn(obs, reference=reference)
path = OUT / f"rollout-{name}.png"
fig.savefig(path, dpi=150, bbox_inches="tight")
print(f"wrote {path}")
print("DONE")