diff --git a/giant/geometry.py b/giant/geometry.py index 1ca6303..ffe27da 100644 --- a/giant/geometry.py +++ b/giant/geometry.py @@ -3,13 +3,22 @@ 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. +learn it from positions sampled from a real steps dataset. -scikit-learn / joblib are an optional dependency (the `geometry` extra); they are -imported lazily so the core install stays lean. +miniCaloSim's detector is a stack of planar layer slabs along one axis (see +`physics/detector-design/minicalosim-geometry.md`), so `material`/`layer_id` +are a pure function of depth. The default ("slab") method exploits this: fit a +1D lookup table of depth-axis segment boundaries and do an exact O(log +#segments) binary search per query, with escape decided by depth/transverse +bounds — far cheaper per call than a nearest-neighbour search over hundreds of +thousands of reference points, which matters because this oracle is queried on +every autoregressive step of a shower rollout. "knn"/"svm" remain as generic +fallbacks (a classifier over 3D positions, escape decided by distance to the +nearest reference point) for geometries that aren't simple slab stacks. + +scikit-learn / joblib are an optional dependency (the `geometry` extra) needed +by "knn"/"svm" and by `save`/`load`; they are imported lazily so the core +install stays lean. """ from __future__ import annotations @@ -19,6 +28,7 @@ from pathlib import Path from typing import Any, Iterable import numpy as np +import pandas as pd import pyarrow.parquet as pq _INSTALL_HINT = ( @@ -36,14 +46,58 @@ def _require_sklearn(): raise ImportError(_INSTALL_HINT) from exc +@dataclass +class _SlabLookup: + """Fast path for a detector that is a stack of planar layer slabs along one + axis (miniCaloSim's actual geometry — see `giant/geometry.py` module docstring + and `physics/detector-design/minicalosim-geometry.md`). `material`/`layer_id` + are then a pure function of depth, found by binary search over `z_edges` + instead of a nearest-neighbour search over the whole reference point cloud — + O(log(#segments)) instead of O(log(#reference points)), with a far smaller + constant factor, and exact rather than approximate. + """ + + axis: int # which of the 3 position components is the depth axis + z_edges: np.ndarray # (n_segments + 1,) sorted boundaries between segments + materials: np.ndarray # (n_segments,) object, material of each segment + layer_ids: np.ndarray # (n_segments,) int64, layer_id of each segment + radius_max: float # largest transverse radius seen in training data + + def query( + self, pos: np.ndarray, margin: float + ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + other = [i for i in range(3) if i != self.axis] + z = pos[:, self.axis] + radius = np.sqrt(pos[:, other[0]] ** 2 + pos[:, other[1]] ** 2) + + idx = np.searchsorted(self.z_edges, z, side="right") - 1 + idx = np.clip(idx, 0, len(self.materials) - 1) + material = self.materials[idx] + layer_id = self.layer_ids[idx] + + escaped = ( + (z < self.z_edges[0] - margin) + | (z > self.z_edges[-1] + margin) + | (radius > self.radius_max + margin) + ) + return material, layer_id, escaped + + @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`. + Two lookup strategies are supported (`metadata["method"]`): + - `"slab"`: exact O(log #segments) binary search exploiting the known + layered-slab detector geometry (see `_SlabLookup`). Fast and preferred. + - `"knn"` / `"svm"`: a generic sklearn classifier over 3D positions, + predicting a class index into `classes` (a list of (material, layer_id) + pairs). Kept as a fallback for geometries that aren't simple slab stacks. + + `escape_threshold` is a distance in position units (mm). For knn/svm it's + compared against the nearest training reference point. For slab it's the + slack allowed beyond the observed depth range / transverse radius before a + point is flagged `escaped`. """ estimator: Any @@ -53,10 +107,10 @@ class GeometryOracle: # Only populated for non-neighbour estimators (SVM) to answer the escape # distance query; KNeighborsClassifier answers it directly. _ref_tree: Any = field(default=None) + # Only populated when metadata["method"] == "slab". + _slab: _SlabLookup | None = field(default=None) - def query( - self, pos: np.ndarray - ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + 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: @@ -68,6 +122,9 @@ class GeometryOracle: np.empty(0, dtype=bool), ) + if self._slab is not None: + return self._slab.query(pos, self.escape_threshold) + # 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. @@ -101,6 +158,7 @@ class GeometryOracle: "escape_threshold": self.escape_threshold, "metadata": self.metadata, "ref_tree": getattr(self, "_ref_tree", None), + "slab": getattr(self, "_slab", None), }, path, ) @@ -118,6 +176,7 @@ class GeometryOracle: metadata=d.get("metadata", {}), ) obj._ref_tree = d.get("ref_tree") + obj._slab = d.get("slab") return obj @@ -189,14 +248,18 @@ def _collect_points( # 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, + 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 - ) + _reservoir_replace(res_pos, res_mat, res_lay, pos, mat, lay, seen, rng) seen += m n = min(seen, subsample) @@ -218,33 +281,164 @@ def _reservoir_replace(res_pos, res_mat, res_lay, pos, mat, lay, seen, rng) -> N res_lay[slots] = lay[keep] +def _fit_slab_lookup( + pos: np.ndarray, + mat: np.ndarray, + lay: np.ndarray, + axis: int, + n_bins: int, +) -> tuple[_SlabLookup, dict]: + """Fit a `_SlabLookup` assuming material/layer_id are a function of depth + (`pos[:, axis]`) alone — true for a stack of planar layer slabs. + + Bins the depth axis into `n_bins` equal-width bins, takes the majority + (material, layer_id) label per bin (robust to the handful of points near a + boundary whose true label is ambiguous at bin resolution), fills any empty + bins from the nearest populated bin, then run-length-encodes consecutive + bins sharing a label into segments. Binary search over the segment + boundaries then answers a query in O(log #segments). + """ + other = [i for i in range(3) if i != axis] + z = pos[:, axis].astype(np.float64) + radius = np.sqrt( + pos[:, other[0]].astype(np.float64) ** 2 + + pos[:, other[1]].astype(np.float64) ** 2 + ) + + z_min, z_max = float(z.min()), float(z.max()) + if z_min == z_max: + raise ValueError( + "all points share the same depth-axis coordinate — pick a " + "different `depth_axis` or use method='knn'/'svm'" + ) + edges = np.linspace(z_min, z_max, n_bins + 1) + bin_idx = np.clip(np.searchsorted(edges, z, side="right") - 1, 0, n_bins - 1) + + counts = ( + pd.DataFrame({"bin": bin_idx, "material": mat, "layer_id": lay}) + .groupby(["bin", "material", "layer_id"]) + .size() + .to_frame("n") + .reset_index() + .sort_values("n", ascending=False) + .drop_duplicates("bin") + ) + + bin_material = np.full(n_bins, "", dtype=object) + bin_layer = np.full(n_bins, -1, dtype=np.int64) + has_data = np.zeros(n_bins, dtype=bool) + idx = counts["bin"].to_numpy() + bin_material[idx] = counts["material"].to_numpy() + bin_layer[idx] = counts["layer_id"].to_numpy() + has_data[idx] = True + + # Forward/backward-fill bins with no samples from the nearest populated one. + fill_from = np.where(has_data, np.arange(n_bins), -1) + for b in range(1, n_bins): + if fill_from[b] == -1: + fill_from[b] = fill_from[b - 1] + for b in range(n_bins - 2, -1, -1): + if fill_from[b] == -1: + fill_from[b] = fill_from[b + 1] + bin_material = bin_material[fill_from] + bin_layer = bin_layer[fill_from] + + # Run-length-encode consecutive bins sharing a label into segments. + changed = ( + np.flatnonzero( + (bin_material[1:] != bin_material[:-1]) | (bin_layer[1:] != bin_layer[:-1]) + ) + + 1 + ) + seg_starts = np.concatenate([[0], changed]) + z_edges = np.concatenate([edges[seg_starts], edges[-1:]]) + materials = bin_material[seg_starts] + layer_ids = bin_layer[seg_starts] + + unique_z = np.unique(z) + median_spacing = float(np.median(np.diff(unique_z))) if len(unique_z) > 1 else 1.0 + radius_max = float(radius.max()) + + slab = _SlabLookup( + axis=axis, + z_edges=z_edges, + materials=materials, + layer_ids=layer_ids, + radius_max=radius_max, + ) + info = { + "n_segments": int(len(materials)), + "z_range": (z_min, z_max), + "median_z_spacing": median_spacing, + "radius_max": radius_max, + } + return slab, info + + def build_geometry_oracle( files: list[Path], - method: str = "knn", + method: str = "slab", k: int = 1, subsample: int = 500_000, escape_factor: float = 5.0, seed: int = 0, + depth_axis: int = 2, + n_bins: int = 2000, ) -> GeometryOracle: """Fit a position -> (material, layer_id) classifier from steps files. - method: "knn" (KNeighborsClassifier, default) or "svm" (SVC). - k: neighbours for the knn classifier. + method: "slab" (default-recommended fast path exploiting the known + layered-slab detector geometry — see `_SlabLookup`), "knn" + (KNeighborsClassifier), or "svm" (SVC). "slab" is O(log #segments) per + query and exact; "knn"/"svm" are generic fallbacks for geometries that + aren't simple slab stacks, at the cost of a much slower query (a + nearest-neighbour or kernel evaluation against up to `subsample` + reference points) and, for "svm", occasional misclassification. + k: neighbours for the knn classifier (ignored otherwise). 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. + escape_factor: escape_threshold = escape_factor * median spacing of the + reference points along the relevant axis/axes, so it scales with the + sampling density of the data. + depth_axis: index (0/1/2 -> x/y/z) of the position component that layers + stack along. Only used by method="slab"; default 2 (z) matches + miniCaloSim's beam-axis-aligned layer stack. + n_bins: depth-axis resolution for method="slab" — should be finer than the + thinnest layer. """ - _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. + # Combined (material, layer_id) class label, used for `classes` regardless + # of method (informational for slab; the actual classifier index for + # knn/svm). pairs = list(zip((str(m) for m in mat), (int(v) for v in lay))) classes = sorted(set(pairs)) + + if method == "slab": + slab, info = _fit_slab_lookup(pos, mat, lay, axis=depth_axis, n_bins=n_bins) + escape_threshold = escape_factor * info["median_z_spacing"] + oracle = GeometryOracle( + estimator=None, + classes=classes, + escape_threshold=escape_threshold, + metadata={ + "method": "slab", + "depth_axis": depth_axis, + "n_bins": n_bins, + "n_reference_points": int(len(pos)), + "escape_factor": escape_factor, + "n_files": len(files), + **info, + }, + ) + oracle._slab = slab + return oracle + + _require_sklearn() + from sklearn.neighbors import KNeighborsClassifier, NearestNeighbors + from sklearn.svm import SVC + 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) @@ -261,14 +455,16 @@ def build_geometry_oracle( # 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'") + raise ValueError(f"unknown method {method!r}; use 'slab', '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) - ] + 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 diff --git a/scripts/dwarf.py b/scripts/dwarf.py index 3d91001..f728bcf 100644 --- a/scripts/dwarf.py +++ b/scripts/dwarf.py @@ -387,6 +387,7 @@ def make_root( class OracleMethod(str, Enum): + slab = "slab" knn = "knn" svm = "svm" @@ -396,13 +397,18 @@ 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") - ], + 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, + typer.Option( + "--method", + help=( + "Lookup strategy: slab (default; exact O(log #segments) fast " + "path for the layered-slab detector geometry), knn, or svm " + "(generic fallbacks for non-slab geometries)" + ), + ), + ] = OracleMethod.slab, k: Annotated[ int, typer.Option("--k", help="Neighbours for the knn classifier") ] = 1, @@ -414,10 +420,24 @@ def build_geometry_oracle( float, typer.Option( "--escape-factor", - help="escape_threshold = this x median NN spacing of reference points", + help="escape_threshold = this x median spacing of reference points", ), ] = 5.0, seed: Annotated[int, typer.Option("--seed", help="Sampling seed")] = 0, + depth_axis: Annotated[ + int, + typer.Option( + "--depth-axis", + help="0/1/2 -> x/y/z axis the layers stack along (method=slab only)", + ), + ] = 2, + n_bins: Annotated[ + int, + typer.Option( + "--n-bins", + help="Depth-axis resolution, finer than the thinnest layer (method=slab only)", + ), + ] = 2000, ) -> None: """Fit a position -> (material, layer_id) oracle for `giant rollout`.""" run_build_geometry_oracle( @@ -428,6 +448,8 @@ def build_geometry_oracle( subsample=subsample, escape_factor=escape_factor, seed=seed, + depth_axis=depth_axis, + n_bins=n_bins, ) diff --git a/scripts/geometry_oracle.py b/scripts/geometry_oracle.py index 1e8517e..e5a36fb 100644 --- a/scripts/geometry_oracle.py +++ b/scripts/geometry_oracle.py @@ -16,11 +16,13 @@ from giant.geometry import build_geometry_oracle def run_build_geometry_oracle( data: Path, out: Path, - method: str = "knn", + method: str = "slab", k: int = 1, subsample: int = 500_000, escape_factor: float = 5.0, seed: int = 0, + depth_axis: int = 2, + n_bins: int = 2000, ) -> None: files = find_parquet_files(data) print(f"found {len(files)} parquet file(s); sampling up to {subsample:,} points") @@ -32,17 +34,33 @@ def run_build_geometry_oracle( subsample=subsample, escape_factor=escape_factor, seed=seed, + depth_axis=depth_axis, + n_bins=n_bins, ) - print(f"method: {method} reference points: {oracle.metadata['n_reference_points']:,}") + 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 method == "slab": + z_lo, z_hi = oracle.metadata["z_range"] + print( + f"depth axis: {'xyz'[depth_axis]} segments: {oracle.metadata['n_segments']} " + f"z range: [{z_lo:.3f}, {z_hi:.3f}] radius_max: {oracle.metadata['radius_max']:.3f}" + ) + print( + f"median depth spacing: {oracle.metadata['median_z_spacing']:.3f} " + f"escape_threshold: {oracle.escape_threshold:.3f} (= {escape_factor}x spacing)" + ) + else: + 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) — " diff --git a/tests/test_geometry.py b/tests/test_geometry.py index 0f1d993..d967315 100644 --- a/tests/test_geometry.py +++ b/tests/test_geometry.py @@ -26,7 +26,9 @@ def _build(subsample=30000, method="knn", escape_factor=5.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, + [Path("x")], + method=method, + subsample=subsample, escape_factor=escape_factor, ) @@ -84,3 +86,93 @@ def test_svm_method_has_escape_tree(): assert orc._ref_tree is not None _, _, escaped = orc.query(np.array([[1e5, 0.0, 0.0]])) assert escaped[0] + + +def _layer_batch(n, rng): + """Two 100mm slabs along z (with an air gap between/around them), bounded + to a 100x100mm transverse footprint — miniCaloSim's actual layer-stack + shape.""" + z = rng.uniform(-20.0, 220.0, n).astype(np.float32) + x = rng.uniform(-50.0, 50.0, n).astype(np.float32) + y = rng.uniform(-50.0, 50.0, n).astype(np.float32) + material = np.full(n, "G4_AIR", dtype=object) + layer_id = np.full(n, -1, dtype=np.int64) + in_l0 = (z >= 0.0) & (z < 100.0) + in_l1 = (z >= 110.0) & (z < 210.0) + material[in_l0] = "G4_PbWO4" + layer_id[in_l0] = 0 + material[in_l1] = "G4_W" + layer_id[in_l1] = 1 + pos = np.stack([x, y, z], axis=1) + return pos, material, layer_id + + +def _build_slab(subsample=60000, n_bins=500, escape_factor=5.0): + rng = np.random.default_rng(0) + batches = [_layer_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="slab", + subsample=subsample, + escape_factor=escape_factor, + depth_axis=2, + n_bins=n_bins, + ) + + +def test_slab_is_default_method(): + rng = np.random.default_rng(0) + batches = [_layer_batch(20000, rng)] + with patch.object(g, "_iter_point_batches", lambda p: iter(batches)): + orc = g.build_geometry_oracle([Path("x")], subsample=20000) + assert orc.metadata["method"] == "slab" + assert orc._slab is not None + + +def test_slab_classes_discovered(): + orc = _build_slab() + assert set(orc.classes) == {("G4_PbWO4", 0), ("G4_W", 1), ("G4_AIR", -1)} + + +def test_slab_query_labels_by_depth(): + orc = _build_slab() + pos = np.array( + [[0.0, 0.0, 50.0], [0.0, 0.0, 105.0], [0.0, 0.0, 150.0]] + ) # layer 0, gap, layer 1 + material, layer_id, escaped = orc.query(pos) + assert list(material) == ["G4_PbWO4", "G4_AIR", "G4_W"] + assert list(layer_id) == [0, -1, 1] + assert not escaped.any() + + +def test_slab_escape_beyond_depth_range(): + orc = _build_slab() + pos = np.array([[0.0, 0.0, 50.0], [0.0, 0.0, 1e5]]) + _, _, escaped = orc.query(pos) + assert not escaped[0] + assert escaped[1] + + +def test_slab_escape_beyond_transverse_radius(): + orc = _build_slab() + pos = np.array([[0.0, 0.0, 50.0], [1e5, 1e5, 50.0]]) + _, _, escaped = orc.query(pos) + assert not escaped[0] + assert escaped[1] + + +def test_slab_save_load_roundtrip(tmp_path): + orc = _build_slab() + p = tmp_path / "slab_oracle.pkl" + orc.save(p) + loaded = g.GeometryOracle.load(p) + + pos = np.array( + [[0.0, 0.0, 50.0], [0.0, 0.0, 105.0], [0.0, 0.0, 150.0], [0.0, 0.0, 1e5]] + ) + 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._slab is not None diff --git a/tests/test_rollout.py b/tests/test_rollout.py index 0fc252a..4674beb 100644 --- a/tests/test_rollout.py +++ b/tests/test_rollout.py @@ -39,7 +39,11 @@ def _oracle(): 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) + # Pinned to "knn" explicitly: this test's escape-threshold semantics + # (tiny threshold -> escape even at a valid interior point, because no + # training point is that close) are KNN-specific, and the fixture's + # box geometry isn't a layer stack the "slab" method could fit anyway. + return g.build_geometry_oracle([Path("x")], method="knn", subsample=20000) def _seeds(n=6):