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
@@ -38,9 +38,7 @@ def per_event(file: str) -> dict:
E0 = float(primary_E[0])
real = pe["real_total_edep"].to_numpy()
gen = pe["gen_total_edep"].to_numpy()
n_steps = (
pl.scan_parquet(file).select(pl.len()).collect(engine="streaming").item()
)
n_steps = pl.scan_parquet(file).select(pl.len()).collect(engine="streaming").item()
return {
"E0": E0,
"n_events": pe.height,
@@ -78,6 +76,8 @@ fmt_row("gen mean/E0", lambda r: f"{r['gen'].mean() / r['E0']:.4f}")
fmt_row("gen max/E0", lambda r: f"{r['gen'].max() / r['E0']:.4f}")
fmt_row("gen p99/E0", lambda r: f"{np.quantile(r['gen'], 0.99) / r['E0']:.4f}")
fmt_row("frac events gen>E0", lambda r: f"{np.mean(r['gen'] > r['E0']):.4f}")
def disp_ratio(r):
return (r["gen"].std() / r["gen"].mean()) / (r["real"].std() / r["real"].mean())
@@ -92,21 +92,31 @@ E0, real_tot, gen_tot = r["E0"], r["real"], r["gen"]
fig, ax = plt.subplots(figsize=(6, 4))
edges = _hist_edges(real_tot, gen_tot, bins=50).tolist()
ax.hist(
real_tot, bins=edges, density=True, histtype="step",
real_tot,
bins=edges,
density=True,
histtype="step",
label=f"real (σ/μ={real_tot.std() / real_tot.mean():.3f})",
)
ax.hist(
gen_tot, bins=edges, density=True, histtype="step",
gen_tot,
bins=edges,
density=True,
histtype="step",
label=f"generated (σ/μ={gen_tot.std() / gen_tot.mean():.3f}, "
f"{np.mean(gen_tot > E0):.1%} > E0)",
)
ax.axvline(E0, color="k", linestyle="--", linewidth=1, label=f"incident energy E0={E0:.0f} MeV")
ax.axvline(
E0, color="k", linestyle="--", linewidth=1, label=f"incident energy E0={E0:.0f} MeV"
)
ax.set_yscale("log")
ax.set_xlabel("total deposited energy per event [MeV]")
ax.set_title("20 ODE steps")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX20}-event-total-energy-vs-E0.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX20}-event-total-energy-vs-E0.png", dpi=150, bbox_inches="tight"
)
fig, ax = plt.subplots(figsize=(6, 4))
ratio_real = real_tot / E0
@@ -129,12 +139,19 @@ all_arrays = [results["10-step (baseline)"]["real"]] + [
]
edges = _hist_edges(*all_arrays, bins=60).tolist()
ax.hist(
results["20-step"]["real"], bins=edges, density=True, histtype="step",
color="k", label="real",
results["20-step"]["real"],
bins=edges,
density=True,
histtype="step",
color="k",
label="real",
)
for name, r2 in results.items():
ax.hist(
r2["gen"], bins=edges, density=True, histtype="step",
r2["gen"],
bins=edges,
density=True,
histtype="step",
label=f"gen {name} ({np.mean(r2['gen'] > r2['E0']):.1%} > E0)",
)
ax.axvline(E0, color="gray", linestyle="--", linewidth=1, label=f"E0={E0:.0f} MeV")
@@ -143,6 +160,8 @@ ax.set_xlabel("total deposited energy per event [MeV]")
ax.set_title("Generated event energy: 10 vs 20 ODE steps")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX20}-compare-event-total-energy.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX20}-compare-event-total-energy.png", dpi=150, bbox_inches="tight"
)
print("DONE")
+1 -3
View File
@@ -66,6 +66,4 @@ print(
f"{'SUM':<14}{tot['10-step']:>14.5f}{tot['20-step']:>14.5f}"
f"{tot['20-step'] / tot['10-step']:>14.2f}"
)
print(
f"{'MEAN':<14}{tot['10-step'] / 9:>14.5f}{tot['20-step'] / 9:>14.5f}"
)
print(f"{'MEAN':<14}{tot['10-step'] / 9:>14.5f}{tot['20-step'] / 9:>14.5f}")
+6 -2
View File
@@ -77,12 +77,16 @@ ax.hist(
label=f"generated (σ/μ={gen_tot.std() / gen_tot.mean():.3f}, "
f"{np.mean(gen_tot > E0):.1%} > E0)",
)
ax.axvline(E0, color="k", linestyle="--", linewidth=1, label=f"incident energy E0={E0:.0f} MeV")
ax.axvline(
E0, color="k", linestyle="--", linewidth=1, label=f"incident energy E0={E0:.0f} MeV"
)
ax.set_yscale("log")
ax.set_xlabel("total deposited energy per event [MeV]")
ax.legend(fontsize=8)
fig.tight_layout()
fig.savefig(OUT / f"{PREFIX}-event-total-energy-vs-E0.png", dpi=150, bbox_inches="tight")
fig.savefig(
OUT / f"{PREFIX}-event-total-energy-vs-E0.png", dpi=150, bbox_inches="tight"
)
print("=== plot: total edep / incident energy ratio ===")
fig, ax = plt.subplots(figsize=(6, 4))
+3 -1
View File
@@ -13,7 +13,9 @@ from pathlib import Path
import giant.analysis as a
rollout_file = sys.argv[1] if len(sys.argv) > 1 else "rollout.parquet"
reference_file = sys.argv[2] if len(sys.argv) > 2 and sys.argv[2] not in ("", "-") else None
reference_file = (
sys.argv[2] if len(sys.argv) > 2 and sys.argv[2] not in ("", "-") else None
)
OUT = Path(sys.argv[3]) if len(sys.argv) > 3 else Path(".")
print(f"=== computing rollout observables: {rollout_file} ===")
File diff suppressed because one or more lines are too long
+525 -65
View File
@@ -51,6 +51,21 @@ how much of the total energy/length, see `pdg_contribution_table_pl`::
plot_pdg_energy_share(table)
plot_pdg_length_share(table)
To run the tiers 1-3 checks above against a full `giant rollout` shower
instead of one-step-ahead predict output, use `load_rollout_vs_truth` in place
of `load_predicted_local` — it builds the same `SampleCollection` from a
`giant rollout` output file ("generated") and any held-out file sharing
`giant train`'s input schema ("real"); the two are independent, unpaired
files (a rollout doesn't replay real events row-for-row), unlike the other
loader's paired pred_*/true_* columns::
from giant.analysis import load_rollout_vs_truth
samples = load_rollout_vs_truth("path/to/rollout.parquet", "path/to/val.parquet")
plot_marginals(samples, group_by="energy")
# ... same plot_kl_bars / plot_correlation_matrices / plot_pairwise /
# plot_direction_alignment / plot_constraint_violations as above.
Four tiers of checks, building on the aggregate marginal/KL check in
`giant.validate.validate_marginals`:
@@ -92,12 +107,17 @@ from giant.constants import (
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
ROLLOUT_COORD_VALUE,
TERM_ENERGY_CUTOFF,
TERM_ESCAPED,
TERM_MAX_STEPS,
TERM_UNKNOWN_PDG,
)
from giant.data.transforms import (
energy_simplex_decode,
inv_log_transform,
local_frame_rotation,
reconstruct_post_pos,
travel_direction,
)
from giant.validate import _histogram_kl
@@ -179,7 +199,16 @@ class SampleCollection:
pdg: np.ndarray # (N,) raw PDG codes
material: np.ndarray # (N,) raw material names
real_raw: np.ndarray # (N, 9) denormalized + delogged real targets
gen_raw: np.ndarray # (N, 9) denormalized + delogged generated targets
gen_raw: np.ndarray # (M, 9) denormalized + delogged generated targets
# Set these three when real_raw/gen_raw are *unpaired* — independent files
# with their own row counts and conditioning (e.g. `load_rollout_vs_truth`,
# comparing a `giant rollout` shower against a held-out truth file) — rather
# than the row-for-row pred_*/true_* pairing `load_predicted_local` produces.
# None (the default) means "same as the real-side field above", which
# reproduces the original paired behavior exactly.
cond_cont_raw_gen: np.ndarray | None = None
pdg_gen: np.ndarray | None = None
material_gen: np.ndarray | None = None
def _check_predict_metadata(path: Path) -> None:
@@ -294,28 +323,66 @@ def load_predicted_local(
# ---------------------------------------------------------------------------
def _gen_pdg(collection: SampleCollection) -> np.ndarray:
return collection.pdg if collection.pdg_gen is None else collection.pdg_gen
def _gen_material(collection: SampleCollection) -> np.ndarray:
return (
collection.material
if collection.material_gen is None
else collection.material_gen
)
def _gen_cond(collection: SampleCollection) -> np.ndarray:
return (
collection.cond_cont_raw
if collection.cond_cont_raw_gen is None
else collection.cond_cont_raw_gen
)
def _group_labels(
collection: SampleCollection,
group_by: str | None,
n_energy_bins: int,
) -> list[tuple[str, np.ndarray]]:
n = len(collection.pdg)
) -> list[tuple[str, np.ndarray, np.ndarray]]:
"""Per-group `(label, mask_real, mask_gen)` triples.
`mask_real` indexes `collection.real_raw` (via `pdg`/`material`/
`cond_cont_raw`); `mask_gen` indexes `collection.gen_raw` via the `*_gen`
fields when set (unpaired real/gen — see `SampleCollection`), or the same
real-side arrays otherwise, which collapses to a single shared mask — the
original paired behavior (real/gen same length, row-for-row).
"""
gen_pdg, gen_material, gen_cond = (
_gen_pdg(collection),
_gen_material(collection),
_gen_cond(collection),
)
n_real, n_gen = len(collection.pdg), len(gen_pdg)
if group_by is None:
return [("all", np.ones(n, dtype=bool))]
return [("all", np.ones(n_real, dtype=bool), np.ones(n_gen, dtype=bool))]
if group_by == "pdg":
return [(f"pdg={v}", collection.pdg == v) for v in np.unique(collection.pdg)]
values = np.unique(np.concatenate([collection.pdg, gen_pdg]))
return [(f"pdg={v}", collection.pdg == v, gen_pdg == v) for v in values]
if group_by == "material":
values = np.unique(np.concatenate([collection.material, gen_material]))
return [
(f"material={v}", collection.material == v)
for v in np.unique(collection.material)
(f"material={v}", collection.material == v, gen_material == v)
for v in values
]
if group_by == "energy":
pre_E = collection.cond_cont_raw[:, 3]
edges = np.quantile(pre_E, np.linspace(0, 1, n_energy_bins + 1))
real_E, gen_E = collection.cond_cont_raw[:, 3], gen_cond[:, 3]
edges = np.quantile(
np.concatenate([real_E, gen_E]), np.linspace(0, 1, n_energy_bins + 1)
)
edges[-1] += 1e-6
bin_idx = np.digitize(pre_E, edges[1:-1])
real_bin = np.digitize(real_E, edges[1:-1])
gen_bin = np.digitize(gen_E, edges[1:-1])
return [
(f"E∈[{edges[i]:.3g},{edges[i + 1]:.3g})", bin_idx == i)
(f"E∈[{edges[i]:.3g},{edges[i + 1]:.3g})", real_bin == i, gen_bin == i)
for i in range(n_energy_bins)
]
raise ValueError(f"unknown group_by={group_by!r}")
@@ -334,16 +401,19 @@ def marginal_table(
failure modes hidden by the aggregate surface at the top.
"""
rows = []
for label, mask in _group_labels(collection, group_by, n_energy_bins):
if mask.sum() < 2:
for label, mask_real, mask_gen in _group_labels(
collection, group_by, n_energy_bins
):
if mask_real.sum() < 2 or mask_gen.sum() < 2:
continue
real, gen = collection.real_raw[mask], collection.gen_raw[mask]
real, gen = collection.real_raw[mask_real], collection.gen_raw[mask_gen]
for j, name in enumerate(RAW_TARGET_NAMES):
rows.append(
{
"group": label,
"dim": name,
"n": int(mask.sum()),
"n": int(mask_real.sum()),
"n_gen": int(mask_gen.sum()),
"real_mean": real[:, j].mean(),
"gen_mean": gen[:, j].mean(),
"real_std": real[:, j].std(),
@@ -649,8 +719,8 @@ def plot_marginals(
squeeze=False,
figsize=(figsize_per_axis[0] * n_cols, figsize_per_axis[1] * n_rows),
)
for row, (label, mask) in enumerate(groups):
real, gen = collection.real_raw[mask], collection.gen_raw[mask]
for row, (label, mask_real, mask_gen) in enumerate(groups):
real, gen = collection.real_raw[mask_real], collection.gen_raw[mask_gen]
for col, j in enumerate(dim_idx):
ax = axes[row][col]
edges = _hist_edges(real[:, j], gen[:, j], bins=bins)
@@ -818,8 +888,6 @@ def plot_pairwise(
"""
pairs = pairs or _DEFAULT_PAIRS
rng = np.random.default_rng(seed)
n = len(collection.pdg)
idx = rng.choice(n, size=min(n_sample, n), replace=False)
fig, axes = plt.subplots(2, len(pairs), squeeze=False, figsize=(4 * len(pairs), 7))
for col, (a, b) in enumerate(pairs):
@@ -827,6 +895,8 @@ def plot_pairwise(
for row, (data, title) in enumerate(
[(collection.real_raw, "real"), (collection.gen_raw, "generated")]
):
n = len(data)
idx = rng.choice(n, size=min(n_sample, n), replace=False)
ax = axes[row][col]
ax.scatter(data[idx, ia], data[idx, ib], s=3, alpha=0.3)
ax.set_xlabel(a)
@@ -1706,16 +1776,43 @@ def plot_pdg_length_share(table: pl.DataFrame, max_slices: int = 6):
# same event_ids can be overlaid against a real reference computed elsewhere.
_ROLLOUT_COLS = [
"event_id", "track_id", "termination_reason",
"pre_x", "pre_y", "pre_z", "pre_dx", "pre_dy", "pre_dz", "pre_E",
"post_x", "post_y", "post_z", "edep",
"event_id",
"track_id",
"termination_reason",
"pre_x",
"pre_y",
"pre_z",
"pre_dx",
"pre_dy",
"pre_dz",
"pre_E",
"post_x",
"post_y",
"post_z",
"edep",
]
def _check_rollout_metadata(path: Path) -> None:
"""Raise if `path` carries coord metadata that isn't `ROLLOUT_COORD_VALUE`.
A missing tag (older rollout output, predating tagging) is let through
silently, matching `giant predict`/`giant rollout`'s own leniency; a tag
that's present but wrong is a real mismatch.
"""
metadata = pq.read_schema(path).metadata or {}
coord = metadata.get(PREDICT_COORD_METADATA_KEY.encode())
if coord is not None and coord.decode() != ROLLOUT_COORD_VALUE:
raise ValueError(
f"{path} is not a rollout file (coord={coord.decode()!r}); "
"expected a `giant rollout` output"
)
@dataclass
class RolloutObservables:
event_table: pd.DataFrame # one row per event_id (mm/MeV)
depth_edges: np.ndarray # (depth_bins+1,) mm, along shower axis
depth_edges: np.ndarray # (depth_bins+1,) mm, along shower axis
transverse_edges: np.ndarray # (transverse_bins+1,) mm, perpendicular
depth_profile: np.ndarray # (depth_bins,) mean edep/event/bin, MeV
depth_profile_std: np.ndarray
@@ -1723,31 +1820,21 @@ class RolloutObservables:
transverse_profile_std: np.ndarray
def compute_rollout_observables(
path: str | Path,
depth_bins: int = 20,
transverse_bins: int = 20,
) -> RolloutObservables:
"""Compute shower observables from a `giant rollout` steps parquet.
def _event_axis_depth_transverse(
df: pd.DataFrame, depth_bins: int, transverse_bins: int
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Shared core of `compute_rollout_observables`/`compute_truth_observables`.
Per event the primary entry step (highest-`pre_E` row) fixes the shower axis
and entry point; every step's deposit (`edep` at `post_pos`) is projected onto
depth-along-axis and transverse-distance-from-axis, then binned. Returns
per-event totals plus dataset-mean longitudinal/transverse profiles.
`df` needs `event_id`/`pre_x,y,z`/`pre_dx,dy,dz`/`pre_E`/`post_x,y,z`/`edep`
— both a `giant rollout` file and a raw truth-schema steps file carry these.
Per event, the highest-`pre_E` row fixes the shower axis/entry point; every
row's `edep` at `post_pos` is projected onto depth-along-axis and
transverse-distance-from-axis, then binned.
Returns `(ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev)`
— `row_ev` maps each row of `df` to its event's index in `ev_ids`;
`depth_ev`/`trans_ev` are `(n_events, bins)` per-event-per-bin edep sums.
"""
pf = pq.ParquetFile(Path(path))
coord = (pf.schema_arrow.metadata or {}).get(
PREDICT_COORD_METADATA_KEY.encode()
)
if coord is not None and coord.decode() != ROLLOUT_COORD_VALUE:
raise ValueError(
f"{path} is not a rollout file (coord={coord.decode()!r}); "
"expected a `giant rollout` output"
)
df = pd.read_parquet(path, columns=_ROLLOUT_COLS)
# Per-event entry point + shower axis from the highest-pre_E row.
entry_idx = df.groupby("event_id")["pre_E"].idxmax()
entry = df.loc[entry_idx].set_index("event_id")
ax = entry[["pre_dx", "pre_dy", "pre_dz"]].to_numpy(dtype=np.float64).copy()
@@ -1768,11 +1855,14 @@ def compute_rollout_observables(
if not (depth_hi - depth_lo > 1e-6 * max(abs(depth_hi), 1.0)):
depth_lo, depth_hi = depth_lo - 0.5, depth_hi + 0.5
depth_edges = np.linspace(depth_lo, depth_hi, depth_bins + 1)
transverse_edges = np.linspace(0.0, max(np.quantile(transverse, 0.999), 1e-6),
transverse_bins + 1)
transverse_edges = np.linspace(
0.0, max(np.quantile(transverse, 0.999), 1e-6), transverse_bins + 1
)
d_bin = np.clip(np.digitize(depth, depth_edges) - 1, 0, depth_bins - 1)
t_bin = np.clip(np.digitize(transverse, transverse_edges) - 1, 0, transverse_bins - 1)
t_bin = np.clip(
np.digitize(transverse, transverse_edges) - 1, 0, transverse_bins - 1
)
# Per-(event, bin) edep sums, then mean/std across events.
depth_ev = np.zeros((n_events, depth_bins))
@@ -1780,6 +1870,40 @@ def compute_rollout_observables(
np.add.at(depth_ev, (row_ev, d_bin), edep)
np.add.at(trans_ev, (row_ev, t_bin), edep)
return ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev
def _centroid_depth(depth_ev: np.ndarray, depth_edges: np.ndarray) -> np.ndarray:
"""Energy-weighted centroid depth per event, from the binned edep sums."""
bin_centers = 0.5 * (depth_edges[:-1] + depth_edges[1:])
tot = depth_ev.sum(axis=1)
centroid = (depth_ev * bin_centers).sum(axis=1) / np.where(tot > 0, tot, 1.0)
return np.where(tot > 0, centroid, 0.0)
def compute_rollout_observables(
path: str | Path,
depth_bins: int = 20,
transverse_bins: int = 20,
) -> RolloutObservables:
"""Compute shower observables from a `giant rollout` steps parquet.
Per event the primary entry step (highest-`pre_E` row) fixes the shower axis
and entry point; every step's deposit (`edep` at `post_pos`) is projected onto
depth-along-axis and transverse-distance-from-axis, then binned. Returns
per-event totals plus dataset-mean longitudinal/transverse profiles.
See `compute_truth_observables` for the truth-schema counterpart, shaped to
plug straight into this function's own output as a `plot_rollout_*`
`reference=` overlay.
"""
_check_rollout_metadata(Path(path))
df = pd.read_parquet(path, columns=_ROLLOUT_COLS)
ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev = (
_event_axis_depth_transverse(df, depth_bins, transverse_bins)
)
# Per-event scalar table.
is_leak = df["termination_reason"].to_numpy() == TERM_ESCAPED
per_ev = df.groupby("event_id").agg(
@@ -1789,14 +1913,11 @@ def compute_rollout_observables(
per_ev["n_tracks"] = df.groupby("event_id")["track_id"].nunique()
per_ev["leaked_E"] = (
df.assign(_leak=np.where(is_leak, df["pre_E"], 0.0))
.groupby("event_id")["_leak"].sum()
.groupby("event_id")["_leak"]
.sum()
)
# Energy-weighted centroid depth per event (from the binned sums).
bin_centers = 0.5 * (depth_edges[:-1] + depth_edges[1:])
tot = depth_ev.sum(axis=1)
centroid = (depth_ev * bin_centers).sum(axis=1) / np.where(tot > 0, tot, 1.0)
per_ev = per_ev.reindex(ev_ids)
per_ev["centroid_depth"] = np.where(tot > 0, centroid, 0.0)
per_ev["centroid_depth"] = _centroid_depth(depth_ev, depth_edges)
per_ev.index.name = "event_id"
return RolloutObservables(
@@ -1810,18 +1931,105 @@ def compute_rollout_observables(
)
_TRUTH_EVENT_COLS = [
"event_id",
"pre_x",
"pre_y",
"pre_z",
"pre_dx",
"pre_dy",
"pre_dz",
"pre_E",
"post_x",
"post_y",
"post_z",
"edep",
"step_length",
]
@dataclass
class TruthObservables:
"""Truth-schema event-level observables, shaped for `plot_rollout_*`'s `reference=`.
The truth-side counterpart to `RolloutObservables`, computed directly from
a raw truth-schema steps file (the same file `load_rollout_vs_truth` takes
as `truth_path`) rather than needing a separate paired `giant predict
--coord local` file covering the same events. Field names mirror
`EventObservables`'s `real_*` convention — `plot_rollout_longitudinal`/
`plot_rollout_transverse`/`plot_rollout_total_energy` already read exactly
these names off their `reference` argument — since this object is only
ever used as a reference overlay, never plotted standalone.
"""
event_table: pd.DataFrame # one row per event_id (mm/MeV)
depth_edges: np.ndarray
transverse_edges: np.ndarray
real_depth_profile: np.ndarray
real_depth_profile_std: np.ndarray
real_transverse_profile: np.ndarray
real_transverse_profile_std: np.ndarray
def compute_truth_observables(
path: str | Path,
depth_bins: int = 20,
transverse_bins: int = 20,
) -> TruthObservables:
"""Compute shower observables from a raw truth-schema steps parquet.
Lets the same held-out truth file used by `load_rollout_vs_truth` (Tier
1-3) also supply the Tier 4 real-shower overlay — pass the result as
`reference=` to `plot_rollout_longitudinal`/`plot_rollout_transverse`/
`plot_rollout_total_energy` — without needing a separate paired `giant
predict --coord local` file for the same events. Same per-event
axis/depth/transverse construction as `compute_rollout_observables` (see
`_event_axis_depth_transverse`), just without the `track_id`/
`termination_reason` columns a rollout file (but not a truth file) has, so
there's no `n_tracks`/`leaked_E` in `event_table`.
"""
df = pd.read_parquet(path, columns=_TRUTH_EVENT_COLS)
ev_ids, row_ev, depth_edges, transverse_edges, depth_ev, trans_ev = (
_event_axis_depth_transverse(df, depth_bins, transverse_bins)
)
per_ev = df.groupby("event_id").agg(
n_steps=("edep", "size"),
real_total_edep=("edep", "sum"),
real_total_length=("step_length", "sum"),
)
per_ev = per_ev.reindex(ev_ids)
per_ev["real_centroid_depth"] = _centroid_depth(depth_ev, depth_edges)
per_ev.index.name = "event_id"
return TruthObservables(
event_table=per_ev.reset_index(),
depth_edges=depth_edges,
transverse_edges=transverse_edges,
real_depth_profile=depth_ev.mean(0),
real_depth_profile_std=depth_ev.std(0),
real_transverse_profile=trans_ev.mean(0),
real_transverse_profile_std=trans_ev.std(0),
)
def plot_rollout_longitudinal(obs: RolloutObservables, reference=None):
"""Mean edep vs depth along the shower axis; optional real-reference overlay.
`reference` may be an `EventObservables` (its `real_depth_profile`) to overlay
the real showers seeded from the same events.
`reference` may be an `EventObservables` or a `TruthObservables` (either
exposes `real_depth_profile`) to overlay the real showers seeded from the
same events — `compute_truth_observables` builds the latter directly from
a raw truth-schema file, with no paired predict-schema file needed.
"""
centers = 0.5 * (obs.depth_edges[:-1] + obs.depth_edges[1:])
fig, ax = plt.subplots(figsize=(7, 4))
ax.plot(centers, obs.depth_profile, label="rollout", color="C0")
ax.fill_between(
centers, obs.depth_profile - obs.depth_profile_std,
obs.depth_profile + obs.depth_profile_std, alpha=0.2, color="C0",
centers,
obs.depth_profile - obs.depth_profile_std,
obs.depth_profile + obs.depth_profile_std,
alpha=0.2,
color="C0",
)
if reference is not None:
rc = 0.5 * (reference.depth_edges[:-1] + reference.depth_edges[1:])
@@ -1840,8 +2048,11 @@ def plot_rollout_transverse(obs: RolloutObservables, reference=None):
fig, ax = plt.subplots(figsize=(7, 4))
ax.plot(centers, obs.transverse_profile, label="rollout", color="C0")
ax.fill_between(
centers, obs.transverse_profile - obs.transverse_profile_std,
obs.transverse_profile + obs.transverse_profile_std, alpha=0.2, color="C0",
centers,
obs.transverse_profile - obs.transverse_profile_std,
obs.transverse_profile + obs.transverse_profile_std,
alpha=0.2,
color="C0",
)
if reference is not None:
rc = 0.5 * (reference.transverse_edges[:-1] + reference.transverse_edges[1:])
@@ -1858,14 +2069,263 @@ def plot_rollout_transverse(obs: RolloutObservables, reference=None):
def plot_rollout_total_energy(obs: RolloutObservables, bins: int = 50, reference=None):
"""Distribution of total deposited energy per shower."""
fig, ax = plt.subplots(figsize=(7, 4))
ax.hist(obs.event_table["total_edep"], bins=bins, histtype="step",
label="rollout", color="C0")
ax.hist(
obs.event_table["total_edep"],
bins=bins,
histtype="step",
label="rollout",
color="C0",
)
if reference is not None:
ax.hist(reference.event_table["real_total_edep"].to_numpy(), bins=bins,
histtype="step", label="real", color="k", ls="--")
ax.hist(
reference.event_table["real_total_edep"].to_numpy(),
bins=bins,
histtype="step",
label="real",
color="k",
ls="--",
)
ax.set_xlabel("total E_dep / event [MeV]")
ax.set_ylabel("events")
ax.set_title("Total deposited energy")
ax.legend()
fig.tight_layout()
return fig
# ── Rollout vs. held-out truth (Tier 1-3, unpaired) ────────────────────────────
#
# `compute_rollout_observables` above covers Tier 4 (event-level) comparisons.
# This builds a `SampleCollection` instead, so the Tier 1-3 diagnostics
# (plot_marginals, plot_kl_bars, correlation_matrices, plot_pairwise,
# direction_alignment, constraint_report) also work on a rollout: `giant
# rollout` output and a training-input-schema truth file (see
# `giant.data.loader.load_steps`) both carry pre_*/post_*/edep/step_length in
# the same physical, world-frame units, so both sides decode into
# RAW_TARGET_NAMES space via the same local_frame_rotation/travel_direction
# construction `giant.data.transforms.build_features` uses for `target_s1` —
# just skipping the log/ALR encode step, since neither file needs it decoded.
#
# Unlike `load_predicted_local`, the two files are independent rather than
# row-for-row paired (a rollout doesn't replay real events step-by-step), so
# real_raw/gen_raw may have different lengths; `SampleCollection`'s `*_gen`
# fields carry the rollout side's own pdg/material/conditioning for grouping.
_WORLD_FRAME_STEP_COLS = [
"pdg",
"material",
"layer_id",
"pre_x",
"pre_y",
"pre_z",
"pre_E",
"pre_dx",
"pre_dy",
"pre_dz",
"post_x",
"post_y",
"post_z",
"post_E",
"post_dx",
"post_dy",
"post_dz",
"edep",
"step_length",
]
def _world_frame_raw_targets(batch: pl.DataFrame) -> np.ndarray:
"""(N, 9) RAW_TARGET_NAMES array from a world-frame steps batch.
See the module note above `_WORLD_FRAME_STEP_COLS` — both `load_rollout_vs_truth`
inputs share this schema.
"""
pre_pos = batch.select(["pre_x", "pre_y", "pre_z"]).to_numpy().astype(np.float32)
pre_dir = batch.select(["pre_dx", "pre_dy", "pre_dz"]).to_numpy().astype(np.float32)
post_pos = (
batch.select(["post_x", "post_y", "post_z"]).to_numpy().astype(np.float32)
)
post_dir = (
batch.select(["post_dx", "post_dy", "post_dz"]).to_numpy().astype(np.float32)
)
pre_E = batch["pre_E"].to_numpy().astype(np.float32)
post_E = batch["post_E"].to_numpy().astype(np.float32)
post_dir_local = local_frame_rotation(pre_dir, post_dir)
travel_dir_local = local_frame_rotation(
pre_dir, travel_direction(pre_pos, post_pos)
)
return np.column_stack(
[
batch["step_length"].to_numpy().astype(np.float32),
pre_E - post_E,
batch["edep"].to_numpy().astype(np.float32),
post_dir_local,
travel_dir_local,
]
).astype(np.float32)
def _world_frame_cond(batch: pl.DataFrame) -> np.ndarray:
"""(N, 9) `_COND_CONT_COLS`-layout array from a world-frame steps batch.
`n_sec` comes from `child_track_ids` (truth schema) or `n_sec_pred`
(rollout schema) — whichever the batch has; kept only for shape parity
with `load_predicted_local`'s `cond_cont_raw` — nothing downstream in this
module groups by it, only `pre_E` (index 3, for `group_by="energy"`).
"""
if "n_sec_pred" in batch.columns:
n_sec = batch["n_sec_pred"].to_numpy().astype(np.float32)
elif "child_track_ids" in batch.columns:
n_sec = batch["child_track_ids"].list.len().to_numpy().astype(np.float32)
else:
n_sec = np.zeros(batch.height, dtype=np.float32)
return np.column_stack(
[
batch.select(["pre_x", "pre_y", "pre_z"]).to_numpy(),
batch["pre_E"].to_numpy(),
batch.select(["pre_dx", "pre_dy", "pre_dz"]).to_numpy(),
batch["layer_id"].to_numpy().astype(np.float32),
n_sec,
]
).astype(np.float32)
def _existing_columns(
source: str | Path | pl.LazyFrame, wanted: list[str]
) -> list[str]:
if isinstance(source, pl.LazyFrame):
available = set(source.collect_schema().names())
else:
available = set(pq.ParquetFile(Path(source)).schema_arrow.names)
return [c for c in wanted if c in available]
# `rollout.py`'s `_terminal_rows` writes one synthetic bookkeeping row per track
# for these termination reasons (escaped/unknown_pdg/energy_cutoff/max_steps):
# step_length=0, post_pos=pre_pos, and — for every reason but escaped — the
# track's *entire remaining pre_E* dumped into `edep` in one row, so the shower's
# total energy still conserves. These aren't steps in any physical sense (truth
# data has no equivalent), so mixing them into a per-step real-vs-generated
# comparison would inject a spurious step_length=0 spike and roughly double the
# apparent mean edep purely from bookkeeping, not model behavior. Real generated
# steps carry "" (continuing) or `TERM_NATURAL_END` (the track's last real step,
# which does have genuine step_length/edep) and are kept.
_SYNTHETIC_ROLLOUT_TERMINATION_REASONS = frozenset(
{TERM_ESCAPED, TERM_UNKNOWN_PDG, TERM_ENERGY_CUTOFF, TERM_MAX_STEPS}
)
def _load_world_frame_side(
source: str | Path | pl.LazyFrame,
sample_frac: float,
seed: int,
batch_size: int,
extra_cols: list[str],
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""Stream one world-frame steps file/LazyFrame into (raw9, cond9, pdg, material).
`extra_cols` (e.g. `["n_sec_pred", "termination_reason"]` or
`["child_track_ids"]`) are included only when present, so the same helper
serves both the rollout and truth schemas without either needing the
other's columns. When `termination_reason` is present (the rollout side),
rows carrying one of `_SYNTHETIC_ROLLOUT_TERMINATION_REASONS` are dropped
before decoding — see that constant's docstring for why.
"""
columns = _WORLD_FRAME_STEP_COLS + _existing_columns(source, extra_cols)
threshold = int(sample_frac * 2**32) if sample_frac < 1.0 else None
raw_parts, cond_parts, pdg_parts, mat_parts = [], [], [], []
offset = 0
for batch in _iter_predicted_local_batches(source, columns, batch_size):
n = batch.height
if threshold is not None:
row_idx = pl.arange(offset, offset + n, eager=True).cast(pl.UInt32)
batch = batch.filter((row_idx.hash(seed=seed) % 2**32) < threshold)
offset += n
if "termination_reason" in batch.columns:
batch = batch.filter(
~pl.col("termination_reason").is_in(
list(_SYNTHETIC_ROLLOUT_TERMINATION_REASONS)
)
)
if batch.height == 0:
continue
raw_parts.append(_world_frame_raw_targets(batch))
cond_parts.append(_world_frame_cond(batch))
pdg_parts.append(batch["pdg"].to_numpy())
mat_parts.append(batch["material"].to_numpy())
if not raw_parts:
raise ValueError(f"{source}: no rows survived (sample_frac={sample_frac})")
return (
np.concatenate(raw_parts, axis=0),
np.concatenate(cond_parts, axis=0),
np.concatenate(pdg_parts, axis=0),
np.concatenate(mat_parts, axis=0),
)
def load_rollout_vs_truth(
rollout_path: str | Path | pl.LazyFrame,
truth_path: str | Path | pl.LazyFrame,
sample_frac: float = 1.0,
seed: int = 0,
batch_size: int = 1_000_000,
) -> SampleCollection:
"""Build a `SampleCollection` comparing a `giant rollout` shower to a truth file.
`truth_path` — any file sharing `giant train`'s input schema (real
miniCaloSim steps, e.g. a held-out/val parquet) — is treated as "real";
`rollout_path` (`giant rollout` output) is treated as "generated". The two
are independent files (not row-for-row paired, since a rollout doesn't
replay real events step-by-step): `real_raw`/`gen_raw` may have different
lengths, and `pdg`/`material`/`cond_cont_raw` are computed separately per
side (`SampleCollection`'s `*_gen` fields) — `_group_labels` (used by
`marginal_table`/`plot_marginals`/`plot_kl_bars`) builds independent masks
for each. `correlation_matrices`, `direction_alignment`, and
`constraint_report` never paired real/gen row-for-row to begin with, so
they need no special handling here.
Every raw target dim is decoded from the world-frame `pre_*`/`post_*`/
`edep`/`step_length` columns both files share — see the module note above
`_WORLD_FRAME_STEP_COLS`. `sample_frac`/`seed`/`batch_size` behave as in
`load_predicted_local`, applied independently to each file.
The rollout side drops synthetic termination-bookkeeping rows (see
`_SYNTHETIC_ROLLOUT_TERMINATION_REASONS`) before decoding, since those
aren't real generated steps. One remaining case where `edep` still isn't
perfectly analogous between the two files: `decode_secondaries` rescales
the valid secondary slots to sum to exactly `e_sec` whenever `n_sec > 0`
(see that function's docstring), but when Stage 1 predicts a nonzero
`e_sec` while Stage 2's `n_sec` head predicts 0 secondaries, there's no
slot to carry that budget at all — `rollout.py` deposits it into that
step's `edep` instead, which truth's Geant4-recorded `edep` never does.
That specific disagreement between the two Stage-1/Stage-2 heads is rare
but not otherwise fixable at decode time.
"""
if not (0 < sample_frac <= 1):
raise ValueError(f"sample_frac must be in (0, 1], got {sample_frac}")
if not isinstance(rollout_path, pl.LazyFrame):
_check_rollout_metadata(Path(rollout_path))
gen_raw, gen_cond, gen_pdg, gen_material = _load_world_frame_side(
rollout_path,
sample_frac,
seed,
batch_size,
extra_cols=["n_sec_pred", "termination_reason"],
)
real_raw, real_cond, real_pdg, real_material = _load_world_frame_side(
truth_path, sample_frac, seed, batch_size, extra_cols=["child_track_ids"]
)
return SampleCollection(
cond_cont_raw=real_cond,
pdg=real_pdg,
material=real_material,
real_raw=real_raw,
gen_raw=gen_raw,
cond_cont_raw_gen=gen_cond,
pdg_gen=gen_pdg,
material_gen=gen_material,
)
+13 -11
View File
@@ -664,9 +664,7 @@ def predict(
typer.echo(f"wrote {total:,} rows → {out}")
def _seed_from_data(
files: list[Path], n_events: int | None
) -> dict[str, np.ndarray]:
def _seed_from_data(files: list[Path], n_events: int | None) -> dict[str, np.ndarray]:
"""Pick each event's primary entry state (argmax-pre_E row) as a shower seed.
Streams conditioning columns and keeps the highest-pre_E step per event_id
@@ -710,25 +708,30 @@ def rollout(
Path, typer.Argument(help="Parquet file/dir to seed showers from (real events)")
],
checkpoint: Annotated[
Path, typer.Option("--checkpoint", "-c", help="Checkpoint .pt (best.pt/last.pt)")
Path,
typer.Option("--checkpoint", "-c", help="Checkpoint .pt (best.pt/last.pt)"),
],
geometry: Annotated[
Path,
typer.Option(
"--geometry", "-g", help="Geometry oracle .pkl (dwarf build-geometry-oracle)"
"--geometry",
"-g",
help="Geometry oracle .pkl (dwarf build-geometry-oracle)",
),
],
energy_cutoff: Annotated[
float,
typer.Option(
"--energy-cutoff", help="Stop a track when its energy drops below this [MeV]"
"--energy-cutoff",
help="Stop a track when its energy drops below this [MeV]",
),
] = 0.1,
max_steps: Annotated[
int, typer.Option("--max-steps", help="Max steps per individual track")
] = 1000,
steps: Annotated[
int, typer.Option("--steps", "-s", help="Flow matching ODE steps per model call")
int,
typer.Option("--steps", "-s", help="Flow matching ODE steps per model call"),
] = 10,
batch_size: Annotated[
int, typer.Option("--batch-size", "-b", help="Tracks stepped per model forward")
@@ -757,7 +760,8 @@ def rollout(
Optional[Path], typer.Option("--out", "-o", help="Output steps parquet")
] = None,
seed: Annotated[
Optional[int], typer.Option("--seed", help="Torch/numpy seed for reproducibility")
Optional[int],
typer.Option("--seed", help="Torch/numpy seed for reproducibility"),
] = None,
) -> None:
"""Roll the surrogate forward into full showers (autoregressive)."""
@@ -845,9 +849,7 @@ def rollout(
ref_path.write_text(yaml.dump(ref, default_flow_style=False, sort_keys=False))
n_rows = len(records["event_id"])
reasons = Counter(
r for r in records["termination_reason"].tolist() if r
)
reasons = Counter(r for r in records["termination_reason"].tolist() if r)
typer.echo(f"wrote {n_rows:,} step rows → {out}")
typer.echo(f"terminations: {dict(reasons)}")
typer.echo(f"reference: {ref_path}")
+1 -1
View File
@@ -13,7 +13,7 @@ K_MAX = 15
# EMB_DIM (continuous type embedding). EMB_DIM must match DenoisingMLP.emb_dim.
# Default emb_dim=16 → SEC_SLOT_DIM=20.
SEC_SLOT_DIM = 20 # 1 + 3 + 16
EMB_DIM = 16 # must match model emb_dim default
EMB_DIM = 16 # must match model emb_dim default
# Per-slot continuous (non-embedding) width: stick-breaking logit + local dir.
CONT_SLOT_DIM = SEC_SLOT_DIM - EMB_DIM # 4
+1 -3
View File
@@ -60,9 +60,7 @@ def _pad_list_col_int(series: pd.Series, K: int, fill: int = 0) -> np.ndarray:
return out
def _pad_dir_col(
dx: pd.Series, dy: pd.Series, dz: pd.Series, K: int
) -> np.ndarray:
def _pad_dir_col(dx: pd.Series, dy: pd.Series, dz: pd.Series, K: int) -> np.ndarray:
"""Pad three list-valued direction columns → (N, K, 3) float32.
Padding direction defaults to (0,0,1) (forward) so it is a valid unit vector.
+26 -4
View File
@@ -360,6 +360,8 @@ def decode_secondaries(
pdg_map_inv: maps model index PDG code
Returns (sec_E, sec_dir_world, sec_pdg_code, sec_valid) each shape (N, K_MAX).
The valid slots' energies (`sec_E[sec_valid]`, per row) always sum to
exactly `e_sec` see the rescaling below.
"""
N, K, _ = sec_cont.shape
stick_logits = sec_cont[:, :, 0] # (N, K)
@@ -372,15 +374,33 @@ def decode_secondaries(
fractions = 1.0 / (1.0 + np.exp(-stick_logits.astype(np.float64)))
sec_E = np.zeros((N, K), dtype=np.float32)
sec_E = np.zeros((N, K), dtype=np.float64)
e_sec = np.asarray(e_sec, dtype=np.float64)
remaining = e_sec.copy()
for i in range(K):
sec_E[:, i] = (fractions[:, i] * remaining).astype(np.float32)
remaining = np.maximum(remaining - sec_E[:, i].astype(np.float64), 0.0)
sec_E[:, i] = fractions[:, i] * remaining
remaining = np.maximum(remaining - sec_E[:, i], 0.0)
sec_valid = np.arange(K)[None, :] < n_sec[:, None] # (N, K)
# Stick-breaking guarantees sum(sec_E[valid]) <= e_sec (each fraction is in
# [0,1] of an already-shrinking remainder) but rarely hits it exactly, so
# rescale the valid slots by one common per-row factor to close that gap —
# rather than dumping the shortfall into whichever slot happens to be last
# by energy rank, which would let one low-energy secondary balloon and
# distort the shower's topology. This preserves each row's relative split
# across its secondaries and only ever scales up (valid_sum <= e_sec).
# Rows where every valid slot decoded to ~zero (scale undefined) fall back
# to an even split of e_sec across the n_sec valid slots.
sec_E = sec_E * sec_valid
valid_sum = sec_E.sum(axis=1)
degenerate = (valid_sum <= _EPS) & (n_sec > 0)
scale = np.where(valid_sum > _EPS, e_sec / np.maximum(valid_sum, _EPS), 0.0)
sec_E = sec_E * scale[:, None]
even_share = e_sec / np.maximum(n_sec, 1).astype(np.float64)
sec_E = np.where(degenerate[:, None] & sec_valid, even_share[:, None], sec_E)
sec_E = sec_E.astype(np.float32)
sec_dir_world = np.zeros((N, K, 3), dtype=np.float32)
for i in range(K):
valid = sec_valid[:, i]
@@ -496,7 +516,9 @@ def build_features(
mat_idx = np.array([mat_map[str(m)] for m in data["material"]], dtype=np.int64)
cond_cat = np.column_stack([pdg_idx, mat_idx]) # (N, 2)
n_sec_raw = data["n_sec"].astype(np.int64) # (N,) unclamped, for the valid-slot mask
n_sec_raw = data["n_sec"].astype(
np.int64
) # (N,) unclamped, for the valid-slot mask
# Clamp the classification label to K_MAX: the head only has K_MAX+1 classes
# (0..K_MAX), and truncating here mirrors the K_MAX-slot truncation already
# applied to sec_cont/sec_pdg_idx by the loader's list padding. Without this,
+4 -4
View File
@@ -123,7 +123,7 @@ class DenoisingMLP(nn.Module):
cond_cont: torch.Tensor,
cond_cat: torch.Tensor,
) -> torch.Tensor:
t_emb = self.time_emb(t) # (B, time_dim)
t_emb = self.time_emb(t) # (B, time_dim)
c_emb = self.cond_enc(cond_cont, cond_cat) # (B, cond_out_dim)
cond = torch.cat([t_emb, c_emb], dim=-1)
x = self.input_proj(x_t)
@@ -178,9 +178,9 @@ class SecondaryConditionEncoder(nn.Module):
cond_cat: torch.Tensor,
stage1_out: torch.Tensor,
) -> torch.Tensor:
base = self.base(cond_cont, cond_cat) # (B, cond_out_dim)
s1 = self.stage1_proj(stage1_out).tanh() # (B, stage1_proj_dim)
return self.fuse(torch.cat([base, s1], dim=-1)) # (B, out_dim)
base = self.base(cond_cont, cond_cat) # (B, cond_out_dim)
s1 = self.stage1_proj(stage1_out).tanh() # (B, stage1_proj_dim)
return self.fuse(torch.cat([base, s1], dim=-1)) # (B, out_dim)
class SecondaryDecoder(nn.Module):
+1 -1
View File
@@ -123,7 +123,7 @@ def run_train_job(
emb_dim = m.get("emb_dim", EMB_DIM)
# SEC_SLOT_DIM must match constants (1 stick + 3 dir + emb_dim)
assert SEC_SLOT_DIM == 1 + 3 + emb_dim, (
f"SEC_SLOT_DIM={SEC_SLOT_DIM} must equal 1+3+emb_dim={1+3+emb_dim}; "
f"SEC_SLOT_DIM={SEC_SLOT_DIM} must equal 1+3+emb_dim={1 + 3 + emb_dim}; "
"update giant/constants.py if emb_dim changed"
)
+183 -52
View File
@@ -39,10 +39,31 @@ from giant.sample import sample_flow, sample_secondaries, snap_type_to_pdg_idx
# Record columns produced per step / per terminal marker.
_RECORD_KEYS = [
"event_id", "track_id", "parent_id", "generation", "step_no", "pdg",
"pre_x", "pre_y", "pre_z", "pre_E", "pre_dx", "pre_dy", "pre_dz",
"post_x", "post_y", "post_z", "post_E", "post_dx", "post_dy", "post_dz",
"edep", "step_length", "material", "layer_id", "n_sec_pred",
"event_id",
"track_id",
"parent_id",
"generation",
"step_no",
"pdg",
"pre_x",
"pre_y",
"pre_z",
"pre_E",
"pre_dx",
"pre_dy",
"pre_dz",
"post_x",
"post_y",
"post_z",
"post_E",
"post_dx",
"post_dy",
"post_dz",
"edep",
"step_length",
"material",
"layer_id",
"n_sec_pred",
"termination_reason",
]
@@ -88,7 +109,12 @@ class _Recorder:
if chunks:
out[k] = np.concatenate(chunks, axis=0)
else:
out[k] = np.empty(0, dtype=object if k in ("material", "termination_reason") else np.float64)
out[k] = np.empty(
0,
dtype=object
if k in ("material", "termination_reason")
else np.float64,
)
return out
@@ -136,13 +162,26 @@ def _terminal_rows(tr: dict[str, np.ndarray], sel: np.ndarray, reason: str, edep
dir_ = tr["pre_dir"][sel]
n = int(sel.sum())
return dict(
event_id=tr["event_id"][sel], track_id=tr["track_id"][sel],
parent_id=tr["parent_id"][sel], generation=tr["generation"][sel],
step_no=tr["step_in_track"][sel], pdg=tr["pdg"][sel],
pre_x=pos[:, 0], pre_y=pos[:, 1], pre_z=pos[:, 2], pre_E=tr["pre_E"][sel],
pre_dx=dir_[:, 0], pre_dy=dir_[:, 1], pre_dz=dir_[:, 2],
post_x=pos[:, 0], post_y=pos[:, 1], post_z=pos[:, 2],
post_E=np.zeros(n), post_dx=dir_[:, 0], post_dy=dir_[:, 1], post_dz=dir_[:, 2],
event_id=tr["event_id"][sel],
track_id=tr["track_id"][sel],
parent_id=tr["parent_id"][sel],
generation=tr["generation"][sel],
step_no=tr["step_in_track"][sel],
pdg=tr["pdg"][sel],
pre_x=pos[:, 0],
pre_y=pos[:, 1],
pre_z=pos[:, 2],
pre_E=tr["pre_E"][sel],
pre_dx=dir_[:, 0],
pre_dy=dir_[:, 1],
pre_dz=dir_[:, 2],
post_x=pos[:, 0],
post_y=pos[:, 1],
post_z=pos[:, 2],
post_E=np.zeros(n),
post_dx=dir_[:, 0],
post_dy=dir_[:, 1],
post_dz=dir_[:, 2],
edep=np.asarray(edep, dtype=np.float64).reshape(n),
step_length=np.zeros(n),
material=tr.get("_material", np.full(len(sel), "", dtype=object))[sel],
@@ -182,8 +221,11 @@ def rollout(
pdg_emb_weight = stage1_model.pdg_embedding_weight()
frontier, counts = make_seed_frontier(
seeds["event_id"], seeds["pdg"], seeds["pre_pos"],
seeds["pre_E"], seeds["pre_dir"],
seeds["event_id"],
seeds["pdg"],
seeds["pre_pos"],
seeds["pre_E"],
seeds["pre_dir"],
)
rec = _Recorder()
@@ -191,14 +233,26 @@ def rollout(
next_parts: list[dict[str, np.ndarray]] = []
n_total = len(frontier["event_id"])
for start in range(0, n_total, batch_size):
chunk = {
k: v[start : start + batch_size] for k, v in frontier.items()
}
chunk = {k: v[start : start + batch_size] for k, v in frontier.items()}
next_parts.append(
_step_chunk(
chunk, stage1_model, sec_decoder, oracle, cond_norm, tgt_norm,
pdg_map, mat_map, pdg_map_inv, pdg_emb_weight, rec, counts,
energy_cutoff, max_steps, steps, device, max_tracks_per_event,
chunk,
stage1_model,
sec_decoder,
oracle,
cond_norm,
tgt_norm,
pdg_map,
mat_map,
pdg_map_inv,
pdg_emb_weight,
rec,
counts,
energy_cutoff,
max_steps,
steps,
device,
max_tracks_per_event,
)
)
frontier = _concat_frontiers(next_parts)
@@ -207,9 +261,23 @@ def rollout(
def _step_chunk(
tr, stage1_model, sec_decoder, oracle, cond_norm, tgt_norm, pdg_map, mat_map,
pdg_map_inv, pdg_emb_weight, rec, counts, energy_cutoff, max_steps, steps,
device, max_tracks_per_event,
tr,
stage1_model,
sec_decoder,
oracle,
cond_norm,
tgt_norm,
pdg_map,
mat_map,
pdg_map_inv,
pdg_emb_weight,
rec,
counts,
energy_cutoff,
max_steps,
steps,
device,
max_tracks_per_event,
) -> dict[str, np.ndarray]:
"""Advance one chunk of tracks by a single step; return the next frontier."""
n = len(tr["event_id"])
@@ -225,19 +293,33 @@ def _step_chunk(
# --- Pre-step termination gates (in priority order; each track picks one) ---
stop = np.zeros(n, dtype=bool)
escaped_sel = escaped & ~stop
rec.add(**_terminal_rows(tr, escaped_sel, TERM_ESCAPED, edep=np.zeros(int(escaped_sel.sum()))))
rec.add(
**_terminal_rows(
tr, escaped_sel, TERM_ESCAPED, edep=np.zeros(int(escaped_sel.sum()))
)
)
stop |= escaped_sel
unknown_sel = ~known_pdg & ~stop
rec.add(**_terminal_rows(tr, unknown_sel, TERM_UNKNOWN_PDG, edep=tr["pre_E"][unknown_sel]))
rec.add(
**_terminal_rows(
tr, unknown_sel, TERM_UNKNOWN_PDG, edep=tr["pre_E"][unknown_sel]
)
)
stop |= unknown_sel
cutoff_sel = (tr["pre_E"] < energy_cutoff) & ~stop
rec.add(**_terminal_rows(tr, cutoff_sel, TERM_ENERGY_CUTOFF, edep=tr["pre_E"][cutoff_sel]))
rec.add(
**_terminal_rows(
tr, cutoff_sel, TERM_ENERGY_CUTOFF, edep=tr["pre_E"][cutoff_sel]
)
)
stop |= cutoff_sel
maxstep_sel = (tr["step_in_track"] >= max_steps) & ~stop
rec.add(**_terminal_rows(tr, maxstep_sel, TERM_MAX_STEPS, edep=tr["pre_E"][maxstep_sel]))
rec.add(
**_terminal_rows(tr, maxstep_sel, TERM_MAX_STEPS, edep=tr["pre_E"][maxstep_sel])
)
stop |= maxstep_sel
active = ~stop
@@ -250,8 +332,12 @@ def _step_chunk(
# --- Build conditioning and run the two stages ---
cond_dict = {
"pre_pos": tr["pre_pos"], "pre_E": tr["pre_E"], "pre_dir": tr["pre_dir"],
"layer_id": layer_id, "material": material, "pdg": tr["pdg"],
"pre_pos": tr["pre_pos"],
"pre_E": tr["pre_E"],
"pre_dir": tr["pre_dir"],
"layer_id": layer_id,
"material": material,
"pdg": tr["pdg"],
}
cond_cont, cond_cat = build_cond_features(cond_dict, pdg_map, mat_map, cond_norm)
cc = torch.from_numpy(cond_cont).float().to(device)
@@ -264,12 +350,18 @@ def _step_chunk(
edep, e_sec, post_E, _delta = energy_simplex_decode(raw[:, 1:3], tr["pre_E"])
post_dir_local = raw[:, 3:6].copy()
post_dir_local /= np.clip(np.linalg.norm(post_dir_local, axis=1, keepdims=True), 1e-8, None)
post_dir_local /= np.clip(
np.linalg.norm(post_dir_local, axis=1, keepdims=True), 1e-8, None
)
post_dir_world = inv_local_frame_rotation(tr["pre_dir"], post_dir_local)
travel_dir_local = raw[:, 6:9].copy()
travel_dir_local /= np.clip(np.linalg.norm(travel_dir_local, axis=1, keepdims=True), 1e-8, None)
post_pos = reconstruct_post_pos(tr["pre_pos"], tr["pre_dir"], step_length, travel_dir_local)
travel_dir_local /= np.clip(
np.linalg.norm(travel_dir_local, axis=1, keepdims=True), 1e-8, None
)
post_pos = reconstruct_post_pos(
tr["pre_pos"], tr["pre_dir"], step_length, travel_dir_local
)
n_sec_np = n_sec_pred.cpu().numpy().astype(np.int64)
@@ -279,8 +371,12 @@ def _step_chunk(
)
sec_pdg_idx = snap_type_to_pdg_idx(sec_type_emb, pdg_emb_weight)
sec_E, sec_dir_world, sec_pdg_code, sec_valid = decode_secondaries(
sec_cont.cpu().numpy(), sec_pdg_idx.cpu().numpy(), n_sec_np,
e_sec, tr["pre_dir"], pdg_map_inv,
sec_cont.cpu().numpy(),
sec_pdg_idx.cpu().numpy(),
n_sec_np,
e_sec,
tr["pre_dir"],
pdg_map_inv,
)
edep = edep.astype(np.float64)
@@ -288,13 +384,21 @@ def _step_chunk(
# --- Spawn secondaries (with per-event track cap) ---
new_tracks, dropped_edep = _spawn_secondaries(
tr, post_pos, sec_valid, sec_E, sec_dir_world, sec_pdg_code,
counts, max_tracks_per_event,
tr,
post_pos,
sec_valid,
sec_E,
sec_dir_world,
sec_pdg_code,
counts,
max_tracks_per_event,
)
# Energy bookkeeping so each step conserves exactly (edep + carried + post_E
# == pre_E): the primary lost `e_sec` to secondaries, but the decoded
# secondaries only carry `sec_E[valid].sum()`. Deposit the unallocated
# residual locally, plus the energy of any sub-cap secondaries we dropped.
# == pre_E): `decode_secondaries` already rescales valid slots to sum to
# exactly `e_sec` whenever n_sec > 0, so `residual` here is ~0 except when
# n_sec == 0 (no secondary to carry the budget at all — the whole `e_sec`
# becomes residual). Also deposit the energy of any sub-cap secondaries
# we dropped for hitting `max_tracks_per_event`.
sec_E_valid_sum = (sec_E * sec_valid).sum(axis=1)
residual = np.maximum(e_sec - sec_E_valid_sum, 0.0)
edep = edep + residual + dropped_edep
@@ -303,31 +407,58 @@ def _step_chunk(
natural = post_E <= 0.0
reason = np.where(natural, TERM_NATURAL_END, "").astype(object)
rec.add(
event_id=tr["event_id"], track_id=tr["track_id"], parent_id=tr["parent_id"],
generation=tr["generation"], step_no=tr["step_in_track"], pdg=tr["pdg"],
pre_x=tr["pre_pos"][:, 0], pre_y=tr["pre_pos"][:, 1], pre_z=tr["pre_pos"][:, 2],
pre_E=tr["pre_E"], pre_dx=tr["pre_dir"][:, 0], pre_dy=tr["pre_dir"][:, 1],
event_id=tr["event_id"],
track_id=tr["track_id"],
parent_id=tr["parent_id"],
generation=tr["generation"],
step_no=tr["step_in_track"],
pdg=tr["pdg"],
pre_x=tr["pre_pos"][:, 0],
pre_y=tr["pre_pos"][:, 1],
pre_z=tr["pre_pos"][:, 2],
pre_E=tr["pre_E"],
pre_dx=tr["pre_dir"][:, 0],
pre_dy=tr["pre_dir"][:, 1],
pre_dz=tr["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=material, layer_id=layer_id, n_sec_pred=n_sec_np,
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=material,
layer_id=layer_id,
n_sec_pred=n_sec_np,
termination_reason=reason,
)
# --- Continue surviving primaries ---
cont = ~natural
cont_frontier = {
"event_id": tr["event_id"][cont], "track_id": tr["track_id"][cont],
"parent_id": tr["parent_id"][cont], "generation": tr["generation"][cont],
"step_in_track": tr["step_in_track"][cont] + 1, "pdg": tr["pdg"][cont],
"pre_pos": post_pos[cont], "pre_E": post_E[cont], "pre_dir": post_dir_world[cont],
"event_id": tr["event_id"][cont],
"track_id": tr["track_id"][cont],
"parent_id": tr["parent_id"][cont],
"generation": tr["generation"][cont],
"step_in_track": tr["step_in_track"][cont] + 1,
"pdg": tr["pdg"][cont],
"pre_pos": post_pos[cont],
"pre_E": post_E[cont],
"pre_dir": post_dir_world[cont],
}
return _concat_frontiers([cont_frontier, new_tracks])
def _spawn_secondaries(
tr, post_pos, sec_valid, sec_E, sec_dir_world, sec_pdg_code, counts,
tr,
post_pos,
sec_valid,
sec_E,
sec_dir_world,
sec_pdg_code,
counts,
max_tracks_per_event,
) -> tuple[dict[str, np.ndarray], np.ndarray]:
"""Turn valid secondaries into new tracks; return (frontier, per-parent dropped edep).
+2 -2
View File
@@ -63,8 +63,8 @@ def sample_secondaries(
sec_cont = x_slots[:, :, :4]
sec_type_emb = x_slots[:, :, 4:]
sec_valid = (
torch.arange(K_MAX, device=device).unsqueeze(0) < n_sec_pred.unsqueeze(1)
sec_valid = torch.arange(K_MAX, device=device).unsqueeze(0) < n_sec_pred.unsqueeze(
1
)
return sec_cont, sec_type_emb, sec_valid
+41 -23
View File
@@ -3,15 +3,17 @@ run_pbwo4, run_sampling) and filing the output into the dataset's raw/ tree:
raw/<kind>/<gen>/<detector>/shard-NNN.root
These executables take `[configName] nEvents` and always write a fixed-name
*.root file into the current directory so running several in parallel
needs separate working directories, and the output filename has to be
discovered rather than assumed (it differs per executable: run_pbwo4 writes
pbwo4_<n>events_hits.root, run_sampling writes sampling_<config>_<n>events_hits.root,
others may differ again). This script gives each run its own scratch
directory under <dataset-root>/.sim-tmp/, requires exactly one *.root to
appear there, and moves it to the next free shard index for that detector
(existing shards are never overwritten).
These executables take `[configName] nEvents [energy_GeV]` (configName is
only accepted by executables with a config selector, e.g. run_sampling;
energy_GeV defaults to 1.0 in the executable itself if omitted here) and
always write a fixed-name *.root file into the current directory so
running several in parallel needs separate working directories, and the
output filename has to be discovered rather than assumed (it differs per
executable: run_pbwo4 writes pbwo4_<n>events_hits.root, run_sampling writes
sampling_<config>_<n>events_hits.root, others may differ again). This script
gives each run its own scratch directory under <dataset-root>/.sim-tmp/,
requires exactly one *.root to appear there, and moves it to the next free
shard index for that detector (existing shards are never overwritten).
--gen must already exist under raw/<kind>/ create one first with
`dwarf bump-gen`.
@@ -105,8 +107,8 @@ def plan_jobs(
return jobs
def job_seed(kind: str, gen: str, job: SimJob) -> int:
"""Deterministic RNG seed for one sim job, unique per (kind, gen, detector, config, shard).
def job_seed(kind: str, gen: str, job: SimJob, energy_gev: float | None) -> int:
"""Deterministic RNG seed for one sim job, unique per (kind, gen, detector, config, shard, energy).
Jobs run concurrently (ThreadPoolExecutor below) and can start within the
same wall-clock second; minicalosim's default seed falls back to
@@ -115,14 +117,31 @@ def job_seed(kind: str, gen: str, job: SimJob) -> int:
landing in separate shard files. Deriving the seed from the full job
identity instead keeps it both unique and reproducible.
"""
key = f"{kind}|{gen}|{job.detector}|{job.config or ''}|{job.shard_index}"
key = (
f"{kind}|{gen}|{job.detector}|{job.config or ''}|{job.shard_index}"
f"|{energy_gev if energy_gev is not None else ''}"
)
return zlib.crc32(key.encode()) & 0x7FFFFFFF
def build_cmd(
executable: Path, job: SimJob, events_per_file: int, energy_gev: float | None
) -> list[str]:
"""minicalosim executables take positional `[configName] nEvents [energy_GeV]`."""
cmd = [str(executable)]
if job.config:
cmd.append(job.config)
cmd.append(str(events_per_file))
if energy_gev is not None:
cmd.append(str(energy_gev))
return cmd
def run_job(
job: SimJob,
executable: Path,
events_per_file: int,
energy_gev: float | None,
dataset_root: Path,
kind: str,
gen: str,
@@ -134,12 +153,9 @@ def run_job(
)
workdir.mkdir(parents=True)
cmd = [str(executable)]
if job.config:
cmd.append(job.config)
cmd.append(str(events_per_file))
cmd = build_cmd(executable, job, events_per_file, energy_gev)
env = dict(os.environ, MINICALOSIM_SEED=str(job_seed(kind, gen, job)))
env = dict(os.environ, MINICALOSIM_SEED=str(job_seed(kind, gen, job, energy_gev)))
result = subprocess.run(cmd, cwd=workdir, capture_output=True, text=True, env=env)
if result.returncode != 0:
@@ -194,6 +210,7 @@ def run_all(
jobs: list[SimJob],
executable: Path,
events_per_file: int,
energy_gev: float | None,
dataset_root: Path,
kind: str,
gen: str,
@@ -208,6 +225,7 @@ def run_all(
job,
executable,
events_per_file,
energy_gev,
dataset_root,
kind,
gen,
@@ -238,6 +256,7 @@ def run_make_root(
dataset_root: str,
jobs: int,
execute: bool,
energy_gev: float | None = None,
) -> None:
if jobs < 1:
raise SystemExit("error: --jobs must be >= 1")
@@ -245,6 +264,8 @@ def run_make_root(
raise SystemExit("error: --num-files must be >= 1")
if events_per_file < 1:
raise SystemExit("error: --events-per-file must be >= 1")
if energy_gev is not None and energy_gev <= 0:
raise SystemExit("error: --energy-gev must be > 0")
if not executable.is_file() or not os.access(executable, os.X_OK):
raise SystemExit(f"error: {executable} is not an executable file")
@@ -257,11 +278,7 @@ def run_make_root(
print(f"=== {'EXECUTING' if execute else 'DRY RUN'} ===")
print(f"executable: {executable}")
for job in planned_jobs:
cmd = (
[str(executable)]
+ ([job.config] if job.config else [])
+ [str(events_per_file)]
)
cmd = build_cmd(executable, job, events_per_file, energy_gev)
dest = (
dataset_root_path
/ "raw"
@@ -270,7 +287,7 @@ def run_make_root(
/ job.detector
/ f"shard-{job.shard_index:03d}.root"
)
seed = job_seed(kind, gen, job)
seed = job_seed(kind, gen, job, energy_gev)
print(f" MINICALOSIM_SEED={seed} {' '.join(cmd)} -> {dest}")
if not execute:
@@ -283,6 +300,7 @@ def run_make_root(
planned_jobs,
executable,
events_per_file,
energy_gev,
dataset_root_path,
kind,
gen,
+11
View File
@@ -363,6 +363,16 @@ def make_root(
gen: Annotated[
str, typer.Option("--gen", help="Existing gen tag under raw/<kind>/, e.g. gen1")
],
energy_gev: Annotated[
float | None,
typer.Option(
"--energy-gev",
help="energy_GeV passed to the executable (default: executable's own "
"default, currently 1.0). Note the dataset detector label is not "
"derived from this — e.g. use '--detector pbwo4_10gev --energy-gev 10' "
"to name the dataset accordingly.",
),
] = None,
kind: Annotated[
str, typer.Option("--kind", help="steps | hits | ... (default: steps)")
] = "steps",
@@ -390,6 +400,7 @@ def make_root(
dataset_root=str(dataset_root),
jobs=jobs,
execute=execute,
energy_gev=energy_gev,
)
+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(