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:
+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