Add event-level shower observables to giant.analysis

Aggregates giant predict --coord local output per event_id into total
deposited energy, longitudinal/transverse shower profiles, and shower-max
depth, reconstructed into world-frame physical units (mm, MeV). Streams the
file in two polars passes rather than building a SampleCollection, since
per-event sums would be corrupted by row subsampling on these large files.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
2026-06-22 15:52:47 +02:00
co-authored by Claude Sonnet 4.6
parent 6fc68fe1aa
commit f81263fa6e
2 changed files with 568 additions and 3 deletions
+191 -1
View File
@@ -3,6 +3,7 @@ import matplotlib
matplotlib.use("Agg") # no display needed for plot smoke tests
import numpy as np
import polars as pl
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
@@ -10,6 +11,7 @@ import pytest
from giant.analysis import (
RAW_TARGET_NAMES,
SampleCollection,
compute_event_observables_pl,
constraint_report,
constraint_report_pl,
correlation_matrices,
@@ -22,8 +24,12 @@ from giant.analysis import (
plot_direction_alignment,
plot_kl_bars,
plot_kl_bars_pl,
plot_longitudinal_profile,
plot_marginals,
plot_pairwise,
plot_shower_max_depth,
plot_total_energy,
plot_transverse_profile,
)
from giant.constants import (
LOCAL_TARGET_NAMES,
@@ -31,7 +37,7 @@ from giant.constants import (
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
)
from giant.data.transforms import log_transform
from giant.data.transforms import inv_log_transform, log_transform, reconstruct_post_pos
def _unit_vectors(rng, n):
@@ -352,3 +358,187 @@ def test_constraint_report_pl_matches_numpy_version(tmp_path):
actual["mean_abs_error"].to_numpy(),
atol=1e-4,
)
# ---------------------------------------------------------------------------
# Tier 4: event-level (shower) observables
# ---------------------------------------------------------------------------
def _make_event_level_arrays(rng):
"""3 events (3/2/4 steps), each with an unambiguous highest-pre_E row.
The forced max-pre_E rows (indices 1, 3, 7) fix a known shower axis/entry
point per event, so the expected event_table can be re-derived
independently in the test without depending on compute_event_observables_pl.
"""
event_id = np.array([0, 0, 0, 1, 1, 2, 2, 2, 2], dtype=np.int64)
n = len(event_id)
pre_pos = rng.uniform(-5.0, 5.0, (n, 3)).astype(np.float32)
pre_dir = _unit_vectors(rng, n)
pre_E = rng.uniform(1.0, 50.0, n).astype(np.float32)
pre_E[1] = 100.0 # event 0's entry step
pre_E[3] = 100.0 # event 1's entry step
pre_E[7] = 100.0 # event 2's entry step
def _local_block():
block = rng.standard_normal((n, 9)).astype(np.float32)
block[:, :3] = log_transform(rng.uniform(0.1, 5.0, (n, 3)).astype(np.float32))
block[:, 3:6] = _unit_vectors(rng, n)
block[:, 6:9] = _unit_vectors(rng, n)
return block
true_log_local = _local_block()
pred_log_local = _local_block()
return event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
def _write_event_level_parquet(path, rng=None):
rng = rng or np.random.default_rng(7)
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local = (
_make_event_level_arrays(rng)
)
n = len(event_id)
table = pa.table(
{
"event_id": event_id,
"pdg": rng.choice([11, -11, 22], n),
"pre_x": pre_pos[:, 0],
"pre_y": pre_pos[:, 1],
"pre_z": pre_pos[:, 2],
"pre_E": pre_E,
"pre_dx": pre_dir[:, 0],
"pre_dy": pre_dir[:, 1],
"pre_dz": pre_dir[:, 2],
"material": rng.choice(["W", "Pb"], n),
"layer_id": rng.integers(0, 10, n).astype(np.int32),
"n_sec": rng.integers(0, 3, n).astype(np.int32),
**{
f"pred_{name}": pred_log_local[:, j]
for j, name in enumerate(LOCAL_TARGET_NAMES)
},
**{
f"true_{name}": true_log_local[:, j]
for j, name in enumerate(LOCAL_TARGET_NAMES)
},
}
)
table = table.replace_schema_metadata(
{
PREDICT_COORD_METADATA_KEY: "local",
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
}
)
pq.write_table(table, path)
return event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
def _expected_event_table(
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
):
"""Independent re-derivation of total/centroid/RMS per event, for comparison."""
expected = {}
for e in sorted(np.unique(event_id).tolist()):
mask = event_id == e
entry_idx = np.where(mask)[0][np.argmax(pre_E[mask])]
entry_pos = pre_pos[entry_idx]
axis_dir = pre_dir[entry_idx]
def agg(log_local):
step_length = inv_log_transform(log_local[mask, 0])
edep = inv_log_transform(log_local[mask, 2])
travel_dir_local = log_local[mask, 6:9]
post_pos = reconstruct_post_pos(
pre_pos[mask], pre_dir[mask], step_length, travel_dir_local
)
disp = post_pos - entry_pos
depth = disp @ axis_dir
transverse = np.linalg.norm(disp - depth[:, None] * axis_dir, axis=1)
total = float(edep.sum())
centroid = float((edep * depth).sum() / total)
rms = float(np.sqrt((edep * transverse**2).sum() / total))
return total, centroid, rms
real_total, real_centroid, real_rms = agg(true_log_local)
gen_total, gen_centroid, gen_rms = agg(pred_log_local)
expected[e] = (
real_total,
gen_total,
real_centroid,
gen_centroid,
real_rms,
gen_rms,
)
return expected
def test_compute_event_observables_pl_matches_manual_reconstruction(tmp_path):
path = tmp_path / "event_level.parquet"
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local = (
_write_event_level_parquet(path)
)
expected = _expected_event_table(
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
)
obs = compute_event_observables_pl(path, depth_bins=5, transverse_bins=5)
table = obs.event_table.sort("event_id")
for i, eid in enumerate(table["event_id"].to_list()):
real_total, gen_total, real_centroid, gen_centroid, real_rms, gen_rms = (
expected[eid]
)
np.testing.assert_allclose(table["real_total_edep"][i], real_total, rtol=1e-4)
np.testing.assert_allclose(table["gen_total_edep"][i], gen_total, rtol=1e-4)
np.testing.assert_allclose(
table["real_centroid_depth"][i], real_centroid, rtol=1e-3, atol=1e-4
)
np.testing.assert_allclose(
table["gen_centroid_depth"][i], gen_centroid, rtol=1e-3, atol=1e-4
)
np.testing.assert_allclose(
table["real_transverse_rms"][i], real_rms, rtol=1e-3, atol=1e-4
)
np.testing.assert_allclose(
table["gen_transverse_rms"][i], gen_rms, rtol=1e-3, atol=1e-4
)
def test_compute_event_observables_pl_profile_shapes(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
obs = compute_event_observables_pl(path, depth_bins=7, transverse_bins=4)
assert obs.depth_edges.shape == (8,)
assert obs.transverse_edges.shape == (5,)
assert obs.real_depth_profile.shape == (7,)
assert obs.gen_depth_profile.shape == (7,)
assert obs.real_transverse_profile.shape == (4,)
assert obs.gen_transverse_profile.shape == (4,)
assert len(obs.event_table) == 3
def test_compute_event_observables_pl_accepts_lazyframe(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
obs_from_path = compute_event_observables_pl(path, depth_bins=5, transverse_bins=5)
obs_from_lf = compute_event_observables_pl(
pl.scan_parquet(path), depth_bins=5, transverse_bins=5
)
np.testing.assert_allclose(
obs_from_lf.event_table.sort("event_id")["real_total_edep"].to_numpy(),
obs_from_path.event_table.sort("event_id")["real_total_edep"].to_numpy(),
)
def test_event_level_plots_run_without_error(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
obs = compute_event_observables_pl(path, depth_bins=5, transverse_bins=5)
assert plot_total_energy(obs) is not None
assert plot_longitudinal_profile(obs) is not None
assert plot_transverse_profile(obs) is not None
assert plot_shower_max_depth(obs) is not None