Fix rollout edep mismatch and add truth overlay to Tier 4 observables

load_rollout_vs_truth was including rollout.py's synthetic termination-
bookkeeping rows (escaped/unknown_pdg/energy_cutoff/max_steps) unfiltered:
these carry step_length=0 and edep=pre_E dumped in one row for shower-level
energy conservation, not a real per-step value, and nearly doubled the
apparent mean edep in a repro. _load_world_frame_side now drops them, keeping
only real generated steps (continuing or natural_end). Also documents the
remaining, unfixable difference: rollout's edep on real steps absorbs any
secondary-energy budget Stage 2 didn't allocate, which truth's edep never does.

Adds compute_truth_observables, the truth-schema counterpart to
compute_rollout_observables, so the Tier 4 event-level plots
(plot_rollout_longitudinal/transverse/total_energy) can overlay a real
reference computed directly from load_rollout_vs_truth's own truth file,
without needing a separate paired giant predict --coord local file. Shares
the depth/transverse binning core with compute_rollout_observables via a new
_event_axis_depth_transverse helper.

Updates rollout_validation.ipynb's Tier 4 section to use this reference and
points ROLLOUT_FILE/TRUTH_FILE at a real prediction/shard pair.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
2026-07-13 12:38:44 +02:00
parent 36bdbf7cc1
commit 4c1250e246
3 changed files with 534 additions and 35 deletions
+6 -8
View File
@@ -43,10 +43,10 @@
"from giant.analysis import load_rollout_vs_truth, plot_kl_bars\n",
"\n",
"# `giant rollout` output for the shower(s) under test.\n",
"ROLLOUT_FILE = \"/home/lars/Programming/giant/rollout.parquet\"\n",
"ROLLOUT_FILE = \"/ceph/lbogner/geant_steps/predictions/e7ee73b7-70c9-4390-83e6-a31e50a7d3cb.parquet\"\n",
"# Any held-out file sharing giant train's input schema (real miniCaloSim\n",
"# steps) \u2014 e.g. the val split the rollout's seed events were drawn from.\n",
"TRUTH_FILE = \"/home/lars/Programming/giant/val.parquet\"\n",
"TRUTH_FILE = \"/ceph/lbogner/geant_steps/processed/steps/gen3/schema2/pbwo4/shard-009.parquet\"\n",
"\n",
"# sample_frac subsamples each file independently (kept memory-bounded for\n",
"# large files); both default to every row when omitted.\n",
@@ -188,9 +188,9 @@
"source": [
"## Tier 4: event-level (shower) observables\n",
"\n",
"Built on `compute_rollout_observables`, not `compute_event_observables_pl` \u2014 the rollout file carries its own `track_id`/`termination_reason` columns that the event-level aggregation needs, and the shower here already *is* a full autoregressive rollout rather than one-step generations re-aggregated by event.\n",
"Built on `compute_rollout_observables` for the rollout side, not `compute_event_observables_pl` \u2014 the rollout file carries its own `track_id`/`termination_reason` columns that the event-level aggregation needs, and the shower here already *is* a full autoregressive rollout rather than one-step generations re-aggregated by event.\n",
"\n",
"To overlay a real reference profile, pass a *paired* `giant predict --coord local` file for the same held-out events as `reference_path` below (see `analysis/export_rollout_observables.py`) \u2014 `TRUTH_FILE` above can't serve as that reference directly, since it's the raw training-input schema, not predict output. Leave `reference_path = None` to skip the overlay."
"The real-shower overlay comes from `compute_truth_observables(TRUTH_FILE)` \u2014 the truth-side counterpart, computed directly from the same raw truth-schema file used for Tiers 1-3 above (no separate paired `giant predict --coord local` file needed, unlike `analysis/export_rollout_observables.py`'s `reference_path`)."
]
},
{
@@ -200,14 +200,12 @@
"metadata": {},
"outputs": [],
"source": [
"from giant.analysis import compute_rollout_observables, compute_event_observables_pl\n",
"from giant.analysis import compute_rollout_observables, compute_truth_observables\n",
"from giant.analysis import plot_rollout_longitudinal, plot_rollout_transverse\n",
"from giant.analysis import plot_rollout_total_energy\n",
"\n",
"obs = compute_rollout_observables(ROLLOUT_FILE)\n",
"\n",
"reference_path = None # optional: a `giant predict --coord local` file, see above\n",
"reference = compute_event_observables_pl(reference_path) if reference_path else None"
"reference = compute_truth_observables(TRUTH_FILE)"
]
},
{
+179 -26
View File
@@ -107,7 +107,10 @@ from giant.constants import (
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
ROLLOUT_COORD_VALUE,
TERM_ENERGY_CUTOFF,
TERM_ESCAPED,
TERM_MAX_STEPS,
TERM_UNKNOWN_PDG,
)
from giant.data.transforms import (
energy_simplex_decode,
@@ -1817,23 +1820,21 @@ class RolloutObservables:
transverse_profile_std: np.ndarray
def compute_rollout_observables(
path: str | Path,
depth_bins: int = 20,
transverse_bins: int = 20,
) -> RolloutObservables:
"""Compute shower observables from a `giant rollout` steps parquet.
def _event_axis_depth_transverse(
df: pd.DataFrame, depth_bins: int, transverse_bins: int
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Shared core of `compute_rollout_observables`/`compute_truth_observables`.
Per event the primary entry step (highest-`pre_E` row) fixes the shower axis
and entry point; every step's deposit (`edep` at `post_pos`) is projected onto
depth-along-axis and transverse-distance-from-axis, then binned. Returns
per-event totals plus dataset-mean longitudinal/transverse profiles.
`df` needs `event_id`/`pre_x,y,z`/`pre_dx,dy,dz`/`pre_E`/`post_x,y,z`/`edep`
— both a `giant rollout` file and a raw truth-schema steps file carry these.
Per event, the highest-`pre_E` row fixes the shower axis/entry point; every
row's `edep` at `post_pos` is projected onto depth-along-axis and
transverse-distance-from-axis, then binned.
Returns `(ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev)`
— `row_ev` maps each row of `df` to its event's index in `ev_ids`;
`depth_ev`/`trans_ev` are `(n_events, bins)` per-event-per-bin edep sums.
"""
_check_rollout_metadata(Path(path))
df = pd.read_parquet(path, columns=_ROLLOUT_COLS)
# Per-event entry point + shower axis from the highest-pre_E row.
entry_idx = df.groupby("event_id")["pre_E"].idxmax()
entry = df.loc[entry_idx].set_index("event_id")
ax = entry[["pre_dx", "pre_dy", "pre_dz"]].to_numpy(dtype=np.float64).copy()
@@ -1869,6 +1870,40 @@ def compute_rollout_observables(
np.add.at(depth_ev, (row_ev, d_bin), edep)
np.add.at(trans_ev, (row_ev, t_bin), edep)
return ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev
def _centroid_depth(depth_ev: np.ndarray, depth_edges: np.ndarray) -> np.ndarray:
"""Energy-weighted centroid depth per event, from the binned edep sums."""
bin_centers = 0.5 * (depth_edges[:-1] + depth_edges[1:])
tot = depth_ev.sum(axis=1)
centroid = (depth_ev * bin_centers).sum(axis=1) / np.where(tot > 0, tot, 1.0)
return np.where(tot > 0, centroid, 0.0)
def compute_rollout_observables(
path: str | Path,
depth_bins: int = 20,
transverse_bins: int = 20,
) -> RolloutObservables:
"""Compute shower observables from a `giant rollout` steps parquet.
Per event the primary entry step (highest-`pre_E` row) fixes the shower axis
and entry point; every step's deposit (`edep` at `post_pos`) is projected onto
depth-along-axis and transverse-distance-from-axis, then binned. Returns
per-event totals plus dataset-mean longitudinal/transverse profiles.
See `compute_truth_observables` for the truth-schema counterpart, shaped to
plug straight into this function's own output as a `plot_rollout_*`
`reference=` overlay.
"""
_check_rollout_metadata(Path(path))
df = pd.read_parquet(path, columns=_ROLLOUT_COLS)
ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev = (
_event_axis_depth_transverse(df, depth_bins, transverse_bins)
)
# Per-event scalar table.
is_leak = df["termination_reason"].to_numpy() == TERM_ESCAPED
per_ev = df.groupby("event_id").agg(
@@ -1881,12 +1916,8 @@ def compute_rollout_observables(
.groupby("event_id")["_leak"]
.sum()
)
# Energy-weighted centroid depth per event (from the binned sums).
bin_centers = 0.5 * (depth_edges[:-1] + depth_edges[1:])
tot = depth_ev.sum(axis=1)
centroid = (depth_ev * bin_centers).sum(axis=1) / np.where(tot > 0, tot, 1.0)
per_ev = per_ev.reindex(ev_ids)
per_ev["centroid_depth"] = np.where(tot > 0, centroid, 0.0)
per_ev["centroid_depth"] = _centroid_depth(depth_ev, depth_edges)
per_ev.index.name = "event_id"
return RolloutObservables(
@@ -1900,11 +1931,95 @@ def compute_rollout_observables(
)
_TRUTH_EVENT_COLS = [
"event_id",
"pre_x",
"pre_y",
"pre_z",
"pre_dx",
"pre_dy",
"pre_dz",
"pre_E",
"post_x",
"post_y",
"post_z",
"edep",
"step_length",
]
@dataclass
class TruthObservables:
"""Truth-schema event-level observables, shaped for `plot_rollout_*`'s `reference=`.
The truth-side counterpart to `RolloutObservables`, computed directly from
a raw truth-schema steps file (the same file `load_rollout_vs_truth` takes
as `truth_path`) rather than needing a separate paired `giant predict
--coord local` file covering the same events. Field names mirror
`EventObservables`'s `real_*` convention — `plot_rollout_longitudinal`/
`plot_rollout_transverse`/`plot_rollout_total_energy` already read exactly
these names off their `reference` argument — since this object is only
ever used as a reference overlay, never plotted standalone.
"""
event_table: pd.DataFrame # one row per event_id (mm/MeV)
depth_edges: np.ndarray
transverse_edges: np.ndarray
real_depth_profile: np.ndarray
real_depth_profile_std: np.ndarray
real_transverse_profile: np.ndarray
real_transverse_profile_std: np.ndarray
def compute_truth_observables(
path: str | Path,
depth_bins: int = 20,
transverse_bins: int = 20,
) -> TruthObservables:
"""Compute shower observables from a raw truth-schema steps parquet.
Lets the same held-out truth file used by `load_rollout_vs_truth` (Tier
1-3) also supply the Tier 4 real-shower overlay — pass the result as
`reference=` to `plot_rollout_longitudinal`/`plot_rollout_transverse`/
`plot_rollout_total_energy` — without needing a separate paired `giant
predict --coord local` file for the same events. Same per-event
axis/depth/transverse construction as `compute_rollout_observables` (see
`_event_axis_depth_transverse`), just without the `track_id`/
`termination_reason` columns a rollout file (but not a truth file) has, so
there's no `n_tracks`/`leaked_E` in `event_table`.
"""
df = pd.read_parquet(path, columns=_TRUTH_EVENT_COLS)
ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev = (
_event_axis_depth_transverse(df, depth_bins, transverse_bins)
)
per_ev = df.groupby("event_id").agg(
n_steps=("edep", "size"),
real_total_edep=("edep", "sum"),
real_total_length=("step_length", "sum"),
)
per_ev = per_ev.reindex(ev_ids)
per_ev["real_centroid_depth"] = _centroid_depth(depth_ev, depth_edges)
per_ev.index.name = "event_id"
return TruthObservables(
event_table=per_ev.reset_index(),
depth_edges=depth_edges,
transverse_edges=transverse_edges,
real_depth_profile=depth_ev.mean(0),
real_depth_profile_std=depth_ev.std(0),
real_transverse_profile=trans_ev.mean(0),
real_transverse_profile_std=trans_ev.std(0),
)
def plot_rollout_longitudinal(obs: RolloutObservables, reference=None):
"""Mean edep vs depth along the shower axis; optional real-reference overlay.
`reference` may be an `EventObservables` (its `real_depth_profile`) to overlay
the real showers seeded from the same events.
`reference` may be an `EventObservables` or a `TruthObservables` (either
exposes `real_depth_profile`) to overlay the real showers seeded from the
same events — `compute_truth_observables` builds the latter directly from
a raw truth-schema file, with no paired predict-schema file needed.
"""
centers = 0.5 * (obs.depth_edges[:-1] + obs.depth_edges[1:])
fig, ax = plt.subplots(figsize=(7, 4))
@@ -2087,6 +2202,21 @@ def _existing_columns(
return [c for c in wanted if c in available]
# `rollout.py`'s `_terminal_rows` writes one synthetic bookkeeping row per track
# for these termination reasons (escaped/unknown_pdg/energy_cutoff/max_steps):
# step_length=0, post_pos=pre_pos, and — for every reason but escaped — the
# track's *entire remaining pre_E* dumped into `edep` in one row, so the shower's
# total energy still conserves. These aren't steps in any physical sense (truth
# data has no equivalent), so mixing them into a per-step real-vs-generated
# comparison would inject a spurious step_length=0 spike and roughly double the
# apparent mean edep purely from bookkeeping, not model behavior. Real generated
# steps carry "" (continuing) or `TERM_NATURAL_END` (the track's last real step,
# which does have genuine step_length/edep) and are kept.
_SYNTHETIC_ROLLOUT_TERMINATION_REASONS = frozenset(
{TERM_ESCAPED, TERM_UNKNOWN_PDG, TERM_ENERGY_CUTOFF, TERM_MAX_STEPS}
)
def _load_world_frame_side(
source: str | Path | pl.LazyFrame,
sample_frac: float,
@@ -2096,9 +2226,12 @@ def _load_world_frame_side(
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Stream one world-frame steps file/LazyFrame into (raw9, cond9, pdg, material).
`extra_cols` (e.g. `["n_sec_pred"]` or `["child_track_ids"]`) are included
only when present, so the same helper serves both the rollout and truth
schemas without either needing the other's columns.
`extra_cols` (e.g. `["n_sec_pred", "termination_reason"]` or
`["child_track_ids"]`) are included only when present, so the same helper
serves both the rollout and truth schemas without either needing the
other's columns. When `termination_reason` is present (the rollout side),
rows carrying one of `_SYNTHETIC_ROLLOUT_TERMINATION_REASONS` are dropped
before decoding — see that constant's docstring for why.
"""
columns = _WORLD_FRAME_STEP_COLS + _existing_columns(source, extra_cols)
threshold = int(sample_frac * 2**32) if sample_frac < 1.0 else None
@@ -2110,6 +2243,12 @@ def _load_world_frame_side(
row_idx = pl.arange(offset, offset + n, eager=True).cast(pl.UInt32)
batch = batch.filter((row_idx.hash(seed=seed) % 2**32) < threshold)
offset += n
if "termination_reason" in batch.columns:
batch = batch.filter(
~pl.col("termination_reason").is_in(
list(_SYNTHETIC_ROLLOUT_TERMINATION_REASONS)
)
)
if batch.height == 0:
continue
raw_parts.append(_world_frame_raw_targets(batch))
@@ -2151,6 +2290,16 @@ def load_rollout_vs_truth(
`edep`/`step_length` columns both files share — see the module note above
`_WORLD_FRAME_STEP_COLS`. `sample_frac`/`seed`/`batch_size` behave as in
`load_predicted_local`, applied independently to each file.
The rollout side drops synthetic termination-bookkeeping rows (see
`_SYNTHETIC_ROLLOUT_TERMINATION_REASONS`) before decoding, since those
aren't real generated steps. Even among the real steps that remain,
`edep` isn't perfectly analogous between the two files: `rollout.py`
tops up a step's `edep` with any secondary-energy budget Stage 2 didn't
allocate to an actual spawned secondary (so every step still conserves
energy exactly), which truth's Geant4-recorded `edep` never does. A
generated `edep` that runs a bit high relative to truth can be this
bookkeeping, not necessarily a Stage-1/Stage-2 miscalibration.
"""
if not (0 < sample_frac <= 1):
raise ValueError(f"sample_frac must be in (0, 1], got {sample_frac}")
@@ -2158,7 +2307,11 @@ def load_rollout_vs_truth(
_check_rollout_metadata(Path(rollout_path))
gen_raw, gen_cond, gen_pdg, gen_material = _load_world_frame_side(
rollout_path, sample_frac, seed, batch_size, extra_cols=["n_sec_pred"]
rollout_path,
sample_frac,
seed,
batch_size,
extra_cols=["n_sec_pred", "termination_reason"],
)
real_raw, real_cond, real_pdg, real_material = _load_world_frame_side(
truth_path, sample_frac, seed, batch_size, extra_cols=["child_track_ids"]
+349 -1
View File
@@ -12,6 +12,8 @@ from giant.analysis import (
RAW_TARGET_NAMES,
SampleCollection,
compute_event_observables_pl,
compute_rollout_observables,
compute_truth_observables,
constraint_report,
constraint_report_pl,
correlation_matrices,
@@ -31,6 +33,9 @@ from giant.analysis import (
plot_pairwise,
plot_pdg_energy_share,
plot_pdg_length_share,
plot_rollout_longitudinal,
plot_rollout_total_energy,
plot_rollout_transverse,
plot_shower_max_depth,
plot_total_energy,
plot_total_length,
@@ -767,7 +772,12 @@ def _write_rollout_parquet(path, n=150, seed=1, coord=ROLLOUT_COORD_VALUE):
"material": rng.choice(["W", "Pb"], n),
"layer_id": rng.integers(0, 10, n).astype(np.int32),
"n_sec_pred": rng.integers(0, 3, n).astype(np.int32),
"termination_reason": rng.choice(["natural_end", "energy_cutoff"], n),
# "" / "natural_end" mark a real generated step (the latter just
# additionally being a track's last); every row here is a real step,
# so all `n` should survive `load_rollout_vs_truth`'s filtering — see
# `test_load_rollout_vs_truth_drops_synthetic_termination_rows` for
# the escaped/unknown_pdg/energy_cutoff/max_steps bookkeeping rows.
"termination_reason": rng.choice(["", "natural_end"], n),
}
)
if coord is not None:
@@ -830,6 +840,80 @@ def test_load_rollout_vs_truth_joint_and_constraint_checks_run(tmp_path):
assert constraint_report(samples) is not None
def test_load_rollout_vs_truth_drops_synthetic_termination_rows(tmp_path):
"""Bookkeeping rows for escaped/unknown_pdg/energy_cutoff/max_steps aren't steps.
`rollout.py._terminal_rows` writes one such row per track termination, with
`step_length=0`/`post_pos=pre_pos` and for every reason but "escaped"
the track's entire remaining `pre_E` dumped into `edep` so the shower's
total energy still conserves. Mixing these into the per-step comparison
would inject a spurious step_length=0 spike and roughly double the
apparent mean edep from bookkeeping alone, not model behavior (see the
`load_rollout_vs_truth` docstring). This checks they're excluded and the
real steps are decoded unaffected by their presence in the same file.
"""
rng = np.random.default_rng(3)
n_real, n_marker = 60, 40
real_fields = _make_world_frame_physical(rng, n_real)
pre_pos, pre_dir, pre_E, step_length, post_pos, post_dir_world, post_E, edep = (
real_fields
)
marker_pre_pos = rng.uniform(-5.0, 5.0, (n_marker, 3)).astype(np.float32)
marker_pre_dir = _unit_vectors(rng, n_marker)
marker_pre_E = rng.uniform(1.0, 100.0, n_marker).astype(np.float32)
def cat(real_col, marker_col):
return np.concatenate([real_col, marker_col])
n = n_real + n_marker
table = pa.table(
{
"event_id": rng.integers(0, 20, n),
"track_id": rng.integers(0, 3, n),
"parent_id": np.full(n, -1, dtype=np.int64),
"generation": np.zeros(n, dtype=np.int64),
"step_no": np.zeros(n, dtype=np.int64),
"pdg": rng.choice([11, -11, 22], n),
"pre_x": cat(pre_pos[:, 0], marker_pre_pos[:, 0]),
"pre_y": cat(pre_pos[:, 1], marker_pre_pos[:, 1]),
"pre_z": cat(pre_pos[:, 2], marker_pre_pos[:, 2]),
"pre_E": cat(pre_E, marker_pre_E),
"pre_dx": cat(pre_dir[:, 0], marker_pre_dir[:, 0]),
"pre_dy": cat(pre_dir[:, 1], marker_pre_dir[:, 1]),
"pre_dz": cat(pre_dir[:, 2], marker_pre_dir[:, 2]),
# Terminal markers: post_pos == pre_pos (zero-length "step").
"post_x": cat(post_pos[:, 0], marker_pre_pos[:, 0]),
"post_y": cat(post_pos[:, 1], marker_pre_pos[:, 1]),
"post_z": cat(post_pos[:, 2], marker_pre_pos[:, 2]),
"post_E": cat(post_E, np.zeros(n_marker, dtype=np.float32)),
"post_dx": cat(post_dir_world[:, 0], marker_pre_dir[:, 0]),
"post_dy": cat(post_dir_world[:, 1], marker_pre_dir[:, 1]),
"post_dz": cat(post_dir_world[:, 2], marker_pre_dir[:, 2]),
# Terminal markers dump the full remaining pre_E into edep.
"edep": cat(edep, marker_pre_E),
"step_length": cat(step_length, np.zeros(n_marker, dtype=np.float32)),
"material": rng.choice(["W", "Pb"], n),
"layer_id": rng.integers(0, 10, n).astype(np.int32),
"n_sec_pred": rng.integers(0, 3, n).astype(np.int32),
"termination_reason": ([""] * n_real + ["energy_cutoff"] * n_marker),
}
)
table = table.replace_schema_metadata(
{PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE}
)
rollout_path = tmp_path / "rollout_with_markers.parquet"
pq.write_table(table, rollout_path)
truth_path = tmp_path / "truth.parquet"
_write_truth_parquet(truth_path, n=20)
samples = load_rollout_vs_truth(rollout_path, truth_path)
assert samples.gen_raw.shape == (n_real, 9)
np.testing.assert_allclose(samples.gen_raw, _expected_raw9(*real_fields), atol=1e-4)
def test_load_rollout_vs_truth_rejects_wrong_coord_metadata(tmp_path):
truth_path = tmp_path / "truth.parquet"
rollout_path = tmp_path / "rollout.parquet"
@@ -838,3 +922,267 @@ def test_load_rollout_vs_truth_rejects_wrong_coord_metadata(tmp_path):
with pytest.raises(ValueError, match="not a rollout file"):
load_rollout_vs_truth(rollout_path, truth_path)
# ---------------------------------------------------------------------------
# Tier 4: compute_rollout_observables / compute_truth_observables (unpaired)
# ---------------------------------------------------------------------------
def _make_rollout_style_event_arrays(rng):
"""3 events (3/2/4 steps), each with an unambiguous highest-pre_E row.
Same forced-max-pre_E-row construction as `_make_event_level_arrays`
above, but for the plain world-frame rollout/truth schema, where
edep/step_length are used directly (no energy-simplex decode needed).
"""
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
step_length = rng.uniform(0.1, 5.0, n).astype(np.float32)
travel_dir = _unit_vectors(rng, n)
post_pos = pre_pos + step_length[:, None] * travel_dir
edep = rng.uniform(0.1, 5.0, n).astype(np.float32)
return event_id, pre_pos, pre_dir, pre_E, post_pos, edep, step_length
def _expected_rollout_style_table(event_id, pre_pos, pre_dir, pre_E, post_pos, edep):
"""Independent re-derivation of total_edep/centroid_depth per event."""
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]
disp = post_pos[mask] - entry_pos
depth = disp @ axis_dir
total_edep = float(edep[mask].sum())
centroid = float((edep[mask] * depth).sum() / total_edep)
expected[e] = (total_edep, centroid, int(mask.sum()))
return expected
def _write_rollout_style_event_parquet(path, rng):
event_id, pre_pos, pre_dir, pre_E, post_pos, edep, step_length = (
_make_rollout_style_event_arrays(rng)
)
n = len(event_id)
table = pa.table(
{
"event_id": event_id,
"track_id": np.zeros(n, dtype=np.int64),
"parent_id": np.full(n, -1, dtype=np.int64),
"generation": np.zeros(n, dtype=np.int64),
"step_no": np.arange(n, dtype=np.int64),
"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],
"post_x": post_pos[:, 0],
"post_y": post_pos[:, 1],
"post_z": post_pos[:, 2],
"post_E": np.zeros(n, dtype=np.float32),
"post_dx": pre_dir[:, 0],
"post_dy": pre_dir[:, 1],
"post_dz": pre_dir[:, 2],
"edep": edep,
"step_length": step_length,
"material": rng.choice(["W", "Pb"], n),
"layer_id": rng.integers(0, 10, n).astype(np.int32),
"n_sec_pred": np.zeros(n, dtype=np.int32),
"termination_reason": [""] * n,
}
)
table = table.replace_schema_metadata(
{PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE}
)
pq.write_table(table, path)
return _expected_rollout_style_table(
event_id, pre_pos, pre_dir, pre_E, post_pos, edep
)
def _write_truth_style_event_parquet(path, rng):
event_id, pre_pos, pre_dir, pre_E, post_pos, edep, step_length = (
_make_rollout_style_event_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),
"child_track_ids": [[] for _ in range(n)],
"e_sec": np.zeros(n, dtype=np.float32),
"step_length": step_length,
"post_E": np.zeros(n, dtype=np.float32),
"edep": edep,
"post_dx": pre_dir[:, 0],
"post_dy": pre_dir[:, 1],
"post_dz": pre_dir[:, 2],
"post_x": post_pos[:, 0],
"post_y": post_pos[:, 1],
"post_z": post_pos[:, 2],
}
)
pq.write_table(table, path)
return _expected_rollout_style_table(
event_id, pre_pos, pre_dir, pre_E, post_pos, edep
)
def test_compute_rollout_observables_matches_manual_reconstruction(tmp_path):
path = tmp_path / "rollout_events.parquet"
# `centroid_depth` is weighted by *binned* depth (bin centers), not the raw
# continuous depth `_expected_rollout_style_table` computes, so it isn't
# checked here — see test_compute_rollout_and_truth_observables_agree_on_identical_data
# for a same-binning cross-check instead.
expected = _write_rollout_style_event_parquet(path, np.random.default_rng(11))
obs = compute_rollout_observables(path, depth_bins=5, transverse_bins=5)
table = obs.event_table.set_index("event_id")
for eid, (total_edep, _centroid, n_steps) in expected.items():
np.testing.assert_allclose(table.loc[eid, "total_edep"], total_edep, rtol=1e-4)
assert table.loc[eid, "n_steps"] == n_steps
def test_compute_truth_observables_matches_manual_reconstruction(tmp_path):
path = tmp_path / "truth_events.parquet"
expected = _write_truth_style_event_parquet(path, np.random.default_rng(12))
obs = compute_truth_observables(path, depth_bins=5, transverse_bins=5)
table = obs.event_table.set_index("event_id")
for eid, (total_edep, _centroid, n_steps) in expected.items():
np.testing.assert_allclose(
table.loc[eid, "real_total_edep"], total_edep, rtol=1e-4
)
assert table.loc[eid, "n_steps"] == n_steps
def test_compute_rollout_and_truth_observables_agree_on_identical_data(tmp_path):
"""`compute_rollout_observables`/`compute_truth_observables` must treat edep
identically: fed the exact same underlying step data (just written once
through each file's own schema), their depth/transverse profiles and total
per-event edep should come out numerically identical.
"""
rng = np.random.default_rng(13)
event_id, pre_pos, pre_dir, pre_E, post_pos, edep, step_length = (
_make_rollout_style_event_arrays(rng)
)
n = len(event_id)
rollout_table = pa.table(
{
"event_id": event_id,
"track_id": np.zeros(n, dtype=np.int64),
"parent_id": np.full(n, -1, dtype=np.int64),
"generation": np.zeros(n, dtype=np.int64),
"step_no": np.arange(n, dtype=np.int64),
"pdg": np.full(n, 11),
"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],
"post_x": post_pos[:, 0],
"post_y": post_pos[:, 1],
"post_z": post_pos[:, 2],
"post_E": np.zeros(n, dtype=np.float32),
"post_dx": pre_dir[:, 0],
"post_dy": pre_dir[:, 1],
"post_dz": pre_dir[:, 2],
"edep": edep,
"step_length": step_length,
"material": np.full(n, "W"),
"layer_id": np.zeros(n, dtype=np.int32),
"n_sec_pred": np.zeros(n, dtype=np.int32),
"termination_reason": [""] * n,
}
)
rollout_table = rollout_table.replace_schema_metadata(
{PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE}
)
rollout_path = tmp_path / "rollout.parquet"
pq.write_table(rollout_table, rollout_path)
truth_table = pa.table(
{
"event_id": event_id,
"pdg": np.full(n, 11),
"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": np.full(n, "W"),
"layer_id": np.zeros(n, dtype=np.int32),
"child_track_ids": [[] for _ in range(n)],
"e_sec": np.zeros(n, dtype=np.float32),
"step_length": step_length,
"post_E": np.zeros(n, dtype=np.float32),
"edep": edep,
"post_dx": pre_dir[:, 0],
"post_dy": pre_dir[:, 1],
"post_dz": pre_dir[:, 2],
"post_x": post_pos[:, 0],
"post_y": post_pos[:, 1],
"post_z": post_pos[:, 2],
}
)
truth_path = tmp_path / "truth.parquet"
pq.write_table(truth_table, truth_path)
rollout_obs = compute_rollout_observables(
rollout_path, depth_bins=5, transverse_bins=5
)
truth_obs = compute_truth_observables(truth_path, depth_bins=5, transverse_bins=5)
np.testing.assert_allclose(rollout_obs.depth_edges, truth_obs.depth_edges)
np.testing.assert_allclose(rollout_obs.transverse_edges, truth_obs.transverse_edges)
np.testing.assert_allclose(rollout_obs.depth_profile, truth_obs.real_depth_profile)
np.testing.assert_allclose(
rollout_obs.transverse_profile, truth_obs.real_transverse_profile
)
rollout_table_sorted = rollout_obs.event_table.sort_values("event_id")
truth_table_sorted = truth_obs.event_table.sort_values("event_id")
np.testing.assert_allclose(
rollout_table_sorted["total_edep"].to_numpy(),
truth_table_sorted["real_total_edep"].to_numpy(),
)
def test_plot_rollout_functions_accept_truth_observables_reference(tmp_path):
rollout_path = tmp_path / "rollout_events.parquet"
truth_path = tmp_path / "truth_events.parquet"
_write_rollout_style_event_parquet(rollout_path, np.random.default_rng(21))
_write_truth_style_event_parquet(truth_path, np.random.default_rng(22))
obs = compute_rollout_observables(rollout_path, depth_bins=5, transverse_bins=5)
reference = compute_truth_observables(truth_path, depth_bins=5, transverse_bins=5)
assert plot_rollout_longitudinal(obs, reference=reference) is not None
assert plot_rollout_transverse(obs, reference=reference) is not None
assert plot_rollout_total_energy(obs, reference=reference) is not None