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:
@@ -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")
|
||||
|
||||
@@ -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}")
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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