Add autoregressive shower rollout driver

Closes the loop from single-step prediction into full showers:

- giant/geometry.py + `dwarf build-geometry-oracle`: learn position ->
  (material, layer_id) from data (KNN/SVM) to supply the conditioning the
  surrogate does not predict; flag detector escape by NN distance.
- giant/rollout.py: breadth-first batched frontier that steps all active
  tracks, spawns secondaries as new tracks, and terminates on energy cutoff,
  per-track max steps, escape, or natural end. Energy is deposited locally on
  every stop except escape (leakage), so showers conserve energy exactly.
- `giant rollout` CLI: seed from real events (argmax pre_E), load checkpoint,
  write a world-frame steps parquet + YAML sidecar.
- giant/analysis.py: compute_rollout_observables + plot_rollout_* for
  single-sided longitudinal/transverse/total-energy shower profiles;
  analysis/export_rollout_observables.py driver.
- scikit-learn added as an optional `geometry` extra (lazy-imported).
- Tests: tests/test_geometry.py, tests/test_rollout.py.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
2026-07-08 10:10:36 +02:00
parent 549c051417
commit 26a176aeaa
13 changed files with 1566 additions and 4 deletions
+4 -1
View File
@@ -8,6 +8,7 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co
uv sync --extra cpu # install dependencies with CPU-only torch (standard/default)
uv sync --extra cuda # install dependencies with CUDA 11.8 torch
uv sync --extra cpu --extra dev # add dev extras (pytest, etc.)
uv sync --extra cpu --extra geometry # add scikit-learn for the geometry oracle (giant rollout)
pytest # run tests
giant train path/to/steps.parquet --mode flow # train (flow matching)
giant train path/to/steps.parquet --mode ddpm # train (DDPM baseline)
@@ -42,7 +43,9 @@ GIANT is a conditional generative surrogate for the Geant4 step function. It rep
**Samplers** (`giant/sample.py`): DDPM, DDIM, and flow matching (ODE integration, ~10 steps). Flow matching is the primary mode.
**Validation** (`giant/validate.py`): step-level marginal comparisons; shower-level rollout validation is planned.
**Validation** (`giant/validate.py`): step-level marginal comparisons. Shower-level (rollout) observables live in `giant/analysis.py` (`compute_rollout_observables` + `plot_rollout_*`), fed by `giant rollout` output.
**Shower rollout** (`giant/rollout.py`, `giant rollout` CLI): autoregressively steps the two-stage model into a full shower — each primary post-step becomes the next pre-step, secondaries are pushed as new tracks, and per-step `material`/`layer_id` come from a `GeometryOracle` (`giant/geometry.py`, built via `dwarf build-geometry-oracle`) that learns position → (material, layer_id) from data and flags detector escape by nearest-neighbour distance. Tracks terminate on energy cutoff, per-track max steps, escape, or natural end; energy is deposited locally on every stop except escape (leakage), so showers conserve energy by construction.
## Roadmap
+46
View File
@@ -0,0 +1,46 @@
"""Export shower-level plots from a `giant rollout` steps parquet.
Not part of the package; run manually. Optionally overlays the real showers
seeded from the same events (a `giant predict --coord local` file) by passing a
reference path. Usage::
python analysis/export_rollout_observables.py ROLLOUT.parquet [REFERENCE_local.parquet]
"""
import sys
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
OUT = Path(sys.argv[3]) if len(sys.argv) > 3 else Path(".")
print(f"=== computing rollout observables: {rollout_file} ===")
obs = a.compute_rollout_observables(rollout_file)
tbl = obs.event_table
print(f"n events: {len(tbl)}")
print(
f"total_edep/event: mean={tbl['total_edep'].mean():.4g} MeV "
f"leaked_E/event: mean={tbl['leaked_E'].mean():.4g} MeV "
f"n_tracks/event: mean={tbl['n_tracks'].mean():.1f} "
f"n_steps/event: mean={tbl['n_steps'].mean():.1f}"
)
reference = None
if reference_file is not None:
print(f"=== computing real reference: {reference_file} ===")
reference = a.compute_event_observables_pl(reference_file)
print("=== plots ===")
for name, fn in [
("longitudinal", a.plot_rollout_longitudinal),
("transverse", a.plot_rollout_transverse),
("total-energy", a.plot_rollout_total_energy),
]:
fig = fn(obs, reference=reference)
path = OUT / f"rollout-{name}.png"
fig.savefig(path, dpi=150, bbox_inches="tight")
print(f"wrote {path}")
print("DONE")
+176
View File
@@ -91,6 +91,8 @@ from giant.constants import (
PREDICT_COORD_METADATA_KEY,
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
ROLLOUT_COORD_VALUE,
TERM_ESCAPED,
)
from giant.data.transforms import (
energy_simplex_decode,
@@ -1693,3 +1695,177 @@ def plot_pdg_length_share(table: pl.DataFrame, max_slices: int = 6):
"length traveled share by particle type",
max_slices,
)
# ── Shower-rollout observables (single-sided; `giant rollout` output) ──────────
#
# The rollout writes a world-frame steps parquet (post_x/y/z + edep already in
# physical units), unlike the local predict schema these other diagnostics read.
# So this is a standalone, generated-only path: no paired `true_*` columns, no
# energy-simplex decode. Because rollout showers are seeded from real events, the
# 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",
]
@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
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
transverse_profile: np.ndarray
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.
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.
"""
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()
ax /= np.clip(np.linalg.norm(ax, axis=1, keepdims=True), 1e-12, None)
ent = entry[["pre_x", "pre_y", "pre_z"]].to_numpy(dtype=np.float64)
ev_ids = entry.index.to_numpy()
ev_row = {int(e): i for i, e in enumerate(ev_ids)}
row_ev = df["event_id"].map(ev_row).to_numpy()
post = df[["post_x", "post_y", "post_z"]].to_numpy(dtype=np.float64)
rel = post - ent[row_ev]
depth = np.einsum("ij,ij->i", rel, ax[row_ev])
transverse = np.linalg.norm(rel - depth[:, None] * ax[row_ev], axis=1)
edep = df["edep"].to_numpy(dtype=np.float64)
n_events = len(ev_ids)
depth_lo, depth_hi = np.quantile(depth, [0.001, 0.999])
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)
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)
# Per-(event, bin) edep sums, then mean/std across events.
depth_ev = np.zeros((n_events, depth_bins))
trans_ev = np.zeros((n_events, transverse_bins))
np.add.at(depth_ev, (row_ev, d_bin), edep)
np.add.at(trans_ev, (row_ev, t_bin), edep)
# Per-event scalar table.
is_leak = df["termination_reason"].to_numpy() == TERM_ESCAPED
per_ev = df.groupby("event_id").agg(
n_steps=("edep", "size"),
total_edep=("edep", "sum"),
)
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()
)
# 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.index.name = "event_id"
return RolloutObservables(
event_table=per_ev.reset_index(),
depth_edges=depth_edges,
transverse_edges=transverse_edges,
depth_profile=depth_ev.mean(0),
depth_profile_std=depth_ev.std(0),
transverse_profile=trans_ev.mean(0),
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.
"""
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",
)
if reference is not None:
rc = 0.5 * (reference.depth_edges[:-1] + reference.depth_edges[1:])
ax.plot(rc, reference.real_depth_profile, label="real", color="k", ls="--")
ax.set_xlabel("depth along axis [mm]")
ax.set_ylabel("mean E_dep / event [MeV]")
ax.set_title("Longitudinal shower profile")
ax.legend()
fig.tight_layout()
return fig
def plot_rollout_transverse(obs: RolloutObservables, reference=None):
"""Mean edep vs transverse distance from the shower axis (Molière-style)."""
centers = 0.5 * (obs.transverse_edges[:-1] + obs.transverse_edges[1:])
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",
)
if reference is not None:
rc = 0.5 * (reference.transverse_edges[:-1] + reference.transverse_edges[1:])
ax.plot(rc, reference.real_transverse_profile, label="real", color="k", ls="--")
ax.set_xlabel("transverse distance [mm]")
ax.set_ylabel("mean E_dep / event [MeV]")
ax.set_title("Transverse shower profile")
ax.set_yscale("log")
ax.legend()
fig.tight_layout()
return fig
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")
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.set_xlabel("total E_dep / event [MeV]")
ax.set_ylabel("events")
ax.set_title("Total deposited energy")
ax.legend()
fig.tight_layout()
return fig
+195
View File
@@ -21,6 +21,7 @@ from giant.constants import (
PREDICT_COORD_METADATA_KEY,
PREDICT_SCHEMA_VERSION,
PREDICT_SCHEMA_VERSION_KEY,
ROLLOUT_COORD_VALUE,
)
from giant.data.loader import (
find_parquet_files,
@@ -37,8 +38,10 @@ from giant.data.transforms import (
reconstruct_post_pos,
Normalizer,
)
from giant.geometry import GeometryOracle
from giant.model.network import DenoisingMLP, SecondaryDecoder
from giant.pipeline import run_train_job
from giant.rollout import rollout as run_rollout
from giant.sample import sample_flow, sample_secondaries, snap_type_to_pdg_idx
_STAGE1_MODEL_KEYS = {
@@ -637,5 +640,197 @@ def predict(
typer.echo(f"wrote {total:,} rows → {out}")
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 —
the codebase's convention for the primary (a secondary always carries less
energy than its parent). See giant/analysis.py:_entry_axis_and_bin_edges.
"""
best_E: dict[int, float] = {}
best: dict[int, tuple] = {}
for path in files:
for chunk in iter_cond_chunks(path):
ev = chunk["event_id"]
pe = chunk["pre_E"]
for i in range(len(ev)):
e = int(ev[i])
if pe[i] > best_E.get(e, -np.inf):
best_E[e] = float(pe[i])
best[e] = (
int(chunk["pdg"][i]),
chunk["pre_pos"][i].astype(np.float64),
float(pe[i]),
chunk["pre_dir"][i].astype(np.float64),
)
event_ids = sorted(best)
if n_events is not None:
event_ids = event_ids[:n_events]
if not event_ids:
raise ValueError("no events found to seed from")
return {
"event_id": np.array(event_ids, dtype=np.int64),
"pdg": np.array([best[e][0] for e in event_ids], dtype=np.int64),
"pre_pos": np.stack([best[e][1] for e in event_ids]),
"pre_E": np.array([best[e][2] for e in event_ids], dtype=np.float64),
"pre_dir": np.stack([best[e][3] for e in event_ids]),
}
@app.command()
def rollout(
data: Annotated[
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)")
],
geometry: Annotated[
Path,
typer.Option(
"--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]"
),
] = 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")
] = 10,
batch_size: Annotated[
int, typer.Option("--batch-size", "-b", help="Tracks stepped per model forward")
] = 4096,
max_tracks_per_event: Annotated[
Optional[int],
typer.Option(
"--max-tracks-per-event",
help="Safety cap on tracks per shower (sub-cap secondaries deposit in place)",
),
] = None,
escape_threshold: Annotated[
Optional[float],
typer.Option(
"--escape-threshold",
help="Override the oracle's NN-distance escape threshold [mm]",
),
] = None,
n_events: Annotated[
Optional[int], typer.Option("--n-events", help="Cap number of seed events")
] = None,
device: Annotated[
Optional[str], typer.Option("--device", "-d", help="cpu | cuda | mps (auto)")
] = None,
out: Annotated[
Optional[Path], typer.Option("--out", "-o", help="Output steps parquet")
] = None,
seed: Annotated[
Optional[int], typer.Option("--seed", help="Torch/numpy seed for reproducibility")
] = None,
) -> None:
"""Roll the surrogate forward into full showers (autoregressive)."""
if seed is not None:
torch.manual_seed(seed)
np.random.seed(seed)
_device = torch.device(device) if device else gconfig.auto_device()
typer.echo(f"device: {_device}")
ckpt = torch.load(checkpoint, map_location="cpu", weights_only=False)
for key in ("model_config", "sec_decoder"):
if key not in ckpt:
typer.echo(
f"error: checkpoint has no {key} — retrain with the current code",
err=True,
)
raise typer.Exit(1)
model_cfg = ckpt["model_config"]
pdg_map = {int(k): v for k, v in ckpt["pdg_map"].items()}
mat_map = {str(k): v for k, v in ckpt["mat_map"].items()}
cond_norm = Normalizer.from_dict(ckpt["normalizer"]["cond"])
tgt_norm = Normalizer.from_dict(ckpt["normalizer"]["target"])
model = DenoisingMLP(**{k: v for k, v in model_cfg.items() if k in _STAGE1_MODEL_KEYS})
model.load_state_dict(ckpt["model"])
model.to(_device).eval()
sec_decoder = SecondaryDecoder(
**{k: v for k, v in model_cfg.items() if k in _SEC_DECODER_MODEL_KEYS}
)
sec_decoder.load_state_dict(ckpt["sec_decoder"])
sec_decoder.to(_device).eval()
typer.echo(f"loaded checkpoint: {checkpoint}")
oracle = GeometryOracle.load(geometry)
typer.echo(
f"loaded geometry oracle: {geometry} "
f"(escape_threshold={oracle.escape_threshold:.3f})"
)
files = find_parquet_files(data)
seeds = _seed_from_data(files, n_events)
typer.echo(f"seeded {len(seeds['event_id']):,} shower(s)")
records = run_rollout(
model,
sec_decoder,
oracle,
seeds,
cond_norm,
tgt_norm,
pdg_map,
mat_map,
energy_cutoff=energy_cutoff,
max_steps=max_steps,
steps=steps,
batch_size=batch_size,
device=_device,
max_tracks_per_event=max_tracks_per_event,
escape_threshold=escape_threshold,
)
out, dataset_path, pred_uuid = _resolve_prediction_output(data, out)
out.parent.mkdir(parents=True, exist_ok=True)
table = pa.table(records).replace_schema_metadata(
{
PREDICT_COORD_METADATA_KEY: ROLLOUT_COORD_VALUE,
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
}
)
pq.write_table(table, out)
ref_path = _write_prediction_ref(checkpoint, pred_uuid, out, dataset_path)
ref = yaml.safe_load(ref_path.read_text())
ref.update(
{
"kind": "rollout",
"geometry_oracle": str(geometry.resolve()),
"energy_cutoff": energy_cutoff,
"max_steps": max_steps,
"steps": steps,
"max_tracks_per_event": max_tracks_per_event,
"n_seed_events": int(len(seeds["event_id"])),
}
)
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
)
typer.echo(f"wrote {n_rows:,} step rows → {out}")
typer.echo(f"terminations: {dict(reasons)}")
typer.echo(f"reference: {ref_path}")
if __name__ == "__main__":
app()
+14
View File
@@ -37,3 +37,17 @@ LOCAL_TARGET_NAMES = [
PREDICT_COORD_METADATA_KEY = "giant.predict.coord"
PREDICT_SCHEMA_VERSION_KEY = "giant.predict.schema_version"
PREDICT_SCHEMA_VERSION = "2"
# Coord-metadata value tagging a `giant rollout` steps parquet (world frame,
# autoregressive shower output). Distinct from predict's "global"/"local".
ROLLOUT_COORD_VALUE = "rollout"
# Per-track termination reasons recorded by the rollout driver. "escaped"
# energy is treated as leakage (not deposited); every other stop deposits the
# track's remaining energy locally so total energy is conserved.
TERM_ESCAPED = "escaped"
TERM_ENERGY_CUTOFF = "energy_cutoff"
TERM_MAX_STEPS = "max_steps"
TERM_NATURAL_END = "natural_end"
TERM_UNKNOWN_PDG = "unknown_pdg"
TERM_MAX_TRACKS = "max_tracks"
+290
View File
@@ -0,0 +1,290 @@
"""Geometry oracle: learn a position -> (material, layer_id) map from step data.
The surrogate conditions on `material` and `layer_id`, but does not predict
them — during a shower rollout they must be looked up from the new position.
There is no in-repo detector geometry (it lives in external miniCaloSim), so we
approximate it with a nearest-neighbour classifier fit on positions sampled from
a real steps dataset. A position whose nearest training neighbour is farther than
a threshold is treated as having escaped the detector (out-of-world), which the
rollout driver uses as a hard track-termination condition.
scikit-learn / joblib are an optional dependency (the `geometry` extra); they are
imported lazily so the core install stays lean.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Iterable
import numpy as np
import pyarrow.parquet as pq
_INSTALL_HINT = (
"the geometry oracle needs scikit-learn — install it with "
"`uv sync --extra cpu --extra geometry`"
)
def _require_sklearn():
try:
import joblib # noqa: F401
from sklearn.neighbors import KNeighborsClassifier # noqa: F401
from sklearn.svm import SVC # noqa: F401
except ImportError as exc: # pragma: no cover - exercised only without extra
raise ImportError(_INSTALL_HINT) from exc
@dataclass
class GeometryOracle:
"""Maps world-frame position -> (material, layer_id, escaped).
`estimator` is a fitted sklearn classifier over 3D positions predicting a
class index into `classes` (a list of (material, layer_id) pairs).
`escape_threshold` is a distance in position units (mm): a query point whose
nearest training reference point is farther than this is flagged `escaped`.
"""
estimator: Any
classes: list[tuple[str, int]]
escape_threshold: float
metadata: dict
# Only populated for non-neighbour estimators (SVM) to answer the escape
# distance query; KNeighborsClassifier answers it directly.
_ref_tree: Any = field(default=None)
def query(
self, pos: np.ndarray
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Return (material (N,) str, layer_id (N,) int, escaped (N,) bool)."""
pos = np.ascontiguousarray(np.asarray(pos, dtype=np.float64))
if pos.ndim != 2 or pos.shape[1] != 3:
raise ValueError(f"pos must be (N, 3), got {pos.shape}")
if len(pos) == 0:
return (
np.empty(0, dtype=object),
np.empty(0, dtype=np.int64),
np.empty(0, dtype=bool),
)
# kneighbors gives the distance to the nearest reference point, which is
# what the escape test needs; it exists on both KNeighborsClassifier and
# (via a stored reference tree) our SVM wrapper below.
dist = self._nearest_distance(pos)
escaped = dist > self.escape_threshold
cls_idx = self.estimator.predict(pos).astype(np.int64)
material = np.array([self.classes[i][0] for i in cls_idx], dtype=object)
layer_id = np.array([self.classes[i][1] for i in cls_idx], dtype=np.int64)
return material, layer_id, escaped
def _nearest_distance(self, pos: np.ndarray) -> np.ndarray:
from sklearn.neighbors import KNeighborsClassifier
if isinstance(self.estimator, KNeighborsClassifier):
dist, _ = self.estimator.kneighbors(pos, n_neighbors=1)
return dist[:, 0]
# SVM (or any non-neighbour estimator): use the separately stored
# NearestNeighbors index purely for the escape distance.
dist, _ = self._ref_tree.kneighbors(pos, n_neighbors=1)
return dist[:, 0]
def save(self, path: str | Path) -> None:
_require_sklearn()
import joblib
joblib.dump(
{
"estimator": self.estimator,
"classes": self.classes,
"escape_threshold": self.escape_threshold,
"metadata": self.metadata,
"ref_tree": getattr(self, "_ref_tree", None),
},
path,
)
@classmethod
def load(cls, path: str | Path) -> "GeometryOracle":
_require_sklearn()
import joblib
d = joblib.load(path)
obj = cls(
estimator=d["estimator"],
classes=[tuple(c) for c in d["classes"]],
escape_threshold=float(d["escape_threshold"]),
metadata=d.get("metadata", {}),
)
obj._ref_tree = d.get("ref_tree")
return obj
_PRE_COLS = ["pre_x", "pre_y", "pre_z"]
_POST_COLS = ["post_x", "post_y", "post_z"]
_LABEL_COLS = ["material", "layer_id"]
def _iter_point_batches(path: Path, batch_size: int = 1_000_000):
"""Yield (pos (M,3), material (M,), layer_id (M,)) from any parquet with a
position + material + layer_id schema (raw steps *or* predict output).
Only the needed columns are read. post_pos points are included when present
(they share their step's label) so boundary regions are densely sampled.
"""
pf = pq.ParquetFile(path)
have = set(pf.schema_arrow.names)
missing = [c for c in (*_PRE_COLS, *_LABEL_COLS) if c not in have]
if missing:
raise ValueError(
f"{path} is missing columns {missing} needed to build a geometry "
"oracle (expected pre_x/y/z, material, layer_id)"
)
has_post = all(c in have for c in _POST_COLS)
cols = [*_PRE_COLS, *_LABEL_COLS] + (_POST_COLS if has_post else [])
for batch in pf.iter_batches(batch_size=batch_size, columns=cols):
d = batch.to_pydict()
pre = np.array([d[c] for c in _PRE_COLS], dtype=np.float32).T
mat = np.array([str(m) for m in d["material"]], dtype=object)
lay = np.asarray(d["layer_id"], dtype=np.int64)
if has_post:
post = np.array([d[c] for c in _POST_COLS], dtype=np.float32).T
yield (
np.concatenate([pre, post], axis=0),
np.concatenate([mat, mat], axis=0),
np.concatenate([lay, lay], axis=0),
)
else:
yield pre, mat, lay
def _collect_points(
files: Iterable[Path],
subsample: int,
seed: int,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""Stream files, reservoir-sample (pos, material, layer_id) points.
Reservoir sampling keeps memory bounded regardless of total file size.
"""
rng = np.random.default_rng(seed)
res_pos = np.empty((subsample, 3), dtype=np.float32)
res_mat = np.empty(subsample, dtype=object)
res_lay = np.empty(subsample, dtype=np.int64)
seen = 0
for path in files:
for pos, mat, lay in _iter_point_batches(path):
m = len(pos)
if seen < subsample:
take = min(subsample - seen, m)
res_pos[seen : seen + take] = pos[:take]
res_mat[seen : seen + take] = mat[:take]
res_lay[seen : seen + take] = lay[:take]
seen += take
if take < m:
# Reservoir is now full; run the standard replacement rule
# on the remainder of this chunk.
_reservoir_replace(
res_pos, res_mat, res_lay, pos[take:], mat[take:],
lay[take:], seen, rng,
)
seen += m - take
else:
_reservoir_replace(
res_pos, res_mat, res_lay, pos, mat, lay, seen, rng
)
seen += m
n = min(seen, subsample)
return res_pos[:n], res_mat[:n], res_lay[:n]
def _reservoir_replace(res_pos, res_mat, res_lay, pos, mat, lay, seen, rng) -> None:
"""Vectorized reservoir replacement for a batch of incoming points."""
m = len(pos)
k = res_pos.shape[0]
# For incoming global index j (seen..seen+m-1), keep with prob k/(j+1),
# replacing a uniformly-chosen reservoir slot.
idx = seen + np.arange(m)
j = rng.integers(0, idx + 1) # j in [0, global_index]
keep = j < k
slots = j[keep]
res_pos[slots] = pos[keep]
res_mat[slots] = mat[keep]
res_lay[slots] = lay[keep]
def build_geometry_oracle(
files: list[Path],
method: str = "knn",
k: int = 1,
subsample: int = 500_000,
escape_factor: float = 5.0,
seed: int = 0,
) -> GeometryOracle:
"""Fit a position -> (material, layer_id) classifier from steps files.
method: "knn" (KNeighborsClassifier, default) or "svm" (SVC).
k: neighbours for the knn classifier.
subsample: max reference points held in memory / used for the fit.
escape_factor: escape_threshold = escape_factor * median 1-NN spacing of the
reference points, so it scales with the sampling density of the data.
"""
_require_sklearn()
from sklearn.neighbors import KNeighborsClassifier, NearestNeighbors
from sklearn.svm import SVC
pos, mat, lay = _collect_points(files, subsample, seed)
if len(pos) == 0:
raise ValueError("no points collected — are these steps parquet files?")
# Combined (material, layer_id) class label -> contiguous index.
pairs = list(zip((str(m) for m in mat), (int(v) for v in lay)))
classes = sorted(set(pairs))
class_to_idx = {c: i for i, c in enumerate(classes)}
y = np.array([class_to_idx[p] for p in pairs], dtype=np.int64)
X = pos.astype(np.float64)
if method == "knn":
estimator = KNeighborsClassifier(n_neighbors=k)
estimator.fit(X, y)
ref_tree = None
elif method == "svm":
estimator = SVC(kernel="rbf")
estimator.fit(X, y)
# SVM cannot answer nearest-neighbour distance queries, so keep a light
# reference tree alongside it purely for the escape test.
ref_tree = NearestNeighbors(n_neighbors=1).fit(X)
else:
raise ValueError(f"unknown method {method!r}; use 'knn' or 'svm'")
# Escape threshold from the reference point spacing. Sample a subset for the
# median 2-NN distance (the 1st neighbour of a training point is itself).
nn = NearestNeighbors(n_neighbors=2).fit(X)
probe = X if len(X) <= 20_000 else X[
np.random.default_rng(seed).choice(len(X), 20_000, replace=False)
]
d2, _ = nn.kneighbors(probe, n_neighbors=2)
median_nn = float(np.median(d2[:, 1]))
escape_threshold = escape_factor * median_nn
oracle = GeometryOracle(
estimator=estimator,
classes=classes,
escape_threshold=escape_threshold,
metadata={
"method": method,
"k": k,
"n_reference_points": int(len(X)),
"median_nn_dist": median_nn,
"escape_factor": escape_factor,
"n_files": len(files),
},
)
oracle._ref_tree = ref_tree
return oracle
+381
View File
@@ -0,0 +1,381 @@
"""Autoregressive shower rollout driver.
Steps the two-stage GIANT surrogate forward into a full particle shower: each
primary post-step becomes the next pre-step, secondaries are pushed as new tracks,
and the material/layer_id conditioning at every step comes from a `GeometryOracle`
(the surrogate does not predict them).
Tracks are advanced breadth-first: every sweep steps all currently-active tracks
once (in `batch_size` chunks), so many tracks share each model forward pass. A
track terminates on one of the recorded `termination_reason`s in constants.py.
Energy accounting: on every terminal stop except escape, the track's remaining
energy is deposited locally so the shower conserves energy; escaped energy is
treated as detector leakage and not deposited.
"""
from __future__ import annotations
import numpy as np
import torch
from giant.constants import (
TERM_ENERGY_CUTOFF,
TERM_ESCAPED,
TERM_MAX_STEPS,
TERM_NATURAL_END,
TERM_UNKNOWN_PDG,
)
from giant.data.transforms import (
Normalizer,
build_cond_features,
decode_secondaries,
energy_simplex_decode,
inv_local_frame_rotation,
inv_log_transform,
reconstruct_post_pos,
)
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",
"termination_reason",
]
def _empty_frontier() -> dict[str, np.ndarray]:
return {
"event_id": np.empty(0, dtype=np.int64),
"track_id": np.empty(0, dtype=np.int64),
"parent_id": np.empty(0, dtype=np.int64),
"generation": np.empty(0, dtype=np.int64),
"step_in_track": np.empty(0, dtype=np.int64),
"pdg": np.empty(0, dtype=np.int64),
"pre_pos": np.empty((0, 3), dtype=np.float64),
"pre_E": np.empty(0, dtype=np.float64),
"pre_dir": np.empty((0, 3), dtype=np.float64),
}
def _concat_frontiers(parts: list[dict[str, np.ndarray]]) -> dict[str, np.ndarray]:
parts = [p for p in parts if len(p["event_id"]) > 0]
if not parts:
return _empty_frontier()
return {k: np.concatenate([p[k] for p in parts], axis=0) for k in parts[0]}
class _Recorder:
"""Accumulates per-step rows into column lists, materialised at the end."""
def __init__(self) -> None:
self._cols: dict[str, list] = {k: [] for k in _RECORD_KEYS}
def add(self, **cols) -> None:
n = len(cols["event_id"])
if n == 0:
return
for k in _RECORD_KEYS:
v = cols[k]
self._cols[k].append(np.asarray(v).reshape(n))
def to_dict(self) -> dict[str, np.ndarray]:
out = {}
for k, chunks in self._cols.items():
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)
return out
def make_seed_frontier(
event_id: np.ndarray,
pdg: np.ndarray,
pre_pos: np.ndarray,
pre_E: np.ndarray,
pre_dir: np.ndarray,
) -> tuple[dict[str, np.ndarray], dict[int, int]]:
"""Build the initial frontier from primary entry states.
Returns (frontier, event_track_count) where the latter tracks the next
unused track_id per event (each primary gets a fresh id starting from 0).
"""
event_id = np.asarray(event_id, dtype=np.int64)
n = len(event_id)
track_id = np.empty(n, dtype=np.int64)
counts: dict[int, int] = {}
for i, ev in enumerate(event_id.tolist()):
c = counts.get(ev, 0)
track_id[i] = c
counts[ev] = c + 1
dir_ = np.asarray(pre_dir, dtype=np.float64)
dir_ = dir_ / np.clip(np.linalg.norm(dir_, axis=1, keepdims=True), 1e-12, None)
frontier = {
"event_id": event_id,
"track_id": track_id,
"parent_id": np.full(n, -1, dtype=np.int64),
"generation": np.zeros(n, dtype=np.int64),
"step_in_track": np.zeros(n, dtype=np.int64),
"pdg": np.asarray(pdg, dtype=np.int64),
"pre_pos": np.asarray(pre_pos, dtype=np.float64),
"pre_E": np.asarray(pre_E, dtype=np.float64),
"pre_dir": dir_,
}
return frontier, counts
def _terminal_rows(tr: dict[str, np.ndarray], sel: np.ndarray, reason: str, edep):
"""Assemble terminal-marker record columns for the selected tracks."""
pos = tr["pre_pos"][sel]
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],
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],
layer_id=tr.get("_layer_id", np.zeros(len(sel), dtype=np.int64))[sel],
n_sec_pred=np.zeros(n, dtype=np.int64),
termination_reason=np.full(n, reason, dtype=object),
)
@torch.no_grad()
def rollout(
stage1_model: torch.nn.Module,
sec_decoder: torch.nn.Module,
oracle,
seeds: dict[str, np.ndarray],
cond_norm: Normalizer,
tgt_norm: Normalizer,
pdg_map: dict[int, int],
mat_map: dict[str, int],
*,
energy_cutoff: float,
max_steps: int,
steps: int = 10,
batch_size: int = 4096,
device: torch.device | None = None,
max_tracks_per_event: int | None = None,
escape_threshold: float | None = None,
) -> dict[str, np.ndarray]:
"""Run showers to completion; return a step-record dict (see _RECORD_KEYS)."""
device = device or torch.device("cpu")
stage1_model.eval()
sec_decoder.eval()
if escape_threshold is not None:
oracle.escape_threshold = float(escape_threshold)
pdg_map_inv = {v: k for k, v in pdg_map.items()}
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"],
)
rec = _Recorder()
while len(frontier["event_id"]) > 0:
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()
}
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,
)
)
frontier = _concat_frontiers(next_parts)
return rec.to_dict()
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,
) -> dict[str, np.ndarray]:
"""Advance one chunk of tracks by a single step; return the next frontier."""
n = len(tr["event_id"])
# --- Geometry lookup + material/layer conditioning ---
material, layer_id, escaped = oracle.query(tr["pre_pos"])
tr = dict(tr)
tr["_material"] = material
tr["_layer_id"] = layer_id
known_pdg = np.array([int(p) in pdg_map for p in tr["pdg"]], dtype=bool)
# --- 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()))))
stop |= escaped_sel
unknown_sel = ~known_pdg & ~stop
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]))
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]))
stop |= maxstep_sel
active = ~stop
if not active.any():
return _empty_frontier()
tr = {k: v[active] for k, v in tr.items()}
material = tr["_material"]
layer_id = tr["_layer_id"]
# --- 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"],
}
cond_cont, cond_cat = build_cond_features(cond_dict, pdg_map, mat_map, cond_norm)
cc = torch.from_numpy(cond_cont).float().to(device)
ck = torch.from_numpy(cond_cat).long().to(device)
stage1_norm, n_sec_pred = sample_flow(stage1_model, cc, ck, steps=steps)
raw = tgt_norm.inverse_transform(stage1_norm.cpu().numpy())
step_length = inv_log_transform(raw[:, 0])
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_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)
n_sec_np = n_sec_pred.cpu().numpy().astype(np.int64)
# --- Secondaries ---
sec_cont, sec_type_emb, _valid = sample_secondaries(
sec_decoder, cc, ck, stage1_norm, n_sec_pred, steps=steps
)
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,
)
edep = edep.astype(np.float64)
post_E = post_E.astype(np.float64)
# --- 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,
)
# 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.
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
# --- Record the stepped rows; mark natural_end where the primary died ---
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],
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,
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],
}
return _concat_frontiers([cont_frontier, new_tracks])
def _spawn_secondaries(
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).
Secondaries are born at their parent's post_pos. When `max_tracks_per_event`
is set and an event is at its cap, further secondaries are not spawned; their
energy is returned as `dropped_edep` (indexed by parent row) so it is
deposited into the parent step instead of vanishing.
"""
B = len(tr["event_id"])
dropped_edep = np.zeros(B, dtype=np.float64)
pr, sl = np.nonzero(sec_valid) # parent-row idx, slot idx
if len(pr) == 0:
return _empty_frontier(), dropped_edep
# Assign a fresh per-event track_id to each candidate in stable parent order,
# applying the per-event cap. The candidate count per chunk is small
# (<= batch_size * K_MAX), so a plain loop is clear and fast enough.
order = np.lexsort((sl, pr)) # group by parent row, slot ascending
kept_pr, kept_sl, kept_tid = [], [], []
for j in order:
ev = int(tr["event_id"][pr[j]])
cur = counts.get(ev, 0)
if max_tracks_per_event is not None and cur >= max_tracks_per_event:
dropped_edep[pr[j]] += float(sec_E[pr[j], sl[j]])
continue
kept_pr.append(pr[j])
kept_sl.append(sl[j])
kept_tid.append(cur)
counts[ev] = cur + 1
if not kept_pr:
return _empty_frontier(), dropped_edep
pr_k = np.array(kept_pr, dtype=np.int64)
sl_k = np.array(kept_sl, dtype=np.int64)
tid_k = np.array(kept_tid, dtype=np.int64)
frontier = {
"event_id": tr["event_id"][pr_k],
"track_id": tid_k,
"parent_id": tr["track_id"][pr_k],
"generation": tr["generation"][pr_k] + 1,
"step_in_track": np.zeros(len(pr_k), dtype=np.int64),
"pdg": sec_pdg_code[pr_k, sl_k].astype(np.int64),
"pre_pos": post_pos[pr_k],
"pre_E": sec_E[pr_k, sl_k].astype(np.float64),
"pre_dir": sec_dir_world[pr_k, sl_k].astype(np.float64),
}
return frontier, dropped_edep
+4 -1
View File
@@ -24,7 +24,10 @@ dev = [
"pytest>=8,<10",
"ruff>=0.15,<1",
"ty>=0.0.50,<0.1",
"giant[convert,analysis]",
"giant[convert,analysis,geometry]",
]
geometry = [
"scikit-learn>=1.4,<2",
]
convert = [
"uproot>=5.3,<6",
+46
View File
@@ -20,6 +20,7 @@ from scripts.bump_dataset_version import (
run_update_manifest,
)
from scripts.create_root_files import run_make_root
from scripts.geometry_oracle import run_build_geometry_oracle
from scripts.hparam_scan import DATA_DEFAULT, SCAN_DIR_DEFAULT, run_hparam_scan
from scripts.migrate_geant_steps import run_migration
from scripts.steps_to_parquet import convert_steps_to_parquet
@@ -385,6 +386,51 @@ def make_root(
)
class OracleMethod(str, Enum):
knn = "knn"
svm = "svm"
@app.command("build-geometry-oracle")
def build_geometry_oracle(
data: Annotated[
Path, typer.Argument(help="Steps parquet file or directory of steps files")
],
out: Annotated[
Path, typer.Option("--out", "-o", help="Output oracle .pkl path")
],
method: Annotated[
OracleMethod,
typer.Option("--method", help="Classifier: knn (default) or svm"),
] = OracleMethod.knn,
k: Annotated[
int, typer.Option("--k", help="Neighbours for the knn classifier")
] = 1,
subsample: Annotated[
int,
typer.Option("--subsample", help="Max reference points sampled from the data"),
] = 500_000,
escape_factor: Annotated[
float,
typer.Option(
"--escape-factor",
help="escape_threshold = this x median NN spacing of reference points",
),
] = 5.0,
seed: Annotated[int, typer.Option("--seed", help="Sampling seed")] = 0,
) -> None:
"""Fit a position -> (material, layer_id) oracle for `giant rollout`."""
run_build_geometry_oracle(
data=data,
out=out,
method=method.value,
k=k,
subsample=subsample,
escape_factor=escape_factor,
seed=seed,
)
@app.command("hparam-scan")
def hparam_scan(
data: Annotated[str, typer.Option("--data")] = DATA_DEFAULT,
+56
View File
@@ -0,0 +1,56 @@
"""Build a position -> (material, layer_id) geometry oracle from steps parquet.
Backs the `dwarf build-geometry-oracle` subcommand. The oracle is consumed by
`giant rollout` to supply the material/layer conditioning at each step, since the
surrogate does not predict them.
"""
from __future__ import annotations
from pathlib import Path
from giant.data.loader import find_parquet_files
from giant.geometry import build_geometry_oracle
def run_build_geometry_oracle(
data: Path,
out: Path,
method: str = "knn",
k: int = 1,
subsample: int = 500_000,
escape_factor: float = 5.0,
seed: int = 0,
) -> None:
files = find_parquet_files(data)
print(f"found {len(files)} parquet file(s); sampling up to {subsample:,} points")
oracle = build_geometry_oracle(
files,
method=method,
k=k,
subsample=subsample,
escape_factor=escape_factor,
seed=seed,
)
print(f"method: {method} reference points: {oracle.metadata['n_reference_points']:,}")
print("classes (material, layer_id):")
for material, layer_id in oracle.classes:
print(f" {material:<12} layer_id={layer_id}")
print(
f"median NN spacing: {oracle.metadata['median_nn_dist']:.3f} "
f"escape_threshold: {oracle.escape_threshold:.3f} "
f"(= {escape_factor}x spacing)"
)
if oracle.escape_threshold <= 0.0:
print(
"warning: escape_threshold is 0 (reference points are coincident) — "
"every rollout query would be flagged as escaped. Pass "
"`giant rollout --escape-threshold <mm>` to override, or use data "
"with distinct step positions."
)
out.parent.mkdir(parents=True, exist_ok=True)
oracle.save(out)
print(f"wrote oracle -> {out}")
+86
View File
@@ -0,0 +1,86 @@
"""Tests for the geometry oracle (position -> material/layer_id + escape)."""
from pathlib import Path
from unittest.mock import patch
import numpy as np
import pytest
from giant import geometry as g
# scikit-learn is an optional extra; skip the whole module if it's missing.
pytest.importorskip("sklearn")
def _box_batch(n, rng):
"""A labelled point cloud: inside a 100mm box -> PbWO4/0, else AIR/-1."""
pos = rng.uniform(-200, 200, (n, 3)).astype(np.float32)
inside = (np.abs(pos) < 100).all(axis=1)
mat = np.where(inside, "G4_PbWO4", "G4_AIR").astype(object)
lay = np.where(inside, 0, -1).astype(np.int64)
return pos, mat, lay
def _build(subsample=30000, method="knn", escape_factor=5.0):
rng = np.random.default_rng(0)
batches = [_box_batch(20000, rng) for _ in range(3)]
with patch.object(g, "_iter_point_batches", lambda p: iter(batches)):
return g.build_geometry_oracle(
[Path("x")], method=method, subsample=subsample,
escape_factor=escape_factor,
)
def test_classes_discovered():
orc = _build()
assert set(orc.classes) == {("G4_PbWO4", 0), ("G4_AIR", -1)}
def test_query_labels_inside_and_outside():
orc = _build()
pos = np.array([[0.0, 0.0, 0.0], [150.0, 150.0, 150.0]])
material, layer_id, _ = orc.query(pos)
assert material[0] == "G4_PbWO4" and layer_id[0] == 0
assert material[1] == "G4_AIR" and layer_id[1] == -1
def test_escape_flag_fires_far_from_data():
orc = _build()
pos = np.array([[0.0, 0.0, 0.0], [1e5, 0.0, 0.0]])
_, _, escaped = orc.query(pos)
assert not escaped[0]
assert escaped[1]
def test_query_empty():
orc = _build()
material, layer_id, escaped = orc.query(np.empty((0, 3)))
assert len(material) == len(layer_id) == len(escaped) == 0
def test_query_bad_shape_raises():
orc = _build()
with pytest.raises(ValueError):
orc.query(np.zeros((4, 2)))
def test_save_load_roundtrip(tmp_path):
orc = _build()
p = tmp_path / "oracle.pkl"
orc.save(p)
loaded = g.GeometryOracle.load(p)
pos = np.array([[0.0, 0.0, 0.0], [150.0, 150.0, 150.0], [1e5, 0.0, 0.0]])
m0, l0, e0 = orc.query(pos)
m1, l1, e1 = loaded.query(pos)
assert (m0 == m1).all() and (l0 == l1).all() and (e0 == e1).all()
assert loaded.escape_threshold == orc.escape_threshold
assert loaded.classes == orc.classes
def test_svm_method_has_escape_tree():
orc = _build(subsample=4000, method="svm")
# SVM cannot answer NN-distance, so a reference tree backs the escape test.
assert orc._ref_tree is not None
_, _, escaped = orc.query(np.array([[1e5, 0.0, 0.0]]))
assert escaped[0]
+144
View File
@@ -0,0 +1,144 @@
"""Tests for the autoregressive shower rollout driver."""
from pathlib import Path
from unittest.mock import patch
import numpy as np
import pytest
import torch
from giant.constants import TERM_ESCAPED, TERM_MAX_STEPS
from giant.data.transforms import Normalizer
from giant.model.network import DenoisingMLP, SecondaryDecoder
from giant.rollout import make_seed_frontier, rollout
pytest.importorskip("sklearn")
from giant import geometry as g # noqa: E402
PDG_MAP = {22: 0, 11: 1, -11: 2}
MAT_MAP = {"G4_AIR": 0, "G4_PbWO4": 1}
def _models():
s1 = DenoisingMLP(pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2)
s2 = SecondaryDecoder(pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2)
return s1.eval(), s2.eval()
def _norms():
rng = np.random.default_rng(0)
cond = Normalizer().fit(rng.standard_normal((1000, 8)).astype(np.float32))
tgt = Normalizer().fit(rng.standard_normal((1000, 9)).astype(np.float32))
return cond, tgt
def _oracle():
rng = np.random.default_rng(0)
pos = rng.uniform(-200, 200, (20000, 3)).astype(np.float32)
inside = (np.abs(pos) < 100).all(axis=1)
mat = np.where(inside, "G4_PbWO4", "G4_AIR").astype(object)
lay = np.where(inside, 0, -1).astype(np.int64)
with patch.object(g, "_iter_point_batches", lambda p: iter([(pos, mat, lay)])):
return g.build_geometry_oracle([Path("x")], subsample=20000)
def _seeds(n=6):
return {
"event_id": np.arange(n, dtype=np.int64),
"pdg": np.full(n, 11, dtype=np.int64),
"pre_pos": np.zeros((n, 3)),
"pre_E": np.linspace(30.0, 90.0, n),
"pre_dir": np.tile([0.0, 0.0, 1.0], (n, 1)),
}
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,
escape_threshold=escape_threshold,
)
def test_seed_frontier_track_ids():
seeds = _seeds(3)
fr, counts = make_seed_frontier(**seeds)
assert (fr["track_id"] == [0, 0, 0]).all() # one primary per event -> id 0
assert (fr["parent_id"] == -1).all()
assert (fr["generation"] == 0).all()
assert all(counts[e] == 1 for e in range(3))
# pre_dir is normalised.
np.testing.assert_allclose(np.linalg.norm(fr["pre_dir"], axis=1), 1.0, atol=1e-6)
def test_rollout_terminates_and_has_rows():
rec = _run()
assert len(rec["event_id"]) > 0
# Every seed event appears.
assert set(rec["event_id"].tolist()) == set(range(6))
def test_max_steps_respected():
# Disable the energy cutoff so tracks survive long enough to hit the step cap.
rec = _run(max_steps=5, energy_cutoff=0.0)
assert rec["step_no"].max() <= 5
assert (rec["termination_reason"] == TERM_MAX_STEPS).any()
def test_energy_conserved_deposit_plus_leak():
seeds = _seeds()
rec = _run(seeds=seeds)
for i, ev in enumerate(seeds["event_id"]):
m = rec["event_id"] == ev
dep = rec["edep"][m].sum()
leak = rec["pre_E"][m & (rec["termination_reason"] == TERM_ESCAPED)].sum()
assert dep + leak == pytest.approx(seeds["pre_E"][i], rel=1e-4)
def test_secondaries_have_valid_parents():
rec = _run()
orphans = 0
for ev in np.unique(rec["event_id"]):
m = rec["event_id"] == ev
tids = set(rec["track_id"][m].tolist())
for pid in rec["parent_id"][m]:
if pid >= 0 and pid not in tids:
orphans += 1
assert orphans == 0
# At least one secondary (generation > 0) is produced by the tiny model.
assert (rec["generation"] > 0).any()
def test_escape_terminates_immediately():
# A tight escape threshold makes even the seed position (origin) escape.
rec = _run(escape_threshold=1e-3)
assert (rec["termination_reason"] == TERM_ESCAPED).all()
assert rec["step_no"].max() == 0
def test_output_schema_complete():
from giant.rollout import _RECORD_KEYS
rec = _run()
assert set(rec.keys()) == set(_RECORD_KEYS)
n = len(rec["event_id"])
assert all(len(v) == n for v in rec.values())
def test_max_tracks_cap_conserves_energy():
# A very small cap forces sub-cap secondaries to deposit in place; energy
# must still balance.
seeds = _seeds()
rec = _run(seeds=seeds, max_tracks_per_event=3)
for i, ev in enumerate(seeds["event_id"]):
m = rec["event_id"] == ev
dep = rec["edep"][m].sum()
leak = rec["pre_E"][m & (rec["termination_reason"] == TERM_ESCAPED)].sum()
assert dep + leak == pytest.approx(seeds["pre_E"][i], rel=1e-4)
assert len(np.unique(rec["track_id"][m])) <= 3
Generated
+124 -2
View File
@@ -470,14 +470,18 @@ dev = [
{ name = "polars" },
{ name = "pytest" },
{ name = "ruff" },
{ name = "scikit-learn" },
{ name = "ty" },
{ name = "uproot" },
]
geometry = [
{ name = "scikit-learn" },
]
[package.metadata]
requires-dist = [
{ name = "awkward", marker = "extra == 'convert'", specifier = ">=2.6,<3" },
{ name = "giant", extras = ["convert", "analysis"], marker = "extra == 'dev'" },
{ name = "giant", extras = ["convert", "analysis", "geometry"], marker = "extra == 'dev'" },
{ name = "ipykernel", marker = "extra == 'analysis'", specifier = ">=7.3.0" },
{ name = "matplotlib", marker = "extra == 'analysis'", specifier = ">=3.8,<4" },
{ name = "numpy", specifier = ">=1.26,<3" },
@@ -488,6 +492,7 @@ requires-dist = [
{ name = "pytest", marker = "extra == 'dev'", specifier = ">=8,<10" },
{ name = "pyyaml", specifier = ">=6,<7" },
{ name = "ruff", marker = "extra == 'dev'", specifier = ">=0.15,<1" },
{ name = "scikit-learn", marker = "extra == 'geometry'", specifier = ">=1.4,<2" },
{ name = "torch", marker = "extra == 'cpu'", specifier = ">=2.3,<2.4", index = "https://download.pytorch.org/whl/cpu", conflict = { package = "giant", extra = "cpu" } },
{ name = "torch", marker = "extra == 'cuda'", specifier = ">=2.3,<2.4", index = "https://download.pytorch.org/whl/cu118", conflict = { package = "giant", extra = "cuda" } },
{ name = "tqdm", specifier = ">=4.60,<5" },
@@ -495,7 +500,7 @@ requires-dist = [
{ name = "typer", specifier = ">=0.12,<1" },
{ name = "uproot", marker = "extra == 'convert'", specifier = ">=5.3,<6" },
]
provides-extras = ["cpu", "cuda", "dev", "convert", "analysis"]
provides-extras = ["cpu", "cuda", "dev", "geometry", "convert", "analysis"]
[[package]]
name = "iniconfig"
@@ -600,6 +605,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/62/a1/3d680cbfd5f4b8f15abc1d571870c5fc3e594bb582bc3b64ea099db13e56/jinja2-3.1.6-py3-none-any.whl", hash = "sha256:85ece4451f492d0c13c5dd7c13a64681a86afae63a5f347908daf103ce6d2f67", size = 134899, upload-time = "2025-03-05T20:05:00.369Z" },
]
[[package]]
name = "joblib"
version = "1.5.3"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/41/f2/d34e8b3a08a9cc79a50b2208a93dce981fe615b64d5a4d4abee421d898df/joblib-1.5.3.tar.gz", hash = "sha256:8561a3269e6801106863fd0d6d84bb737be9e7631e33aaed3fb9ce5953688da3", size = 331603, upload-time = "2025-12-15T08:41:46.427Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/7b/91/984aca2ec129e2757d1e4e3c81c3fcda9d0f85b74670a094cc443d9ee949/joblib-1.5.3-py3-none-any.whl", hash = "sha256:5fc3c5039fc5ca8c0276333a188bbd59d6b7ab37fe6632daa76bc7f9ec18e713", size = 309071, upload-time = "2025-12-15T08:41:44.973Z" },
]
[[package]]
name = "jupyter-client"
version = "8.9.1"
@@ -891,6 +905,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/43/e3/7d92a15f894aa0c9c4b49b8ee9ac9850d6e63b03c9c32c0367a13ae62209/mpmath-1.3.0-py3-none-any.whl", hash = "sha256:a0b2b9fe80bbcd81a6647ff13108738cfb482d481d826cc0e02f5b35e5c88d2c", size = 536198, upload-time = "2023-03-07T16:47:09.197Z" },
]
[[package]]
name = "narwhals"
version = "2.23.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/e8/ac/66ed1fc6e38a0c0f330627ec5c5d597990d6159b6712b82af0ad2c65f06c/narwhals-2.23.0.tar.gz", hash = "sha256:13e7ff5b4bb4a2f77b907c2e4d8a76e273dfc1323a3c997440a2f9fd26aed408", size = 656209, upload-time = "2026-07-01T11:21:53.278Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/f4/4e/afc8c31605cb8be1d3bb4438c4d979daa104dab6306cd2b87abe9c3a7299/narwhals-2.23.0-py3-none-any.whl", hash = "sha256:769e7b9ab102c93d8fa019f6b4cd1a657909b04a20bf6210e5a35aae06814ae9", size = 458938, upload-time = "2026-07-01T11:21:51.677Z" },
]
[[package]]
name = "nest-asyncio2"
version = "1.7.2"
@@ -1571,6 +1594,96 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/2d/c7/c53e8dbff9c9dc4b7928773421ae294a5d28fcb8dcda1a089579d3a7e510/ruff-0.15.17-py3-none-win_arm64.whl", hash = "sha256:f3be1fbb34bcdfd146240d8fb92a709d4c2c8191348580a3c044ec60fa0b4456", size = 11355275, upload-time = "2026-06-11T17:54:43.635Z" },
]
[[package]]
name = "scikit-learn"
version = "1.9.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "joblib" },
{ name = "narwhals" },
{ name = "numpy" },
{ name = "scipy" },
{ name = "threadpoolctl" },
]
sdist = { url = "https://files.pythonhosted.org/packages/fa/6f/37092bdb25f712817231799fc5674d8e704066a8a70c1d2d40517e18b4ab/scikit_learn-1.9.0.tar.gz", hash = "sha256:8833266989d3a5110178a9fae30783675460724d0e1efb13b14901d2c660c557", size = 7750767, upload-time = "2026-06-02T11:54:32.706Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ac/20/75f915ff375d6249e6550ac740fdbbd66159a068fd3af1400ff62036b07a/scikit_learn-1.9.0-cp312-cp312-macosx_10_13_x86_64.whl", hash = "sha256:2bd41b0d201bc81575531b96b713d3eb5e5f50fb0b82101ff0f92294fdc236ac", size = 8741122, upload-time = "2026-06-02T11:53:24.08Z" },
{ url = "https://files.pythonhosted.org/packages/cc/d5/2b5148f2279196775e1db2aeb85d14b70ac80e7e32b3b28e7ebeafb0901d/scikit_learn-1.9.0-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:5be45aa4a42a68a533913a6ed736cf309de2226411c79ef8d609a5456f1939b1", size = 8261512, upload-time = "2026-06-02T11:53:27.183Z" },
{ url = "https://files.pythonhosted.org/packages/a0/ee/5adbc77656b71f9456a2f5a7a9fdb4bcf9207a6b962889f1c2f9323afa4e/scikit_learn-1.9.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5e50ed4da51974e86e940690e9a3d82e729b62b5a49f7c9bac534d515d39d86f", size = 8837603, upload-time = "2026-06-02T11:53:30.328Z" },
{ url = "https://files.pythonhosted.org/packages/6c/c2/63fdda36c56437eeb44aaf9493c8bcd62ce230ab1598924fc626ffbfa943/scikit_learn-1.9.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:056c92bb67ad4c28463c2f2653d9701449201e7e7a9e94e321be0f71c4fef2b8", size = 9132097, upload-time = "2026-06-02T11:53:33.456Z" },
{ url = "https://files.pythonhosted.org/packages/83/a4/c8e67227c680e2259c8864ae72ff48b06e16a6f51253a22167aa02a8aa4e/scikit_learn-1.9.0-cp312-cp312-win_amd64.whl", hash = "sha256:4306775fad04cc4b472a1b15af1ae9cede1540fbfcc17fbce3767cd8dc7ae283", size = 8211173, upload-time = "2026-06-02T11:53:36.602Z" },
{ url = "https://files.pythonhosted.org/packages/cf/fd/3c0863792e98e67e9184aa4029288a175935eb65443afcd30d4f143450cf/scikit_learn-1.9.0-cp312-cp312-win_arm64.whl", hash = "sha256:26e22435f63bcdcf396b574273f29f13dd531f5ea035801f5be10ba1540a4e60", size = 7867451, upload-time = "2026-06-02T11:53:39.075Z" },
{ url = "https://files.pythonhosted.org/packages/3c/01/cf3310626b6d48d3e9be69a1223f9180360b5e6edb045f50fade723ce494/scikit_learn-1.9.0-cp313-cp313-macosx_10_13_x86_64.whl", hash = "sha256:80746d63bd4b6eaca54d36fe5feaf4d28bb38dc6f9470f81c7cad7c40155f119", size = 8705188, upload-time = "2026-06-02T11:53:41.964Z" },
{ url = "https://files.pythonhosted.org/packages/3e/04/5acd7ae280c5f93b6ac5ef6cdec14eef4c8d1cd91d85b3292989c94d96b1/scikit_learn-1.9.0-cp313-cp313-macosx_12_0_arm64.whl", hash = "sha256:5b934c45c252844a91d69fda3a34cff5e7307e1db10d77cb10a3980312c74713", size = 8228299, upload-time = "2026-06-02T11:53:44.817Z" },
{ url = "https://files.pythonhosted.org/packages/0c/39/ffe829a5b8ecb40a518724a997794657fdc354ada5e8fe8e64d998c0bac9/scikit_learn-1.9.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:38c3dcb9a1ffb85505ec53d54c7b4aea0cff70050425a7760c2af661ac85df05", size = 8789690, upload-time = "2026-06-02T11:53:47.461Z" },
{ url = "https://files.pythonhosted.org/packages/1f/88/8dab5de10c638c083772a6be83a3d8106ced492f74a928c8693638e5bb50/scikit_learn-1.9.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:da76d09304a4706db7cc1e3ebaa3b6b98a67365cc11d2996c4f1e58ba47df714", size = 9087723, upload-time = "2026-06-02T11:53:50.702Z" },
{ url = "https://files.pythonhosted.org/packages/20/3f/7917ca72464038f6240ec70c29f94862d08a34a74291ae4d4ec5eb8186a0/scikit_learn-1.9.0-cp313-cp313-win_amd64.whl", hash = "sha256:5808d98f15c6bf6d9d96d2348c1997392a5888ce7097e664105f930c4bca1277", size = 8184330, upload-time = "2026-06-02T11:53:53.396Z" },
{ url = "https://files.pythonhosted.org/packages/78/c7/15739eb2f61fda3c54639e9942414e5a19ad8a8d1f5a3266afad7cb7df80/scikit_learn-1.9.0-cp313-cp313-win_arm64.whl", hash = "sha256:d77f54c017633791bc0225a43e2f8d03745fdcfe4880268fcc4df15f505dec2e", size = 7840653, upload-time = "2026-06-02T11:53:56.035Z" },
{ url = "https://files.pythonhosted.org/packages/f4/7d/c9a35cf59b20a86fec24d306f1547b78dec194b08d367ce2a3e4854169d9/scikit_learn-1.9.0-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:9656acd4e93f74e0b66c8a36c88830a99252dfa900044d36bc2212ae89a47162", size = 8713289, upload-time = "2026-06-02T11:53:58.788Z" },
{ url = "https://files.pythonhosted.org/packages/3c/a7/552a7821597c632b907f7bfe8f36f9f572777af8ef8a48353041cf8e091a/scikit_learn-1.9.0-cp314-cp314-macosx_12_0_arm64.whl", hash = "sha256:24360002ae845e7866522b0a5bbf690802e7bc388cac8663502e78aa98598aa2", size = 8245141, upload-time = "2026-06-02T11:54:01.694Z" },
{ url = "https://files.pythonhosted.org/packages/7d/79/f4a0c4fe9711154cddabf913471153af79056382ddc612cfe5ee0ff4b72e/scikit_learn-1.9.0-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5162ad10a418c8a282dde04c9aa06965de3e9a65f33c1440c0ae69bb1a09d913", size = 8847671, upload-time = "2026-06-02T11:54:04.448Z" },
{ url = "https://files.pythonhosted.org/packages/f0/af/4d72d9e475ac83719160c662619e4bf7b95c19507cd582e7d0167a3c3dae/scikit_learn-1.9.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1fea2cc5677ab49d6f5bade978c866da44957b712d92e9635e8b4f723013c3cb", size = 9118104, upload-time = "2026-06-02T11:54:07.205Z" },
{ url = "https://files.pythonhosted.org/packages/a2/d5/6a58eea2cb9abbb9b3f2bb8b2cfb3243d1152d69f442d256c7af71304769/scikit_learn-1.9.0-cp314-cp314-win_amd64.whl", hash = "sha256:64fa347efc1c839c487433e40c5144d38c336e8a2b59c81aa8660373945c2673", size = 8290674, upload-time = "2026-06-02T11:54:10.087Z" },
{ url = "https://files.pythonhosted.org/packages/65/5b/d4c879cf358f1187141cf90ced473f087183489090244f50c124a2ee478b/scikit_learn-1.9.0-cp314-cp314-win_arm64.whl", hash = "sha256:1b944b6db288f6b926e3650026ddafb988929de95d11fc2cc5fa117773c9ba42", size = 7978807, upload-time = "2026-06-02T11:54:12.769Z" },
{ url = "https://files.pythonhosted.org/packages/8a/43/bfae3121ec67ae09150d453c442c7c1cc166e9aefe056e6ab3b7728a5cfc/scikit_learn-1.9.0-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:4ccacf04ca5f4b492158a5f28afe0ace43f81b2571e4b9a66d34848b46128949", size = 9031941, upload-time = "2026-06-02T11:54:15.436Z" },
{ url = "https://files.pythonhosted.org/packages/75/b0/20a4546eb17f3b25d3c66df15810411c14ed5065bcfab50b53c96fb627b2/scikit_learn-1.9.0-cp314-cp314t-macosx_12_0_arm64.whl", hash = "sha256:ee1a8db2c18c08e34c7412d4b10be1cac214cd4ea7dc9715a6a327eb49a37c96", size = 8613528, upload-time = "2026-06-02T11:54:18.842Z" },
{ url = "https://files.pythonhosted.org/packages/18/3c/e440e039bb82cd19004edaaad00acbde0fb9b461083c3ecf37941c557312/scikit_learn-1.9.0-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:147e9329ef0e39f75d4cffa02b2aa48d827832684926cd5210d9a2cb5c57246b", size = 8855050, upload-time = "2026-06-02T11:54:21.699Z" },
{ url = "https://files.pythonhosted.org/packages/43/26/b341b8dab5998da6270a3a42c2152c578501354d36f944b5856757035ef8/scikit_learn-1.9.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:5bad8f8b9950321b54c965fdcbac6c6c55e79e16646b49977bcf3668d3870a1a", size = 9097190, upload-time = "2026-06-02T11:54:24.454Z" },
{ url = "https://files.pythonhosted.org/packages/fb/de/b650b4d69b84468cfa2e28a3ff7b8103743029e6446ce1a97fe060ef688c/scikit_learn-1.9.0-cp314-cp314t-win_amd64.whl", hash = "sha256:78fc56eafd4edb9575d2d8950d1dd152061abb573341a1cb7e099fc40f6c6666", size = 8963204, upload-time = "2026-06-02T11:54:27.428Z" },
{ url = "https://files.pythonhosted.org/packages/ee/f3/ff83d76d7418112e5a61326443cdda87be3545dd8d6599c95b2481a4419e/scikit_learn-1.9.0-cp314-cp314t-win_arm64.whl", hash = "sha256:051075bda8b7aab87b1906ab3d4740a1e1224a19d7b3781a576736edc94e76aa", size = 8222661, upload-time = "2026-06-02T11:54:30.192Z" },
]
[[package]]
name = "scipy"
version = "1.18.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "numpy" },
]
sdist = { url = "https://files.pythonhosted.org/packages/a7/25/c2700dfaf6442b4effaa91af24ebce5dc9d31bb4a69706313aae70d72cd0/scipy-1.18.0.tar.gz", hash = "sha256:67b2ad2ad54c72ca6d04975a9b2df8c3638c34ddd5b28738e94fc2b57929d378", size = 30774447, upload-time = "2026-06-19T15:01:43.456Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/6a/19/ca10ead60b0acc80b2b833c2c4a4f2ff753d0f58b811f70d911c7e94a25c/scipy-1.18.0-cp312-cp312-macosx_10_15_x86_64.whl", hash = "sha256:7bd21faaf5a1a3b2eff922d02db5f191b99a6518db9078a8fb23169f6d22259a", size = 31056519, upload-time = "2026-06-19T14:59:45.203Z" },
{ url = "https://files.pythonhosted.org/packages/96/72/1e6442a00cd2924d361aa1b642ab6373ec35c6fabf311a760be9f76e0f13/scipy-1.18.0-cp312-cp312-macosx_12_0_arm64.whl", hash = "sha256:265915e79107de9f946b855e50d7470d5893ec3f54b342e1aa6201cbdcd8bb6b", size = 28681889, upload-time = "2026-06-19T14:59:48.103Z" },
{ url = "https://files.pythonhosted.org/packages/9b/2d/11dd93d21e147a73ba22bd75c0b9208d3a2e0ec76d53170ce7d9029b1015/scipy-1.18.0-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:9ab7b758be6940954a713ee466e2043e9f6e2ed965c1fce5c91039f4be3d90a9", size = 20423580, upload-time = "2026-06-19T14:59:50.665Z" },
{ url = "https://files.pythonhosted.org/packages/9c/01/93552f75e0d2a7dd115a45e59209c51e8d514daff02fc887d2623be06fe1/scipy-1.18.0-cp312-cp312-macosx_14_0_x86_64.whl", hash = "sha256:97b6cddaaee0a779ef6b5ca83c9604b27cc16b2b8fc22c142652df8793319fb8", size = 23054441, upload-time = "2026-06-19T14:59:53.564Z" },
{ url = "https://files.pythonhosted.org/packages/3c/23/21f5e703643d66f21faa6b4c73195bfcad70c55efcb4f1ab327cd7c4101a/scipy-1.18.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:52a96e21517c7292375c0e27dd796a811f03fcea5fd4d108fdfea8145dcf17ab", size = 33968720, upload-time = "2026-06-19T14:59:56.415Z" },
{ url = "https://files.pythonhosted.org/packages/dd/aa/1b939f6c67ed68635bb538e6752d3dacc02f66535182e939a89581a44e9c/scipy-1.18.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:1f55797419e16e7f30cf88ffb3113ce0467f00cfe3f70d5c281730b21769bfc2", size = 35287115, upload-time = "2026-06-19T14:59:59.411Z" },
{ url = "https://files.pythonhosted.org/packages/b6/ff/eec46be7e9234208f801062b53e1983085eddebd693f6c9bfb03b459830d/scipy-1.18.0-cp312-cp312-musllinux_1_2_aarch64.whl", hash = "sha256:ad033410e2e0672ffdc1042110cef20e1c46f8fd0616cee1d44d8d58fad8fc11", size = 35577989, upload-time = "2026-06-19T15:00:02.235Z" },
{ url = "https://files.pythonhosted.org/packages/84/ca/210d4759c7210bb7d269437421959b39a33434e2776b60c5cb8a763bb30a/scipy-1.18.0-cp312-cp312-musllinux_1_2_x86_64.whl", hash = "sha256:4a55985d54c769c872e64b7f4c8a81cc30ef700cc04296abbbf3705439c126de", size = 37421717, upload-time = "2026-06-19T15:00:05.102Z" },
{ url = "https://files.pythonhosted.org/packages/2b/54/9a9edb45345bd6744da5ddfb6628e5d5185920494c6a67ec45b6381004cb/scipy-1.18.0-cp312-cp312-win_amd64.whl", hash = "sha256:71ccc8faa2dd16ac310233203474a8b5cb67f10dedd54a3116d34943f4b19132", size = 36597428, upload-time = "2026-06-19T15:00:08.112Z" },
{ url = "https://files.pythonhosted.org/packages/99/0e/33f32a2a58987e26aec0f7df252cbbad1e90ae77bdbc76f40dd4ed0cf0ea/scipy-1.18.0-cp312-cp312-win_arm64.whl", hash = "sha256:d88363fd9d8fbd3511bd273f1a49efb2a540773ddf92a91d57498ce7dd7f3e76", size = 24351481, upload-time = "2026-06-19T15:00:11.103Z" },
{ url = "https://files.pythonhosted.org/packages/05/52/9c0136c2de7ae0779b7b366447766cec6d9f0702c56bb8ffeb04c8fd3af4/scipy-1.18.0-cp313-cp313-macosx_10_15_x86_64.whl", hash = "sha256:09143f676d157d9f546d663504ef9c1becb819824f1afc018814176411942446", size = 31036107, upload-time = "2026-06-19T15:00:14.03Z" },
{ url = "https://files.pythonhosted.org/packages/02/73/0291a64843270f4efb86cdcf2ee0f2048631b65ec6b405398b2b4dbf11bf/scipy-1.18.0-cp313-cp313-macosx_12_0_arm64.whl", hash = "sha256:5efe260f69417b97ddae455bfb5a95e8359f7f66ad7fa9522a60feb66f169520", size = 28663303, upload-time = "2026-06-19T15:00:16.819Z" },
{ url = "https://files.pythonhosted.org/packages/d3/0f/10ffa0b697a572f4e0d48b92a88895d366422f019f723e7e14a84c050dac/scipy-1.18.0-cp313-cp313-macosx_14_0_arm64.whl", hash = "sha256:68363b7eaacd8b5dd426df56d782cc156468ac79a127a1b87ca597d6e2e82197", size = 20404960, upload-time = "2026-06-19T15:00:19.635Z" },
{ url = "https://files.pythonhosted.org/packages/7e/d2/e896cea21ba8edd6c81d4c55b1ffcc717e79698dcbebf9641b4cfb4c6622/scipy-1.18.0-cp313-cp313-macosx_14_0_x86_64.whl", hash = "sha256:c5557d8be5da8e41353fcd4d21491fdbab83b062fc579e94dc09a7c8ab4f669b", size = 23034074, upload-time = "2026-06-19T15:00:22.107Z" },
{ url = "https://files.pythonhosted.org/packages/ea/b2/e83ea34279a52c03374477c74006256ec78df65fc877baa4617d6de1d202/scipy-1.18.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:0d13bca67c096d89fb95ced0d8921807300fce0275643aef9533cc63a0773468", size = 33942038, upload-time = "2026-06-19T15:00:24.964Z" },
{ url = "https://files.pythonhosted.org/packages/f6/af/e8fe5fb136f51e2b01678b92cb4106d10d8cd68ec147ead2e7cb0ac75398/scipy-1.18.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:a46f9273dbd0eb1cefba61c9b8648b4dfe3cbc14a080176f9a73e44b8336dc7f", size = 35266390, upload-time = "2026-06-19T15:00:28.059Z" },
{ url = "https://files.pythonhosted.org/packages/3a/49/2c5cbb907b56695fc67517811d1db234dfd83381a84814ec220aded2794d/scipy-1.18.0-cp313-cp313-musllinux_1_2_aarch64.whl", hash = "sha256:5aba46108853ddfc77906b6557aac839d2b52e900c1d72a1180adaaab58d265f", size = 35551324, upload-time = "2026-06-19T15:00:31.014Z" },
{ url = "https://files.pythonhosted.org/packages/bb/73/eda39f7a2d306ff0ffc574afd13c0bbb6d10a603d9a413998ee269487a80/scipy-1.18.0-cp313-cp313-musllinux_1_2_x86_64.whl", hash = "sha256:b6f758e35f12757b5d95c00bc6de2438e229c2664b7a92e96f205959d9f2dfa4", size = 37404785, upload-time = "2026-06-19T15:00:34.072Z" },
{ url = "https://files.pythonhosted.org/packages/b7/d2/ae881ee28d014f38e0ccbfd974a06a919ba9af34f1f74bf42b5301891d63/scipy-1.18.0-cp313-cp313-win_amd64.whl", hash = "sha256:1afac4a847207c7ff8efd321734a50b06d0280b3b2a2c0fc2f413101747ad7c7", size = 36554943, upload-time = "2026-06-19T15:00:36.903Z" },
{ url = "https://files.pythonhosted.org/packages/70/3a/21154e2d54eb3639c6bf4dbae2e531c68356bfe95990daa30df33b30d556/scipy-1.18.0-cp313-cp313-win_arm64.whl", hash = "sha256:c5dbddf60e58c2312316d097271a8e73d40eaf2eabfa4d95ed7d3695bbf2ce7b", size = 24350911, upload-time = "2026-06-19T15:00:40.062Z" },
{ url = "https://files.pythonhosted.org/packages/78/b5/915a19b3de2f7430062b509653563db1633ddbb6f021b06731521115d4e2/scipy-1.18.0-cp314-cp314-macosx_10_15_x86_64.whl", hash = "sha256:4c256ee70c0d1a8a2ace807e199ccd4e3f57037433842abb3fb36bc17eaa9578", size = 31036253, upload-time = "2026-06-19T15:00:43.216Z" },
{ url = "https://files.pythonhosted.org/packages/d7/88/b72def7262e150d16be13fca37a96481138d624e700340bc3362a7588929/scipy-1.18.0-cp314-cp314-macosx_12_0_arm64.whl", hash = "sha256:2ef3abc54a4ffc53765374b0d5728532dfdd2585ed23f6b11c206a1f0b1b9af8", size = 28673758, upload-time = "2026-06-19T15:00:46.663Z" },
{ url = "https://files.pythonhosted.org/packages/91/02/2e636a61a525632c373cf6a9c24442a3ffb79e364d38e98b32042964ac32/scipy-1.18.0-cp314-cp314-macosx_14_0_arm64.whl", hash = "sha256:f2a6af57bd9e4a75d70e4117e78a1bbee84f79ae3fbb6d0111005d6ebcc4cb8d", size = 20415514, upload-time = "2026-06-19T15:00:49.399Z" },
{ url = "https://files.pythonhosted.org/packages/c9/b6/2135974442f6aba159d9d39d774a1c8cb19947016725d69fecc685df45bf/scipy-1.18.0-cp314-cp314-macosx_14_0_x86_64.whl", hash = "sha256:3f1ac564d3bf6c03d861d2cd87a1bea0da2887136f7fb1bf519c05a8971452d6", size = 23034398, upload-time = "2026-06-19T15:00:51.941Z" },
{ url = "https://files.pythonhosted.org/packages/f6/e6/ba89ec5abf6ee9257c0d1ec985573f3ae32742c24bc03e016388a40b1b15/scipy-1.18.0-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:40395a5fcd1abee49a5c7aaa98c29db393eedc835138560a588c47ec16156690", size = 33998032, upload-time = "2026-06-19T15:00:54.838Z" },
{ url = "https://files.pythonhosted.org/packages/7f/c4/bc41eb19b0fd0db868f4132920879019318d80cc522ad8f2bca4611af808/scipy-1.18.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8ca01e8ae69f1b18e9a58d91afead31be3cef0dd905a10249dac559ee15460a0", size = 35283333, upload-time = "2026-06-19T15:00:58.152Z" },
{ url = "https://files.pythonhosted.org/packages/53/a4/cbdeef6eb3830a8462a9d4ada814de5fc984345cc9ecf17cbec51a036f1e/scipy-1.18.0-cp314-cp314-musllinux_1_2_aarch64.whl", hash = "sha256:7a7f3b01647384dbc3a711e8c6778e0aabbe93959249fef5c7393396bcac0867", size = 35610216, upload-time = "2026-06-19T15:01:01.155Z" },
{ url = "https://files.pythonhosted.org/packages/80/4d/b2b82502b65f661d1b789c1665dcdf315d5f12194e06fc0b37946294ebae/scipy-1.18.0-cp314-cp314-musllinux_1_2_x86_64.whl", hash = "sha256:6aa94e78ec192a30063a5e72e561c28af769dc311190b24fe91774eff1969709", size = 37418960, upload-time = "2026-06-19T15:01:04.155Z" },
{ url = "https://files.pythonhosted.org/packages/93/3e/902d836831474b0ab5a37d16404f7bc5fafd9efba632890e271ba952635f/scipy-1.18.0-cp314-cp314-win_amd64.whl", hash = "sha256:2d8bbdc6c817f5b4006a54d799d4f5bab6f910193cbb9a1ff310833d4d270f61", size = 37288845, upload-time = "2026-06-19T15:01:07.822Z" },
{ url = "https://files.pythonhosted.org/packages/b6/43/8d73b337a3bdb14daa0314f0434210747c02d79d729ce1777574a817dcf6/scipy-1.18.0-cp314-cp314-win_arm64.whl", hash = "sha256:18e9575f1569b2c54174e6159d32942e03731177f63dce7975f0a0c88d102f5b", size = 24988971, upload-time = "2026-06-19T15:01:11.076Z" },
{ url = "https://files.pythonhosted.org/packages/b4/b4/f11918b0508a2787031a0499a03fbe3546f3bb5ca05d01038c45b278c09a/scipy-1.18.0-cp314-cp314t-macosx_10_15_x86_64.whl", hash = "sha256:f351e0dd702687d12a402b867a1b4146a256923e1c38317cbc472f6372b94707", size = 31399325, upload-time = "2026-06-19T15:01:13.723Z" },
{ url = "https://files.pythonhosted.org/packages/7b/d1/1f287b57c0ff0ee5185dff3946d92c8017d39b0e431f0ae79a3ff1859512/scipy-1.18.0-cp314-cp314t-macosx_12_0_arm64.whl", hash = "sha256:7c7a51b33ce387193c97f228320cf8e87361daa1bba750638677729598b3e677", size = 29092110, upload-time = "2026-06-19T15:01:16.908Z" },
{ url = "https://files.pythonhosted.org/packages/ff/1a/7b74eb6c392fdcb27d414c0e7558a6d0231eb3b6d73571f479bb81ea8794/scipy-1.18.0-cp314-cp314t-macosx_14_0_arm64.whl", hash = "sha256:84031d7b052a54fae2f8632e0ec802073d385476eb9a63079bce6e23ef9283d4", size = 20833811, upload-time = "2026-06-19T15:01:20.488Z" },
{ url = "https://files.pythonhosted.org/packages/7c/ad/f3941716320a7b9cb4d68734a903b45fe16eff5fb7da7e16f2e619304979/scipy-1.18.0-cp314-cp314t-macosx_14_0_x86_64.whl", hash = "sha256:56abf29a7c067dde59be8b9a22d606a4ea1b2f2a4b756d9d903c62818f5dacce", size = 23396644, upload-time = "2026-06-19T15:01:23.364Z" },
{ url = "https://files.pythonhosted.org/packages/22/22/1446b62ffe07f9719b7d9b1b6a4e05a772833ae8f441fe4c22c34c9b250f/scipy-1.18.0-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:1ad44305cfa24b1ba5803cbbebf033590ccbac1aa5d612d727b785325ab408b0", size = 34079318, upload-time = "2026-06-19T15:01:26.002Z" },
{ url = "https://files.pythonhosted.org/packages/56/3b/b87da667098bb470fa30c7011b0ba351ee976dd395c78798c66e941665a3/scipy-1.18.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:945c1761b93f38d7f99ae81ae80c63e621471608c7eeead563f6df025585cd58", size = 35324320, upload-time = "2026-06-19T15:01:28.881Z" },
{ url = "https://files.pythonhosted.org/packages/f8/a1/c7932f91909759b0267f75fdea34e91309f96b895757534b76a90b6b4344/scipy-1.18.0-cp314-cp314t-musllinux_1_2_aarch64.whl", hash = "sha256:1a4441f15d620578772a49e5ab48c0ee1f7a0220e387110283062729136b2553", size = 35699541, upload-time = "2026-06-19T15:01:31.968Z" },
{ url = "https://files.pythonhosted.org/packages/f7/86/5185061a1fcc41d18c5dc2463969b3a3964b31d9ac67b2fb05d4c7ff7670/scipy-1.18.0-cp314-cp314t-musllinux_1_2_x86_64.whl", hash = "sha256:9aac6192fac56bf2ca534389d24623f07b39ff83317d58287285e7fbd622ff76", size = 37472480, upload-time = "2026-06-19T15:01:35.136Z" },
{ url = "https://files.pythonhosted.org/packages/31/8e/f04c68e39919a010d34f2ee1367fd705b0a25a02f609d755f0bfbc0a15fc/scipy-1.18.0-cp314-cp314t-win_amd64.whl", hash = "sha256:e40baea28ae7f5475c779741e2d90b1247c78531207b49c7030e698ff81cee3f", size = 37365390, upload-time = "2026-06-19T15:01:38.091Z" },
{ url = "https://files.pythonhosted.org/packages/d5/19/969dc072906c84dd0a3b05dcf57ea750936087d7873549e408b35cfc3f97/scipy-1.18.0-cp314-cp314t-win_arm64.whl", hash = "sha256:368e0a705903c466aa5f08eefb39e6b1b6b2d659e7352a31fd9e2438365be0f8", size = 25279661, upload-time = "2026-06-19T15:01:40.817Z" },
]
[[package]]
name = "shellingham"
version = "1.5.4"
@@ -1626,6 +1739,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/9b/24/84ce997e8ae6296168a74d0d9c4dde572d90fb23fd7c0b219c30ff71e00e/tbb-2021.13.1-py3-none-win_amd64.whl", hash = "sha256:cbf024b2463fdab3ebe3fa6ff453026358e6b903839c80d647e08ad6d0796ee9", size = 286908, upload-time = "2024-08-07T15:09:05.677Z" },
]
[[package]]
name = "threadpoolctl"
version = "3.6.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/b7/4d/08c89e34946fce2aec4fbb45c9016efd5f4d7f24af8e5d93296e935631d8/threadpoolctl-3.6.0.tar.gz", hash = "sha256:8ab8b4aa3491d812b623328249fab5302a68d2d71745c8a4c719a2fcaba9f44e", size = 21274, upload-time = "2025-03-13T13:49:23.031Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/32/d5/f9a850d79b0851d1d4ef6456097579a9005b31fea68726a4ae5f2d82ddd9/threadpoolctl-3.6.0-py3-none-any.whl", hash = "sha256:43a0b8fd5a2928500110039e43a5eed8480b918967083ea48dc3ab9f13c4a7fb", size = 18638, upload-time = "2025-03-13T13:49:21.846Z" },
]
[[package]]
name = "torch"
version = "2.3.1"