Diffrax integration - #817
Conversation
a656228 to
b43c245
Compare
| "integrator used in solving the ODE system of the microkinetic " | ||
| "simulation (Kvaerno methods require overreact[fast])" | ||
| ), | ||
| choices=[ |
There was a problem hiding this comment.
It would be nice to build those choices out of simulate._DIFFRAX_METHODS, so that all new methods are listed in a single place.
| method="RK23", | ||
| max_step=np.inf, | ||
| first_step=np.finfo(np.float64).eps, | ||
| first_step=None, |
There was a problem hiding this comment.
This is probably correct. Do you mind commenting?
| First step size. If not given, Diffrax chooses one automatically, | ||
| while the SciPy backend uses `np.finfo(np.float64).eps` for backwards | ||
| compatibility. |
There was a problem hiding this comment.
Now I see. Maybe we leave with None and the first step size gets chosen automatically no matter what?
| if first_step is not None: | ||
| first_step = np.min([first_step, max_step / 2.0]) |
There was a problem hiding this comment.
I'm happy to remove this if it's not needed anymore.
| if first_step is None: | ||
| first_step = np.finfo(np.float64).eps |
There was a problem hiding this comment.
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?
| first_step=first_step, | ||
| rtol=rtol, | ||
| atol=atol, | ||
| # jac=jac, # noqa: ERA001 |
| stepsize_controller = diffrax.PIDController( | ||
| rtol=rtol, | ||
| atol=atol, | ||
| dtmax=max_step, | ||
| ) |
| return y, r | ||
|
|
||
|
|
||
| def _get_y_diffrax(dydt, y0, t_span, method, max_step, first_step, rtol, atol): |
There was a problem hiding this comment.
Maybe add defaults here (in principle the ones in around lines 62-65 above, but I'm open to changing them as well).
| # Avoid differentiating 0**0 for compounds that do not participate in | ||
| # a reaction, this causes NaN fileed jacobians in diffrax otherwise. | ||
| bases = jnp.where(M == 0, 1.0, y) |
| "jax>=0.4", | ||
| "jaxlib>=0.4", |
There was a problem hiding this comment.
Are those required too or diffrax already pulls them in?
This is awesome and would close #771.
Let's put everything behind the I also noted that your change in |
Initial idea for integrating diffrax, currently dispatches to a _get_y_diffrax function if the selected method is one of the provided by this integration
b43c245 to
d600870
Compare
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
fastfeature flag?In this initial implementation I've kept it behind the
fastfeature 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...