Files
giant/tests/test_analysis.py
T
lars 4c1250e246 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>
2026-07-13 12:38:44 +02:00

1189 lines
44 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import matplotlib
matplotlib.use("Agg") # no display needed for plot smoke tests
import numpy as np
import polars as pl
import pyarrow as pa
import pyarrow.parquet as pq
import pytest
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,
direction_alignment,
load_predicted_local,
load_rollout_vs_truth,
marginal_table,
marginal_table_pl,
pdg_contribution_table_pl,
plot_constraint_violations,
plot_correlation_matrices,
plot_direction_alignment,
plot_kl_bars,
plot_kl_bars_pl,
plot_longitudinal_profile,
plot_marginals,
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,
plot_transverse_profile,
)
from giant.constants import (
LOCAL_TARGET_NAMES,
PREDICT_COORD_METADATA_KEY,
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
ROLLOUT_COORD_VALUE,
)
from giant.data.transforms import (
energy_simplex_decode,
inv_log_transform,
local_frame_rotation,
log_transform,
reconstruct_post_pos,
travel_direction,
)
def _unit_vectors(rng, n):
v = rng.standard_normal((n, 3)).astype(np.float32)
return v / np.linalg.norm(v, axis=1, keepdims=True)
def _make_collection(n=200, seed=0, gen_offset=0.0) -> SampleCollection:
rng = np.random.default_rng(seed)
real = np.column_stack(
[
rng.uniform(0.1, 5.0, n), # step_length
rng.uniform(0.1, 5.0, n), # delta_e
rng.uniform(0.1, 5.0, n), # edep
_unit_vectors(rng, n), # post_dir
_unit_vectors(rng, n), # travel_dir
]
).astype(np.float32)
gen = real + gen_offset
pre_E = rng.uniform(1.0, 100.0, n).astype(np.float32)
cond_cont_raw = np.column_stack(
[
rng.standard_normal((n, 3)),
pre_E,
rng.standard_normal((n, 3)),
rng.integers(0, 5, n),
rng.integers(0, 3, n),
]
).astype(np.float32)
return SampleCollection(
cond_cont_raw=cond_cont_raw,
pdg=rng.choice([11, -11, 22], size=n),
material=rng.choice(["W", "Pb"], size=n),
real_raw=real,
gen_raw=gen,
)
def test_marginal_table_aggregate_has_all_dims():
table = marginal_table(_make_collection())
assert set(table["dim"]) == set(RAW_TARGET_NAMES)
assert (table["group"] == "all").all()
@pytest.mark.parametrize("group_by", ["pdg", "material", "energy"])
def test_marginal_table_grouped_covers_all_rows(group_by):
collection = _make_collection()
table = marginal_table(collection, group_by=group_by)
assert table["n"].groupby(table["group"]).first().sum() == len(collection.pdg)
def test_marginal_table_identical_distributions_have_zero_kl():
collection = _make_collection(gen_offset=0.0)
table = marginal_table(collection)
np.testing.assert_allclose(table["kl_real_gen"], 0.0, atol=1e-6)
def test_marginal_table_shifted_distribution_has_positive_kl():
collection = _make_collection(gen_offset=3.0)
table = marginal_table(collection)
assert (table["kl_real_gen"] > 0).all()
def test_correlation_matrices_are_symmetric_unit_diagonal():
real_corr, gen_corr = correlation_matrices(_make_collection())
for corr in (real_corr, gen_corr):
np.testing.assert_allclose(np.diag(corr), 1.0, atol=1e-5)
np.testing.assert_allclose(corr, corr.T, atol=1e-5)
def test_direction_alignment_real_data_is_unit_norm_dot_product():
real_cos, gen_cos = direction_alignment(_make_collection())
assert np.all(real_cos >= -1.0 - 1e-5) and np.all(real_cos <= 1.0 + 1e-5)
assert np.all(gen_cos >= -1.0 - 1e-5) and np.all(gen_cos <= 1.0 + 1e-5)
def test_constraint_report_clean_data_has_no_violations():
report = constraint_report(_make_collection(gen_offset=0.0))
assert (report["violation_rate"] == 0.0).all()
def test_constraint_report_flags_negative_log_dims_and_bad_norms():
collection = _make_collection(gen_offset=0.0)
collection.gen_raw[:, 0] = -1.0 # negative step_length
collection.gen_raw[:, 3:6] *= 2.0 # post_dir no longer unit norm
report = constraint_report(collection)
violations = dict(zip(report["check"], report["violation_rate"]))
assert violations["step_length >= 0"] == 1.0
assert violations["post_dir unit norm"] == 1.0
def test_plot_marginals_runs_without_error():
fig = plot_marginals(_make_collection())
assert fig is not None
def test_plot_marginals_grouped_runs_without_error():
fig = plot_marginals(_make_collection(), group_by="material")
assert fig is not None
def test_plot_correlation_matrices_runs_without_error():
fig = plot_correlation_matrices(_make_collection())
assert fig is not None
def test_plot_pairwise_runs_without_error():
fig = plot_pairwise(_make_collection())
assert fig is not None
def test_plot_direction_alignment_runs_without_error():
fig = plot_direction_alignment(_make_collection())
assert fig is not None
def test_plot_constraint_violations_runs_without_error():
fig = plot_constraint_violations(_make_collection())
assert fig is not None
def test_plot_kl_bars_runs_without_error():
fig = plot_kl_bars(_make_collection())
assert fig is not None
def test_plot_kl_bars_grouped_runs_without_error():
fig = plot_kl_bars(_make_collection(), group_by="material")
assert fig is not None
def test_plot_kl_bars_caps_groups_by_pdg():
collection = _make_collection(n=600)
collection.pdg = np.arange(600) % 8 # 8 distinct pdg values, > max_groups
fig = plot_kl_bars(collection, group_by="pdg", max_groups=3)
ax = fig.axes[0]
assert len({line.get_label() for line in ax.containers}) <= 3
def _write_predicted_local_parquet(path, n=50, metadata=None, rng=None):
"""Mimic `giant predict --coord local`'s output schema for the loader tests.
Column 0 is a log-scaled step_length; columns 12 are the deposit/secondary
ALR energy logits (unconstrained reals, decoded against pre_E); columns 38
are direction components.
"""
rng = rng or np.random.default_rng(0)
true_log_local = rng.standard_normal((n, 9)).astype(np.float32)
true_log_local[:, 0] = log_transform(rng.uniform(0.1, 5.0, n).astype(np.float32))
pred_log_local = true_log_local + rng.normal(0, 0.01, (n, 9)).astype(np.float32)
pre_E = rng.uniform(1.0, 100.0, n).astype(np.float32)
table = pa.table(
{
"event_id": rng.integers(0, 10, n),
"pdg": rng.choice([11, -11, 22], n),
"pre_x": rng.standard_normal(n).astype(np.float32),
"pre_y": rng.standard_normal(n).astype(np.float32),
"pre_z": rng.standard_normal(n).astype(np.float32),
"pre_E": pre_E,
"pre_dx": rng.standard_normal(n).astype(np.float32),
"pre_dy": rng.standard_normal(n).astype(np.float32),
"pre_dz": rng.standard_normal(n).astype(np.float32),
"material": rng.choice(["W", "Pb"], n),
"layer_id": rng.integers(0, 10, n).astype(np.int32),
"n_sec": rng.integers(0, 3, n).astype(np.int32),
**{
f"pred_{name}": pred_log_local[:, j]
for j, name in enumerate(LOCAL_TARGET_NAMES)
},
**{
f"true_{name}": true_log_local[:, j]
for j, name in enumerate(LOCAL_TARGET_NAMES)
},
}
)
if metadata is not None:
table = table.replace_schema_metadata(metadata)
pq.write_table(table, path)
return true_log_local, pred_log_local, pre_E
def test_load_predicted_local_round_trips_values(tmp_path):
path = tmp_path / "predicted_local.parquet"
true_log_local, pred_log_local, pre_E = _write_predicted_local_parquet(
path,
metadata={
PREDICT_COORD_METADATA_KEY: "local",
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
},
)
collection = load_predicted_local(path)
def expected_raw(log_local):
raw = log_local.copy()
raw[:, 0] = np.exp(log_local[:, 0]) - 1e-8
edep, _e_sec, _post_E, delta_e = energy_simplex_decode(log_local[:, 1:3], pre_E)
raw[:, 1] = delta_e
raw[:, 2] = edep
return raw
np.testing.assert_allclose(
collection.real_raw, expected_raw(true_log_local), atol=1e-4
)
np.testing.assert_allclose(
collection.gen_raw, expected_raw(pred_log_local), atol=1e-4
)
def test_load_predicted_local_usable_by_downstream_plots(tmp_path):
path = tmp_path / "predicted_local.parquet"
_write_predicted_local_parquet(
path,
metadata={
PREDICT_COORD_METADATA_KEY: "local",
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
},
)
collection = load_predicted_local(path)
assert marginal_table(collection) is not None
assert plot_marginals(collection) is not None
def test_load_predicted_local_rejects_missing_metadata(tmp_path):
path = tmp_path / "no_metadata.parquet"
_write_predicted_local_parquet(path, metadata=None)
with pytest.raises(ValueError, match="no '.*' parquet metadata"):
load_predicted_local(path)
def test_load_predicted_local_rejects_global_coord(tmp_path):
path = tmp_path / "global.parquet"
_write_predicted_local_parquet(
path,
metadata={
PREDICT_COORD_METADATA_KEY: "global",
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
},
)
with pytest.raises(ValueError, match="coord=local"):
load_predicted_local(path)
def test_load_predicted_local_rejects_mismatched_schema_version(tmp_path):
path = tmp_path / "old_version.parquet"
_write_predicted_local_parquet(
path,
metadata={
PREDICT_COORD_METADATA_KEY: "local",
PREDICT_SCHEMA_VERSION_KEY: "999",
},
)
with pytest.raises(ValueError, match="schema version"):
load_predicted_local(path)
def _predicted_local_path(tmp_path, n=200, seed=0):
path = tmp_path / "predicted_local.parquet"
_write_predicted_local_parquet(
path,
n=n,
rng=np.random.default_rng(seed),
metadata={
PREDICT_COORD_METADATA_KEY: "local",
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
},
)
return path
@pytest.mark.parametrize("group_by", [None, "pdg", "material", "energy"])
def test_marginal_table_pl_matches_numpy_version(tmp_path, group_by):
path = _predicted_local_path(tmp_path)
collection = load_predicted_local(path)
expected = marginal_table(collection, group_by=group_by).sort_values(
["group", "dim"]
)
actual = (
marginal_table_pl(path, group_by=group_by).sort(["group", "dim"]).to_pandas()
)
assert list(expected["group"]) == list(actual["group"])
assert list(expected["n"]) == list(actual["n"])
for col in ["real_mean", "gen_mean", "real_std", "gen_std"]:
np.testing.assert_allclose(
expected[col].to_numpy(),
actual[col].to_numpy(),
atol=1e-4,
rtol=1e-4,
)
# KL uses np.histogram (numpy path) vs polars Series.hist (lazy path); the two
# backends bin the boundary (min/max) sample differently, so allow a small
# absolute discrepancy rather than requiring bit-identical estimates.
np.testing.assert_allclose(
expected["kl_real_gen"].to_numpy(),
actual["kl_real_gen"].to_numpy(),
atol=2e-2,
)
@pytest.mark.parametrize("group_by", [None, "pdg", "material", "energy"])
def test_plot_kl_bars_pl_runs_without_error(tmp_path, group_by):
path = _predicted_local_path(tmp_path)
fig = plot_kl_bars_pl(path, group_by=group_by)
assert fig is not None
def test_marginal_table_pl_rejects_missing_metadata(tmp_path):
path = tmp_path / "no_metadata.parquet"
_write_predicted_local_parquet(path, metadata=None)
with pytest.raises(ValueError, match="no '.*' parquet metadata"):
marginal_table_pl(path)
def test_constraint_report_pl_matches_numpy_version(tmp_path):
path = _predicted_local_path(tmp_path)
collection = load_predicted_local(path)
expected = constraint_report(collection)
actual = constraint_report_pl(path).to_pandas()
assert list(expected["check"]) == list(actual["check"])
np.testing.assert_allclose(
expected["violation_rate"].to_numpy(),
actual["violation_rate"].to_numpy(),
atol=1e-6,
)
np.testing.assert_allclose(
expected["mean_abs_error"].to_numpy(),
actual["mean_abs_error"].to_numpy(),
atol=1e-4,
)
# ---------------------------------------------------------------------------
# Tier 4: event-level (shower) observables
# ---------------------------------------------------------------------------
def _make_event_level_arrays(rng):
"""3 events (3/2/4 steps), each with an unambiguous highest-pre_E row.
The forced max-pre_E rows (indices 1, 3, 7) fix a known shower axis/entry
point per event, so the expected event_table can be re-derived
independently in the test without depending on compute_event_observables_pl.
"""
event_id = np.array([0, 0, 0, 1, 1, 2, 2, 2, 2], dtype=np.int64)
n = len(event_id)
pre_pos = rng.uniform(-5.0, 5.0, (n, 3)).astype(np.float32)
pre_dir = _unit_vectors(rng, n)
pre_E = rng.uniform(1.0, 50.0, n).astype(np.float32)
pre_E[1] = 100.0 # event 0's entry step
pre_E[3] = 100.0 # event 1's entry step
pre_E[7] = 100.0 # event 2's entry step
def _local_block():
block = rng.standard_normal((n, 9)).astype(np.float32)
block[:, 0] = log_transform(rng.uniform(0.1, 5.0, n).astype(np.float32))
# cols 12 stay as random ALR energy logits (decoded against pre_E)
block[:, 3:6] = _unit_vectors(rng, n)
block[:, 6:9] = _unit_vectors(rng, n)
return block
true_log_local = _local_block()
pred_log_local = _local_block()
return event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
def _write_event_level_parquet(path, rng=None):
rng = rng or np.random.default_rng(7)
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local = (
_make_event_level_arrays(rng)
)
n = len(event_id)
pdg = rng.choice([11, -11, 22], n)
table = pa.table(
{
"event_id": event_id,
"pdg": pdg,
"pre_x": pre_pos[:, 0],
"pre_y": pre_pos[:, 1],
"pre_z": pre_pos[:, 2],
"pre_E": pre_E,
"pre_dx": pre_dir[:, 0],
"pre_dy": pre_dir[:, 1],
"pre_dz": pre_dir[:, 2],
"material": rng.choice(["W", "Pb"], n),
"layer_id": rng.integers(0, 10, n).astype(np.int32),
"n_sec": rng.integers(0, 3, n).astype(np.int32),
**{
f"pred_{name}": pred_log_local[:, j]
for j, name in enumerate(LOCAL_TARGET_NAMES)
},
**{
f"true_{name}": true_log_local[:, j]
for j, name in enumerate(LOCAL_TARGET_NAMES)
},
}
)
table = table.replace_schema_metadata(
{
PREDICT_COORD_METADATA_KEY: "local",
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
}
)
pq.write_table(table, path)
return event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local, pdg
def _expected_event_table(
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
):
"""Independent re-derivation of total/centroid/RMS per event, for comparison."""
expected = {}
for e in sorted(np.unique(event_id).tolist()):
mask = event_id == e
entry_idx = np.where(mask)[0][np.argmax(pre_E[mask])]
entry_pos = pre_pos[entry_idx]
axis_dir = pre_dir[entry_idx]
def agg(log_local):
step_length = inv_log_transform(log_local[mask, 0])
edep, _e_sec, _post_E, _delta_e = energy_simplex_decode(
log_local[mask, 1:3], pre_E[mask]
)
travel_dir_local = log_local[mask, 6:9]
post_pos = reconstruct_post_pos(
pre_pos[mask], pre_dir[mask], step_length, travel_dir_local
)
disp = post_pos - entry_pos
depth = disp @ axis_dir
transverse = np.linalg.norm(disp - depth[:, None] * axis_dir, axis=1)
total_edep = float(edep.sum())
total_length = float(step_length.sum())
centroid = float((edep * depth).sum() / total_edep)
rms = float(np.sqrt((edep * transverse**2).sum() / total_edep))
return total_edep, total_length, centroid, rms
real_total_edep, real_total_length, real_centroid, real_rms = agg(
true_log_local
)
gen_total_edep, gen_total_length, gen_centroid, gen_rms = agg(pred_log_local)
expected[e] = (
real_total_edep,
gen_total_edep,
real_total_length,
gen_total_length,
real_centroid,
gen_centroid,
real_rms,
gen_rms,
)
return expected
def test_compute_event_observables_pl_matches_manual_reconstruction(tmp_path):
path = tmp_path / "event_level.parquet"
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local, _pdg = (
_write_event_level_parquet(path)
)
expected = _expected_event_table(
event_id, pre_pos, pre_dir, pre_E, true_log_local, pred_log_local
)
obs = compute_event_observables_pl(path, depth_bins=5, transverse_bins=5)
table = obs.event_table.sort("event_id")
for i, eid in enumerate(table["event_id"].to_list()):
(
real_total_edep,
gen_total_edep,
real_total_length,
gen_total_length,
real_centroid,
gen_centroid,
real_rms,
gen_rms,
) = expected[eid]
np.testing.assert_allclose(
table["real_total_edep"][i], real_total_edep, rtol=1e-4
)
np.testing.assert_allclose(
table["gen_total_edep"][i], gen_total_edep, rtol=1e-4
)
np.testing.assert_allclose(
table["real_total_length"][i], real_total_length, rtol=1e-4
)
np.testing.assert_allclose(
table["gen_total_length"][i], gen_total_length, rtol=1e-4
)
np.testing.assert_allclose(
table["real_centroid_depth"][i], real_centroid, rtol=1e-3, atol=1e-4
)
np.testing.assert_allclose(
table["gen_centroid_depth"][i], gen_centroid, rtol=1e-3, atol=1e-4
)
np.testing.assert_allclose(
table["real_transverse_rms"][i], real_rms, rtol=1e-3, atol=1e-4
)
np.testing.assert_allclose(
table["gen_transverse_rms"][i], gen_rms, rtol=1e-3, atol=1e-4
)
def test_compute_event_observables_pl_profile_shapes(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
obs = compute_event_observables_pl(path, depth_bins=7, transverse_bins=4)
assert obs.depth_edges.shape == (8,)
assert obs.transverse_edges.shape == (5,)
assert obs.real_depth_profile.shape == (7,)
assert obs.gen_depth_profile.shape == (7,)
assert obs.real_transverse_profile.shape == (4,)
assert obs.gen_transverse_profile.shape == (4,)
assert len(obs.event_table) == 3
def test_compute_event_observables_pl_accepts_lazyframe(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
obs_from_path = compute_event_observables_pl(path, depth_bins=5, transverse_bins=5)
obs_from_lf = compute_event_observables_pl(
pl.scan_parquet(path), depth_bins=5, transverse_bins=5
)
np.testing.assert_allclose(
obs_from_lf.event_table.sort("event_id")["real_total_edep"].to_numpy(),
obs_from_path.event_table.sort("event_id")["real_total_edep"].to_numpy(),
)
def test_event_level_plots_run_without_error(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
obs = compute_event_observables_pl(path, depth_bins=5, transverse_bins=5)
assert plot_total_energy(obs) is not None
assert plot_total_length(obs) is not None
assert plot_longitudinal_profile(obs) is not None
assert plot_transverse_profile(obs) is not None
assert plot_shower_max_depth(obs) is not None
def test_pdg_contribution_table_pl_matches_manual_sums(tmp_path):
path = tmp_path / "event_level.parquet"
_, _, _, pre_E, true_log_local, pred_log_local, pdg = _write_event_level_parquet(
path
)
real_edep = energy_simplex_decode(true_log_local[:, 1:3], pre_E)[0]
gen_edep = energy_simplex_decode(pred_log_local[:, 1:3], pre_E)[0]
real_length = inv_log_transform(true_log_local[:, 0])
gen_length = inv_log_transform(pred_log_local[:, 0])
expected = {}
for p in np.unique(pdg):
mask = pdg == p
expected[int(p)] = (
float(real_edep[mask].sum()),
float(gen_edep[mask].sum()),
float(real_length[mask].sum()),
float(gen_length[mask].sum()),
)
table = pdg_contribution_table_pl(path).sort("pdg")
for i, p in enumerate(table["pdg"].to_list()):
real_e, gen_e, real_l, gen_l = expected[int(p)]
np.testing.assert_allclose(table["real_total_edep"][i], real_e, rtol=1e-4)
np.testing.assert_allclose(table["gen_total_edep"][i], gen_e, rtol=1e-4)
np.testing.assert_allclose(table["real_total_length"][i], real_l, rtol=1e-4)
np.testing.assert_allclose(table["gen_total_length"][i], gen_l, rtol=1e-4)
def test_pdg_contribution_table_pl_accepts_lazyframe(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
from_path = pdg_contribution_table_pl(path).sort("pdg")
from_lf = pdg_contribution_table_pl(pl.scan_parquet(path)).sort("pdg")
np.testing.assert_allclose(
from_lf["real_total_edep"].to_numpy(), from_path["real_total_edep"].to_numpy()
)
def test_pdg_pie_plots_run_without_error(tmp_path):
path = tmp_path / "event_level.parquet"
_write_event_level_parquet(path)
table = pdg_contribution_table_pl(path)
assert plot_pdg_energy_share(table) is not None
assert plot_pdg_length_share(table) is not None
def test_plot_pdg_energy_share_caps_slices():
table = pl.DataFrame(
{
"pdg": list(range(10)),
"real_total_edep": [float(10 - i) for i in range(10)],
"gen_total_edep": [float(10 - i) for i in range(10)],
"real_total_length": [float(10 - i) for i in range(10)],
"gen_total_length": [float(10 - i) for i in range(10)],
}
)
fig = plot_pdg_energy_share(table, max_slices=4)
for ax in fig.axes:
assert len(ax.patches) == 4
# ---------------------------------------------------------------------------
# load_rollout_vs_truth: unpaired rollout-vs-truth SampleCollection
# ---------------------------------------------------------------------------
def _make_world_frame_physical(rng, n):
"""Random-but-physical pre/post step fields shared by the truth/rollout schemas."""
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, 100.0, n).astype(np.float32)
step_length = rng.uniform(0.1, 5.0, n).astype(np.float32)
travel_dir_world = _unit_vectors(rng, n)
post_pos = pre_pos + step_length[:, None] * travel_dir_world
post_dir_world = _unit_vectors(rng, n)
delta_e = (rng.uniform(0.0, 1.0, n) * pre_E).astype(np.float32)
post_E = pre_E - delta_e
edep = (delta_e * rng.uniform(0.0, 1.0, n)).astype(np.float32)
return pre_pos, pre_dir, pre_E, step_length, post_pos, post_dir_world, post_E, edep
def _expected_raw9(
pre_pos, pre_dir, pre_E, step_length, post_pos, post_dir_world, post_E, edep
):
post_dir_local = local_frame_rotation(pre_dir, post_dir_world)
travel_dir_local = local_frame_rotation(
pre_dir, travel_direction(pre_pos, post_pos)
)
return np.column_stack(
[step_length, pre_E - post_E, edep, post_dir_local, travel_dir_local]
).astype(np.float32)
def _write_truth_parquet(path, n=200, seed=0):
rng = np.random.default_rng(seed)
fields = _make_world_frame_physical(rng, n)
pre_pos, pre_dir, pre_E, step_length, post_pos, post_dir_world, post_E, edep = (
fields
)
table = pa.table(
{
"event_id": rng.integers(0, 20, n),
"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": [list(range(int(k))) for k in rng.integers(0, 3, n)],
"e_sec": rng.uniform(0.0, 1.0, n).astype(np.float32),
"step_length": step_length,
"post_E": post_E,
"edep": edep,
"post_dx": post_dir_world[:, 0],
"post_dy": post_dir_world[:, 1],
"post_dz": post_dir_world[:, 2],
"post_x": post_pos[:, 0],
"post_y": post_pos[:, 1],
"post_z": post_pos[:, 2],
}
)
pq.write_table(table, path)
return _expected_raw9(*fields)
def _write_rollout_parquet(path, n=150, seed=1, coord=ROLLOUT_COORD_VALUE):
rng = np.random.default_rng(seed)
fields = _make_world_frame_physical(rng, n)
pre_pos, pre_dir, pre_E, step_length, post_pos, post_dir_world, post_E, edep = (
fields
)
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": 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": post_E,
"post_dx": post_dir_world[:, 0],
"post_dy": post_dir_world[:, 1],
"post_dz": post_dir_world[:, 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": rng.integers(0, 3, n).astype(np.int32),
# "" / "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:
table = table.replace_schema_metadata({PREDICT_COORD_METADATA_KEY: coord})
pq.write_table(table, path)
return _expected_raw9(*fields)
def test_load_rollout_vs_truth_decodes_raw_targets_correctly(tmp_path):
truth_path = tmp_path / "truth.parquet"
rollout_path = tmp_path / "rollout.parquet"
expected_real = _write_truth_parquet(truth_path, n=200, seed=0)
expected_gen = _write_rollout_parquet(rollout_path, n=150, seed=1)
samples = load_rollout_vs_truth(rollout_path, truth_path)
np.testing.assert_allclose(samples.real_raw, expected_real, atol=1e-4)
np.testing.assert_allclose(samples.gen_raw, expected_gen, atol=1e-4)
def test_load_rollout_vs_truth_allows_unpaired_lengths(tmp_path):
truth_path = tmp_path / "truth.parquet"
rollout_path = tmp_path / "rollout.parquet"
_write_truth_parquet(truth_path, n=200)
_write_rollout_parquet(rollout_path, n=150)
samples = load_rollout_vs_truth(rollout_path, truth_path)
assert samples.real_raw.shape == (200, 9)
assert samples.gen_raw.shape == (150, 9)
assert samples.pdg.shape == (200,)
assert samples.pdg_gen.shape == (150,)
@pytest.mark.parametrize("group_by", [None, "pdg", "material", "energy"])
def test_load_rollout_vs_truth_downstream_plots_run_without_error(tmp_path, group_by):
truth_path = tmp_path / "truth.parquet"
rollout_path = tmp_path / "rollout.parquet"
_write_truth_parquet(truth_path, n=200)
_write_rollout_parquet(rollout_path, n=150)
samples = load_rollout_vs_truth(rollout_path, truth_path)
table = marginal_table(samples, group_by=group_by)
assert set(table["dim"]) == set(RAW_TARGET_NAMES)
assert plot_marginals(samples, group_by=group_by) is not None
assert plot_kl_bars(samples, group_by=group_by) is not None
def test_load_rollout_vs_truth_joint_and_constraint_checks_run(tmp_path):
truth_path = tmp_path / "truth.parquet"
rollout_path = tmp_path / "rollout.parquet"
_write_truth_parquet(truth_path, n=200)
_write_rollout_parquet(rollout_path, n=150)
samples = load_rollout_vs_truth(rollout_path, truth_path)
assert plot_correlation_matrices(samples) is not None
assert plot_pairwise(samples) is not None
assert plot_direction_alignment(samples) is not None
assert plot_constraint_violations(samples) is not None
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"
_write_truth_parquet(truth_path)
_write_rollout_parquet(rollout_path, coord="local")
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