# Parameter Types in DiffEqFlux.jl versus DifferentialEquations.jl

**URL:** <https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635>\
**Category:** Modelling & Simulations\
**Tags:** diffeq\
**Created:** [April 18, 2022, 1:39pm UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635 "2022-04-18T13:39:05Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![jbu](https://avatars.discourse-cdn.com/v4/letter/j/a8b319/32.png) [@jbu](https://discourse.julialang.org/u/jbu)\
**Post date:** [April 18, 2022, 1:39pm UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/1 "2022-04-18T13:39:05Z")

</div>

What are the differences between allowed parameter types in DiffEqFlux.jl versus DifferentialEquations.jl?

The documentation for DifferentialEquations [states the following](https://diffeq.sciml.ai/stable/tutorials/ode_example/#Defining-Parameterized-Functions):

> Note that the type for the parameters `p` can be anything: you can use arrays, static arrays, named tuples, etc. to enclose your parameters in a way that is sensible for your problem.

For DiffEqFlux though, this doesn’t seem to be the case. Based on [this issue](https://github.com/SciML/DiffEqFlux.jl/issues/178#issuecomment-603141784) (and my own testing), it appears that automatic differentiation with DiffEqFlux only works when `p` is an N-dimensional array. In particular, making `p` an array of arrays or an array of mutable structs seems to cause errors when trying to compute gradients via Zygote or ForwardDiff.

I couldn’t find anything in the DiffEqFlux documentation directly stating that `p` was restricted to being an N-dimensional array. I was wondering if someone could confirm this or clarify how to make additional types of `p` work with DiffEqFlux.

If it’s helpful, here’s a MWE (adapted from [this issue post](https://github.com/SciML/DiffEqFlux.jl/issues/178#issuecomment-603141784)):

> ****
>
> Software versions:
> 
> Julia: v1.7.2  
> DifferentialEquations: v7.1.0  
> DiffEqFlux: v1.45.3  
> Zygote: v0.6.38  
> ForwardDiff: v0.10.25
> 
> ```julia
> using Flux, DiffEqFlux, OrdinaryDiffEq, ForwardDiff
> 
> # ODE functions
> 
> function f1!(dx, x, p, t)
> dx[1] = p[1, 1]
> dx[2] = p[2, 1]
> end
> p1 = [1. 5.; 5. 1.]
> 
> function f2!(dx, x, p, t)
> dx[1] = p[1][1]
> dx[2] = p[2][1]
> end
> p2 = [[1.], [5.]]
> 
> function f3!(dx, x, p, t)
> dx[1] = p.one
> dx[2] = p.two
> end
> 
> mutable struct two_p
> one::Float64
> two::Float64
> end
> 
> # Define length() and iterate() to remove (some) Zygote errors
> Base.length(X::two_p) = 2
> Base.getindex(X::two_p, i::Int) = begin
> if i == 1
> return X.one
> elseif i == 2
> return X.two
> else
> BoundsError()
> end
> end
> Base.iterate(X::two_p, state=1) = begin
> if 1 <= state <= 2
> return (X[state], state+1)
> elseif state > 2
> return nothing
> else
> error()
> end
> end
> 
> p3 = two_p(1.,5.)
> 
> # ODE Problem setup
> 
> x0 = [1., 2.]
> tspan = (0., 2.)
> prob1 = ODEProblem(f1!, x0, tspan, p1)
> prob2 = ODEProblem(f2!, x0, tspan, p2)
> prob3 = ODEProblem(f3!, x0, tspan, p3)
> 
> function predict_adjoint1(p) 
> Array(concrete_solve(prob1, Tsit5(), x0, p))
> end
> 
> function predict_adjoint2(p) 
> Array(concrete_solve(prob2, Tsit5(), x0, p))
> end
> 
> function predict_adjoint3(p) 
> Array(concrete_solve(prob3, Tsit5(), x0, p))
> end
> 
> # Loss functions
> 
> function loss_adjoint1(p)
> prediction = predict_adjoint1(p)
> loss = sum(abs2, prediction[:,end] .-1)
> loss
> end
> 
> function loss_adjoint2(p)
> prediction = predict_adjoint2(p)
> loss = sum(abs2, prediction[:,end] .-1)
> loss
> end
> 
> function loss_adjoint3(p)
> prediction = predict_adjoint3(p)
> loss = sum(abs2, prediction[:,end] .-1)
> loss
> end
> 
> # Compute Gradients
> 
> Zygote.gradient(loss_adjoint1,p1) # Works
> ForwardDiff.gradient(loss_adjoint1,p1) # Works
> 
> Zygote.gradient(loss_adjoint2,p2) # ERROR: MethodError: no method matching Float64(::Vector{Float64})
> ForwardDiff.gradient(loss_adjoint2,p2) # ERROR: MethodError: no method matching one(::Type{Vector{Float64}})
> 
> Zygote.gradient(loss_adjoint3,p3) # ERROR: type Array has no field one 
> ForwardDiff.gradient(loss_adjoint3,p3) # ERROR: MethodError: no method matching gradient(::typeof(loss_adjoint3), ::two_p)
> 
> ```

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [April 18, 2022, 7:40pm UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/2 "2022-04-18T19:40:53Z")

</div>

[Related](https://discourse.julialang.org/t/using-flux-gradient-on-differentialequations-solve-results-in-an-error/79564)?

---

<div class="post-metadata">

**Author:** ![00krishna](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/00krishna/32/8843_2.png) [@00krishna](https://discourse.julialang.org/u/00krishna)\
**Post date:** [April 20, 2022, 5:34am UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/3 "2022-04-20T05:34:57Z")

</div>

If I remember correctly, I think that `DiffEqFlux` works with `ComponentArrays` for the parameters. I don’t think that a struct would work, but I am also pretty sure that a struct won’t work in DifferentialEquations.jl for the parameters. I would have to think about it, but I believe the gradient tape has difficulty tracing back through operations on the struct.

---

<div class="post-metadata">

**Author:** ![jonniedie](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jonniedie/32/12842_2.png) [@jonniedie](https://discourse.julialang.org/u/jonniedie)\
**Post date:** [April 20, 2022, 7:52pm UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/4 "2022-04-20T19:52:58Z")

</div>

Yep, `ComponentArrays` should work here. If it doesn’t, please open an issue!

---

<div class="post-metadata">

**Author:** ![goerch](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/goerch/32/29122_2.png) [@goerch](https://discourse.julialang.org/u/goerch)\
**Post date:** [April 20, 2022, 9:01pm UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/5 "2022-04-20T21:01:14Z")

</div>

Sorry to insist, is this the proposed solution for the issue referenced above too? Then I’ll try that. Thanks!

---

<div class="post-metadata">

**Author:** ![jbu](https://avatars.discourse-cdn.com/v4/letter/j/a8b319/32.png) [@jbu](https://discourse.julialang.org/u/jbu)\
**Post date:** [April 21, 2022, 4:12am UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/6 "2022-04-21T04:12:11Z")

</div>

Thanks for this information–this is helpful to know! That makes sense that it’s not possible to set `p` equal to a struct.

I guess what’s confused me is that DifferentialEquations.jl works perfectly fine with `p` being a `Vector{Any}` containing a mix of scalars, arrays, structs, and even functions. But taking the gradient when `p` is a `Vector{Any}` doesn’t seem to work.

As an example, let’s say we have the following code:

```julia
using DifferentialEquations, DiffEqFlux

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

    du[1] = -p[1]'*u
    du[2] = (p[2].a + p[2].b)u[2]
    du[3] = p[3](u,t)
    return nothing
end

struct mystruct
    a
    b
end

function control(u,t)
    return -exp(-t)*u[3]
end

u0 = [10,15,20]
p = [[1;2;3], mystruct(-1,-2), control]
tspan = (0.0,10.0)

prob = ODEProblem(f!,u0, tspan, p)

sol = solve(prob, Tsit5()) # Solves without errors

```

This code runs without errors, which is great! The parameter vector `p` contains an array, a struct, and even a function, and everything works perfectly fine.

However, let’s say we want to define a loss function with respect to the first entry of `p` (e.g. the entry `[1;2;3]`) and take the gradient:

```julia
function loss(p1)
    sol = solve(prob, Tsit5(), p=[p1, mystruct(-1,-2), control])
    return sum(abs2, sol)
end

grad(p) = Zygote.gradient(loss, p)

p2 = [4;5;6]
grad(p2) # ERROR: MethodError: no method matching Int64(::Vector{Int64})

```

Even though DifferentialEquations handled the ODE solving just fine, Zygote crashes when taking the gradient for the first entry of `p`. Using ForwardDiff also results in an error:

```julia
gradF(p) = ForwardDiff.gradient(loss,p)
gradF(p2) # ERROR: TypeError: in typeassert, expected Float64, got a value of type ForwardDiff.Dual{Nothing, Float64, 3}

```

Since DifferentialEquations works with a `p` of type `Vector{Any}`, it’s difficult to tell whether the Zygote/ForwardDiff errors are due to the type of `p` or some other problem with the function `f!`.

(Granted, I could be doing something wrong–I’d love to know if I’m missing something here)

---

<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:** [June 4, 2022, 3:38pm UTC](https://discourse.julialang.org/t/parameter-types-in-diffeqflux-jl-versus-differentialequations-jl/79635/7 "2022-06-04T15:38:16Z")

</div>

I added much better error messages in this PR:

[https://github.com/SciML/DiffEqSensitivity.jl/pull/596](https://github.com/SciML/DiffEqSensitivity.jl/pull/596)

That should give a lot more clarity here.

Note that the current interface is kind of “requires AbstractArray”, but in reality there’s a bit more generality than it could have with a `SciMLParameters Interface`, which just hasn’t been written down and fully described. I plan to create this interface package rather soon though.
