5 ms·
You're discussing switching algorithms like LSODA? This is a really good question that I don't have a neat answer to. You can actually also hit similar issues
by patrickkidger 4y ago
You're discussing switching algorithms like LSODA?
This is a really good question that I don't have a neat answer to. You can actually also hit similar issues when naively vectorising some stiff algorithms that detect when to recalculate Jacobians: at each timestep, some batch elements will want the recalculation whilst some won't.
In both cases, something like what you suggest might be a possible solution. And thankfully support for such custom batching will Soon (tm) be coming to JAX.
In the mean time, Diffrax actually side-steps the problem by not implementing those kinds of solvers. New solvers usually get added on a by-request basis, so this just hasn't been an issue yet.
So:
1. Do you know what DifferentialEquations.jl does in this scenario when running on the GPU? (On the CPU needing to sychronise batch elements usually isn't such a concern.)
2. If the price of using JAX is giving up a bit of performance in this one use case, then I think I'm okay with that. The clear contenders in this space are JAX and Julia, and right now it's a concious choice to be using JAX as a better fit for the problems I'm tackling. This comparison is something I've written about before: https://discourse.julialang.org/t/state-of-machine-learning-in-julia/74385/4 https://discourse.julialang.org/t/state-of-machine-learning-....
- adgjlsfhk1 4y agoOne of the best general purpose solvers (which DifferentialEquations.jl often uses by default) is AutoTsit5. This is a solver that detect stiffness per timestep and uses either Tsit5 (a 5th order RK solver) or Rosenbrock23 depending on the stiffness. These solvers are great because they typically give really good results for a wide variety of problems. The main way DifferentialEquations.jl deals with this currently is being faster on the CPU than other solvers are on GPU. For small models, GPU has too much overhead, and for large models, you can do linear algebra within timesteps on the GPU without compromising solver efficiency.