4 ms·
Taking a union of differential equations is a pretty bad idea for exactly the reasons you describe. But it's absolutely possible to parallelise multiple diffeq
by patrickkidger 4y ago
Taking a union of differential equations is a pretty bad idea for exactly the reasons you describe. But it's absolutely possible to parallelise multiple diffeq solves without needing to use the same time steps for each solve. That is, the time steps can be vectorised, rather than broadcast, over the batch axis.
So the only thing you actually need are the same number of steps for each solve, which can be easily accomplished by just padding out the solves that finish with slightly fewer steps. In practice this ends up introducing negligible overhead whilst solving the above issue very neatly. For example this is precisely what Diffrax (https://github.com/patrick-kidger/diffrax https://github.com/patrick-kidger/diffrax) does under `jax.vmap`.
I've not dug into what Julia does here; is this not already done when broadcasting `DifferentialEquations.solve`?
- adgjlsfhk1 4y agoHow does this work with stiffness adaptive solvers? The computation you need to do for when an equation is stiff looks pretty different from when it's not stiff, and I'm having a hard time seeing how that would work if you vectorize the timesteps. Do you need to synchronize after every timestep to split into separate batches for stiff and non-stiff?
- patrickkidger 4y agoYou'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.
- ChrisRackauckas 4y agoYeah, that's why I don't like vmap so much. It just never gets you an optimal algorithm. On CPUs, you don't need to do the same number of steps each solve: you might as well fully decouple all of the solves and max out the threads. Broadcasting DiffEq.solve will do this fully decoupled version, but it's better to parallelize it either with multithreading and distributed (or polyester). That's of course the better thing to do on CPU, and then you don't need to worry about filling kernels or whatnot because you can just oversubscribe if you have to. On GPU vmap isn't as bad at face value, I thought it would be fairly close to optimal, but when we dug in we found it wasn't the right approach for the ODEs we were looking at either. Julia has KernelAbstractions.jl which can do similar code transformations as vmap, but when digging into the CUDA profiler we found it takes a pretty substantial number of trajectories to fill the kernels. There were a few more optimizations we could do to it, but it was reaching its limit rather quickly. Meanwhile, it turns out that if you can just generate .ptx kernels for the ODE, minimize the usage of global memory by keeping all ODE parts locally in registers, and then only couple within warps, you can beat that by about 100x. So some of our other GPU stuff has gone stale for a bit while we've been working to get this approach completed. It was somewhat of a disappointing conclusion because it means we have to really specialize some codes for GPU, but anyways, full details will come probably around the end of summer.