From 0d3ec06d7c750abe341c6aa03eedbf4bdec4c469 Mon Sep 17 00:00:00 2001 From: Etienne Ferrier Date: Tue, 29 Sep 2026 19:57:55 +0000 Subject: [PATCH] Fix NaN step size after a failed step with complex y When an implicit step's root find fails, the solver reports `y_error = inf`, which is `inf + 0j` for complex `y`. `PIDController` then divides it by the real tolerance scale `atol + rtol * |y|`. JAX promotes the scale to complex, and complex division computes `inf * 0 = nan` for the imaginary part, so the scaled error is `inf + nanj`. Every lineax norm of that is NaN (`jnp.abs(inf + nanj)` is NaN in JAX), so the next step size is NaN and the solve loops until `max_steps`. With real `y` the same path gives `inf`: the step is rejected and dt shrinks, as intended. Divide the real and imaginary parts separately instead. This only changes complex `y`, and gives the same values as before whenever the error is finite. The fix is in the controller rather than where `y_error = inf` is set, because the error must keep the dtype of `y`, and any complex `inf` divided by a promoted real gives a NaN. The controller also covers every source of `inf`: the Runge-Kutta and implicit Euler failure paths and the NaN-to-inf mapping in `_integrate.py`. This affects every implicit solver (ImplicitEuler, Kvaerno*, KenCarp*) with complex `y`: any step whose root find fails ends the solve. In v0.7.2, where VeryChord gives up after two iterations (before #754), this already happens with the default dt0, e.g. Kvaerno5 on `-1j * (sz + cos(t) * sx) @ y`. Fixes #774. Co-Authored-By: Claude Opus 5.5 --- diffrax/_step_size_controller/pid.py | 10 ++++++- test/test_solver.py | 43 ++++++++++++++++++++++++++++ 2 files changed, 52 insertions(+), 1 deletion(-) diff --git a/diffrax/_step_size_controller/pid.py b/diffrax/_step_size_controller/pid.py index 752b9bed..aa87bc39 100644 --- a/diffrax/_step_size_controller/pid.py +++ b/diffrax/_step_size_controller/pid.py @@ -487,7 +487,15 @@ def _scale(_y0, _y1_candidate, _y_error): _y1_candidate = jnp.where(_nan, _y0, _y1_candidate) _y = jnp.maximum(jnp.abs(_y0), jnp.abs(_y1_candidate)) with jax.numpy_dtype_promotion("standard"): - return _y_error / (self.atol + _y * self.rtol) + _scale = self.atol + _y * self.rtol + if jnp.iscomplexobj(_y_error): + # Divide the real and imaginary parts separately. A failed step + # (e.g. an implicit solve that did not converge) reports + # `y_error = inf`, i.e. `inf + 0j`. Complex division by `_scale` + # would give `inf + nanj`, hence a NaN norm and a NaN step size, + # and the solve would never recover. + return lax.complex(_y_error.real / _scale, _y_error.imag / _scale) + return _y_error / _scale scaled_error = self.norm(jtu.tree_map(_scale, y0, y1_candidate, y_error)) keep_step = scaled_error < 1 diff --git a/test/test_solver.py b/test/test_solver.py index dfff3f6b..b7adad93 100644 --- a/test/test_solver.py +++ b/test/test_solver.py @@ -54,6 +54,49 @@ def test_implicit_euler_adaptive(): assert out2.result == diffrax.RESULTS.successful +@pytest.mark.parametrize( + "solver", (diffrax.ImplicitEuler(), diffrax.Kvaerno3(), diffrax.Kvaerno5()) +) +@pytest.mark.parametrize( + "dtype, complex_dtype", + ((jnp.float64, jnp.complex128), (jnp.float32, jnp.complex64)), +) +def test_implicit_adaptive_complex_failed_step(solver, dtype, complex_dtype): + # `dt0=1` is too large for the root find, so the first step fails and reports + # `y_error=inf`. The step must be rejected and retried with a smaller `dt`. For + # complex `y` this used to give a NaN scaled error, hence a NaN `dt`, and the + # solve never finished. The complex solve should match the real one. + term = diffrax.ODETerm(lambda t, y, args: -10 * y**3) + t0 = jnp.array(0, dtype) + t1 = jnp.array(1, dtype) + dt0 = jnp.array(1, dtype) + tol = 1e-5 if dtype == jnp.float64 else 1e-3 + stepsize_controller = diffrax.PIDController(rtol=tol, atol=tol) + sols = [] + for y_dtype in (dtype, complex_dtype): + sol = diffrax.diffeqsolve( + term, + solver, + t0, + t1, + dt0, + jnp.array(1, y_dtype), + stepsize_controller=stepsize_controller, + max_steps=1000, + throw=False, + ) + assert sol.result == diffrax.RESULTS.successful + assert sol.stats["num_rejected_steps"] > 0 + sols.append(sol) + real_sol, complex_sol = sols + assert complex_sol.ys is not None and real_sol.ys is not None + assert complex_sol.stats["num_steps"] == real_sol.stats["num_steps"] + assert tree_allclose(complex_sol.ys.real, real_sol.ys) + assert tree_allclose(complex_sol.ys.imag, jnp.zeros_like(real_sol.ys)) + true_y1 = jnp.array([1 / 21**0.5], dtype) + assert tree_allclose(real_sol.ys, true_y1, rtol=100 * tol, atol=100 * tol) + + class _DoubleDopri5(diffrax.AbstractRungeKutta): tableau: ClassVar[diffrax.MultiButcherTableau] = diffrax.MultiButcherTableau( diffrax.Dopri5.tableau, diffrax.Dopri5.tableau