Skip to content

Fix NaN step size after a failed step with complex y - #775

Open
etiferrier wants to merge 1 commit into
patrick-kidger:mainfrom
etiferrier:fix-complex-inf-error
Open

etiferrier wants to merge 1 commit into
patrick-kidger:mainfrom
etiferrier:fix-complex-inf-error

Conversation

@etiferrier

Copy link
Copy Markdown

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.

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 patrick-kidger#754),
this already happens with the default dt0, e.g. Kvaerno5 on
`-1j * (sz + cos(t) * sx) @ y`.

Fixes patrick-kidger#774.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Implicit solvers never recover from a failed step when y is complex (NaN step size)

1 participant