Fix NaN step size after a failed step with complex y - #775
Open
etiferrier wants to merge 1 commit into
Open
etiferrier wants to merge 1 commit into
etiferrier wants to merge 1 commit into
Conversation
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
When an implicit step's root find fails, the solver reports
y_error = inf, which isinf + 0jfor complexy.PIDControllerthen divides it by the real tolerance scaleatol + rtol * |y|. JAX promotes the scale to complex, and complex division computesinf * 0 = nanfor the imaginary part, so the scaled error isinf + 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 untilmax_steps. With realythe same path givesinf: 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 wherey_error = infis set, because the error must keep the dtype ofy, and any complexinfdivided by a promoted real gives a NaN. The controller also covers every source ofinf: 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.