26a176aeaa
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>
47 lines
1.6 KiB
Python
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")
|