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