Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 9 additions & 1 deletion diffrax/_step_size_controller/pid.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
43 changes: 43 additions & 0 deletions test/test_solver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down