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:
+191
-1
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user