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 1–2 are the deposit/secondary ALR energy logits (unconstrained reals, decoded against pre_E); columns 3–8 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 1–2 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