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:
+377
-2
@@ -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
@@ -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