Skip to content

Diffrax integration - #817

Open
caprilesport wants to merge 11 commits into
geem-lab:mainfrom
caprilesport:diffrax
Open

Diffrax integration#817
caprilesport wants to merge 11 commits into
geem-lab:mainfrom
caprilesport:diffrax

Conversation

@caprilesport

@caprilesport caprilesport commented Aug 7, 2026

Copy link
Copy Markdown
Member

Initial integration of the Diffrax solvers, mainly to add the stiff solvers it provides.

The diffrax solver requires Jax, so the first question is, we now require Jax as a core feature or hide diffrax behind the fast feature flag?

In this initial implementation I've kept it behind the fast feature flag, but an option would be to replace the scipy solvers with the diffrax ones, as I think they are much more solid than the ones scipy exposes.

The only problem is that if someone wants to reproduce old simulated system this would be a bit of a breaking update, so I think it would be better to still keep the old solvers in the current proposed change? Open to feedback regarding this.

I haven't yet throughtly tested this in a "real world scenario", but the plan is to try to get some stiff system this week to compare the scipy solvers and the newly added Kvaerno solvers to see if they can handle it better.

I've also noted that I accidently linked the diffrax issue (#771) when doing the uv migration (which should have been #759, oopsie) instead of linking the correct issue. So if you want to re-open that one and close the uv issue...

@schneiderfelipe schneiderfelipe linked an issue Aug 14, 2026 that may be closed by this pull request
3 tasks

@schneiderfelipe schneiderfelipe left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Good stuff!

Comment thread overreact/_cli.py Outdated
Comment thread overreact/simulate.py
Comment thread overreact/simulate.py
Comment thread overreact/simulate.py Outdated
Comment thread overreact/simulate.py Outdated
Comment on lines +188 to +189
if first_step is None:
first_step = np.finfo(np.float64).eps

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Let's either pass a given first step through or let the underlying algorithms decide it. I think it is simpler and better matches diffrax?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yeah I've made a bit of a mess of it I think. Now it's either set by the user, kept as None (in the case of diffrax, then it chooses one) or the old heuristic is applied. I think keeping the old behaviour at least for now is the call, but if you want to change this I'm happy to do it in this PR.

Comment thread overreact/simulate.py Outdated
Comment thread overreact/simulate.py
Comment thread overreact/simulate.py Outdated
Comment thread overreact/simulate.py
Comment thread pyproject.toml Outdated
@schneiderfelipe

Copy link
Copy Markdown
Member

Initial integration of the Diffrax solvers, mainly to add the stiff solvers it provides.

This is awesome and would close #771.

The diffrax solver requires Jax, so the first question is, we now require Jax as a core feature or hide diffrax behind the fast feature flag?

In this initial implementation I've kept it behind the fast feature flag, but an option would be to replace the scipy solvers with the diffrax ones, as I think they are much more solid than the ones scipy exposes.

The only problem is that if someone wants to reproduce old simulated system this would be a bit of a breaking update, so I think it would be better to still keep the old solvers in the current proposed change? Open to feedback regarding this.

I haven't yet throughtly tested this in a "real world scenario", but the plan is to try to get some stiff system this week to compare the scipy solvers and the newly added Kvaerno solvers to see if they can handle it better.

Let's put everything behind the fast flag for now and leave having Diffrax being default a decision for the future.

I also noted that your change in _dydt might solve #422. Would you like to try solving it in this PR as well?

Uses the constants defined in overreact.simulate to sync the possible
choices with the simulate module.

Improves the error message to also print the available scipy solvers
that don't require jax.

A supported_methods constant also is exposed by the module.
After avoiding differentiating 0**0 we basically solved the issue where
we couldn't use jax in solve_ivp.
@caprilesport

Copy link
Copy Markdown
Member Author

I think all the comments are adressed, If there's anything else you think is needed I'm happy to include in this PR.

I also noted that your change in _dydt might solve #422. Would you like to try solving it in this PR as well?

Yes it does! I didn't quite know why we left it disabled, but just allowing it and running the test suite runs smoothly. So apparently after we merge we can also close #422 :)

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.

Consider migrating to diffrax

2 participants