# Is it possible to do Nested AD ~elegantly~ in Julia? (PINNs)

**URL:** <https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888>\
**Category:** General Usage\
**Tags:** machine-learning\
**Created:** [May 15, 2023, 5:12pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888 "2023-05-15T17:12:01Z")\
**Posts on this page:** 20\
**Page:** 1

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 15, 2023, 5:12pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/1 "2023-05-15T17:12:01Z")

</div>

To differentiate a loss function defined in terms of a network’s derivatives seems to have been an issue since forever, specially when Zygote is involved; see [[1]](https://discourse.julialang.org/t/how-to-use-gradient-of-neural-network-as-the-loss-function/50569/11) [[2]](https://discourse.julialang.org/t/current-status-of-nested-ad/70753) [[3]](https://discourse.julialang.org/t/flux-pinn-1d-burgers/93262) and many more over at Zygote’s git.

In many of these threads (specially pre 2022) it is said that [the release of Diffractor.jl would correct this issue](https://discourse.julialang.org/t/gradient-calculation-in-pinn/61525/3). Now that it seems that [Diffractor is mostly dead](https://discourse.julialang.org/t/state-of-diffractor-jl/92959), what is left?

In my case, because I am working with exotic architectures, it is not possible to use @ChrisRackauckas’s NeuralPDE library. I have also been unable to reproduce the “ReverseDiff-Over-Zygote” hack that is commonly thrown around in these discussions.

Is there any other experimental AD library that is capable of this? Should I wait for it to be released? Or should I appeal to sketchy finite differences under the hood? That wouldn’t exactly thrill reviewers…

---

<div class="post-metadata">

**Author:** ![tim.holy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tim.holy/32/52_2.png) [@tim.holy](https://discourse.julialang.org/u/tim.holy)\
**Post date:** [May 16, 2023, 12:34pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/2 "2023-05-16T12:34:02Z")

</div>

I can’t answer the question, but note that Diffractor.jl has seen \>40 commits since Jan 1. That’s more than some packages that you would probably attest are very much alive.

 ![image](https://global.discourse-cdn.com/julialang/original/3X/5/2/52283a372982869475f01c432036bde47573bbe6.jpeg)

---

<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:** [May 16, 2023, 1:40pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/4 "2023-05-16T13:40:22Z")

</div>

Nested AD for this kind of thing is asymptotically much worse than numerical. If you’re going to use anything else, the thing to try is:

> **[GitHub - JuliaDiff/TaylorDiff.jl: Taylor-mode automatic differentiation for...](https://github.com/JuliaDiff/TaylorDiff.jl)**
>
> Taylor-mode automatic differentiation for higher-order derivatives - GitHub - JuliaDiff/TaylorDiff.jl: Taylor-mode automatic differentiation for higher-order derivatives

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [May 16, 2023, 2:35pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/5 "2023-05-16T14:35:14Z")

</div>

On the flipside, it relies heavily on nightly compiler features and has yet to conquer issues such as [Even very basic broadcasting breaks inferability · Issue #147 · JuliaDiff/Diffractor.jl · GitHub](https://github.com/JuliaDiff/Diffractor.jl/issues/147). To my knowledge only forward mode is coming any time soon, so even if all that is addressed Diffractor may not be fit for the purpose in this thread. Given the history of promotion around this project (if anyone is curious, search for “PINN” and “Diffractor” on Discourse) and the still completely unspecified timelines, I think it’d be very hard to argue the community is overcompensating on the expectation management front.

Digressing a bit, I wonder how much crossover this thread has with [Difficulties writing a program that computes PDEs involving Laplacians with AD](https://discourse.julialang.org/t/difficulties-writing-a-program-that-computes-pdes-involving-laplacians-with-ad/98834), which doesn’t appear to have any answers yet.

---

<div class="post-metadata">

**Author:** ![tim.holy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tim.holy/32/52_2.png) [@tim.holy](https://discourse.julialang.org/u/tim.holy)\
**Post date:** [May 16, 2023, 2:37pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/6 "2023-05-16T14:37:48Z")

</div>

Sure, I’m not saying it’s near-ready, just that pronouncing something “dead” is very different than the question of “is it ready yet?”

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [May 16, 2023, 2:54pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/7 "2023-05-16T14:54:24Z")

</div>

Oh, I’m not arguing either way about that. The point was to add more context since a lot of people are interested in Diffractor but don’t know where it currently sits readiness-wise. Seeing a bunch of commit activity in isolation doesn’t provide much signal (positive or negative) for that 🙂

---

<div class="post-metadata">

**Author:** ![Tomas\_Pevny](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomas_pevny/32/25466_2.png) [@Tomas\_Pevny](https://discourse.julialang.org/u/Tomas_Pevny)\
**Post date:** [May 16, 2023, 3:53pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/8 "2023-05-16T15:53:16Z")

</div>

I have been doing higher order diff with Zygote. It required redefining rules of some function, because they were not AD friendly (contained mutation), but it was doable. Time to first gradient was long, sometimes even half an hour. But it worked.

---

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 16, 2023, 4:39pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/9 "2023-05-16T16:39:03Z")

</div>

Thanks Chris. Is this the route you went for when making NeuralPDE? I have gone over the repo and the associated article a few times but just couldn’t find where and how the numerical derivatives were calculated. It does work brilliantly however, great work.

---

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 16, 2023, 4:43pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/10 "2023-05-16T16:43:19Z")

</div>

That sound quite lengthy. What is the size of the network? Have you tried Reverse-over-Zygote?

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [May 16, 2023, 5:31pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/11 "2023-05-16T17:31:02Z")

</div>

Enzyme does nested AD.

---

<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:** [May 16, 2023, 5:34pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/12 "2023-05-16T17:34:05Z")

</div>

The point though is that nested AD is not what you want to do here. Even if you can do nested AD, it’s still not really the solution so I don’t know why people keep mentioning it.

> [@Bizzi](#):
>
> Is this the route you went for when making NeuralPDE?

No, because it wasn’t ready when we built NeuralPDE, so we did a mixture of numerical and reverse mode to hit the asymptotically optimal form, and will be (over the summer) replacing the numerical parts with this TaylorDiff form to have an optimal all AD solution.

---

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 16, 2023, 7:37pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/13 "2023-05-16T19:37:37Z")

</div>

I see. Then what is the intuition behind using TaylorDiff?

Let’s take your PINNs tutorial as a MWE: Flux can peek inside the finite difference net\_xFD and train the network:

```julia
using Flux, Statistics #, TaylorDiff, ForwardDiff, Statistics, Plots

NN = Chain(Dense(1 => 12,tanh),
           Dense(12 => 12,tanh),
           Dense(12 => 12,tanh),
           Dense(12 => 1))
net(x) = x*first(NN([x]))
net(1) #Works

ϵ = Float32((eps(Float32))^(1/2)) #Naive Finite Difference
net_xFD(x) = (net(x+ϵ)-net(x))/(ϵ)
net_xFD(1) #Works

ts = 1f-2:1f-2:1f0 #Training set and loss function
loss() = mean(abs2(net_xFD(t)-cos(t)) for t in ts) 

#Training Loop
opt = Flux.Adam()
data = Iterators.repeated((), 500)
iter = 0
cb = function ()
  global iter += 1
  if iter % 50 == 0
    display(loss())
  end
end
Flux.train!(loss, Flux.params(NN), data, opt; cb=cb) #Works 

```

Now, obviously if I just replaced the naive finite difference net\_xFD for the TaylorDiff derivative it wouldn’t work. In fact, it seems like TaylorDiff can’t even take derivatives of Flux networks natively:

```julia
using TaylorDiff, ForwardDiff

net_xAD(x) = ForwardDiff.derivative(net,x) 
net_xAD(1) #Works

net_xTD(x) = TaylorDiff.derivative(net,x,1) 
net_xTD(1) #Doesn't work

```

So how would one go about training using TaylorDiff? Should I forego Flux completely? This is all very confusing.

---

<div class="post-metadata">

**Author:** ![Tomas\_Pevny](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomas_pevny/32/25466_2.png) [@Tomas\_Pevny](https://discourse.julialang.org/u/Tomas_Pevny)\
**Post date:** [May 16, 2023, 7:42pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/14 "2023-05-16T19:42:37Z")

</div>

Yes, I did reverse Zygote over Zygote. The real problem was only the operators. I had to have my own version of `logitcrossentropy` for example. But it was not bad, I think it was done in less than a day, though I had a prior experience with writing custom rules.

---

<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:** [May 16, 2023, 9:55pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/15 "2023-05-16T21:55:55Z")

</div>

> [@Bizzi](#):
>
> So how would one go about training using TaylorDiff? Should I forego Flux completely? This is all very confusing.

Yes use Lux as the tests show.

> <https://github.com/JuliaDiff/TaylorDiff.jl/blob/main/test/lux.jl>

There’s a PINN example with Flux though:

> <https://github.com/JuliaDiff/TaylorDiff.jl/blob/main/benchmark/pinn.jl>

---

<div class="post-metadata">

**Author:** ![patrick-kidger](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/patrick-kidger/32/20378_2.png) [@patrick-kidger](https://discourse.julialang.org/u/patrick-kidger)\
**Post date:** [May 17, 2023, 1:52am UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/16 "2023-05-17T01:52:46Z")

</div>

So the above Flux example does actually use nested AD. It has a `gradient(loss_by_taylordiff, ...)`, and `loss_by_taylordiff` itself calls `derivative`. Moreover I don’t think there’s anything wrong with that – nested AD is the appropriate thing to do here.

I believe (please correct me if need be) that Chris’ admonition against nested AD, and preference for Taylor-mode AD, is specifically when computing second derivatives directly, e.g. when you’re directly computing some d2y/dt2.

That’s not what @Bizzi’s example appears to require. The second derivative is “indirect”: `loss` contains a derivative, but must itself also be differentiated.

Assuming I’ve got all that correct – Bizzi, the appropriate (asymptotically correct) thing to do here is exactly what your example above is already doing. Use forward-mode autodiff (`ForwardDiff`) to compute `net_xAD` as the input `t` is a scalar, then use reverse-mode autodiff (`Flux`) to optimise the overall problem, as the output of `loss` is a scalar.

---

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 17, 2023, 4:04pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/17 "2023-05-17T16:04:35Z")

</div>

Thanks for taking the time to clarify these bits, Patrick.

I’m afraid I don’t understand your last paragraph, however. Whenever I use, say, ForwardDiff to calculate the derivative of the network, it looks like Flux can no longer “look inside” the loss function and see its dependence on the network’s parameters. As a result, the gradient returns 0 and the training loss never decreases. Using the example above:

```julia
#PINNs Example ForwardDiff
using Flux, Statistics, ForwardDiff

NN = Chain(Dense(1 => 12,tanh),
           Dense(12 => 12,tanh),
           Dense(12 => 12,tanh),
           Dense(12 => 1))
net(x) = x*first(NN([x]))
net(1) #Works

net_xAD(x) = ForwardDiff.derivative(net,x) 
net_xAD(1) #Works

ts = 1f-2:1f-2:1f0
loss() = mean(abs2(net_xAD(t)-cos(t)) for t in ts) 

opt = Flux.Adam()
data = Iterators.repeated((), 500)
iter = 0
cb = function ()
  global iter += 1
  if iter % 50 == 0
    display(loss())
  end
end
Flux.train!(loss, Flux.params(NN), data, opt; cb=cb) #Runs, but does not decrease the loss

```

The displayed losses are:

```julia
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0
1.3544431f0

```

Am I missing something? I feel like this could be related to the @functor macro, but I can’t really see how. Why should the gradient work with the finite differences net\_xFD but not with net\_xAD? My only explanation for this was that nested AD was not supported.

---

<div class="post-metadata">

**Author:** ![patrick-kidger](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/patrick-kidger/32/20378_2.png) [@patrick-kidger](https://discourse.julialang.org/u/patrick-kidger)\
**Post date:** [May 17, 2023, 4:27pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/18 "2023-05-17T16:27:05Z")

</div>

I guess I should replace my last paragraph with just “Use forward-mode autodiff to compute `net_xAD`, then use reverse-mode autodiff to optimise the overall problem.” I believe the statements I made about autodiff are correct, and that what you’re seeing here is a Julia bug: that these packages are silently failing to compose.

Honestly, I don’t actually use Julia in my work, as I’ve ran into _far_ too many issues exactly like what you’re seeing here. Try using JAX instead, which has done nested AD for years without difficulty. ([Shameless advert](http://github.com/patrick-kidger/equinox).)

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [May 17, 2023, 5:32pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/19 "2023-05-17T17:32:14Z")

</div>

Did you happen to see the warning from [Add warnings to ForwardDiff functions by mcabbott · Pull Request #1224 · FluxML/Zygote.jl · GitHub](https://github.com/FluxML/Zygote.jl/pull/1224)? If not, that’s probably a bug. Patrick’s suspicion is basically right though: Zygote (the AD Flux uses by default) over ForwardDiff will not work for your code as-is. For more on how this affects PINNs specifically, see [PINN loss doesn't converge to 0? · Issue #1966 · FluxML/Flux.jl · GitHub](https://github.com/FluxML/Flux.jl/issues/1966).

The examples Chris posted above avoid this by replacing one or both of the ADs involved with TaylorDiff. You could try other ADs as well (e.g. Enzyme as suggested above). I’ll defer to the people who actually work in this area to make suggestions though.

---

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 17, 2023, 5:32pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/20 "2023-05-17T17:32:34Z")

</div>

I see. Although I have considered dropping Julia multiple times at this point, my immense appreciation for the language’s concept drives me to try a little bit more. If I am unable to overcome this issue by the end of the week, however, I am probably going back to Python. If I do, I’ll certainly take a look at JAX.

---

<div class="post-metadata">

**Author:** ![Bizzi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bizzi/32/51484_2.png) [@Bizzi](https://discourse.julialang.org/u/Bizzi)\
**Post date:** [May 17, 2023, 5:46pm UTC](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888/21 "2023-05-17T17:46:02Z")

</div>

I see. So this is a limitation of Zygote + ForwardDiff specifically? I’ll try the implementation with TaylorDiff next. I’m aware that Lux is better behaved than Flux for a variety of applications, but it shouldn’t make a difference in this case, correct? (Given that Lux is also built on top of Zygote).

[Next page](https://discourse.julialang.org/t/is-it-possible-to-do-nested-ad-elegantly-in-julia-pinns/98888.md?page=2)
