# Differentiate through solve for a subset of parameters

**URL:** <https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869>\
**Category:** Modelling & Simulations\
**Tags:** diffeq\
**Created:** [November 3, 2021, 12:54pm UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869 "2021-11-03T12:54:34Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![hexaeder](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/hexaeder/32/24403_2.png) [@hexaeder](https://discourse.julialang.org/u/hexaeder)\
**Post date:** [November 3, 2021, 12:54pm UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869/1 "2021-11-03T12:54:34Z")

</div>

If I pass a vector of dual numbers as parameters, the solver will recocnize the dual numbers and provide a vector `dx<:Vector{Dual...}` to my diff eq function.

In my problem I have several parameters I want to tune but also additional static parameters. I tried to wrap them into a Tuple. But now DiffEq won’t provide `dx` as a vector of dual numbers anymore which leads to errors. Can I somehow tell the solve manually that I wan’t to diff?

```julia
using Pkg
Pkg.activate(temp=true)
Pkg.add("OrdinaryDiffEq")

using OrdinaryDiffEq
using OrdinaryDiffEq.ForwardDiff

function diffeq(du, u, (p, fixed_p), t)
    p1, p2 = p
    du[1] = -p1*u[1] + p2 + fixed_p
end

function loss(p)
    prob = ODEProblem(diffeq, ones(1), (0.0, 1.0), (p, 10.0))
    sol = solve(prob, Tsit5())
    return sol[end][1]
end

loss([1.0, 2.0])

ForwardDiff.gradient(loss2, [1.0, 2.0])

```

Error: `du<:Vector{Float64}` and can’t take any dual numbers.

---

<div class="post-metadata">

**Author:** ![lindnemi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lindnemi/32/23421_2.png) [@lindnemi](https://discourse.julialang.org/u/lindnemi)\
**Post date:** [November 3, 2021, 1:23pm UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869/2 "2021-11-03T13:23:04Z")

</div>

This came up before: [https://github.com/SciML/DiffEqFlux.jl/issues/178](https://github.com/SciML/DiffEqFlux.jl/issues/178)

What worked for me was to use a wrapper:

```julia
function wrapper(dx, x, p, t)
  diffeq(dx, x, (p, fixed_p) ,t)
end

```

---

<div class="post-metadata">

**Author:** ![tomaklutfu](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomaklutfu/32/2411_2.png) [@tomaklutfu](https://discourse.julialang.org/u/tomaklutfu)\
**Post date:** [November 4, 2021, 4:11am UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869/3 "2021-11-04T04:11:40Z")

</div>

> [@hexaeder](#):
>
> `ones(1)`

> [@hexaeder](#):
>
> `prob = ODEProblem(diffeq, ones(1), (0.0, 1.0), (p, 10.0))`

`ones(1)` should be `ones(eltype(p), 1)`.

---

<div class="post-metadata">

**Author:** ![hexaeder](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/hexaeder/32/24403_2.png) [@hexaeder](https://discourse.julialang.org/u/hexaeder)\
**Post date:** [November 4, 2021, 9:33am UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869/4 "2021-11-04T09:33:13Z")

</div>

Thanks for the answers! I guess I’ll go with the wrapper for now.

> [@tomaklutfu](#):
>
> `ones(1)` should be `ones(eltype(p), 1)`

Doesn’t that imply that I’m differentiating with respect to the initial condition?

EDIT:  
Well I can’t achieve a allocation free wrapper in my real example. However your proposed solution

```julia
function loss(p)
    x0 = ones(1)
    prob = ODEProblem(diffeq, convert(Vector{eltype(p)}, x0), (0.0, 1.0), (p, 10.0))
    sol = solve(prob, Tsit5())
    return sol[end][1]
end

```

seems to work just fine, at least for the mini example. I’d still be interested in the implications of passing the initial state as a vector of dual numbers…

---

<div class="post-metadata">

**Author:** ![tomaklutfu](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomaklutfu/32/2411_2.png) [@tomaklutfu](https://discourse.julialang.org/u/tomaklutfu)\
**Post date:** [November 4, 2021, 10:59am UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869/5 "2021-11-04T10:59:10Z")

</div>

> [@hexaeder](#):
>
> Doesn’t that imply that I’m differentiating with respect to the initial condition?

No, it does not. `ones` gives zero differential to each element and `ForwardDiff.gradient` sets `1` to corresponding duals in the elements of `p`.

In general, you cannot have zero allocation in solving `ODEProblem`. You can minimize allocations by using `StaticArray` type and not storing every step in the solution giving `save_on=false` to `solve`.

---

<div class="post-metadata">

**Author:** ![hexaeder](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/hexaeder/32/24403_2.png) [@hexaeder](https://discourse.julialang.org/u/hexaeder)\
**Post date:** [November 4, 2021, 11:07am UTC](https://discourse.julialang.org/t/differentiate-through-solve-for-a-subset-of-parameters/70869/6 "2021-11-04T11:07:33Z")

</div>

yeah I know I am not aiming for zero allocations for everything. But my real world problem is more complex than that, I have some base parameters `p_base` and additional parameters `p_tune`

```julia
N = length(p_base)
# we reshape the array such that each vertex gets its own column
p_tune = reshape(p_additional, :, N)
# each vertex gets a tuple of parameters, P_base first and than the others
p = [(p_base[i], p_tune[:, i]...) for i in 1:N]

```

which creates a new array `p` containing of Tuples which are mixed from base and tune parameters.  
I don’t want to perform this reshuffling of my parameters at every single call of my diffeq function, therefore I’d prefer to not use a wrapper.

But converting `x0 = convert(Vector{eltype(p)}, orig_x0`) seems to tell the solver just fine to hand down `Dual` caches to my mutating function instead of `Vector{Float64}`.
