From 57680645debbc4f647c95c9290d5ee2089823c23 Mon Sep 17 00:00:00 2001 From: nstarman Date: Wed, 19 Aug 2026 23:46:33 -0400 Subject: [PATCH 1/7] =?UTF-8?q?=F0=9F=90=9B=20fix(curveframes):=20two=20gu?= =?UTF-8?q?ards=20admitted=20NaN=20and=20returned=20NaN=20silently?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Found by auditing what each guard admits rather than what it claims. `x <= tol` and `x > hi` are both False for a NaN, so a guard written as a direct comparison lets one through and returns a NaN result with nothing raised -- worse than the case the guard exists for, which at least errors. _orthonormalize([1, nan, 0], T0) -> NaN triad, no error ArcLength(..., s_max=5km)(Q(nan,"km")) -> NaN position, no error `TubularChart`'s reach guard already avoids this, written `~(f > 0)` in #699 after the same bug. These two were still in the direct form: bishop.py norm <= 1e-12 * |v| -> ~(norm > tol) arclength.py (s < -margin) | (s > s_max + margin) -> ~in_domain The cases each guard was written for are unaffected: an exactly parallel normal, an all-zero normal, and an `s` genuinely outside the solved domain all still raise, and a well-conditioned normal and an in-domain `s` still pass. Four tests, mutation-verified against reverting both predicates. Co-Authored-By: Claude Opus 5 --- .../coordinaxs/curveframes/_src/arclength.py | 6 +- .../src/coordinaxs/curveframes/_src/bishop.py | 6 +- .../tests/unit/test_guards_reject_nan.py | 63 +++++++++++++++++++ 3 files changed, 73 insertions(+), 2 deletions(-) create mode 100644 packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py diff --git a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py index 67d99acf9..572ddf2b5 100644 --- a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py +++ b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py @@ -212,7 +212,11 @@ def _eval_tau_dense( s_max_val = jnp.asarray(s_max.ustrip(s_unit)) margin = _S_MAX_MARGIN * jnp.abs(s_max_val) - out_of_domain = (s_val < -margin) | (s_val > s_max_val + margin) + # Negated `in-domain` rather than the two out-of-domain comparisons: a NaN + # `s` compares False against both, so the direct form admits it and hands + # back a NaN position with nothing raised. + in_domain = (s_val >= -margin) & (s_val <= s_max_val + margin) + out_of_domain = ~in_domain # No clip: every point of the solved range is real data, so the only # thing to do with a genuine overshoot is refuse it. diff --git a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py index 78a778e6e..85556025b 100644 --- a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py +++ b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py @@ -160,7 +160,11 @@ def _orthonormalize(v: Any, T0_val: Any) -> Any: w = v - jnp.dot(v, T0_val) * T0_val norm = jnp.linalg.norm(w) # `<=`, not `<`, so that an all-zero `v` (threshold 0) still raises. - w = eqx.error_if(w, norm <= 1e-12 * jnp.linalg.norm(v), _MSG_PARALLEL_NORMAL) + # `~(norm > tol)`, not `norm <= tol`: a NaN compares False against both, so + # the `<=` form admits a NaN `v` and returns a NaN triad with nothing + # raised. Same reason `TubularChart`'s reach guard is written `~(f > 0)`. + tol = 1e-12 * jnp.linalg.norm(v) + w = eqx.error_if(w, ~(norm > tol), _MSG_PARALLEL_NORMAL) return w / norm diff --git a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py new file mode 100644 index 000000000..6c6e3f0f6 --- /dev/null +++ b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py @@ -0,0 +1,63 @@ +"""A NaN must not walk through a guard that claims to reject bad input. + +`x <= tol` and `x > hi` are both False for a NaN, so a guard written as a +direct comparison admits it and returns a NaN result with nothing raised -- +worse than the case the guard was written for, which at least errors. + +`TubularChart`'s reach guard already avoids this by testing `~(f > 0)`; these +pin the same property for the two guards that were still written the direct +way. +""" + +import jax.numpy as jnp +import pytest + +import unxt as u + +import coordinaxs.curveframes as cxfc +from coordinaxs.curveframes._src.bishop import _orthonormalize + + +def helix(tau: u.AbstractQuantity) -> u.AbstractQuantity: + t = tau.ustrip("s") + return u.Q(jnp.stack([jnp.cos(t), jnp.sin(t), 0.3 * t]), "km") + + +@pytest.mark.parametrize("bad", [jnp.nan, jnp.inf], ids=["nan", "inf"]) +def test_a_non_finite_initial_normal_is_rejected(bad: float) -> None: + """`_orthonormalize` returned a NaN triad instead of raising. + + The rejection was `norm <= 1e-12 * |v|`; with a NaN `v` both sides are NaN + and the comparison is False. Fails if it goes back to the direct form. + """ + T0 = jnp.array([1.0, 0.0, 0.0]) + with pytest.raises(Exception, match="parallel"): + _orthonormalize(jnp.array([1.0, bad, 0.0]), T0) + + +def test_the_legitimate_degenerate_cases_still_raise() -> None: + """The case the guard was written for keeps working.""" + T0 = jnp.array([1.0, 0.0, 0.0]) + for v in (jnp.array([1.0, 0.0, 0.0]), jnp.array([0.0, 0.0, 0.0])): + with pytest.raises(Exception, match="parallel"): + _orthonormalize(v, T0) + + # and a well-conditioned normal is untouched + out = _orthonormalize(jnp.array([0.0, 2.0, 0.0]), T0) + assert jnp.allclose(jnp.asarray(out), jnp.array([0.0, 1.0, 0.0])) + + +def test_a_nan_arc_length_is_rejected() -> None: + """The domain guard returned a NaN position instead of raising. + + The test was `(s < -margin) | (s > s_max + margin)`; a NaN is False for + both, so it fell through to `diffrax`, which happily interpolated a NaN. + """ + fast = cxfc.ArcLength(helix, "s", s_max=u.Q(5.0, "km")) + with pytest.raises(Exception, match="solved domain"): + fast(u.Q(jnp.nan, "km")) + + # in-domain and genuinely-outside both behave as before + assert jnp.isfinite(fast(u.Q(2.0, "km")).ustrip("km")).all() + with pytest.raises(Exception, match="solved domain"): + fast(u.Q(99.0, "km")) From 94adcd0643ad411dd0bab21c5910ffe9bbd5f289 Mon Sep 17 00:00:00 2001 From: nstarman Date: Wed, 19 Aug 2026 23:55:36 -0400 Subject: [PATCH 2/7] =?UTF-8?q?=F0=9F=90=9B=20fix(transforms,charts):=20ra?= =?UTF-8?q?nge=20guards=20admitted=20NaN=20and=20leaked=20it?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Same class as the previous commit, found by continuing the audit into core. `x >= hi` and `x > hi` are both False for a NaN, so a guard written as a direct comparison admits it: LorentzBoost(Q([nan,0,0], "")).gamma -> non-finite, no error leq(Q(nan, "m"), Q(2, "m")) -> admits geq(Q(nan, "m"), Q(2, "m")) -> admits `lorentz.py`'s own comment states the intent it was failing: guard "so every derived quantity reports the same superluminal error rather than one of them leaking a non-finite value". Four predicates negated: beta_sq >= 1.0 -> ~(beta_sq < 1.0) speed >= 1.0 -> ~(speed < 1.0) jnp.any(x > max_val) -> jnp.any(~(x <= max_val)) jnp.any(x < min_val) -> jnp.any(~(x >= min_val)) Valid inputs are unaffected: gamma(beta=0.5) = 1.154701, and in-range values still pass `leq`/`geq`. Not changed, and worth a separate decision: the zero/degeneracy guards (`norm == 0`, `jnp.isclose(s, 0)`, `jnp.allclose(norm, 0)` in `builders.py`, `scale.py`, `reflect.py`, `lorentz.py`) admit NaN for the same reason. A NaN is not zero, so whether they should reject it is a question about intent rather than a mechanical fix. Co-Authored-By: Claude Opus 5 --- .../tests/unit/test_guards_reject_nan.py | 36 +++++++++++++++++++ src/coordinax/_src/charts/checks.py | 7 ++-- .../transforms/_src/actions/lorentz.py | 7 ++-- 3 files changed, 46 insertions(+), 4 deletions(-) diff --git a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py index 6c6e3f0f6..d15971972 100644 --- a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py +++ b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py @@ -61,3 +61,39 @@ def test_a_nan_arc_length_is_rejected() -> None: assert jnp.isfinite(fast(u.Q(2.0, "km")).ustrip("km")).all() with pytest.raises(Exception, match="solved domain"): fast(u.Q(99.0, "km")) + + +# -------------------------------------------------------------------------- +# The same class in core: range guards written as a direct comparison. + + +def test_a_nan_boost_velocity_is_rejected() -> None: + """`LorentzBoost` returned a non-finite gamma instead of raising. + + The guard was `beta_sq >= 1.0`, and its own comment says it exists so that + no derived quantity "leaks a non-finite value" -- which a NaN did, being + False for that comparison. + """ + import coordinax.transforms as cxfm + + bad = cxfm.LorentzBoost(u.Q(jnp.array([jnp.nan, 0.0, 0.0]), "")) + with pytest.raises(Exception, match="subluminal"): + _ = bad.gamma + + # subluminal is unaffected + ok = cxfm.LorentzBoost(u.Q(jnp.array([0.5, 0.0, 0.0]), "")) + assert jnp.isfinite(ok.gamma) + + +def test_nan_fails_the_coordinate_bounds_checks() -> None: + """`leq`/`geq` admitted a NaN, so an out-of-range coordinate slipped by.""" + from coordinax._src.charts.checks import geq, leq + + with pytest.raises(Exception, match="less than or equal"): + leq(u.Q(jnp.nan, "m"), u.Q(2, "m")) + with pytest.raises(Exception, match="greater than or equal"): + geq(u.Q(jnp.nan, "m"), u.Q(2, "m")) + + # in-range values still pass + leq(u.Q(1.0, "m"), u.Q(2, "m")) + geq(u.Q(3.0, "m"), u.Q(2, "m")) diff --git a/src/coordinax/_src/charts/checks.py b/src/coordinax/_src/charts/checks.py index 461eb9871..61ba8354a 100644 --- a/src/coordinax/_src/charts/checks.py +++ b/src/coordinax/_src/charts/checks.py @@ -115,7 +115,9 @@ def leq( """ name = f" {name}" if name else name msg = f"The input{name} must be less than or equal to {comp_name}." - return eqx.error_if(x, u.ustrip("", jnp.any(x > max_val)), msg) + # `~(x <= max)` rather than `x > max`: a NaN is False for both, so the + # direct form admits it silently. + return eqx.error_if(x, u.ustrip("", jnp.any(~(x <= max_val))), msg) def geq( @@ -144,7 +146,8 @@ def geq( """ name = f" {name}" if name else name msg = f"The input{name} must be greater than or equal to {comp_name}." - return eqx.error_if(x, u.ustrip("", jnp.any(x < min_val)), msg) + # `~(x >= min)`, for the NaN reason given in `leq`. + return eqx.error_if(x, u.ustrip("", jnp.any(~(x >= min_val))), msg) def check_manifolds_match_charts( diff --git a/src/coordinax/transforms/_src/actions/lorentz.py b/src/coordinax/transforms/_src/actions/lorentz.py index e98ed4a05..36c4c075d 100644 --- a/src/coordinax/transforms/_src/actions/lorentz.py +++ b/src/coordinax/transforms/_src/actions/lorentz.py @@ -277,7 +277,10 @@ def gamma(self) -> Array: """ beta_sq = jnp.sum(self.beta**2) - beta_sq = eqx.error_if(beta_sq, beta_sq >= 1.0, _MSG_SUPERLUMINAL) + # `~(x < 1)`, not `x >= 1`: a NaN is False for both comparisons, so the + # direct form admits it and leaks the non-finite value this guard exists + # to stop. + beta_sq = eqx.error_if(beta_sq, ~(beta_sq < 1.0), _MSG_SUPERLUMINAL) return 1.0 / jnp.sqrt(1.0 - beta_sq) @property @@ -297,7 +300,7 @@ def rapidity(self) -> Array: # same condition as `gamma`, so every derived quantity reports the same # superluminal error rather than one of them leaking a non-finite value. speed = self.speed - speed = eqx.error_if(speed, speed >= 1.0, _MSG_SUPERLUMINAL) + speed = eqx.error_if(speed, ~(speed < 1.0), _MSG_SUPERLUMINAL) return jnp.arctanh(speed) # ----------------------------------------------------- From 63fc1e2255666267d7ef2d778a654a45a4a5eed7 Mon Sep 17 00:00:00 2001 From: nstarman Date: Thu, 20 Aug 2026 11:45:44 -0400 Subject: [PATCH 3/7] =?UTF-8?q?=E2=9C=85=20test(charts,transforms):=20put?= =?UTF-8?q?=20the=20core=20NaN=20guards=20in=20the=20core=20suites?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `nox -s "pytest(package='coordinax')"` collects only the root `tests/` tree, so the `LorentzBoost.gamma` and `leq`/`geq` regressions did not run when core was tested alone. They now live beside the tests for the guards they cover: the boost one folds into the existing subluminal check as a second parameter, and `TestLeq`/`TestGeq` each gain a NaN case. Narrow the remaining curveframes guards from `Exception` to `eqx.EquinoxRuntimeError`, which is what all of them raise. Co-Authored-By: Claude Opus 5 --- .../tests/unit/test_guards_reject_nan.py | 45 +++---------------- tests/unit/charts/test_checks.py | 15 +++++++ tests/unit/transforms/test_lorentz_boost.py | 10 +++-- 3 files changed, 27 insertions(+), 43 deletions(-) diff --git a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py index d15971972..3869f07b9 100644 --- a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py +++ b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py @@ -9,6 +9,7 @@ way. """ +import equinox as eqx import jax.numpy as jnp import pytest @@ -31,7 +32,7 @@ def test_a_non_finite_initial_normal_is_rejected(bad: float) -> None: and the comparison is False. Fails if it goes back to the direct form. """ T0 = jnp.array([1.0, 0.0, 0.0]) - with pytest.raises(Exception, match="parallel"): + with pytest.raises(eqx.EquinoxRuntimeError, match="parallel"): _orthonormalize(jnp.array([1.0, bad, 0.0]), T0) @@ -39,7 +40,7 @@ def test_the_legitimate_degenerate_cases_still_raise() -> None: """The case the guard was written for keeps working.""" T0 = jnp.array([1.0, 0.0, 0.0]) for v in (jnp.array([1.0, 0.0, 0.0]), jnp.array([0.0, 0.0, 0.0])): - with pytest.raises(Exception, match="parallel"): + with pytest.raises(eqx.EquinoxRuntimeError, match="parallel"): _orthonormalize(v, T0) # and a well-conditioned normal is untouched @@ -54,46 +55,10 @@ def test_a_nan_arc_length_is_rejected() -> None: both, so it fell through to `diffrax`, which happily interpolated a NaN. """ fast = cxfc.ArcLength(helix, "s", s_max=u.Q(5.0, "km")) - with pytest.raises(Exception, match="solved domain"): + with pytest.raises(eqx.EquinoxRuntimeError, match="solved domain"): fast(u.Q(jnp.nan, "km")) # in-domain and genuinely-outside both behave as before assert jnp.isfinite(fast(u.Q(2.0, "km")).ustrip("km")).all() - with pytest.raises(Exception, match="solved domain"): + with pytest.raises(eqx.EquinoxRuntimeError, match="solved domain"): fast(u.Q(99.0, "km")) - - -# -------------------------------------------------------------------------- -# The same class in core: range guards written as a direct comparison. - - -def test_a_nan_boost_velocity_is_rejected() -> None: - """`LorentzBoost` returned a non-finite gamma instead of raising. - - The guard was `beta_sq >= 1.0`, and its own comment says it exists so that - no derived quantity "leaks a non-finite value" -- which a NaN did, being - False for that comparison. - """ - import coordinax.transforms as cxfm - - bad = cxfm.LorentzBoost(u.Q(jnp.array([jnp.nan, 0.0, 0.0]), "")) - with pytest.raises(Exception, match="subluminal"): - _ = bad.gamma - - # subluminal is unaffected - ok = cxfm.LorentzBoost(u.Q(jnp.array([0.5, 0.0, 0.0]), "")) - assert jnp.isfinite(ok.gamma) - - -def test_nan_fails_the_coordinate_bounds_checks() -> None: - """`leq`/`geq` admitted a NaN, so an out-of-range coordinate slipped by.""" - from coordinax._src.charts.checks import geq, leq - - with pytest.raises(Exception, match="less than or equal"): - leq(u.Q(jnp.nan, "m"), u.Q(2, "m")) - with pytest.raises(Exception, match="greater than or equal"): - geq(u.Q(jnp.nan, "m"), u.Q(2, "m")) - - # in-range values still pass - leq(u.Q(1.0, "m"), u.Q(2, "m")) - geq(u.Q(3.0, "m"), u.Q(2, "m")) diff --git a/tests/unit/charts/test_checks.py b/tests/unit/charts/test_checks.py index 84777bd0f..208a5ec28 100644 --- a/tests/unit/charts/test_checks.py +++ b/tests/unit/charts/test_checks.py @@ -164,6 +164,13 @@ def test_array_with_value_above_max_raises(self) -> None: ): checks.leq(x, max_q) + def test_nan_raises(self) -> None: + """A NaN is False for ``x <= max``, so the direct form admitted it.""" + with pytest.raises( + (eqx.EquinoxRuntimeError, ValueError), match="must be less than or equal to" + ): + checks.leq(u.Q(jnp.nan, "m"), u.Q(5, "m")) + class TestGeq: """Tests for geq (greater than or equal) check.""" @@ -206,3 +213,11 @@ def test_array_with_value_below_min_raises(self) -> None: match="must be greater than or equal to", ): checks.geq(x, min_q) + + def test_nan_raises(self) -> None: + """A NaN is False for ``x >= min``, so the direct form admitted it.""" + with pytest.raises( + (eqx.EquinoxRuntimeError, ValueError), + match="must be greater than or equal to", + ): + checks.geq(u.Q(jnp.nan, "m"), u.Q(5, "m")) diff --git a/tests/unit/transforms/test_lorentz_boost.py b/tests/unit/transforms/test_lorentz_boost.py index 7e2e78d7d..a64b6d5d0 100644 --- a/tests/unit/transforms/test_lorentz_boost.py +++ b/tests/unit/transforms/test_lorentz_boost.py @@ -160,14 +160,18 @@ def test_from_velocity_accepts_other_speed_units(self): assert float(in_kms.speed) == pytest.approx(0.5, abs=1e-4) @pytest.mark.parametrize("attr", ["gamma", "rapidity"]) - def test_superluminal_boost_is_rejected(self, attr): + @pytest.mark.parametrize("beta", [1.5, jnp.nan], ids=["superluminal", "nan"]) + def test_a_non_subluminal_boost_is_rejected(self, attr, beta): """Every derived quantity guards, not just ``gamma``. ``rapidity`` used to reach ``arctanh(|beta| >= 1)`` and hand back - ``inf``/``nan`` while ``gamma`` on the same object raised. + ``inf``/``nan`` while ``gamma`` on the same object raised. A ``nan`` + beta needs the guard written as ``~(beta_sq < 1)``: it is False for + ``beta_sq >= 1`` too, so the direct form let it through and returned + the non-finite value the guard exists to prevent. """ with pytest.raises(eqx.EquinoxRuntimeError, match="subluminal"): - _ = getattr(cxfm.LorentzBoost([1.5, 0.0, 0.0]), attr) + _ = getattr(cxfm.LorentzBoost([beta, 0.0, 0.0]), attr) def test_zero_direction_is_rejected(self): """A zero ``direction`` has no axis to normalise onto. From 7137171085a2ee3c30dd039b5c1845b062d40720 Mon Sep 17 00:00:00 2001 From: nstarman Date: Thu, 20 Aug 2026 15:52:50 -0400 Subject: [PATCH 4/7] =?UTF-8?q?=F0=9F=90=9B=20fix(transforms):=20four=20de?= =?UTF-8?q?generacy=20guards=20admitted=20NaN=20and=20inf?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `norm == 0`, `allclose(norm, 0)` and `isclose(s, 0)` each test one point of their own precondition. `axis / norm` is a unit vector only where the norm is finite and positive -- it is NaN iff a component is, `inf` iff a component is, and 0 iff the vector is -- so the equality form caught the zero case and let the other two normalise to exactly the silent NaN each guard's own comment says it exists to prevent: nine NaN entries in `R`, in `H`, and three in `beta`. `Scale`'s `inf` factor was the quietest of the lot. Its reciprocal is 0.0, so `from_factors([2, inf]).inverse` came back finite, singular, and with nothing to notice. Each guard now tests the precondition itself, and the messages say so. The four existing tests gain `nan` and `inf` parameters; reverting any one guard fails both of its new cases and neither of the old. Co-Authored-By: Claude Opus 5 --- .../transforms/_src/actions/builders.py | 16 ++++++++++++---- .../transforms/_src/actions/lorentz.py | 15 ++++++++++----- .../transforms/_src/actions/reflect.py | 12 ++++++++---- src/coordinax/transforms/_src/actions/scale.py | 15 +++++++++++---- tests/unit/transforms/test_builders.py | 17 +++++++++++++---- tests/unit/transforms/test_lorentz_boost.py | 15 +++++++++------ tests/unit/transforms/test_reflect.py | 11 ++++++++--- .../test_spatial_linear_transforms.py | 12 +++++++++--- 8 files changed, 80 insertions(+), 33 deletions(-) diff --git a/src/coordinax/transforms/_src/actions/builders.py b/src/coordinax/transforms/_src/actions/builders.py index cae4d35c7..479ddc921 100644 --- a/src/coordinax/transforms/_src/actions/builders.py +++ b/src/coordinax/transforms/_src/actions/builders.py @@ -15,7 +15,10 @@ from .rotate import Rotate from .translate import Translate -_MSG_ZERO_AXIS = "`RotationAboutAxis.axis` must be non-zero; got a zero-length axis." +_MSG_ZERO_AXIS = ( + "`RotationAboutAxis.axis` must be non-zero and finite; normalising anything " + "else gives a NaN rotation matrix." +) def _as_axis(axis: Any, /) -> Shaped[Array, "3"]: @@ -75,9 +78,14 @@ def __call__(self, tau: Any, /) -> Rotate: """Build the `Rotate` operator at time parameter ``tau``.""" theta = jnp.asarray(u.ustrip("rad", self.omega * tau + self.phase)) norm = jnp.linalg.vector_norm(self.axis) - # A zero-length axis defines no rotation; normalizing it would give a - # silently NaN `R`. `error_if` also fires under `jit`. - axis = eqx.error_if(self.axis, norm == 0, _MSG_ZERO_AXIS) + # `axis / norm` is a unit vector only where `axis` is finite and non-zero, + # which is exactly a finite positive `norm`: the norm is NaN iff a component + # is, `inf` iff a component is, and 0 iff the axis is. Testing `norm == 0` + # alone caught the last of the three and let the other two normalise to the + # silently NaN `R` this guard exists to prevent. `error_if` fires under `jit`. + axis = eqx.error_if( + self.axis, ~((norm > 0) & jnp.isfinite(norm)), _MSG_ZERO_AXIS + ) n = axis / norm # Rodrigues' formula: R = I cos(th) + sin(th) [n]_x + (1-cos th) n n^T # The `0.0 * n[0]` terms keep the zero entries as functions of `axis` diff --git a/src/coordinax/transforms/_src/actions/lorentz.py b/src/coordinax/transforms/_src/actions/lorentz.py index 36c4c075d..e5be21546 100644 --- a/src/coordinax/transforms/_src/actions/lorentz.py +++ b/src/coordinax/transforms/_src/actions/lorentz.py @@ -31,8 +31,8 @@ _MSG_ZERO_DIRECTION = ( - "LorentzBoost.from_rapidity requires a non-zero `direction`; the zero " - "vector has no boost axis to normalise onto." + "LorentzBoost.from_rapidity requires a non-zero, finite `direction`; " + "anything else has no boost axis to normalise onto." ) @@ -244,9 +244,14 @@ def from_rapidity( """ d = _float(direction) norm = jnp.linalg.norm(d) - # A zero direction has no boost axis to normalise onto; dividing would - # give `nan` betas that then propagate silently into every matrix entry. - norm = eqx.error_if(norm, norm == 0.0, _MSG_ZERO_DIRECTION) + # `d / norm` is a unit vector only where `norm` is finite and positive: it + # is NaN iff a component of `d` is, `inf` iff a component is, and 0 iff `d` + # is. Testing `norm == 0.0` alone caught only the last, so a NaN direction + # built the `nan` betas this guard exists to prevent -- reported later by + # `gamma`'s subluminal check, which names the wrong cause. + norm = eqx.error_if( + norm, ~((norm > 0.0) & jnp.isfinite(norm)), _MSG_ZERO_DIRECTION + ) return cls(jnp.tanh(_float(rapidity)) * (d / norm)) # ----------------------------------------------------- diff --git a/src/coordinax/transforms/_src/actions/reflect.py b/src/coordinax/transforms/_src/actions/reflect.py index f8929090e..48cf884f0 100644 --- a/src/coordinax/transforms/_src/actions/reflect.py +++ b/src/coordinax/transforms/_src/actions/reflect.py @@ -21,7 +21,9 @@ HMatrix: TypeAlias = Shaped[Array, " N N"] -_MSG_ZERO_NORMAL: Final = "Reflect.from_normal requires a nonzero normal vector." +_MSG_ZERO_NORMAL: Final = ( + "Reflect.from_normal requires a nonzero normal vector that is finite." +) @final @@ -76,9 +78,11 @@ def from_normal(cls: type["Reflect"], normal: Any, /) -> "Reflect": raise ValueError(msg) norm = jnp.linalg.norm(n) - # Defer the zero-normal check so it survives jit (a plain `bool` on a - # traced value raises TracerBoolConversionError). - n = eqx.error_if(n, jnp.allclose(norm, 0), _MSG_ZERO_NORMAL) + # Deferred so it survives jit (a plain `bool` on a traced value raises + # TracerBoolConversionError). `n / norm` is a unit vector only where `norm` + # is finite and positive; `allclose(norm, 0)` alone was False for a NaN or + # `inf` norm, so those normalised to a silently NaN `H`. + n = eqx.error_if(n, ~((norm > 0) & jnp.isfinite(norm)), _MSG_ZERO_NORMAL) n_hat = n / norm H = jnp.eye(n.shape[0], dtype=n_hat.dtype) - 2 * jnp.outer(n_hat, n_hat) diff --git a/src/coordinax/transforms/_src/actions/scale.py b/src/coordinax/transforms/_src/actions/scale.py index 27b92d918..7b2e25fdb 100644 --- a/src/coordinax/transforms/_src/actions/scale.py +++ b/src/coordinax/transforms/_src/actions/scale.py @@ -23,7 +23,9 @@ SMatrix: TypeAlias = Shaped[Array, " N N"] SFactors: TypeAlias = Shaped[Array, " N"] -_MSG_SINGULAR: Final = "Scale matrix must be invertible." +_MSG_SINGULAR: Final = ( + "Scale matrix must be invertible: every factor finite and non-zero." +) _MSG_NOT_DIAGONAL: Final = ( "Scale requires a diagonal matrix -- it scales the axes, and nothing else. " "For a general linear map use `Linear`." @@ -122,9 +124,14 @@ def from_factors(cls: type["Scale"], factors: Any, /) -> "Scale": if s.ndim != 1: msg = f"Scale.from_factors requires a vector; got shape={s.shape!r}." raise ValueError(msg) - # Defer the singular check so it survives jit (a plain `bool` on a - # traced value raises TracerBoolConversionError). - s = eqx.error_if(s, jnp.any(jnp.isclose(s, 0)), _MSG_SINGULAR) + # Deferred so it survives jit (a plain `bool` on a traced value raises + # TracerBoolConversionError). A factor must also be finite: `isclose(s, 0)` + # is False for both NaN and `inf`, and an `inf` factor was the worst of the + # three -- its reciprocal is 0.0, so `inverse` came back finite, singular, + # and silent. + s = eqx.error_if( + s, jnp.any(jnp.isclose(s, 0) | ~jnp.isfinite(s)), _MSG_SINGULAR + ) return cls._from_diagonal(s) @property diff --git a/tests/unit/transforms/test_builders.py b/tests/unit/transforms/test_builders.py index 762c98ed4..00f489380 100644 --- a/tests/unit/transforms/test_builders.py +++ b/tests/unit/transforms/test_builders.py @@ -1,5 +1,6 @@ """Tests for the built-in TimeDep builders.""" +import equinox as eqx import jax import jax.numpy as jnp import pytest @@ -83,10 +84,18 @@ def y(rate_x): assert jnp.allclose(jax.grad(y)(3.0), 2000.0, atol=1e-6) -def test_rotation_about_axis_zero_axis_raises(): - """A zero-length axis must fail loudly, not normalize to a NaN `R`.""" - b = cxfm.builders.RotationAboutAxis(u.Q(1, "rad/s"), axis=jnp.zeros(3)) - with pytest.raises(Exception, match="must be non-zero"): +@pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) +def test_rotation_about_axis_unnormalisable_axis_raises(bad): + """An axis that cannot be normalised must fail loudly, not give a NaN `R`. + + `axis / norm` is a unit vector only for a finite positive norm. A NaN or + `inf` axis is False for `norm == 0` and used to normalise straight through + to nine NaN entries in `R`. + """ + b = cxfm.builders.RotationAboutAxis( + u.Q(1, "rad/s"), axis=jnp.array([bad, 0.0, 0.0]) + ) + with pytest.raises(eqx.EquinoxRuntimeError, match="must be non-zero"): b(u.Q(1.0, "s")) diff --git a/tests/unit/transforms/test_lorentz_boost.py b/tests/unit/transforms/test_lorentz_boost.py index a64b6d5d0..fcabc0661 100644 --- a/tests/unit/transforms/test_lorentz_boost.py +++ b/tests/unit/transforms/test_lorentz_boost.py @@ -173,14 +173,17 @@ def test_a_non_subluminal_boost_is_rejected(self, attr, beta): with pytest.raises(eqx.EquinoxRuntimeError, match="subluminal"): _ = getattr(cxfm.LorentzBoost([beta, 0.0, 0.0]), attr) - def test_zero_direction_is_rejected(self): - """A zero ``direction`` has no axis to normalise onto. - - Dividing by its zero norm produced ``nan`` betas that then propagated - silently into every entry of the matrix. + @pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) + def test_an_unnormalisable_direction_is_rejected(self, bad): + """A ``direction`` that cannot be normalised has no boost axis. + + Dividing by its norm produced ``nan`` betas that then propagated + silently into every entry of the matrix. ``norm == 0.0`` caught only + the zero case; the other two were reported later by ``gamma``'s + subluminal check, which names the wrong cause. """ with pytest.raises(eqx.EquinoxRuntimeError, match="non-zero"): - cxfm.LorentzBoost.from_rapidity(0.5, (0.0, 0.0, 0.0)) + cxfm.LorentzBoost.from_rapidity(0.5, (bad, 0.0, 0.0)) class TestPhysicalPredictions: diff --git a/tests/unit/transforms/test_reflect.py b/tests/unit/transforms/test_reflect.py index 1dff5def8..bff517a61 100644 --- a/tests/unit/transforms/test_reflect.py +++ b/tests/unit/transforms/test_reflect.py @@ -17,11 +17,16 @@ from .conftest import EXPECTED_IDENTITY, EXPECTED_REFLECT -def test_reflect_from_normal_zero_raises_under_jit() -> None: - """A zero normal is rejected even under jit (no tracer bool).""" +@pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) +def test_reflect_from_normal_unnormalisable_raises_under_jit(bad: float) -> None: + """A normal that cannot be normalised is rejected, even under jit. + + Deferred so no tracer bool is taken. `allclose(norm, 0)` was False for a + NaN or `inf` norm, so those reached `n / norm` and gave a NaN `H`. + """ build = eqx.filter_jit(cxfm.Reflect.from_normal) with pytest.raises(eqx.EquinoxRuntimeError, match="nonzero normal"): - jax.block_until_ready(build(jnp.asarray([0.0, 0.0, 0.0])).H) + jax.block_until_ready(build(jnp.asarray([bad, 0.0, 0.0])).H) def _extract_xyz(result: Any) -> tuple[float, float, float]: diff --git a/tests/unit/transforms/test_spatial_linear_transforms.py b/tests/unit/transforms/test_spatial_linear_transforms.py index b63d19014..a9159b3d2 100644 --- a/tests/unit/transforms/test_spatial_linear_transforms.py +++ b/tests/unit/transforms/test_spatial_linear_transforms.py @@ -22,11 +22,17 @@ def _to_np(x: object, unit: str) -> np.ndarray: return np.asarray(u.ustrip(unit, x), dtype=float) -def test_scale_from_factors_singular_raises_under_jit() -> None: - """A zero scale factor is rejected even under jit (no tracer bool).""" +@pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) +def test_scale_from_factors_singular_raises_under_jit(bad: float) -> None: + """A non-invertible scale factor is rejected, even under jit. + + `isclose(s, 0)` is False for both NaN and `inf`. The `inf` case was the + quietest: its reciprocal is 0.0, so `inverse` came back finite, singular, + and with nothing to notice. + """ build = eqx.filter_jit(cxfm.Scale.from_factors) with pytest.raises(eqx.EquinoxRuntimeError, match="invertible"): - jax.block_until_ready(build(jnp.asarray([2.0, 0.0, 4.0])).s) + jax.block_until_ready(build(jnp.asarray([2.0, bad, 4.0])).s) def test_scale_from_factors_nonsingular_jits() -> None: From 8b307f476404bbcc8b5f3593d2bac101c3001e91 Mon Sep 17 00:00:00 2001 From: nstarman Date: Thu, 20 Aug 2026 16:03:32 -0400 Subject: [PATCH 5/7] =?UTF-8?q?=E2=99=BB=EF=B8=8F=20refactor(transforms):?= =?UTF-8?q?=20one=20`=5Funnormalisable`,=20not=20three=20copies=20of=20it?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `~((norm > 0) & jnp.isfinite(norm))` stood verbatim in `builders.py`, `lorentz.py` and `reflect.py`, each under its own restatement of why. Three copies of a predicate can drift, and this PR exists because a guard drifted from its own stated intent. It moves to `actions/utils.py` with the reasoning in its docstring; the three call sites become one line each. Also inline two one-use locals (`tol`, `in_domain`), fold `bishop.py`'s two comments into one that matches the code now that `norm <= tol` is gone, and cut the test docstrings back to what each rejects rather than re-deriving the mechanism the source already states. Net -24 lines. No behaviour change: all twelve bad inputs still raise, all four valid ones still pass, and reverting the helper alone fails all six `nan`/`inf` cases across the three modules. Co-Authored-By: Claude Opus 5 --- .../coordinaxs/curveframes/_src/arclength.py | 3 +-- .../src/coordinaxs/curveframes/_src/bishop.py | 10 +++---- .../tests/unit/test_guards_reject_nan.py | 27 ++++++++++--------- .../transforms/_src/actions/builders.py | 12 +++------ .../transforms/_src/actions/lorentz.py | 13 ++++----- .../transforms/_src/actions/reflect.py | 12 ++++----- .../transforms/_src/actions/scale.py | 15 ++++------- .../transforms/_src/actions/utils.py | 13 +++++++++ tests/unit/transforms/test_builders.py | 7 +---- tests/unit/transforms/test_lorentz_boost.py | 13 ++------- tests/unit/transforms/test_reflect.py | 6 +---- .../test_spatial_linear_transforms.py | 7 +---- 12 files changed, 57 insertions(+), 81 deletions(-) diff --git a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py index 572ddf2b5..75f7bb4fd 100644 --- a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py +++ b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py @@ -215,8 +215,7 @@ def _eval_tau_dense( # Negated `in-domain` rather than the two out-of-domain comparisons: a NaN # `s` compares False against both, so the direct form admits it and hands # back a NaN position with nothing raised. - in_domain = (s_val >= -margin) & (s_val <= s_max_val + margin) - out_of_domain = ~in_domain + out_of_domain = ~((s_val >= -margin) & (s_val <= s_max_val + margin)) # No clip: every point of the solved range is real data, so the only # thing to do with a genuine overshoot is refuse it. diff --git a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py index 85556025b..9f00cd018 100644 --- a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py +++ b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py @@ -159,12 +159,12 @@ def _orthonormalize(v: Any, T0_val: Any) -> Any: """ w = v - jnp.dot(v, T0_val) * T0_val norm = jnp.linalg.norm(w) - # `<=`, not `<`, so that an all-zero `v` (threshold 0) still raises. # `~(norm > tol)`, not `norm <= tol`: a NaN compares False against both, so - # the `<=` form admits a NaN `v` and returns a NaN triad with nothing - # raised. Same reason `TubularChart`'s reach guard is written `~(f > 0)`. - tol = 1e-12 * jnp.linalg.norm(v) - w = eqx.error_if(w, ~(norm > tol), _MSG_PARALLEL_NORMAL) + # the `<=` form admits a NaN `v` and returns a NaN triad with nothing raised. + # Same reason `TubularChart`'s reach guard is written `~(f > 0)`. Negating + # `>` keeps the non-strict sense, so an all-zero `v` (threshold 0) still + # raises. + w = eqx.error_if(w, ~(norm > 1e-12 * jnp.linalg.norm(v)), _MSG_PARALLEL_NORMAL) return w / norm diff --git a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py index 3869f07b9..610f50a8b 100644 --- a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py +++ b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py @@ -24,26 +24,27 @@ def helix(tau: u.AbstractQuantity) -> u.AbstractQuantity: return u.Q(jnp.stack([jnp.cos(t), jnp.sin(t), 0.3 * t]), "km") -@pytest.mark.parametrize("bad", [jnp.nan, jnp.inf], ids=["nan", "inf"]) -def test_a_non_finite_initial_normal_is_rejected(bad: float) -> None: +T0 = jnp.array([1.0, 0.0, 0.0]) + + +@pytest.mark.parametrize( + "v", + [[1.0, 0.0, 0.0], [0.0, 0.0, 0.0], [1.0, jnp.nan, 0.0], [1.0, jnp.inf, 0.0]], + ids=["parallel", "zero", "nan", "inf"], +) +def test_an_unusable_initial_normal_is_rejected(v: list[float]) -> None: """`_orthonormalize` returned a NaN triad instead of raising. The rejection was `norm <= 1e-12 * |v|`; with a NaN `v` both sides are NaN - and the comparison is False. Fails if it goes back to the direct form. + and the comparison is False. The first two cases are what the guard was + written for and must keep raising. """ - T0 = jnp.array([1.0, 0.0, 0.0]) with pytest.raises(eqx.EquinoxRuntimeError, match="parallel"): - _orthonormalize(jnp.array([1.0, bad, 0.0]), T0) - + _orthonormalize(jnp.array(v), T0) -def test_the_legitimate_degenerate_cases_still_raise() -> None: - """The case the guard was written for keeps working.""" - T0 = jnp.array([1.0, 0.0, 0.0]) - for v in (jnp.array([1.0, 0.0, 0.0]), jnp.array([0.0, 0.0, 0.0])): - with pytest.raises(eqx.EquinoxRuntimeError, match="parallel"): - _orthonormalize(v, T0) - # and a well-conditioned normal is untouched +def test_a_well_conditioned_normal_is_untouched() -> None: + """The guard costs the valid case nothing.""" out = _orthonormalize(jnp.array([0.0, 2.0, 0.0]), T0) assert jnp.allclose(jnp.asarray(out), jnp.array([0.0, 1.0, 0.0])) diff --git a/src/coordinax/transforms/_src/actions/builders.py b/src/coordinax/transforms/_src/actions/builders.py index 479ddc921..61eab2d2a 100644 --- a/src/coordinax/transforms/_src/actions/builders.py +++ b/src/coordinax/transforms/_src/actions/builders.py @@ -14,6 +14,7 @@ from .custom_types import CDict from .rotate import Rotate from .translate import Translate +from .utils import _unnormalisable _MSG_ZERO_AXIS = ( "`RotationAboutAxis.axis` must be non-zero and finite; normalising anything " @@ -78,14 +79,9 @@ def __call__(self, tau: Any, /) -> Rotate: """Build the `Rotate` operator at time parameter ``tau``.""" theta = jnp.asarray(u.ustrip("rad", self.omega * tau + self.phase)) norm = jnp.linalg.vector_norm(self.axis) - # `axis / norm` is a unit vector only where `axis` is finite and non-zero, - # which is exactly a finite positive `norm`: the norm is NaN iff a component - # is, `inf` iff a component is, and 0 iff the axis is. Testing `norm == 0` - # alone caught the last of the three and let the other two normalise to the - # silently NaN `R` this guard exists to prevent. `error_if` fires under `jit`. - axis = eqx.error_if( - self.axis, ~((norm > 0) & jnp.isfinite(norm)), _MSG_ZERO_AXIS - ) + # Anything but a finite positive norm normalises to the silently NaN `R` + # this guard exists to prevent. `error_if` fires under `jit`. + axis = eqx.error_if(self.axis, _unnormalisable(norm), _MSG_ZERO_AXIS) n = axis / norm # Rodrigues' formula: R = I cos(th) + sin(th) [n]_x + (1-cos th) n n^T # The `0.0 * n[0]` terms keep the zero entries as functions of `axis` diff --git a/src/coordinax/transforms/_src/actions/lorentz.py b/src/coordinax/transforms/_src/actions/lorentz.py index e5be21546..a27b2f1d6 100644 --- a/src/coordinax/transforms/_src/actions/lorentz.py +++ b/src/coordinax/transforms/_src/actions/lorentz.py @@ -17,6 +17,7 @@ from .base import AbstractTransform from .identity import identity from .linear import AbstractLinearTransform +from .utils import _unnormalisable from coordinax.transforms._src import groups #: Speed of light, used only to convert a velocity into a dimensionless beta. @@ -244,14 +245,10 @@ def from_rapidity( """ d = _float(direction) norm = jnp.linalg.norm(d) - # `d / norm` is a unit vector only where `norm` is finite and positive: it - # is NaN iff a component of `d` is, `inf` iff a component is, and 0 iff `d` - # is. Testing `norm == 0.0` alone caught only the last, so a NaN direction - # built the `nan` betas this guard exists to prevent -- reported later by - # `gamma`'s subluminal check, which names the wrong cause. - norm = eqx.error_if( - norm, ~((norm > 0.0) & jnp.isfinite(norm)), _MSG_ZERO_DIRECTION - ) + # Anything but a finite positive norm builds the `nan` betas this guard + # exists to prevent -- reported later by `gamma`'s subluminal check, + # which names the wrong cause. + norm = eqx.error_if(norm, _unnormalisable(norm), _MSG_ZERO_DIRECTION) return cls(jnp.tanh(_float(rapidity)) * (d / norm)) # ----------------------------------------------------- diff --git a/src/coordinax/transforms/_src/actions/reflect.py b/src/coordinax/transforms/_src/actions/reflect.py index 48cf884f0..8779762c7 100644 --- a/src/coordinax/transforms/_src/actions/reflect.py +++ b/src/coordinax/transforms/_src/actions/reflect.py @@ -17,13 +17,12 @@ from .base import AbstractTransform from .identity import identity from .linear import AbstractLinearTransform +from .utils import _unnormalisable from coordinax.transforms._src import groups HMatrix: TypeAlias = Shaped[Array, " N N"] -_MSG_ZERO_NORMAL: Final = ( - "Reflect.from_normal requires a nonzero normal vector that is finite." -) +_MSG_ZERO_NORMAL: Final = "Reflect.from_normal needs a finite, nonzero normal." @final @@ -79,10 +78,9 @@ def from_normal(cls: type["Reflect"], normal: Any, /) -> "Reflect": norm = jnp.linalg.norm(n) # Deferred so it survives jit (a plain `bool` on a traced value raises - # TracerBoolConversionError). `n / norm` is a unit vector only where `norm` - # is finite and positive; `allclose(norm, 0)` alone was False for a NaN or - # `inf` norm, so those normalised to a silently NaN `H`. - n = eqx.error_if(n, ~((norm > 0) & jnp.isfinite(norm)), _MSG_ZERO_NORMAL) + # TracerBoolConversionError). Anything but a finite positive norm + # normalises to a silently NaN `H`. + n = eqx.error_if(n, _unnormalisable(norm), _MSG_ZERO_NORMAL) n_hat = n / norm H = jnp.eye(n.shape[0], dtype=n_hat.dtype) - 2 * jnp.outer(n_hat, n_hat) diff --git a/src/coordinax/transforms/_src/actions/scale.py b/src/coordinax/transforms/_src/actions/scale.py index 7b2e25fdb..fc539b909 100644 --- a/src/coordinax/transforms/_src/actions/scale.py +++ b/src/coordinax/transforms/_src/actions/scale.py @@ -23,9 +23,7 @@ SMatrix: TypeAlias = Shaped[Array, " N N"] SFactors: TypeAlias = Shaped[Array, " N"] -_MSG_SINGULAR: Final = ( - "Scale matrix must be invertible: every factor finite and non-zero." -) +_MSG_SINGULAR: Final = "Scale matrix must be invertible: factors finite, non-zero." _MSG_NOT_DIAGONAL: Final = ( "Scale requires a diagonal matrix -- it scales the axes, and nothing else. " "For a general linear map use `Linear`." @@ -125,13 +123,10 @@ def from_factors(cls: type["Scale"], factors: Any, /) -> "Scale": msg = f"Scale.from_factors requires a vector; got shape={s.shape!r}." raise ValueError(msg) # Deferred so it survives jit (a plain `bool` on a traced value raises - # TracerBoolConversionError). A factor must also be finite: `isclose(s, 0)` - # is False for both NaN and `inf`, and an `inf` factor was the worst of the - # three -- its reciprocal is 0.0, so `inverse` came back finite, singular, - # and silent. - s = eqx.error_if( - s, jnp.any(jnp.isclose(s, 0) | ~jnp.isfinite(s)), _MSG_SINGULAR - ) + # TracerBoolConversionError). An `inf` factor is the quiet one: its + # reciprocal is 0.0, so `inverse` came back finite, singular, and silent. + bad = jnp.isclose(s, 0) | ~jnp.isfinite(s) + s = eqx.error_if(s, jnp.any(bad), _MSG_SINGULAR) return cls._from_diagonal(s) @property diff --git a/src/coordinax/transforms/_src/actions/utils.py b/src/coordinax/transforms/_src/actions/utils.py index 4d35b1663..c50cf5d7f 100644 --- a/src/coordinax/transforms/_src/actions/utils.py +++ b/src/coordinax/transforms/_src/actions/utils.py @@ -12,6 +12,8 @@ from collections.abc import Iterable from typing import Any +import jax.numpy as jnp + import coordinax.representations as cxr from coordinax._src.exceptions import NoGlobalCartesianChartError @@ -69,3 +71,14 @@ def require_matching_keys( + (f"; unexpected {extra}" if extra else "") + "." ) + + +def _unnormalisable(norm: Any, /) -> Any: + """Whether ``v / norm`` fails to be a unit vector, for ``norm = |v|``. + + True unless the norm is finite and strictly positive, which is exactly the + precondition: a norm is NaN iff a component of ``v`` is, ``inf`` iff a + component is, and ``0`` iff ``v`` is. Testing ``norm == 0`` alone catches + the last of those three and lets the other two normalise to a silent NaN. + """ + return ~((norm > 0) & jnp.isfinite(norm)) diff --git a/tests/unit/transforms/test_builders.py b/tests/unit/transforms/test_builders.py index 00f489380..cbf3dba55 100644 --- a/tests/unit/transforms/test_builders.py +++ b/tests/unit/transforms/test_builders.py @@ -86,12 +86,7 @@ def y(rate_x): @pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) def test_rotation_about_axis_unnormalisable_axis_raises(bad): - """An axis that cannot be normalised must fail loudly, not give a NaN `R`. - - `axis / norm` is a unit vector only for a finite positive norm. A NaN or - `inf` axis is False for `norm == 0` and used to normalise straight through - to nine NaN entries in `R`. - """ + """An axis that cannot be normalised must fail loudly, not give a NaN `R`.""" b = cxfm.builders.RotationAboutAxis( u.Q(1, "rad/s"), axis=jnp.array([bad, 0.0, 0.0]) ) diff --git a/tests/unit/transforms/test_lorentz_boost.py b/tests/unit/transforms/test_lorentz_boost.py index fcabc0661..799e7f886 100644 --- a/tests/unit/transforms/test_lorentz_boost.py +++ b/tests/unit/transforms/test_lorentz_boost.py @@ -165,23 +165,14 @@ def test_a_non_subluminal_boost_is_rejected(self, attr, beta): """Every derived quantity guards, not just ``gamma``. ``rapidity`` used to reach ``arctanh(|beta| >= 1)`` and hand back - ``inf``/``nan`` while ``gamma`` on the same object raised. A ``nan`` - beta needs the guard written as ``~(beta_sq < 1)``: it is False for - ``beta_sq >= 1`` too, so the direct form let it through and returned - the non-finite value the guard exists to prevent. + ``inf``/``nan`` while ``gamma`` on the same object raised. """ with pytest.raises(eqx.EquinoxRuntimeError, match="subluminal"): _ = getattr(cxfm.LorentzBoost([beta, 0.0, 0.0]), attr) @pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) def test_an_unnormalisable_direction_is_rejected(self, bad): - """A ``direction`` that cannot be normalised has no boost axis. - - Dividing by its norm produced ``nan`` betas that then propagated - silently into every entry of the matrix. ``norm == 0.0`` caught only - the zero case; the other two were reported later by ``gamma``'s - subluminal check, which names the wrong cause. - """ + """A ``direction`` that cannot be normalised has no boost axis.""" with pytest.raises(eqx.EquinoxRuntimeError, match="non-zero"): cxfm.LorentzBoost.from_rapidity(0.5, (bad, 0.0, 0.0)) diff --git a/tests/unit/transforms/test_reflect.py b/tests/unit/transforms/test_reflect.py index bff517a61..7e492f454 100644 --- a/tests/unit/transforms/test_reflect.py +++ b/tests/unit/transforms/test_reflect.py @@ -19,11 +19,7 @@ @pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) def test_reflect_from_normal_unnormalisable_raises_under_jit(bad: float) -> None: - """A normal that cannot be normalised is rejected, even under jit. - - Deferred so no tracer bool is taken. `allclose(norm, 0)` was False for a - NaN or `inf` norm, so those reached `n / norm` and gave a NaN `H`. - """ + """A normal that cannot be normalised is rejected, even under jit.""" build = eqx.filter_jit(cxfm.Reflect.from_normal) with pytest.raises(eqx.EquinoxRuntimeError, match="nonzero normal"): jax.block_until_ready(build(jnp.asarray([bad, 0.0, 0.0])).H) diff --git a/tests/unit/transforms/test_spatial_linear_transforms.py b/tests/unit/transforms/test_spatial_linear_transforms.py index a9159b3d2..70ee04b24 100644 --- a/tests/unit/transforms/test_spatial_linear_transforms.py +++ b/tests/unit/transforms/test_spatial_linear_transforms.py @@ -24,12 +24,7 @@ def _to_np(x: object, unit: str) -> np.ndarray: @pytest.mark.parametrize("bad", [0.0, jnp.nan, jnp.inf], ids=["zero", "nan", "inf"]) def test_scale_from_factors_singular_raises_under_jit(bad: float) -> None: - """A non-invertible scale factor is rejected, even under jit. - - `isclose(s, 0)` is False for both NaN and `inf`. The `inf` case was the - quietest: its reciprocal is 0.0, so `inverse` came back finite, singular, - and with nothing to notice. - """ + """A non-invertible scale factor is rejected, even under jit.""" build = eqx.filter_jit(cxfm.Scale.from_factors) with pytest.raises(eqx.EquinoxRuntimeError, match="invertible"): jax.block_until_ready(build(jnp.asarray([2.0, bad, 4.0])).s) From 155001262ff3b31044bd68df12e7e46de01a6a18 Mon Sep 17 00:00:00 2001 From: nstarman Date: Tue, 25 Aug 2026 12:56:38 -0400 Subject: [PATCH 6/7] =?UTF-8?q?=F0=9F=93=9D=20docs(transforms,curveframes)?= =?UTF-8?q?:=20say=20the=20NaN=20reason=20once,=20not=20seven=20times?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `_unnormalisable`'s docstring already states why a guard cannot be written as a direct comparison. Seven comments re-derived it, and four of them opened by restating the predicate the helper had just been extracted to hold. Each site keeps only what is local to it -- a NaN `R`, `nan` betas, a NaN `H`, an `inf` factor whose reciprocal is 0.0 -- and `gamma` points at the comment `rapidity` already carries, the way `geq` points at `leq`. Two error messages also told the user what breaks internally rather than what to pass; both are one line now, and `test_builders.py` matches on the wording that replaced it. Net -20 lines. No behaviour change: the twelve bad inputs still raise with the messages above, and reverting the helper still fails all six `nan`/`inf` cases. Co-Authored-By: Claude Opus 5 --- .../src/coordinaxs/curveframes/_src/arclength.py | 4 +--- .../src/coordinaxs/curveframes/_src/bishop.py | 7 ++----- src/coordinax/_src/charts/checks.py | 3 +-- src/coordinax/transforms/_src/actions/builders.py | 8 ++------ src/coordinax/transforms/_src/actions/lorentz.py | 14 ++++---------- src/coordinax/transforms/_src/actions/reflect.py | 3 +-- src/coordinax/transforms/_src/actions/scale.py | 5 ++--- src/coordinax/transforms/_src/actions/utils.py | 3 +-- tests/unit/transforms/test_builders.py | 2 +- 9 files changed, 15 insertions(+), 34 deletions(-) diff --git a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py index 75f7bb4fd..251b98ec3 100644 --- a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py +++ b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/arclength.py @@ -212,9 +212,7 @@ def _eval_tau_dense( s_max_val = jnp.asarray(s_max.ustrip(s_unit)) margin = _S_MAX_MARGIN * jnp.abs(s_max_val) - # Negated `in-domain` rather than the two out-of-domain comparisons: a NaN - # `s` compares False against both, so the direct form admits it and hands - # back a NaN position with nothing raised. + # A NaN `s` is False for both out-of-domain tests, so negate in-domain. out_of_domain = ~((s_val >= -margin) & (s_val <= s_max_val + margin)) # No clip: every point of the solved range is real data, so the only diff --git a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py index 9f00cd018..0ca1dbc23 100644 --- a/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py +++ b/packages/coordinaxs.curveframes/src/coordinaxs/curveframes/_src/bishop.py @@ -159,11 +159,8 @@ def _orthonormalize(v: Any, T0_val: Any) -> Any: """ w = v - jnp.dot(v, T0_val) * T0_val norm = jnp.linalg.norm(w) - # `~(norm > tol)`, not `norm <= tol`: a NaN compares False against both, so - # the `<=` form admits a NaN `v` and returns a NaN triad with nothing raised. - # Same reason `TubularChart`'s reach guard is written `~(f > 0)`. Negating - # `>` keeps the non-strict sense, so an all-zero `v` (threshold 0) still - # raises. + # `~(norm > tol)`, not `norm <= tol`: a NaN is False for both, so the `<=` + # form returns a NaN triad. Negating `>` keeps an all-zero `v` raising. w = eqx.error_if(w, ~(norm > 1e-12 * jnp.linalg.norm(v)), _MSG_PARALLEL_NORMAL) return w / norm diff --git a/src/coordinax/_src/charts/checks.py b/src/coordinax/_src/charts/checks.py index 61ba8354a..3c19e396d 100644 --- a/src/coordinax/_src/charts/checks.py +++ b/src/coordinax/_src/charts/checks.py @@ -115,8 +115,7 @@ def leq( """ name = f" {name}" if name else name msg = f"The input{name} must be less than or equal to {comp_name}." - # `~(x <= max)` rather than `x > max`: a NaN is False for both, so the - # direct form admits it silently. + # `~(x <= max)`, not `x > max`: a NaN is False for both, admitting it. return eqx.error_if(x, u.ustrip("", jnp.any(~(x <= max_val))), msg) diff --git a/src/coordinax/transforms/_src/actions/builders.py b/src/coordinax/transforms/_src/actions/builders.py index 61eab2d2a..9c134fc9c 100644 --- a/src/coordinax/transforms/_src/actions/builders.py +++ b/src/coordinax/transforms/_src/actions/builders.py @@ -16,10 +16,7 @@ from .translate import Translate from .utils import _unnormalisable -_MSG_ZERO_AXIS = ( - "`RotationAboutAxis.axis` must be non-zero and finite; normalising anything " - "else gives a NaN rotation matrix." -) +_MSG_ZERO_AXIS = "`RotationAboutAxis.axis` must be finite and non-zero." def _as_axis(axis: Any, /) -> Shaped[Array, "3"]: @@ -79,8 +76,7 @@ def __call__(self, tau: Any, /) -> Rotate: """Build the `Rotate` operator at time parameter ``tau``.""" theta = jnp.asarray(u.ustrip("rad", self.omega * tau + self.phase)) norm = jnp.linalg.vector_norm(self.axis) - # Anything but a finite positive norm normalises to the silently NaN `R` - # this guard exists to prevent. `error_if` fires under `jit`. + # Anything else normalises to a silently NaN `R`; `error_if` fires under jit. axis = eqx.error_if(self.axis, _unnormalisable(norm), _MSG_ZERO_AXIS) n = axis / norm # Rodrigues' formula: R = I cos(th) + sin(th) [n]_x + (1-cos th) n n^T diff --git a/src/coordinax/transforms/_src/actions/lorentz.py b/src/coordinax/transforms/_src/actions/lorentz.py index a27b2f1d6..61371982f 100644 --- a/src/coordinax/transforms/_src/actions/lorentz.py +++ b/src/coordinax/transforms/_src/actions/lorentz.py @@ -31,10 +31,7 @@ ) -_MSG_ZERO_DIRECTION = ( - "LorentzBoost.from_rapidity requires a non-zero, finite `direction`; " - "anything else has no boost axis to normalise onto." -) +_MSG_ZERO_DIRECTION = "LorentzBoost.from_rapidity needs a finite, non-zero `direction`." def _float(x: Any, /) -> Array: @@ -245,9 +242,8 @@ def from_rapidity( """ d = _float(direction) norm = jnp.linalg.norm(d) - # Anything but a finite positive norm builds the `nan` betas this guard - # exists to prevent -- reported later by `gamma`'s subluminal check, - # which names the wrong cause. + # Anything else builds `nan` betas -- reported later by `gamma`'s + # subluminal check, which names the wrong cause. norm = eqx.error_if(norm, _unnormalisable(norm), _MSG_ZERO_DIRECTION) return cls(jnp.tanh(_float(rapidity)) * (d / norm)) @@ -279,9 +275,7 @@ def gamma(self) -> Array: """ beta_sq = jnp.sum(self.beta**2) - # `~(x < 1)`, not `x >= 1`: a NaN is False for both comparisons, so the - # direct form admits it and leaks the non-finite value this guard exists - # to stop. + # `~(x < 1)`, not `x >= 1`, for the NaN reason given in `rapidity` below. beta_sq = eqx.error_if(beta_sq, ~(beta_sq < 1.0), _MSG_SUPERLUMINAL) return 1.0 / jnp.sqrt(1.0 - beta_sq) diff --git a/src/coordinax/transforms/_src/actions/reflect.py b/src/coordinax/transforms/_src/actions/reflect.py index 8779762c7..564095c70 100644 --- a/src/coordinax/transforms/_src/actions/reflect.py +++ b/src/coordinax/transforms/_src/actions/reflect.py @@ -78,8 +78,7 @@ def from_normal(cls: type["Reflect"], normal: Any, /) -> "Reflect": norm = jnp.linalg.norm(n) # Deferred so it survives jit (a plain `bool` on a traced value raises - # TracerBoolConversionError). Anything but a finite positive norm - # normalises to a silently NaN `H`. + # TracerBoolConversionError). Anything else normalises to a NaN `H`. n = eqx.error_if(n, _unnormalisable(norm), _MSG_ZERO_NORMAL) n_hat = n / norm diff --git a/src/coordinax/transforms/_src/actions/scale.py b/src/coordinax/transforms/_src/actions/scale.py index fc539b909..37b2514fb 100644 --- a/src/coordinax/transforms/_src/actions/scale.py +++ b/src/coordinax/transforms/_src/actions/scale.py @@ -122,9 +122,8 @@ def from_factors(cls: type["Scale"], factors: Any, /) -> "Scale": if s.ndim != 1: msg = f"Scale.from_factors requires a vector; got shape={s.shape!r}." raise ValueError(msg) - # Deferred so it survives jit (a plain `bool` on a traced value raises - # TracerBoolConversionError). An `inf` factor is the quiet one: its - # reciprocal is 0.0, so `inverse` came back finite, singular, and silent. + # Deferred so it survives jit, as in `Reflect.from_normal`. `inf` is the + # quiet failure: 1/inf = 0.0, so `inverse` came back finite and singular. bad = jnp.isclose(s, 0) | ~jnp.isfinite(s) s = eqx.error_if(s, jnp.any(bad), _MSG_SINGULAR) return cls._from_diagonal(s) diff --git a/src/coordinax/transforms/_src/actions/utils.py b/src/coordinax/transforms/_src/actions/utils.py index c50cf5d7f..e2ec8d631 100644 --- a/src/coordinax/transforms/_src/actions/utils.py +++ b/src/coordinax/transforms/_src/actions/utils.py @@ -78,7 +78,6 @@ def _unnormalisable(norm: Any, /) -> Any: True unless the norm is finite and strictly positive, which is exactly the precondition: a norm is NaN iff a component of ``v`` is, ``inf`` iff a - component is, and ``0`` iff ``v`` is. Testing ``norm == 0`` alone catches - the last of those three and lets the other two normalise to a silent NaN. + component is, and ``0`` iff ``v`` is. """ return ~((norm > 0) & jnp.isfinite(norm)) diff --git a/tests/unit/transforms/test_builders.py b/tests/unit/transforms/test_builders.py index cbf3dba55..f4ee0a5f3 100644 --- a/tests/unit/transforms/test_builders.py +++ b/tests/unit/transforms/test_builders.py @@ -90,7 +90,7 @@ def test_rotation_about_axis_unnormalisable_axis_raises(bad): b = cxfm.builders.RotationAboutAxis( u.Q(1, "rad/s"), axis=jnp.array([bad, 0.0, 0.0]) ) - with pytest.raises(eqx.EquinoxRuntimeError, match="must be non-zero"): + with pytest.raises(eqx.EquinoxRuntimeError, match="finite and non-zero"): b(u.Q(1.0, "s")) From e56385c38171ab55624d9142ad36e3bc34f5f747 Mon Sep 17 00:00:00 2001 From: nstarman Date: Thu, 27 Aug 2026 11:41:04 +0200 Subject: [PATCH 7/7] =?UTF-8?q?=F0=9F=93=9D=20test(curveframes):=20drop=20?= =?UTF-8?q?the=20third=20telling=20of=20the=20NaN=20reason?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The guard file's docstrings still explained the mechanism at length after it was cut from the source comments it duplicates: the module docstring restated what `bishop.py` and `arclength.py` each say beside their own guard. Each test now says what it pins. The arc-length case becomes one parametrize over `nan` and a genuine overshoot, matching the four sibling files and this file's own other test, with the in-domain assertion beside the well- conditioned-normal one that already plays that role. Net -12 lines, same eight cases. Reverting either guard still fails exactly the three non-finite ones. Co-Authored-By: Claude Opus 5 --- .../tests/unit/test_guards_reject_nan.py | 36 +++++++------------ 1 file changed, 12 insertions(+), 24 deletions(-) diff --git a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py index 610f50a8b..2ef6804b1 100644 --- a/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py +++ b/packages/coordinaxs.curveframes/tests/unit/test_guards_reject_nan.py @@ -1,12 +1,7 @@ -"""A NaN must not walk through a guard that claims to reject bad input. +"""The two `curveframes` guards reject what they cannot handle. -`x <= tol` and `x > hi` are both False for a NaN, so a guard written as a -direct comparison admits it and returns a NaN result with nothing raised -- -worse than the case the guard was written for, which at least errors. - -`TubularChart`'s reach guard already avoids this by testing `~(f > 0)`; these -pin the same property for the two guards that were still written the direct -way. +Each covers the degenerate case it was written for and the non-finite ones it +used to admit before being rewritten as a negated positive test. """ import equinox as eqx @@ -33,12 +28,7 @@ def helix(tau: u.AbstractQuantity) -> u.AbstractQuantity: ids=["parallel", "zero", "nan", "inf"], ) def test_an_unusable_initial_normal_is_rejected(v: list[float]) -> None: - """`_orthonormalize` returned a NaN triad instead of raising. - - The rejection was `norm <= 1e-12 * |v|`; with a NaN `v` both sides are NaN - and the comparison is False. The first two cases are what the guard was - written for and must keep raising. - """ + """`parallel` and `zero` are the original cases; the rest returned a NaN triad.""" with pytest.raises(eqx.EquinoxRuntimeError, match="parallel"): _orthonormalize(jnp.array(v), T0) @@ -49,17 +39,15 @@ def test_a_well_conditioned_normal_is_untouched() -> None: assert jnp.allclose(jnp.asarray(out), jnp.array([0.0, 1.0, 0.0])) -def test_a_nan_arc_length_is_rejected() -> None: - """The domain guard returned a NaN position instead of raising. - - The test was `(s < -margin) | (s > s_max + margin)`; a NaN is False for - both, so it fell through to `diffrax`, which happily interpolated a NaN. - """ +@pytest.mark.parametrize("s", [jnp.nan, 99.0], ids=["nan", "outside"]) +def test_an_out_of_domain_arc_length_is_rejected(s: float) -> None: + """A NaN fell through to `diffrax`, which happily interpolated it.""" fast = cxfc.ArcLength(helix, "s", s_max=u.Q(5.0, "km")) with pytest.raises(eqx.EquinoxRuntimeError, match="solved domain"): - fast(u.Q(jnp.nan, "km")) + fast(u.Q(s, "km")) + - # in-domain and genuinely-outside both behave as before +def test_an_in_domain_arc_length_is_untouched() -> None: + """The guard costs the valid case nothing.""" + fast = cxfc.ArcLength(helix, "s", s_max=u.Q(5.0, "km")) assert jnp.isfinite(fast(u.Q(2.0, "km")).ustrip("km")).all() - with pytest.raises(eqx.EquinoxRuntimeError, match="solved domain"): - fast(u.Q(99.0, "km"))