Skip to content
Original file line number Diff line number Diff line change
Expand Up @@ -212,7 +212,8 @@ 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)
# 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
# thing to do with a genuine overshoot is refuse it.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -159,8 +159,9 @@ 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 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


Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
"""The two `curveframes` guards reject what they cannot handle.

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
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")


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:
"""`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)


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]))


@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(s, "km"))


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()
6 changes: 4 additions & 2 deletions src/coordinax/_src/charts/checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,7 +115,8 @@ 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)`, not `x > max`: a NaN is False for both, admitting it.
return eqx.error_if(x, u.ustrip("", jnp.any(~(x <= max_val))), msg)


def geq(
Expand Down Expand Up @@ -144,7 +145,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(
Expand Down
8 changes: 4 additions & 4 deletions src/coordinax/transforms/_src/actions/builders.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,9 @@
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; got a zero-length axis."
_MSG_ZERO_AXIS = "`RotationAboutAxis.axis` must be finite and non-zero."


def _as_axis(axis: Any, /) -> Shaped[Array, "3"]:
Expand Down Expand Up @@ -75,9 +76,8 @@ 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)
# 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
# The `0.0 * n[0]` terms keep the zero entries as functions of `axis`
Expand Down
17 changes: 8 additions & 9 deletions src/coordinax/transforms/_src/actions/lorentz.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -30,10 +31,7 @@
)


_MSG_ZERO_DIRECTION = (
"LorentzBoost.from_rapidity requires a non-zero `direction`; the zero "
"vector has no boost axis to normalise onto."
)
_MSG_ZERO_DIRECTION = "LorentzBoost.from_rapidity needs a finite, non-zero `direction`."


def _float(x: Any, /) -> Array:
Expand Down Expand Up @@ -244,9 +242,9 @@ 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)
# 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))

# -----------------------------------------------------
Expand Down Expand Up @@ -277,7 +275,8 @@ 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`, 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)

@property
Expand All @@ -297,7 +296,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)

# -----------------------------------------------------
Expand Down
9 changes: 5 additions & 4 deletions src/coordinax/transforms/_src/actions/reflect.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +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."
_MSG_ZERO_NORMAL: Final = "Reflect.from_normal needs a finite, nonzero normal."


@final
Expand Down Expand Up @@ -76,9 +77,9 @@ 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). Anything else normalises to a 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)
Expand Down
9 changes: 5 additions & 4 deletions src/coordinax/transforms/_src/actions/scale.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@
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: 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`."
Expand Down Expand Up @@ -122,9 +122,10 @@ 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, 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)

@property
Expand Down
12 changes: 12 additions & 0 deletions src/coordinax/transforms/_src/actions/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -69,3 +71,13 @@ 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.
"""
return ~((norm > 0) & jnp.isfinite(norm))
15 changes: 15 additions & 0 deletions tests/unit/charts/test_checks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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"))
12 changes: 8 additions & 4 deletions tests/unit/transforms/test_builders.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Tests for the built-in TimeDep builders."""

import equinox as eqx
import jax
import jax.numpy as jnp
import pytest
Expand Down Expand Up @@ -83,10 +84,13 @@ 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`."""
b = cxfm.builders.RotationAboutAxis(
u.Q(1, "rad/s"), axis=jnp.array([bad, 0.0, 0.0])
)
with pytest.raises(eqx.EquinoxRuntimeError, match="finite and non-zero"):
b(u.Q(1.0, "s"))


Expand Down
16 changes: 7 additions & 9 deletions tests/unit/transforms/test_lorentz_boost.py
Original file line number Diff line number Diff line change
Expand Up @@ -160,23 +160,21 @@ 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.
"""
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.

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."""
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:
Expand Down
7 changes: 4 additions & 3 deletions tests/unit/transforms/test_reflect.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,12 @@
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."""
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]:
Expand Down
7 changes: 4 additions & 3 deletions tests/unit/transforms/test_spatial_linear_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,11 +22,12 @@ 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."""
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:
Expand Down