From f81263fa6e69f7491017a52d9a0ceaa1369cf291 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Mon, 22 Jun 2026 15:52:47 +0200 Subject: [PATCH] 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 --- giant/analysis.py | 379 ++++++++++++++++++++++++++++++++++++++++- tests/test_analysis.py | 192 ++++++++++++++++++++- 2 files changed, 568 insertions(+), 3 deletions(-) diff --git a/giant/analysis.py b/giant/analysis.py index 9e62903..39393ec 100644 --- a/giant/analysis.py +++ b/giant/analysis.py @@ -27,7 +27,22 @@ skip the checkpoint/model entirely and load the parquet directly:: (`--coord global` output isn't supported here — it has no ground-truth columns to compare against.) -Three tiers of checks, building on the aggregate marginal/KL check in +For event-level (shower) observables on the same `--coord local` predict +output, see `compute_event_observables_pl` — it streams the file directly +(no row subsampling, no `SampleCollection`), since per-event sums would be +corrupted by partial events:: + + from giant.analysis import compute_event_observables_pl + from giant.analysis import plot_total_energy, plot_longitudinal_profile + from giant.analysis import plot_transverse_profile, plot_shower_max_depth + + obs = compute_event_observables_pl("path/to/steps_predicted_local.parquet") + plot_total_energy(obs) + plot_longitudinal_profile(obs) + plot_transverse_profile(obs) + plot_shower_max_depth(obs) + +Four tiers of checks, building on the aggregate marginal/KL check in `giant.validate.validate_marginals`: 1. stratified marginals — per-dimension real-vs-generated comparison, sliced by @@ -39,6 +54,14 @@ Three tiers of checks, building on the aggregate marginal/KL check in step_length/delta_e/edep, checked in denormalized physical units; nothing in the unconstrained MLP output enforces these, so violations are a pure generation artifact. +4. event-level observables — total deposited energy, longitudinal/transverse + shower profiles, and shower-max depth, aggregated per `event_id` in the + world frame with physical units (mm, MeV). This re-aggregates one-step-ahead + generations (each row generated conditioned on the *real* preceding state) + grouped by event — not a full autoregressive shower rollout — so it won't + surface covariate-shift failures that only appear under true rollout, only + how well one-step generation reconstructs aggregate shower structure when + fed real conditioning throughout. `collect_samples` takes a `steps` argument (forwarded to the flow ODE integrator or, in ddim mode, the DDIM substep count) so a later @@ -69,7 +92,12 @@ from giant.constants import ( ) from giant.data.dataset import train_val_split from giant.data.loader import find_parquet_files, load_steps -from giant.data.transforms import Normalizer, build_features, inv_log_transform +from giant.data.transforms import ( + Normalizer, + build_features, + inv_log_transform, + reconstruct_post_pos, +) from giant.model.network import DenoisingMLP from giant.model.schedule import CosineSchedule from giant.sample import sample_ddim, sample_ddpm, sample_flow @@ -925,3 +953,350 @@ def plot_constraint_violations(collection: SampleCollection): ax.set_title(f"generated {name}") fig.tight_layout() return fig + + +# --------------------------------------------------------------------------- +# Tier 4: event-level (shower) observables +# +# Built directly on `giant predict --coord local` parquet output, streamed in +# two passes rather than materialized into a SampleCollection: per-event sums +# (total deposited energy, etc.) would be silently corrupted by the row +# subsampling `load_predicted_local(sample_frac=...)` uses to keep large files +# in memory, and the real files here run into the tens of millions of rows. +# --------------------------------------------------------------------------- + +_PRE_COLS = ["pre_x", "pre_y", "pre_z", "pre_dx", "pre_dy", "pre_dz", "pre_E"] + + +@dataclass +class EventObservables: + event_table: pl.DataFrame # one row per event_id; real_*/gen_* columns, mm/MeV + depth_edges: np.ndarray # (depth_bins+1,) mm, along shower axis + transverse_edges: np.ndarray # (transverse_bins+1,) mm, perpendicular to axis + real_depth_profile: np.ndarray # (depth_bins,) mean edep/event/bin, MeV + gen_depth_profile: np.ndarray + real_depth_profile_std: np.ndarray # event-to-event RMS per bin + gen_depth_profile_std: np.ndarray + real_transverse_profile: np.ndarray + gen_transverse_profile: np.ndarray + real_transverse_profile_std: np.ndarray + gen_transverse_profile_std: np.ndarray + + +def _entry_axis_and_bin_edges( + lf: pl.LazyFrame, depth_bins: int, transverse_bins: int +) -> tuple[pl.DataFrame, np.ndarray, np.ndarray]: + """Per-event shower axis (highest-pre_E row) plus depth/transverse bin edges. + + Bin edges are sized from `pre_pos` alone (no post_pos reconstruction + needed) — pre-step positions already trace the shower's extent closely + enough to pick a sensible range, which avoids a second full streaming pass + just to size the bins. + """ + narrow = lf.select(["event_id", *_PRE_COLS]) + entry = narrow.group_by("event_id").agg( + pl.col("pre_x").get(pl.col("pre_E").arg_max()).alias("entry_x"), + pl.col("pre_y").get(pl.col("pre_E").arg_max()).alias("entry_y"), + pl.col("pre_z").get(pl.col("pre_E").arg_max()).alias("entry_z"), + pl.col("pre_dx").get(pl.col("pre_E").arg_max()).alias("axis_x"), + pl.col("pre_dy").get(pl.col("pre_E").arg_max()).alias("axis_y"), + pl.col("pre_dz").get(pl.col("pre_E").arg_max()).alias("axis_z"), + ) + + joined = narrow.join(entry, on="event_id") + dx = pl.col("pre_x") - pl.col("entry_x") + dy = pl.col("pre_y") - pl.col("entry_y") + dz = pl.col("pre_z") - pl.col("entry_z") + depth = dx * pl.col("axis_x") + dy * pl.col("axis_y") + dz * pl.col("axis_z") + tx = dx - depth * pl.col("axis_x") + ty = dy - depth * pl.col("axis_y") + tz = dz - depth * pl.col("axis_z") + transverse = (tx**2 + ty**2 + tz**2).sqrt() + + stats = ( + joined.select(depth.alias("depth_proxy"), transverse.alias("transverse_proxy")) + .select( + pl.col("depth_proxy").quantile(0.001).alias("depth_lo"), + pl.col("depth_proxy").quantile(0.999).alias("depth_hi"), + pl.col("transverse_proxy").quantile(0.999).alias("transverse_hi"), + ) + .collect() + .row(0, named=True) + ) + + entry_df = entry.collect().sort("event_id") + + depth_lo, depth_hi = stats["depth_lo"], stats["depth_hi"] + if not (depth_hi - depth_lo > 1e-6 * max(abs(depth_hi), 1.0)): + depth_lo, depth_hi = depth_lo - 0.5, depth_hi + 0.5 + depth_edges = np.linspace(depth_lo, depth_hi, depth_bins + 1) + transverse_hi = max(stats["transverse_hi"], 1e-6) + transverse_edges = np.linspace(0.0, transverse_hi, transverse_bins + 1) + return entry_df, depth_edges, transverse_edges + + +def _iter_predicted_local_batches( + source: str | Path | pl.LazyFrame, columns: list[str], batch_size: int +): + """Yield the needed columns in bounded-memory chunks. + + A path is streamed via pyarrow's row-batch reader so the file is never + fully materialized; a `pl.LazyFrame` (the in-memory test-fixture case) is + just collected once, since that data is small by construction. + """ + if isinstance(source, pl.LazyFrame): + yield source.select(columns).collect() + return + pf = pq.ParquetFile(Path(source)) + for batch in pf.iter_batches(batch_size=batch_size, columns=columns): + yield pl.from_arrow(batch) + + +def compute_event_observables_pl( + source: str | Path | pl.LazyFrame, + depth_bins: int = 20, + transverse_bins: int = 20, + batch_size: int = 1_000_000, +) -> EventObservables: + """Stream a `giant predict --coord local` parquet file into event-level observables. + + For each event, the highest-`pre_E` row is taken as the primary's entry + step (secondaries always carry less energy than their parent), fixing a + shower axis/entry point shared by both real and generated rows (both are + conditioned on the same real pre-step state). Every row's `post_pos`/`edep` + is reconstructed into the world frame in physical units (mm, MeV) via + `giant.data.transforms.reconstruct_post_pos`/`inv_log_transform` — the same + functions `giant predict --coord global` uses — and projected onto + depth-along-axis / transverse-distance-from-axis. + + Runs in two passes: a cheap pure-polars pass over `pre_*` columns only + (shower axis + bin-edge sizing), then one streaming pass over the full + file accumulating per-event and per-bin sums in numpy. Never materializes + the file as a `SampleCollection` — see the module-level note above. + """ + lf = _scan_predicted_local(source) + entry_df, depth_edges, transverse_edges = _entry_axis_and_bin_edges( + lf, depth_bins, transverse_bins + ) + + event_ids = entry_df["event_id"].to_numpy() + n_events = len(event_ids) + entry_pos = ( + entry_df.select(["entry_x", "entry_y", "entry_z"]).to_numpy().astype(np.float32) + ) + axis_dir = ( + entry_df.select(["axis_x", "axis_y", "axis_z"]).to_numpy().astype(np.float32) + ) + + n_steps = np.zeros(n_events, dtype=np.int64) + real_total_edep = np.zeros(n_events, dtype=np.float64) + gen_total_edep = np.zeros(n_events, dtype=np.float64) + real_sum_edep_depth = np.zeros(n_events, dtype=np.float64) + gen_sum_edep_depth = np.zeros(n_events, dtype=np.float64) + real_sum_edep_transverse2 = np.zeros(n_events, dtype=np.float64) + gen_sum_edep_transverse2 = np.zeros(n_events, dtype=np.float64) + real_depth_bin_edep = np.zeros((n_events, depth_bins), dtype=np.float64) + gen_depth_bin_edep = np.zeros((n_events, depth_bins), dtype=np.float64) + real_transverse_bin_edep = np.zeros((n_events, transverse_bins), dtype=np.float64) + gen_transverse_bin_edep = np.zeros((n_events, transverse_bins), dtype=np.float64) + + pred_cols = [f"pred_{name}" for name in LOCAL_TARGET_NAMES] + true_cols = [f"true_{name}" for name in LOCAL_TARGET_NAMES] + needed_cols = [ + "event_id", + "pre_x", + "pre_y", + "pre_z", + "pre_dx", + "pre_dy", + "pre_dz", + *pred_cols, + *true_cols, + ] + + for batch_df in _iter_predicted_local_batches(source, needed_cols, batch_size): + eid = batch_df["event_id"].to_numpy() + idx = np.searchsorted(event_ids, eid) + + pre_pos = ( + batch_df.select(["pre_x", "pre_y", "pre_z"]).to_numpy().astype(np.float32) + ) + pre_dir = ( + batch_df.select(["pre_dx", "pre_dy", "pre_dz"]) + .to_numpy() + .astype(np.float32) + ) + + def _reconstruct(cols: list[str]) -> tuple[np.ndarray, np.ndarray]: + raw = batch_df.select(cols).to_numpy().astype(np.float32) + step_length = inv_log_transform(raw[:, 0]) + edep = inv_log_transform(raw[:, 2]) + travel_dir_local = raw[:, 6:9] + post_pos = reconstruct_post_pos( + pre_pos, pre_dir, step_length, travel_dir_local + ) + return post_pos, edep + + real_post_pos, real_edep = _reconstruct(true_cols) + gen_post_pos, gen_edep = _reconstruct(pred_cols) + + e_pos = entry_pos[idx] + a_dir = axis_dir[idx] + + def _depth_transverse(post_pos: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + disp = post_pos - e_pos + depth = np.sum(disp * a_dir, axis=1) + perp = disp - depth[:, None] * a_dir + transverse = np.linalg.norm(perp, axis=1) + return depth, transverse + + real_depth, real_transverse = _depth_transverse(real_post_pos) + gen_depth, gen_transverse = _depth_transverse(gen_post_pos) + + real_depth_bin = np.digitize(real_depth, depth_edges[1:-1]) + gen_depth_bin = np.digitize(gen_depth, depth_edges[1:-1]) + real_transverse_bin = np.digitize(real_transverse, transverse_edges[1:-1]) + gen_transverse_bin = np.digitize(gen_transverse, transverse_edges[1:-1]) + + np.add.at(n_steps, idx, 1) + np.add.at(real_total_edep, idx, real_edep) + np.add.at(gen_total_edep, idx, gen_edep) + np.add.at(real_sum_edep_depth, idx, real_edep * real_depth) + np.add.at(gen_sum_edep_depth, idx, gen_edep * gen_depth) + np.add.at(real_sum_edep_transverse2, idx, real_edep * real_transverse**2) + np.add.at(gen_sum_edep_transverse2, idx, gen_edep * gen_transverse**2) + np.add.at(real_depth_bin_edep, (idx, real_depth_bin), real_edep) + np.add.at(gen_depth_bin_edep, (idx, gen_depth_bin), gen_edep) + np.add.at(real_transverse_bin_edep, (idx, real_transverse_bin), real_edep) + np.add.at(gen_transverse_bin_edep, (idx, gen_transverse_bin), gen_edep) + + safe_real_total = np.where(real_total_edep > 0, real_total_edep, 1.0) + safe_gen_total = np.where(gen_total_edep > 0, gen_total_edep, 1.0) + real_centroid_depth = real_sum_edep_depth / safe_real_total + gen_centroid_depth = gen_sum_edep_depth / safe_gen_total + real_transverse_rms = np.sqrt(real_sum_edep_transverse2 / safe_real_total) + gen_transverse_rms = np.sqrt(gen_sum_edep_transverse2 / safe_gen_total) + + depth_centers = 0.5 * (depth_edges[:-1] + depth_edges[1:]) + real_max_depth = depth_centers[np.argmax(real_depth_bin_edep, axis=1)] + gen_max_depth = depth_centers[np.argmax(gen_depth_bin_edep, axis=1)] + + event_table = pl.DataFrame( + { + "event_id": event_ids, + "n_steps": n_steps, + "real_total_edep": real_total_edep, + "gen_total_edep": gen_total_edep, + "real_centroid_depth": real_centroid_depth, + "gen_centroid_depth": gen_centroid_depth, + "real_transverse_rms": real_transverse_rms, + "gen_transverse_rms": gen_transverse_rms, + "real_max_depth": real_max_depth, + "gen_max_depth": gen_max_depth, + } + ) + + return EventObservables( + event_table=event_table, + depth_edges=depth_edges, + transverse_edges=transverse_edges, + real_depth_profile=real_depth_bin_edep.mean(axis=0), + gen_depth_profile=gen_depth_bin_edep.mean(axis=0), + real_depth_profile_std=real_depth_bin_edep.std(axis=0), + gen_depth_profile_std=gen_depth_bin_edep.std(axis=0), + real_transverse_profile=real_transverse_bin_edep.mean(axis=0), + gen_transverse_profile=gen_transverse_bin_edep.mean(axis=0), + real_transverse_profile_std=real_transverse_bin_edep.std(axis=0), + gen_transverse_profile_std=gen_transverse_bin_edep.std(axis=0), + ) + + +def plot_total_energy(observables: EventObservables, bins: int = 50): + """Real-vs-generated histogram of total deposited energy per event, with resolution.""" + table = observables.event_table + real = table["real_total_edep"].to_numpy() + gen = table["gen_total_edep"].to_numpy() + + fig, ax = plt.subplots(figsize=(6, 4)) + edges = _hist_edges(real, gen, bins=bins) + ax.hist( + real, + bins=edges, + density=True, + histtype="step", + label=f"real (σ/μ={real.std() / real.mean():.3f})", + ) + ax.hist( + gen, + bins=edges, + density=True, + histtype="step", + label=f"generated (σ/μ={gen.std() / gen.mean():.3f})", + ) + ax.set_yscale("log") + ax.set_xlabel("total deposited energy per event [MeV]") + ax.legend(fontsize=8) + fig.tight_layout() + return fig + + +def _plot_profile( + centers: np.ndarray, + real_mean: np.ndarray, + gen_mean: np.ndarray, + real_std: np.ndarray, + gen_std: np.ndarray, + xlabel: str, +): + fig, ax = plt.subplots(figsize=(6, 4)) + ax.errorbar(centers, real_mean, yerr=real_std, fmt="o-", label="real", capsize=2) + ax.errorbar(centers, gen_mean, yerr=gen_std, fmt="s-", label="generated", capsize=2) + ax.set_xlabel(xlabel) + ax.set_ylabel("mean edep per event per bin [MeV]") + ax.legend(fontsize=8) + fig.tight_layout() + return fig + + +def plot_longitudinal_profile(observables: EventObservables): + """E_dep(depth) mean ± event-to-event RMS, real vs generated.""" + centers = 0.5 * (observables.depth_edges[:-1] + observables.depth_edges[1:]) + return _plot_profile( + centers, + observables.real_depth_profile, + observables.gen_depth_profile, + observables.real_depth_profile_std, + observables.gen_depth_profile_std, + "depth along shower axis [mm]", + ) + + +def plot_transverse_profile(observables: EventObservables): + """E_dep(transverse distance) mean ± event-to-event RMS, real vs generated (Molière-style).""" + centers = 0.5 * ( + observables.transverse_edges[:-1] + observables.transverse_edges[1:] + ) + return _plot_profile( + centers, + observables.real_transverse_profile, + observables.gen_transverse_profile, + observables.real_transverse_profile_std, + observables.gen_transverse_profile_std, + "transverse distance from shower axis [mm]", + ) + + +def plot_shower_max_depth(observables: EventObservables, bins: int = 30): + """Real-vs-generated histogram of per-event shower-maximum depth.""" + table = observables.event_table + real = table["real_max_depth"].to_numpy() + gen = table["gen_max_depth"].to_numpy() + + fig, ax = plt.subplots(figsize=(6, 4)) + edges = _hist_edges(real, gen, bins=bins) + ax.hist(real, bins=edges, density=True, histtype="step", label="real") + ax.hist(gen, bins=edges, density=True, histtype="step", label="generated") + ax.set_xlabel("depth of shower maximum [mm]") + ax.legend(fontsize=8) + fig.tight_layout() + return fig diff --git a/tests/test_analysis.py b/tests/test_analysis.py index ad9f611..5525699 100644 --- a/tests/test_analysis.py +++ b/tests/test_analysis.py @@ -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