Encode edep/secondary/post energy as a conservation-constrained simplex

Replaces the independent log_delta_e/log_edep targets with 2 additive-log-ratio
coordinates over the deposit/secondary/post-energy simplex (fractions of pre_E
summing to 1), so edep + e_sec + post_E == pre_E holds by construction after
decoding (softmax) rather than being learned approximately. Requires e_sec
(secondary energy) as a new conditioning input and a steps_to_parquet.py pass
to derive it from child track first-step energies.
This commit is contained in:
2026-06-25 16:01:13 +02:00
parent 64c6bd1cef
commit 8475199609
13 changed files with 415 additions and 83 deletions
+47
View File
@@ -1,5 +1,7 @@
import numpy as np
from giant.data.transforms import (
energy_simplex_decode,
energy_simplex_encode,
inv_log_transform,
local_frame_rotation,
log_transform,
@@ -111,6 +113,51 @@ def test_reconstruct_post_pos_general_roundtrip():
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 = 500
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)