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
parent 6fc68fe1aa
commit f81263fa6e
2 changed files with 568 additions and 3 deletions
+377 -2
View File
@@ -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
+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