diff --git a/CLAUDE.md b/CLAUDE.md index 4214e06..d997e58 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -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 diff --git a/analysis/export_rollout_observables.py b/analysis/export_rollout_observables.py new file mode 100644 index 0000000..f54e143 --- /dev/null +++ b/analysis/export_rollout_observables.py @@ -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") diff --git a/giant/analysis.py b/giant/analysis.py index f2e8944..c8c110b 100644 --- a/giant/analysis.py +++ b/giant/analysis.py @@ -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 diff --git a/giant/cli.py b/giant/cli.py index d9f8208..4b61b10 100644 --- a/giant/cli.py +++ b/giant/cli.py @@ -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() diff --git a/giant/constants.py b/giant/constants.py index c42b8f2..8b42522 100644 --- a/giant/constants.py +++ b/giant/constants.py @@ -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" diff --git a/giant/geometry.py b/giant/geometry.py new file mode 100644 index 0000000..1ca6303 --- /dev/null +++ b/giant/geometry.py @@ -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 diff --git a/giant/rollout.py b/giant/rollout.py new file mode 100644 index 0000000..9faa7af --- /dev/null +++ b/giant/rollout.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index 582f397..cf06f84 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/scripts/dwarf.py b/scripts/dwarf.py index f9f2889..3d91001 100644 --- a/scripts/dwarf.py +++ b/scripts/dwarf.py @@ -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, diff --git a/scripts/geometry_oracle.py b/scripts/geometry_oracle.py new file mode 100644 index 0000000..1e8517e --- /dev/null +++ b/scripts/geometry_oracle.py @@ -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 ` 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}") diff --git a/tests/test_geometry.py b/tests/test_geometry.py new file mode 100644 index 0000000..0f1d993 --- /dev/null +++ b/tests/test_geometry.py @@ -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] diff --git a/tests/test_rollout.py b/tests/test_rollout.py new file mode 100644 index 0000000..0fc252a --- /dev/null +++ b/tests/test_rollout.py @@ -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 diff --git a/uv.lock b/uv.lock index aba0c58..a197032 100644 --- a/uv.lock +++ b/uv.lock @@ -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"