0b3ece52ed
Routes on the physics process (Compton, phot, brems, ...) that ends a step, supervised by a small classifier since process is a post-step outcome unobservable at gate time. Threads a process label end-to-end through the data pipeline (loader, build_features, dataset batches, training loss/checkpointing) alongside the existing EnergyRouter.
283 lines
12 KiB
Python
283 lines
12 KiB
Python
import numpy as np
|
|
import pytest
|
|
from giant.constants import K_MAX
|
|
from giant.data.transforms import (
|
|
build_features,
|
|
energy_simplex_decode,
|
|
energy_simplex_encode,
|
|
inv_local_frame_rotation,
|
|
inv_log_transform,
|
|
local_frame_rotation,
|
|
log_transform,
|
|
Normalizer,
|
|
reconstruct_post_pos,
|
|
travel_direction,
|
|
)
|
|
|
|
|
|
def test_log_transform_invertible():
|
|
x = np.array([0.1, 1.0, 10.0, 1000.0], dtype=np.float32)
|
|
np.testing.assert_allclose(inv_log_transform(log_transform(x)), x, rtol=1e-5)
|
|
|
|
|
|
def test_local_frame_rotation_noop_when_aligned():
|
|
N = 8
|
|
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
|
|
rng = np.random.default_rng(0)
|
|
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
|
|
result = local_frame_rotation(pre_dir, post_dir)
|
|
np.testing.assert_allclose(result, post_dir, atol=1e-5)
|
|
|
|
|
|
def test_local_frame_rotation_preserves_angle():
|
|
"""Angle between pre_dir and post_dir must equal angle between ẑ and rotated."""
|
|
rng = np.random.default_rng(1)
|
|
N = 200
|
|
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
|
|
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
|
|
|
|
rotated = local_frame_rotation(pre_dir, post_dir)
|
|
|
|
cos_before = (pre_dir * post_dir).sum(axis=1)
|
|
cos_after = rotated[:, 2] # dot with ẑ = z-component (unit vectors)
|
|
np.testing.assert_allclose(cos_after, cos_before, atol=1e-5)
|
|
|
|
|
|
def test_local_frame_rotation_preserves_norm():
|
|
rng = np.random.default_rng(2)
|
|
N = 100
|
|
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
|
|
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
|
|
result = local_frame_rotation(pre_dir, post_dir)
|
|
np.testing.assert_allclose(np.linalg.norm(result, axis=1), 1.0, atol=1e-5)
|
|
|
|
|
|
def test_local_frame_rotation_rejects_near_zero_pre_dir():
|
|
"""A degenerate (near-zero-norm) pre_dir has no well-defined frame — must
|
|
raise instead of silently falling back to an arbitrary rotation axis."""
|
|
pre_dir = np.array([[0.0, 0.0, 0.0], [0.0, 0.0, 1.0]], dtype=np.float32)
|
|
post_dir = np.array([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], dtype=np.float32)
|
|
with pytest.raises(ValueError, match="near-zero norm"):
|
|
local_frame_rotation(pre_dir, post_dir)
|
|
with pytest.raises(ValueError, match="near-zero norm"):
|
|
inv_local_frame_rotation(pre_dir, post_dir)
|
|
|
|
|
|
def test_local_frame_rotation_normalizes_non_unit_pre_dir():
|
|
"""A pre_dir with float32-drift norm (not exactly 1) must still produce the
|
|
same result as its exactly-normalized counterpart, not a skewed frame."""
|
|
rng = np.random.default_rng(9)
|
|
N = 50
|
|
pre_dir_unit = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir_unit /= np.linalg.norm(pre_dir_unit, axis=1, keepdims=True)
|
|
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
|
|
|
|
pre_dir_scaled = pre_dir_unit * rng.uniform(0.9, 1.1, size=(N, 1)).astype(
|
|
np.float32
|
|
)
|
|
expected = local_frame_rotation(pre_dir_unit, post_dir)
|
|
result = local_frame_rotation(pre_dir_scaled, post_dir)
|
|
np.testing.assert_allclose(result, expected, atol=1e-4)
|
|
|
|
|
|
def test_travel_direction_is_unit_norm():
|
|
rng = np.random.default_rng(5)
|
|
N = 50
|
|
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
|
|
post_pos = pre_pos + rng.standard_normal((N, 3)).astype(np.float32)
|
|
result = travel_direction(pre_pos, post_pos)
|
|
np.testing.assert_allclose(np.linalg.norm(result, axis=1), 1.0, atol=1e-5)
|
|
|
|
|
|
def test_travel_direction_matches_normalized_displacement():
|
|
rng = np.random.default_rng(6)
|
|
N = 50
|
|
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
|
|
disp = rng.standard_normal((N, 3)).astype(np.float32)
|
|
post_pos = pre_pos + disp
|
|
expected = disp / np.linalg.norm(disp, axis=1, keepdims=True)
|
|
np.testing.assert_allclose(travel_direction(pre_pos, post_pos), expected, atol=1e-5)
|
|
|
|
|
|
def test_reconstruct_post_pos_straight_line():
|
|
"""When post_pos = pre_pos + L * pre_dir, travel_dir equals pre_dir, so its
|
|
local-frame encoding is ẑ — reconstruction must recover post_pos exactly."""
|
|
rng = np.random.default_rng(7)
|
|
N = 20
|
|
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
|
|
step_length = rng.uniform(0.1, 5.0, size=N).astype(np.float32)
|
|
post_pos = pre_pos + step_length[:, None] * pre_dir
|
|
|
|
travel_dir_local = local_frame_rotation(
|
|
pre_dir, travel_direction(pre_pos, post_pos)
|
|
)
|
|
np.testing.assert_allclose(travel_dir_local, np.tile([0, 0, 1], (N, 1)), atol=1e-4)
|
|
|
|
reconstructed = reconstruct_post_pos(
|
|
pre_pos, pre_dir, step_length, travel_dir_local
|
|
)
|
|
np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)
|
|
|
|
|
|
def test_reconstruct_post_pos_general_roundtrip():
|
|
"""Full encode (build_features-style) -> decode (cli.py predict-style) path."""
|
|
rng = np.random.default_rng(8)
|
|
N = 100
|
|
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
|
|
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
|
|
post_pos = pre_pos + rng.standard_normal((N, 3)).astype(np.float32)
|
|
step_length = np.linalg.norm(post_pos - pre_pos, axis=1).astype(np.float32)
|
|
|
|
travel_dir_local = local_frame_rotation(
|
|
pre_dir, travel_direction(pre_pos, post_pos)
|
|
)
|
|
reconstructed = reconstruct_post_pos(
|
|
pre_pos, pre_dir, step_length, travel_dir_local
|
|
)
|
|
np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)
|
|
|
|
|
|
def test_energy_simplex_conservation():
|
|
"""Decoding any ALR coords yields energies that sum to pre_E exactly."""
|
|
rng = np.random.default_rng(11)
|
|
N = 500
|
|
z = rng.standard_normal((N, 2)).astype(np.float32) * 3.0
|
|
pre_E = rng.uniform(1.0, 100.0, N).astype(np.float32)
|
|
edep, e_sec, post_E, delta_e = energy_simplex_decode(z, pre_E)
|
|
np.testing.assert_allclose(edep + e_sec + post_E, pre_E, rtol=1e-5, atol=1e-4)
|
|
np.testing.assert_allclose(delta_e, edep + e_sec, rtol=1e-5, atol=1e-4)
|
|
assert np.all(edep >= 0) and np.all(e_sec >= 0) and np.all(post_E >= 0)
|
|
|
|
|
|
def test_energy_simplex_roundtrip():
|
|
"""Encode → decode recovers energies whose lost part already sums to delta_e."""
|
|
rng = np.random.default_rng(12)
|
|
N = 500000
|
|
pre_E = rng.uniform(1.0, 100.0, N).astype(np.float32)
|
|
post_E = (pre_E * rng.uniform(0.0, 1.0, N)).astype(np.float32)
|
|
delta_e = pre_E - post_E
|
|
g = rng.uniform(0.0, 1.0, N).astype(np.float32)
|
|
edep = (g * delta_e).astype(np.float32)
|
|
e_sec = ((1.0 - g) * delta_e).astype(np.float32)
|
|
|
|
z = energy_simplex_encode(edep, e_sec, post_E, pre_E)
|
|
edep_r, e_sec_r, post_E_r, _ = energy_simplex_decode(z, pre_E)
|
|
# Tolerance reflects the tiny simplex floor (~1e-5 of pre_E).
|
|
np.testing.assert_allclose(edep_r, edep, atol=5e-3)
|
|
np.testing.assert_allclose(e_sec_r, e_sec, atol=5e-3)
|
|
np.testing.assert_allclose(post_E_r, post_E, atol=5e-3)
|
|
|
|
|
|
def test_energy_simplex_handles_boundary_zeros():
|
|
"""e_sec=0 (no secondaries) and post_E=0 (track end) stay finite and decode near 0."""
|
|
pre_E = np.array([10.0, 50.0, 100.0], dtype=np.float32)
|
|
edep = np.array([4.0, 50.0, 0.0], dtype=np.float32)
|
|
e_sec = np.array([0.0, 0.0, 0.0], dtype=np.float32) # no secondaries
|
|
post_E = np.array([6.0, 0.0, 100.0], dtype=np.float32) # row 1: track ends
|
|
|
|
z = energy_simplex_encode(edep, e_sec, post_E, pre_E)
|
|
assert np.all(np.isfinite(z))
|
|
_, e_sec_r, post_E_r, _ = energy_simplex_decode(z, pre_E)
|
|
np.testing.assert_allclose(e_sec_r, 0.0, atol=1e-2)
|
|
assert post_E_r[1] < 1e-2 # the absorbed track decodes to ~0 post energy
|
|
|
|
|
|
def test_normalizer_roundtrip():
|
|
rng = np.random.default_rng(3)
|
|
X = rng.standard_normal((200, 9)).astype(np.float32)
|
|
norm = Normalizer().fit(X)
|
|
np.testing.assert_allclose(norm.inverse_transform(norm.transform(X)), X, atol=1e-5)
|
|
|
|
|
|
def test_normalizer_serialization():
|
|
rng = np.random.default_rng(4)
|
|
X = rng.standard_normal((50, 6)).astype(np.float32)
|
|
norm = Normalizer().fit(X)
|
|
norm2 = Normalizer.from_dict(norm.to_dict())
|
|
assert norm2.mean is not None and norm.mean is not None
|
|
assert norm2.std is not None and norm.std is not None
|
|
np.testing.assert_allclose(norm2.mean, norm.mean)
|
|
np.testing.assert_allclose(norm2.std, norm.std)
|
|
|
|
|
|
def test_build_features_clamps_n_sec_label_to_k_max():
|
|
"""A step with more secondaries than K_MAX must not overflow the
|
|
n_sec classifier's K_MAX+1 classes (regression test: this used to hand
|
|
cross_entropy an out-of-range target and crash CUDA training with
|
|
'unique_by_key: failed to synchronize: cudaErrorAssert')."""
|
|
N = 3
|
|
raw_n_sec = np.array([0, 5, K_MAX + 20], dtype=np.int32)
|
|
rng = np.random.default_rng(0)
|
|
data = {
|
|
"pdg": np.array([11, 11, 11], dtype=np.int32),
|
|
"material": np.array(["PbWO4", "PbWO4", "PbWO4"], dtype=object),
|
|
"pre_pos": rng.standard_normal((N, 3)).astype(np.float32),
|
|
"pre_E": np.full(N, 10.0, dtype=np.float32),
|
|
"pre_dir": np.tile(np.array([0, 0, 1], dtype=np.float32), (N, 1)),
|
|
"layer_id": np.zeros(N, dtype=np.int32),
|
|
"n_sec": raw_n_sec,
|
|
"e_sec": np.full(N, 1.0, dtype=np.float32),
|
|
"step_length": np.full(N, 1.0, dtype=np.float32),
|
|
"post_E": np.full(N, 9.0, dtype=np.float32),
|
|
"edep": np.full(N, 1.0, dtype=np.float32),
|
|
"post_dir": np.tile(np.array([0, 0, 1], dtype=np.float32), (N, 1)),
|
|
"post_pos": rng.standard_normal((N, 3)).astype(np.float32),
|
|
}
|
|
pdg_map = {11: 0}
|
|
mat_map = {"PbWO4": 0}
|
|
|
|
_, _, _, n_sec, _, _, _, _, _ = build_features(data, pdg_map, mat_map)
|
|
|
|
assert n_sec.max() <= K_MAX
|
|
np.testing.assert_array_equal(n_sec, [0, 5, K_MAX])
|
|
|
|
|
|
def _minimal_step_data(N: int, process: np.ndarray | None = None) -> dict:
|
|
rng = np.random.default_rng(0)
|
|
data = {
|
|
"pdg": np.full(N, 11, dtype=np.int32),
|
|
"material": np.full(N, "PbWO4", dtype=object),
|
|
"pre_pos": rng.standard_normal((N, 3)).astype(np.float32),
|
|
"pre_E": np.full(N, 10.0, dtype=np.float32),
|
|
"pre_dir": np.tile(np.array([0, 0, 1], dtype=np.float32), (N, 1)),
|
|
"layer_id": np.zeros(N, dtype=np.int32),
|
|
"n_sec": np.zeros(N, dtype=np.int32),
|
|
"e_sec": np.full(N, 1.0, dtype=np.float32),
|
|
"step_length": np.full(N, 1.0, dtype=np.float32),
|
|
"post_E": np.full(N, 9.0, dtype=np.float32),
|
|
"edep": np.full(N, 1.0, dtype=np.float32),
|
|
"post_dir": np.tile(np.array([0, 0, 1], dtype=np.float32), (N, 1)),
|
|
"post_pos": rng.standard_normal((N, 3)).astype(np.float32),
|
|
}
|
|
if process is not None:
|
|
data["process"] = process
|
|
return data
|
|
|
|
|
|
def test_build_features_proc_idx_zero_without_proc_map():
|
|
data = _minimal_step_data(3, process=np.array(["compt", "phot", "eIoni"], dtype=object))
|
|
pdg_map, mat_map = {11: 0}, {"PbWO4": 0}
|
|
|
|
*_, proc_idx, _, _ = build_features(data, pdg_map, mat_map)
|
|
|
|
np.testing.assert_array_equal(proc_idx, [0, 0, 0])
|
|
|
|
|
|
def test_build_features_proc_idx_looks_up_proc_map():
|
|
data = _minimal_step_data(3, process=np.array(["compt", "phot", "eIoni"], dtype=object))
|
|
pdg_map, mat_map = {11: 0}, {"PbWO4": 0}
|
|
proc_map = {"compt": 0, "phot": 1, "eIoni": 2}
|
|
|
|
*_, proc_idx, _, _ = build_features(data, pdg_map, mat_map, proc_map=proc_map)
|
|
|
|
np.testing.assert_array_equal(proc_idx, [0, 1, 2])
|