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