From 26c1e70cfbbf3bcb7a0969e43a269934bae051d1 Mon Sep 17 00:00:00 2001 From: Gabriel Gerlero Date: Thu, 27 Aug 2026 12:56:31 -0300 Subject: [PATCH] Fix catastrophic cancellation in VanGenuchten near theta_range[0] --- src/frontx/models.py | 3 ++- tests/test_models.py | 37 +++++++++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+), 1 deletion(-) create mode 100644 tests/test_models.py diff --git a/src/frontx/models.py b/src/frontx/models.py index 1faf9a8..610017c 100644 --- a/src/frontx/models.py +++ b/src/frontx/models.py @@ -4,6 +4,7 @@ import equinox as eqx import jax +import jax.numpy as jnp import numpy as np from . import Param @@ -275,7 +276,7 @@ def _kr( /, ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) - return Se**self.l * (1 - (1 - Se ** (1 / self._m)) ** self._m) ** 2 + return Se**self.l * jnp.expm1(self._m * jnp.log1p(-(Se ** (1 / self._m)))) ** 2 class LETxs(_RichardsModel): diff --git a/tests/test_models.py b/tests/test_models.py new file mode 100644 index 0000000..9d25c8c --- /dev/null +++ b/tests/test_models.py @@ -0,0 +1,37 @@ +import jax +import jax.numpy as jnp +import pytest + +from frontx.models import VanGenuchten + +jax.config.update("jax_enable_x64", True) + + +def test_van_genuchten_near_residual() -> None: + D = VanGenuchten(n=1.1) + + # Reference values computed at high precision. + theta = 1e-3 + expected = ( + 2.613452611709472e-36, + 3.0054705034658917e-32, + 3.1557440286391852e-28, + ) + + assert D(theta) == pytest.approx(expected[0], rel=1e-10, abs=0.0) + assert jax.grad(D)(theta) == pytest.approx(expected[1], rel=1e-10, abs=0.0) + assert jax.grad(jax.grad(D))(theta) == pytest.approx( + expected[2], rel=1e-10, abs=0.0 + ) + + +def test_van_genuchten_kr_near_residual_float32() -> None: + model = VanGenuchten(n=1.5) + + theta = jnp.asarray(1e-3, dtype=jnp.float32) + + assert model._kr(theta) == pytest.approx( + 3.5136418469739605e-21, + rel=1e-6, + abs=0.0, + )