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