# ODEProblem(....) vs NeuralODE(....) for neural ODEs

**URL:** <https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670>\
**Category:** Machine Learning\
**Tags:** flux, optimization, neural-network, diffeqflux\
**Created:** [April 8, 2024, 8:35am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670 "2024-04-08T08:35:11Z")\
**Posts on this page:** 15\
**Page:** 1

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 8, 2024, 8:35am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/1 "2024-04-08T08:35:11Z")

</div>

Hello!

I am wondering about the difference when using `ODEProblem(....)` from DifferentialEquations.jl and ` NeuralODE(....)` from DiffEqFlux.jl in terms of time performance.

Consider the code below, thus, definining the RHS explicitly and solving the neural ODE by using `ODEProblem(....)`:

```julia
dudt2 = Lux.Chain(Lux.Dense(6, 8, swish), 
Lux.Dense(8, 8, swish),
Lux.Dense(8, 8, swish),
Lux.Dense(8, 6))

function rhs!(du, u, p, t)

    û = dudt2(u, p, st)[1]
    du[1] = û[1]
    du[2] = û[2]
    du[3] = û[3]
    du[4] = û[4]
    du[5] = û[5]
    du[6] = û[6]

end

function predict_neuralode(θ,st,dudt2,tspan,tsteps,u0)
    prob_neuralode = ODEProblem(rhs!, u0, tspan)
    _prob = remake(prob_neuralode, p = θ)
    Array(solve(_prob, saveat = tsteps)) 
end

```

Now consider the code below, thus, solving the neural ODE by using `NeuralODE(....)` from Diff

```julia
dudt2 = Lux.Chain(Lux.Dense(6, 8, swish), 
Lux.Dense(8, 8, swish),
Lux.Dense(8, 8, swish),
Lux.Dense(8, 6))

function predict_neuralode(p,st,dudt2,tspan,tsteps,u0)
    prob_neuralode = NeuralODE(dudt2, tspan, saveat = tsteps)
    return Array(prob_neuralode(u0, p, st)[1])
end

```

For my specific problem, the simulation time is 3 times slower when using `ODEProblem(....)` and solving the neural ODE than when using ` NeuralODE(....)`.

What is the reason for it being much slower? And is there a way to fix the significantly weaker time performance?

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 8, 2024, 12:44pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/2 "2024-04-08T12:44:11Z")

</div>

Is it the same solver and options? NeuralODE sets a few defaults that make sense for neural ODEs and optimizes a few things based on how it’s normally used. Check the `solve` results.

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 8, 2024, 8:18pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/3 "2024-04-08T20:18:40Z")

</div>

Thanks! So I studied the output of `solve(....)` (for the second case where I don’t define a RHS function I studied the output of `NeuralODE(....)`). The only difference that I found was that for `interp`, the `cache` is different as I have shown in the figure below. The variable in the top, `pred_neuralnew` is the output from `NeuralODE` and `pred_ode` is the output from `solve` for the `ODEProblem`. Is there a way to change it so that the `cache` is the same for both cases?

 ![image](https://global.discourse-cdn.com/julialang/original/3X/0/e/0ea2bc359d6e3046baa5df68f8c56fd3e5b8137f.png)

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 8, 2024, 8:52pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/4 "2024-04-08T20:52:00Z")

</div>

`dense=false`. You cannot adjoint the interpolation so it must set it to false. I can look at the code later and see.

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 9, 2024, 8:01am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/5 "2024-04-09T08:01:43Z")

</div>

I’ve created and uploaded a toy-example for a model of multiple chemical reactions taking place.  
In this specific case, defining the RHS explicitly and using `ODEProblem(....)` and `solve(....)` together with `dense = false` have similar computation times compared to using only `NeuralODE(....)` however there are still some differences in terms of computation time.

I am also working on a larger scale version of this and the computation time is significantly worse when using `ODEProblem(....)` and `solve(....)` compared to only `NeuralODE(....)`. I have also noticed that the activation function has a huge effect on the similarity of the computation times. For instance when using `tanh()`, the computation times are more similar rather than using relu-like activation functions such as `swish()`.  
I am wondering how else one can modify `ODEProblem(....)` and `solve(....)` so that it is equivalent to `NeuralODE(....)` besides using `dense = false` in `solve`.

[test\_node\_vs\_ode.jl](https://discourse.julialang.org/uploads/short-url/4fKVZ7B6OVPU0zvNFl6S8R7eH16.jl) (2.7 KB)

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 9, 2024, 8:08am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/6 "2024-04-09T08:08:03Z")

</div>

It’s just out of place and ZygoteVJP:

> <https://github.com/SciML/DiffEqFlux.jl/blob/master/src/neural_de.jl#L49-L55>

Did you try and out of place definition?

```julia
function rhs!(u, p, t)
  dudt2(u, p, st)[1]
end

```

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 15, 2024, 7:48am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/7 "2024-04-15T07:48:23Z")

</div>

So I have tried to benchmark the code below vs just using `NeuralODE(....)` and when using Adam the computation time is similar. However when using Adam and switching to BFGS when close to the minima, the code below is significantly faster compared to `NeuralODE(.....)`. Do you know why? And how I can modify the code below so that the computation time is the same as when using `NeuralODE(....)` together with Adam + BFGS?

```julia

function rhs!(du, u, p, t)
    
    û = dudt2(u, p, st)[1]
    du[1] = û[1] 
    du[2] = û[2] 
    du[3] = û[3]
    du[4] = û[4] 
    du[5] = û[5] 
    du[6] = û[6] 

end

basic_tgrad(u,p,t) = zeros(GT_data)

function predict_neuralode(θ,st,dudt2,tspan,tsteps,u0)
    ff = ODEFunction{false}(rhs!; tgrad = basic_tgrad)
    prob = ODEProblem{false}(ff, u0, tspan)
    _prob = remake(prob, p = θ)
    Array(solve(_prob, Vern7(), saveat = tsteps, sensealg = InterpolatingAdjoint(; autojacvec = ZygoteVJP())))
end

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 15, 2024, 11:13am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/8 "2024-04-15T11:13:00Z")

</div>

I would be surprised if the code below runs, since `rhs!(du, u, p, t)` plus `ODEFunction{false}(rhs!; tgrad = basic_tgrad)` is contradictory: the `false` directly implies it’s only looking for a dispatch `rhs(u, p, t)` which doesn’t exist.

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 15, 2024, 11:38am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/9 "2024-04-15T11:38:32Z")

</div>

Okay that’s interesting because it actually did run and even converge. Let me try removing the `{false}` and try benchmarking that.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 15, 2024, 11:39am UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/10 "2024-04-15T11:39:11Z")

</div>

Run it in a new REPL, you’ll see that what you have there requires the function I defined above.

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 15, 2024, 12:42pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/11 "2024-04-15T12:42:58Z")

</div>

So when I remove `{false}` from `ODEFunction(...)` I actually get the following error message:

```julia

ERROR: Nonconforming functions detected. If a model function `f` is defined
as in-place, then all constituent functions like `jac` and `paramjac`
must be in-place (and vice versa with out-of-place). Detected that
some overloads did not conform to the same convention as `f`.

Nonconforming functions: ["tgrad"]

```

However when using `{false}` together with `rhs!(du, u, p, t)`, as I had it before it works well.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 15, 2024, 12:47pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/12 "2024-04-15T12:47:47Z")

</div>

Yes, that’s what I said. The out of place definition:

```julia
function rhs!(u, p, t)
  dudt2(u, p, st)[1]
end

```

is required for the `false` version (that’s what it means), and that’s what’s faster for Zygote reverse mode. It should be faster for Adam and BFGS for this use case. I think the code you’re testing with was just mixing this up.

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 15, 2024, 1:24pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/13 "2024-04-15T13:24:33Z")

</div>

Thanks for the clarification. But what if I want to use the function:

```julia
function rhs!(du, u, p, t)
    
    û = dudt2(u, p, st)[1]
    du[1] = û[1] 
    du[2] = û[2] 
    du[3] = û[3]
    du[4] = û[4] 
    du[5] = û[5] 
    du[6] = û[6] 

end

```

This also only works together with {false}.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [April 15, 2024, 1:29pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/14 "2024-04-15T13:29:11Z")

</div>

> [@KianH](#):
>
> This also only works together with {false}.

No, that’s the in-place function. It only works with `{true}`, which is default preferred. What I’m saying is you probably don’t want to do in-place with neural networks: that’s what currently isn’t optimized in reverse mode.

---

<div class="post-metadata">

**Author:** ![KianH](https://avatars.discourse-cdn.com/v4/letter/k/c6cbf5/32.png) [@KianH](https://discourse.julialang.org/u/KianH)\
**Post date:** [April 15, 2024, 2:16pm UTC](https://discourse.julialang.org/t/odeproblem-vs-neuralode-for-neural-odes/112670/15 "2024-04-15T14:16:22Z")

</div>

Alright thanks!
