Also clip the positive tail of raw predicted log_mass in rollout
CI / Format (ruff format) (push) Failing after 29s
CI / Lint (ruff check) (push) Successful in 30s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 27s
CI / Tests (push) Successful in 1m10s

The previous commit clipped log_mass's negative tail (undershooting
_EPS made mass go slightly negative). The mirror case also crashes
rollout: a sufficiently large raw predicted log_mass overflows
exp() in float32, giving mass = inf, which then fails the same
downstream log_transform finiteness check when that mass is fed back
in as conditioning for a further step. Bound the upper tail too, at a
value comfortably below float32's overflow point.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
2026-08-14 12:22:14 +02:00
parent d0381728ee
commit 98b09d2b4b
2 changed files with 43 additions and 9 deletions
+24
View File
@@ -566,3 +566,27 @@ def test_decode_secondaries_extreme_negative_log_mass_stays_nonnegative():
# stay well above -eps so log_transform(mass) stays finite.
assert sec_mass[0, 0] > -1e-8
log_transform(sec_mass[0, 0])
def test_decode_secondaries_extreme_positive_log_mass_stays_finite():
from giant.data.transforms import decode_secondaries, log_transform
N = 1
e_sec = np.array([5.0], dtype=np.float32)
n_sec = np.array([1])
pre_dir = np.array([[0.0, 0.0, 1.0]], dtype=np.float32)
sec_cont = np.zeros((N, K_MAX, 6), dtype=np.float32)
sec_cont[0, 0, 0] = 10.0 # stick logit -> ~all of e_sec
sec_cont[0, 0, 1:4] = [0, 0, 1]
sec_cont[0, 0, 4] = 200.0 # raw model prediction: extremely positive log_mass
sec_cont[0, 0, 5] = 1.0
_, _, sec_mass, _, _ = decode_secondaries(sec_cont, n_sec, e_sec, pre_dir)
# Mirror image of the extreme-negative case above: exp(log_mass)
# overflows float32 to inf for an unclipped raw prediction this large,
# which then crashes the next log_transform call the same way a
# negative mass would.
assert np.isfinite(sec_mass[0, 0])
log_transform(sec_mass[0, 0])