diff --git a/diffrax/_global_interpolation.py b/diffrax/_global_interpolation.py index 278d7c5a..5601405d 100644 --- a/diffrax/_global_interpolation.py +++ b/diffrax/_global_interpolation.py @@ -405,7 +405,7 @@ def _check_ts(ts: Real[Array, " times"]) -> Real[Array, " times"]: raise ValueError(f"`ts` must be of length at least 2; got {ts.shape[0]}") # Also catches any NaN times. ts = eqxi.error_if( - ts, ts[:-1] >= ts[1:], "`ts` must be monotonically strictly increasing." + ts, ~(ts[:-1] < ts[1:]), "`ts` must be monotonically strictly increasing." ) return ts diff --git a/test/test_global_interpolation.py b/test/test_global_interpolation.py index 9c1c236c..30e069bd 100644 --- a/test/test_global_interpolation.py +++ b/test/test_global_interpolation.py @@ -3,6 +3,7 @@ from typing import cast import diffrax +import equinox as eqx import jax import jax.numpy as jnp import jax.random as jr @@ -388,3 +389,28 @@ def test_dense_interpolation_vmap(solver, getkey): diffrax.Ralston: 1e-3, }.get(type(solver), 1e-6) assert tree_allclose(derivs, true_derivs, atol=deriv_tol, rtol=deriv_tol) + + +@pytest.mark.parametrize( + "interpolate", + [ + diffrax.linear_interpolation, + diffrax.rectilinear_interpolation, + diffrax.backward_hermite_coefficients, + ], +) +@pytest.mark.parametrize("dtype", [jnp.float32, jnp.float64]) +@pytest.mark.parametrize("compile", [False, True]) +@pytest.mark.parametrize("nan_index", [0, 1, 2]) +def test_interpolation_rejects_nan_times(interpolate, dtype, compile, nan_index): + """NaN knots violate strict ordering at the beginning, middle and end.""" + ts = jnp.array([0.0, 1.0, 2.0], dtype=dtype).at[nan_index].set(jnp.nan) + ys = jnp.array([0.0, 1.0, 2.0], dtype=dtype) + if compile: + interpolate = eqx.filter_jit(interpolate) + + with pytest.raises( + (ValueError, RuntimeError), match="monotonically strictly increasing" + ): + result = interpolate(ts, ys) + jtu.tree_map(lambda x: x.block_until_ready(), result)