"""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)])): # 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): 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