Merge branch 'phase2-secondary-prediction' into 4-prototype-a-mixture-of-experts-routing-tree-architecture

Brings in the rollout-validation fixes developed alongside Phase 2
(exact e_sec budget rescaling in decode_secondaries, filtering
synthetic termination rows out of load_rollout_vs_truth, Tier 4 truth
overlay, --energy-gev support in dwarf make-root) and reconciles them
with this branch's mixture-of-experts routing work: build_features/
build_models/dataset plumbing keep the ProcessRouter's proc_map/
proc_idx threading, and create_root_files.py's job_seed folds in both
the per-job seed derivation and the new energy_gev component.
This commit is contained in:
2026-07-15 09:59:43 +02:00
21 changed files with 2040 additions and 208 deletions
+526
View File
@@ -12,11 +12,14 @@ 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,
@@ -30,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,
@@ -40,12 +46,15 @@ from giant.constants import (
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,
)
@@ -660,3 +669,520 @@ def test_plot_pdg_energy_share_caps_slices():
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
+41 -14
View File
@@ -139,24 +139,29 @@ def test_plan_jobs_multiple_detectors_each_start_independently(tmp_path):
def test_job_seed_deterministic():
job = SimJob(detector="pbwo4", config=None, shard_index=3)
assert job_seed("steps", "gen1", job) == job_seed("steps", "gen1", job)
assert job_seed("steps", "gen1", job, None) == job_seed("steps", "gen1", job, None)
def test_job_seed_varies_by_shard_index():
a = SimJob(detector="pbwo4", config=None, shard_index=0)
b = SimJob(detector="pbwo4", config=None, shard_index=1)
assert job_seed("steps", "gen1", a) != job_seed("steps", "gen1", b)
assert job_seed("steps", "gen1", a, None) != job_seed("steps", "gen1", b, None)
def test_job_seed_varies_by_detector():
a = SimJob(detector="pbwo4", config=None, shard_index=0)
b = SimJob(detector="sampling_pb_scint", config="pb_scint", shard_index=0)
assert job_seed("steps", "gen1", a) != job_seed("steps", "gen1", b)
assert job_seed("steps", "gen1", a, None) != job_seed("steps", "gen1", b, None)
def test_job_seed_varies_by_gen():
job = SimJob(detector="pbwo4", config=None, shard_index=0)
assert job_seed("steps", "gen1", job) != job_seed("steps", "gen2", job)
assert job_seed("steps", "gen1", job, None) != job_seed("steps", "gen2", job, None)
def test_job_seed_varies_by_energy():
job = SimJob(detector="pbwo4", config=None, shard_index=0)
assert job_seed("steps", "gen1", job, 1.0) != job_seed("steps", "gen1", job, 10.0)
def test_run_job_passes_deterministic_seed_env_var(tmp_path):
@@ -166,11 +171,11 @@ def test_run_job_passes_deterministic_seed_env_var(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=5)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert result.dest is not None
payload = json.loads(result.dest.read_text())
assert payload["seed"] == str(job_seed("steps", "gen1", job))
assert payload["seed"] == str(job_seed("steps", "gen1", job, None))
def test_run_job_moves_output_to_correct_shard_path(tmp_path):
@@ -181,7 +186,7 @@ def test_run_job_moves_output_to_correct_shard_path(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=7)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert result.ok
assert result.dest == gen_dir / "pbwo4" / "shard-007.root"
@@ -197,7 +202,7 @@ def test_run_job_passes_config_arg_and_isolates_cwd(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="sampling_pb_scint", config="pb_scint", shard_index=0)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert result.ok
assert result.dest is not None
@@ -215,13 +220,27 @@ def test_run_job_omits_config_arg_when_none(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=0)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert result.dest is not None
payload = json.loads(result.dest.read_text())
assert payload["argv"] == ["10000"]
def test_run_job_appends_energy_arg_when_given(tmp_path):
fake = _write_fake_executable(tmp_path / "fake_exe.py")
(tmp_path / "raw" / "steps" / "gen1").mkdir(parents=True)
tmp_root = tmp_path / ".sim-tmp"
tmp_root.mkdir()
job = SimJob(detector="pbwo4_10gev", config=None, shard_index=0)
result = run_job(job, fake, 10000, 10.0, tmp_path, "steps", "gen1", tmp_root)
assert result.dest is not None
payload = json.loads(result.dest.read_text())
assert payload["argv"] == ["10000", "10.0"]
def test_run_job_fails_when_executable_errors(tmp_path):
fake = _write_fake_executable(tmp_path / "fake_exe.py", exit_code=1)
(tmp_path / "raw" / "steps" / "gen1").mkdir(parents=True)
@@ -229,7 +248,7 @@ def test_run_job_fails_when_executable_errors(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=0)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert not result.ok
assert "exited 1" in result.message
@@ -242,7 +261,7 @@ def test_run_job_fails_when_no_root_file_produced(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=0)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert not result.ok
assert "found 0" in result.message
@@ -255,7 +274,7 @@ def test_run_job_fails_when_multiple_root_files_produced(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=0)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert not result.ok
assert "found 2" in result.message
@@ -270,7 +289,7 @@ def test_run_job_refuses_to_overwrite_existing_shard(tmp_path):
tmp_root.mkdir()
job = SimJob(detector="pbwo4", config=None, shard_index=0)
result = run_job(job, fake, 10000, tmp_path, "steps", "gen1", tmp_root)
result = run_job(job, fake, 10000, None, tmp_path, "steps", "gen1", tmp_root)
assert not result.ok
assert "overwrite" in result.message
@@ -285,7 +304,15 @@ def test_run_all_caps_concurrency(tmp_path):
jobs = [SimJob(detector="pbwo4", config=None, shard_index=i) for i in range(6)]
results = run_all(
jobs, fake, 10000, tmp_path, "steps", "gen1", max_workers=2, tmp_root=tmp_root
jobs,
fake,
10000,
None,
tmp_path,
"steps",
"gen1",
max_workers=2,
tmp_root=tmp_root,
)
assert all(r.ok and r.dest is not None for r in results)
+141 -3
View File
@@ -1,4 +1,5 @@
"""Tests for Phase 2: secondary particle prediction."""
import numpy as np
import pytest
import torch
@@ -11,6 +12,7 @@ from giant.sample import sample_secondaries, snap_type_to_pdg_idx
# ── helpers ──────────────────────────────────────────────────────────────────
def _stage1(pdg=3, mat=2):
return DenoisingMLP(pdg_vocab=pdg, mat_vocab=mat, hidden_dim=32, n_blocks=2)
@@ -29,6 +31,7 @@ def _cond(B=8, pdg=3, mat=2):
# ── DenoisingMLP Phase-2 additions ───────────────────────────────────────────
def test_predict_n_sec_shape():
B = 8
model = _stage1()
@@ -53,6 +56,7 @@ def test_pdg_embedding_weight_shape():
# ── SecondaryDecoder ──────────────────────────────────────────────────────────
def test_sec_decoder_output_shape():
B = 8
decoder = _sec_decoder()
@@ -89,6 +93,7 @@ def test_sec_decoder_gradients():
# ── masked flow matching loss ─────────────────────────────────────────────────
def test_flow_matching_loss_secondary_scalar():
B, pdg, mat = 8, 3, 2
decoder = _sec_decoder(pdg, mat)
@@ -96,7 +101,9 @@ def test_flow_matching_loss_secondary_scalar():
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
sec_mask = torch.ones(B, K_MAX, dtype=torch.bool)
loss = flow_matching_loss_secondary(decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask)
loss = flow_matching_loss_secondary(
decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask
)
assert loss.shape == ()
assert loss.item() >= 0.0
@@ -109,7 +116,9 @@ def test_flow_matching_loss_secondary_mask_zeros_padding():
cond_cont, cond_cat = _cond(B, pdg, mat)
stage1_out = torch.randn(B, X_DIM)
sec_mask = torch.zeros(B, K_MAX, dtype=torch.bool)
loss = flow_matching_loss_secondary(decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask)
loss = flow_matching_loss_secondary(
decoder, x1, cond_cont, cond_cat, stage1_out, sec_mask
)
assert loss.item() == pytest.approx(0.0, abs=1e-6)
@@ -128,6 +137,7 @@ def test_flow_matching_loss_secondary_has_grad():
# ── sampling ──────────────────────────────────────────────────────────────────
def test_sample_secondaries_shapes():
B, pdg, mat = 6, 3, 2
decoder = _sec_decoder(pdg, mat)
@@ -169,6 +179,7 @@ def test_snap_type_to_pdg_idx_shape():
# ── encode_secondaries round-trip ─────────────────────────────────────────────
def test_encode_secondaries_energy_conservation():
"""Decoded stick-breaking fractions must sum to ≈ e_sec."""
from giant.data.transforms import encode_secondaries
@@ -217,6 +228,133 @@ def test_encode_secondaries_direction_encoding():
sec_cont = encode_secondaries(sec_E_list, sec_dir_list, sec_valid, e_sec, pre_dir)
# dir columns are sec_cont[:, :, 1:4]
local_dirs = sec_cont[:, :2, 1:4] # (N, 2, 3) — valid slots only
local_dirs = sec_cont[:, :2, 1:4] # (N, 2, 3) — valid slots only
norms_out = np.linalg.norm(local_dirs, axis=-1)
np.testing.assert_allclose(norms_out, 1.0, atol=1e-5)
# ── decode_secondaries: exact energy conservation ────────────────────────────
def _random_sec_cont(rng, N, stick_logit_scale=1.0):
sec_cont = rng.standard_normal((N, K_MAX, 4)).astype(np.float32)
sec_cont[:, :, 0] *= stick_logit_scale
dirs = sec_cont[:, :, 1:]
dirs /= np.linalg.norm(dirs, axis=-1, keepdims=True)
return sec_cont
def test_decode_secondaries_valid_slots_sum_to_e_sec():
"""The valid slots' energies must sum to exactly e_sec, not just <= e_sec.
Rows with n_sec=0 are excluded: there's no slot to put the budget in, so
valid_sum is correctly 0 regardless of e_sec there (see
test_decode_secondaries_zero_n_sec_has_zero_energy) the shortfall in
that case is handled downstream (e.g. rollout.py dumps it into edep).
"""
from giant.data.transforms import decode_secondaries
rng = np.random.default_rng(0)
N = 200
sec_cont = _random_sec_cont(rng, N)
sec_pdg_pred = np.zeros((N, K_MAX), dtype=np.int64)
n_sec = rng.integers(0, K_MAX + 1, size=N)
e_sec = rng.uniform(0.0, 50.0, size=N).astype(np.float32)
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
sec_E, _sec_dir, _sec_pdg, sec_valid = decode_secondaries(
sec_cont, sec_pdg_pred, n_sec, e_sec, pre_dir, {0: 22}
)
valid_sum = (sec_E * sec_valid).sum(axis=1)
has_secondaries = n_sec > 0
np.testing.assert_allclose(
valid_sum[has_secondaries],
e_sec[has_secondaries],
atol=1e-3,
rtol=1e-5,
)
def test_decode_secondaries_zero_n_sec_has_zero_energy():
"""n_sec=0 rows get no secondaries and no forced energy assignment."""
from giant.data.transforms import decode_secondaries
rng = np.random.default_rng(1)
N = 10
sec_cont = _random_sec_cont(rng, N)
sec_pdg_pred = np.zeros((N, K_MAX), dtype=np.int64)
n_sec = np.zeros(N, dtype=np.int64)
e_sec = rng.uniform(1.0, 10.0, size=N).astype(np.float32)
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
sec_E, _sec_dir, _sec_pdg, sec_valid = decode_secondaries(
sec_cont, sec_pdg_pred, n_sec, e_sec, pre_dir, {0: 22}
)
assert not sec_valid.any()
np.testing.assert_allclose(sec_E, 0.0)
def test_decode_secondaries_degenerate_row_falls_back_to_even_split():
"""All-zero stick fractions for the valid slots fall back to an even split."""
from giant.data.transforms import decode_secondaries
rng = np.random.default_rng(2)
N = 4
sec_cont = _random_sec_cont(rng, N)
# Drive every valid slot's stick-breaking fraction to ~0 (huge negative logit).
n_sec = np.array([0, 1, 3, K_MAX])
for i, k in enumerate(n_sec):
sec_cont[i, :k, 0] = -80.0
sec_pdg_pred = np.zeros((N, K_MAX), dtype=np.int64)
e_sec = np.array([0.0, 4.0, 9.0, 30.0], dtype=np.float32)
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
sec_E, _sec_dir, _sec_pdg, sec_valid = decode_secondaries(
sec_cont, sec_pdg_pred, n_sec, e_sec, pre_dir, {0: 22}
)
for i, k in enumerate(n_sec):
if k == 0:
continue
np.testing.assert_allclose(sec_E[i, :k], e_sec[i] / k, atol=1e-4)
np.testing.assert_allclose(sec_E[i, :k].sum(), e_sec[i], atol=1e-3)
def test_decode_secondaries_rescale_preserves_relative_shares():
"""Rescaling should keep each valid slot's *share* of the budget unchanged.
A shortfall shouldn't get dumped into whichever slot is last by energy
rank it should be spread proportionally, i.e. sec_E[i] / sec_E[j] for
two valid slots must match before and after the e_sec rescale.
"""
from giant.data.transforms import decode_secondaries
rng = np.random.default_rng(3)
N = 1
sec_cont = _random_sec_cont(rng, N)
n_sec = np.array([4])
pre_dir = np.array([[0.0, 0.0, 1.0]], dtype=np.float32)
sec_pdg_pred = np.zeros((N, K_MAX), dtype=np.int64)
sec_E_small, _, _, sec_valid = decode_secondaries(
sec_cont,
sec_pdg_pred,
n_sec,
np.array([5.0], dtype=np.float32),
pre_dir,
{0: 22},
)
sec_E_large, _, _, _ = decode_secondaries(
sec_cont,
sec_pdg_pred,
n_sec,
np.array([50.0], dtype=np.float32),
pre_dir,
{0: 22},
)
ratio_small = sec_E_small[0, :4] / sec_E_small[0, 0]
ratio_large = sec_E_large[0, :4] / sec_E_large[0, 0]
np.testing.assert_allclose(ratio_small, ratio_large, rtol=1e-4)
+20 -5
View File
@@ -56,16 +56,31 @@ def _seeds(n=6):
}
def _run(escape_threshold=1e9, energy_cutoff=1.0, max_steps=30,
max_tracks_per_event=300, seeds=None):
def _run(
escape_threshold=1e9,
energy_cutoff=1.0,
max_steps=30,
max_tracks_per_event=300,
seeds=None,
):
torch.manual_seed(0)
np.random.seed(0)
s1, s2 = _models()
cond, tgt = _norms()
return rollout(
s1, s2, _oracle(), seeds or _seeds(), cond, tgt, PDG_MAP, MAT_MAP,
energy_cutoff=energy_cutoff, max_steps=max_steps, steps=4,
batch_size=128, max_tracks_per_event=max_tracks_per_event,
s1,
s2,
_oracle(),
seeds or _seeds(),
cond,
tgt,
PDG_MAP,
MAT_MAP,
energy_cutoff=energy_cutoff,
max_steps=max_steps,
steps=4,
batch_size=128,
max_tracks_per_event=max_tracks_per_event,
escape_threshold=escape_threshold,
)
+10 -3
View File
@@ -263,6 +263,14 @@ def _minimal_step_data(N: int, process: np.ndarray | None = None) -> dict:
return data
def _step_data_no_sec_lists(n_sec: np.ndarray) -> dict:
"""Minimal build_features input with n_sec but no per-secondary list columns
(mimics a parquet that skipped the parent->child join)."""
data = _minimal_step_data(len(n_sec))
data["n_sec"] = np.asarray(n_sec, dtype=np.int32)
return data
def test_build_features_proc_idx_zero_without_proc_map():
data = _minimal_step_data(3, process=np.array(["compt", "phot", "eIoni"], dtype=object))
pdg_map, mat_map = {11: 0}, {"PbWO4": 0}
@@ -287,8 +295,7 @@ def test_build_features_require_secondaries_raises_when_lists_missing():
through the parent->child join; require_secondaries must catch it instead of
silently zeroing every Stage-2 target (regression: this collapsed the
secondary species to a single PDG index during training)."""
data = _minimal_step_data(3)
data["n_sec"] = np.array([0, 2, 1], dtype=np.int32) # secondaries, but no lists
data = _step_data_no_sec_lists(np.array([0, 2, 1]))
pdg_map, mat_map = {11: 0}, {"PbWO4": 0}
with pytest.raises(ValueError, match="per-secondary columns"):
@@ -298,7 +305,7 @@ def test_build_features_require_secondaries_raises_when_lists_missing():
def test_build_features_require_secondaries_ok_when_no_secondaries():
"""require_secondaries only fires when secondaries actually exist; a file
with n_sec == 0 everywhere (e.g. Stage-1-only) must still load."""
data = _minimal_step_data(3) # n_sec all zero
data = _step_data_no_sec_lists(np.zeros(3, dtype=np.int32))
pdg_map, mat_map = {11: 0}, {"PbWO4": 0}
_, _, _, _, sec_cont, sec_pdg_idx, *_ = build_features(