From 4c1250e246b4fc9d93ea37b7f6c7d2cba196f42a Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Mon, 13 Jul 2026 12:38:44 +0200 Subject: [PATCH] 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 --- analysis/rollout_validation.ipynb | 14 +- giant/analysis.py | 205 ++++++++++++++--- tests/test_analysis.py | 350 +++++++++++++++++++++++++++++- 3 files changed, 534 insertions(+), 35 deletions(-) diff --git a/analysis/rollout_validation.ipynb b/analysis/rollout_validation.ipynb index 70af3f2..ff5cd69 100644 --- a/analysis/rollout_validation.ipynb +++ b/analysis/rollout_validation.ipynb @@ -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)" ] }, { diff --git a/giant/analysis.py b/giant/analysis.py index d491fb9..0453ac9 100644 --- a/giant/analysis.py +++ b/giant/analysis.py @@ -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"] diff --git a/tests/test_analysis.py b/tests/test_analysis.py index b46e556..2fe4d7e 100644 --- a/tests/test_analysis.py +++ b/tests/test_analysis.py @@ -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