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:
@@ -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
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user