# Issue with Zygote over ForwardDiff.derivative

**URL:** https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824
**Category:** Machine Learning
**Created:** [November 2, 2021, 7:03pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824 "2021-11-02T19:03:08Z")
**Posts on this page:** 11
**Page:** 1

<div class="post-metadata">

### Author: ![jlmaccal](https://avatars.discourse-cdn.com/v4/letter/j/839c29/32.png) [@jlmaccal](https://discourse.julialang.org/u/jlmaccal)
#### Post date: [November 2, 2021, 7:03pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/1 "2021-11-02T19:03:09Z")

</div>

I’m having some trouble getting Zygote over `ForwardDiff.derivative` to work.

I’m going to refer to this [prior post](https://discourse.julialang.org/t/is-it-possible-perform-reverse-mode-differentiation-flux-jl-with-zygote-jl-of-a-forward-mode-differentiation-result-e-g-forwarddiff/30945). The following code used to fail with `ERROR: setindex! not defined for ForwardDiff.Partials{1,Float64}`. The fix was to define some additional adjoints, as @ChrisRackauckas showed in DiffEqFlux.jl.

```julia
using Flux, ForwardDiff

f = Chain(x -> fill(x, 3), Dense(3, 3, softplus))
df(x) = ForwardDiff.derivative(f, x)

x = rand()
f(x) #Works
df(x) #Works
gs = gradient(() -> sum(df(x)), params(f)) #Fails

```

However, the code above now runs without the additional adjoints, but the gradients returned are `nothing`.

One of my codes used a similar Zygote over `ForwardDiff.derivative` idea, but it no longer trains as all of the gradients are `nothing`. Something seems to have changed, but I don’t know where to start. Unfortunately, I don’t have the old Project or Manifest files, so I don’t know what versions I was using.

---

<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: [November 2, 2021, 7:04pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/2 "2021-11-02T19:04:22Z")

</div>

what’s your full MWE?

---

<div class="post-metadata">

### Author: ![jlmaccal](https://avatars.discourse-cdn.com/v4/letter/j/839c29/32.png) [@jlmaccal](https://discourse.julialang.org/u/jlmaccal)
#### Post date: [November 2, 2021, 7:27pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/3 "2021-11-02T19:27:32Z")

</div>

The example above fails in the way I describe. `gs` should have the gradients wrt to `params(f)`, but has `nothing` instead.

---

<div class="post-metadata">

### Author: ![jlmaccal](https://avatars.discourse-cdn.com/v4/letter/j/839c29/32.png) [@jlmaccal](https://discourse.julialang.org/u/jlmaccal)
#### Post date: [November 2, 2021, 7:30pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/4 "2021-11-02T19:30:36Z")

</div>

My example is more along the lines of this, computing the dot product between the gradient and a vector, which is equivalent to the directional derivative.

```julia
using Flux
using ForwardDiff

net = Chain(Dense(2, 128, relu), Dense(128, 128, relu), Dense(128, 1))
p, re = Flux.destructure(net)

x = randn(Float32, 2, 128)
dx = randn(Float32, 2, 128)

grads = Flux.gradient(p -> sum(ForwardDiff.derivative(h -> re(p)(x + h*dx), 0.0f0)), p)

```

This used to fail without defining some extra adjoints as in DiffEqFlux. It now just gives `nothing`

---

<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: [November 2, 2021, 7:48pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/5 "2021-11-02T19:48:09Z")

</div>

If you add the DiffEqFlux adjoints does it work?

---

<div class="post-metadata">

### Author: ![jlmaccal](https://avatars.discourse-cdn.com/v4/letter/j/839c29/32.png) [@jlmaccal](https://discourse.julialang.org/u/jlmaccal)
#### Post date: [November 2, 2021, 7:53pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/6 "2021-11-02T19:53:34Z")

</div>

No, that doesn’t make a difference.

---

<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: [November 2, 2021, 7:59pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/7 "2021-11-02T19:59:46Z")

</div>

Interesting. @mcabbott would you know something about what might’ve changed?

---

<div class="post-metadata">

### Author: ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)
#### Post date: [November 2, 2021, 8:10pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/8 "2021-11-02T20:10:16Z")

</div>

Yes this won’t work, sadly. The warning from Zygote.forwarddiff is:

```julia
Note that the function `f` will *drop gradients* for any closed-over values.

```

and that’s what’s being used [here](https://github.com/FluxML/Zygote.jl/blob/master/src/lib/forward.jl#L145). That is, it’s forward-over-forward, and takes derivatives only with respect to the explicit parameter, not to anything closed over (since ForwardDiff is unaware of those).

Making it give errors when `f` closes over anything would be better. Making it actually work… I’m not sure, might be possible? Does DiffEqFlux.jl have (pirate?) code which handles this?

---

<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: [November 2, 2021, 8:15pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/9 "2021-11-02T20:15:04Z")

</div>

> [@mcabbott](#):
>
> Does DiffEqFlux.jl have (pirate?) code which handles this?

Yes.

```julia
ZygoteRules.@adjoint function ForwardDiff.Dual{T}(x, ẋ::Tuple) where T
  @assert length(ẋ) == 1
  ForwardDiff.Dual{T}(x, ẋ), ḋ -> (ḋ.partials[1], (ḋ.value,))
end

ZygoteRules.@adjoint ZygoteRules.literal_getproperty(d::ForwardDiff.Dual{T}, ::Val{:partials}) where T =
  d.partials, ṗ -> (ForwardDiff.Dual{T}(ṗ[1], 0),)

ZygoteRules.@adjoint ZygoteRules.literal_getproperty(d::ForwardDiff.Dual{T}, ::Val{:value}) where T =
  d.value, ẋ -> (ForwardDiff.Dual{T}(0, ẋ),)

```

All of our pirate code is: [https://github.com/SciML/DiffEqFlux.jl/blob/v1.44.0/src/DiffEqFlux.jl#L60-L74](https://github.com/SciML/DiffEqFlux.jl/blob/v1.44.0/src/DiffEqFlux.jl#L60-L74) and we should upstream some of it.

---

<div class="post-metadata">

### Author: ![jlmaccal](https://avatars.discourse-cdn.com/v4/letter/j/839c29/32.png) [@jlmaccal](https://discourse.julialang.org/u/jlmaccal)
#### Post date: [November 2, 2021, 8:51pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/10 "2021-11-02T20:51:23Z")

</div>

Is there any possibility of a workaround? This used to work ~1 year ago.

The only other approach I have been able to make work is ReverseDiff over Zygote, but for some reason this is super slow (I’ll create another thread about this).

---

<div class="post-metadata">

### Author: ![facusapienza](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/facusapienza/32/24317_2.png) [@facusapienza](https://discourse.julialang.org/u/facusapienza)
#### Post date: [January 21, 2024, 4:18pm UTC](https://discourse.julialang.org/t/issue-with-zygote-over-forwarddiff-derivative/70824/11 "2024-01-21T16:18:25Z")

</div>

Has it been any update or progress in this line? I recently encountered a similar problem that I posted in [Nested and different AD methods altogether: How to add AD calculations inside my loss function when using neural differential equations?](https://discourse.julialang.org/t/nested-and-different-ad-methods-altogether-how-to-add-ad-calculations-inside-my-loss-function-when-using-neural-differential-equations/108985) that I am trying to make work. Thanks!
