# Zygote dozens\* of times slower than manually written function

**URL:** <https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748>\
**Category:** Performance\
**Tags:** zygote, forwarddiff\
**Created:** [April 20, 2022, 4:07pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748 "2022-04-20T16:07:25Z")\
**Posts on this page:** 18\
**Page:** 1

<div class="post-metadata">

**Author:** ![Gabrielmlando](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gabrielmlando/32/20628_2.png) [@Gabrielmlando](https://discourse.julialang.org/u/Gabrielmlando)\
**Post date:** [April 20, 2022, 4:07pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/1 "2022-04-20T16:07:25Z")

</div>

I am trying to use `Zygote` in order to avoid the need of manually passing gradients/Hessians. I started by the simplest possible example for my applications, namely a hydrogen atom. The Hamiltonian and initial point are defined as

```julia
using StaticArrays, ForwardDiff, Zygote

H(X) = X[1]^2/2.0 + X[2]^2/2.0 + X[3]^2/2.0 - ( X[4]^2 + X[5]^2 + X[6]^2 )^-0.5
x0 = @SVector rand(6);

```

I then get

```julia
∇1H(X) = Zygote.gradient(H, X)
@btime ∇1H($x0);
>> 3.451 μs (57 allocations: 2.28 KiB)

```

Which is insanely slow and has unjustifiable allocations. Even if I use `ForwardDiff` directly, the result is far superior:

```julia
∇2H(X) = ForwardDiff.gradient(H, X)
@btime ∇2H($x0);
>> 171.613 ns (0 allocations: 0 bytes)

```

Still, it is twice what I get from a manual evaluation of the gradient:

```julia
∇H(X) = @SVector [
    X[1] ,
    X[2] ,
    X[3] ,
    X[4] * ( X[4]^2 + X[5]^2 + X[6]^2 )^-1.5 , 
    X[5] * ( X[4]^2 + X[5]^2 + X[6]^2 )^-1.5 ,
    X[6] * ( X[4]^2 + X[5]^2 + X[6]^2 )^-1.5 ,
]

@btime ∇H($x0);
>> 71.700 ns (0 allocations: 0 bytes)

```

Am I doing something wrong when using `Zygote`? Should I just keep passing gradients manually?

---

<div class="post-metadata">

**Author:** ![ufechner7](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ufechner7/32/51363_2.png) [@ufechner7](https://discourse.julialang.org/u/ufechner7)\
**Post date:** [April 20, 2022, 5:10pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/2 "2022-04-20T17:10:15Z")

</div>

Well, 3 µs is not that slow… If I would have to do it manually I would probably need 3h…

I guess handwritten derivatives will always be faster to execute than machine generated derivatives, if, and this is the big if it is feasible to write them manually.

Automated differentation uses dual numbers which always causes some overhead.

---

<div class="post-metadata">

**Author:** ![marius311](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marius311/32/3953_2.png) [@marius311](https://discourse.julialang.org/u/marius311)\
**Post date:** [April 20, 2022, 5:36pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/3 "2022-04-20T17:36:08Z")

</div>

I don’t think you’re doing anything wrong, Zygote is known to not be super fast for code which does lots of scalar indexing like yours. All the compiler magic going on often makes inference fail on the gradient calculation, so the alocations you’re seeing are probably because the gradient code ends up being type unstable. For gradients w.r.t. a small number of parameters like this I would just recommend using ForwardDiff. You can also use Zygote.forwarddiff to temporarily switch to ForwardDiff,

```julia
H(X) = Zygote.forwarddiff(X) do X
    X[1]^2/2.0 + X[2]^2/2.0 + X[3]^2/2.0 - ( X[4]^2 + X[5]^2 + X[6]^2 )^-0.5
end

```

now you can embed `H(X)` in some bigger calculation if you wanted to still use Zygote on that, and this piece of it will use ForwardDiff / be fast.

---

<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, 6:54pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/4 "2022-04-20T18:54:15Z")

</div>

> [@Gabrielmlando](#):
>
> Am I doing something wrong when using `Zygote` ? Should I just keep passing gradients manually?

Why did you choose `Zygote` in the first place? If I would be working with `Flux`, I’d be working with `Zygote`, obviously. Otherwise I’d check [AD](https://github.com/JuliaDiff), [Finite Differences](https://github.com/JuliaDiff/FiniteDifferences.jl), [Symbolics.jl](https://github.com/JuliaSymbolics/Symbolics.jl) or manually derive and decide what is appropriate for my use case (I think I’ve seen technical limits for most of the mentioned system in one or the other way, especially w.r.t. complex numbers for example).

Edit: and (obviously?) you can use all these to validate correctness.

---

<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:** [April 20, 2022, 7:56pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/5 "2022-04-20T19:56:11Z")

</div>

Just to hammer this point home, Zygote will allocate at least 1 full-sized array for **every** indexing operation in that function.

---

<div class="post-metadata">

**Author:** ![cjdoris](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cjdoris/32/213133_2.png) [@cjdoris](https://discourse.julialang.org/u/cjdoris)\
**Post date:** [April 20, 2022, 9:20pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/6 "2022-04-20T21:20:36Z")

</div>

Contrary to the title, in the original post Zygote is only 50x slower, not hundreds.

---

<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:43pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/7 "2022-04-20T21:43:38Z")

</div>

We could defuse the title of the posting, should we?

---

<div class="post-metadata">

**Author:** ![DNF](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dnf/32/10191_2.png) [@DNF](https://discourse.julialang.org/u/DNF)\
**Post date:** [April 20, 2022, 10:11pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/8 "2022-04-20T22:11:31Z")

</div>

I don’t know about Zygote, but you should definitely clean up your function, first of all. `X^-0.5` is a really bad operation to do.

Here’s a comparison:

```julia
H(X) = X[1]^2/2.0 + X[2]^2/2.0 + X[3]^2/2.0 - ( X[4]^2 + X[5]^2 + X[6]^2 )^-0.5
G(X) = (X[1]^2 + X[2]^2 + X[3]^2)/2 - 1/sqrt(X[4]^2 + X[5]^2 + X[6]^2)
∇2H(X) = ForwardDiff.gradient(H, X)
∇2G(X) = ForwardDiff.gradient(G, X)

```

Benchmarks:

```julia
1.7.2> x0 = @SVector rand(6);

1.7.2> @btime H($x0)
  30.071 ns (0 allocations: 0 bytes)
-0.2515642893020753

1.7.2> @btime G($x0)
  3.100 ns (0 allocations: 0 bytes)
-0.2515642893020753

1.7.2> @btime ∇2H($x0);
  91.802 ns (0 allocations: 0 bytes)

1.7.2> @btime ∇2G($x0);
  22.211 ns (0 allocations: 0 bytes)

```

This probably won’t help with Zygote, but it’s good to improve the basics first.

---

<div class="post-metadata">

**Author:** ![DNF](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dnf/32/10191_2.png) [@DNF](https://discourse.julialang.org/u/DNF)\
**Post date:** [April 20, 2022, 10:18pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/9 "2022-04-20T22:18:17Z")

</div>

> [@Gabrielmlando](#):
>
> ```julia
> ∇H(X) = @SVector [
> X[1] ,
> X[2] ,
> X[3] ,
> X[4] * ( X[4]^2 + X[5]^2 + X[6]^2 )^-1.5 , 
> X[5] * ( X[4]^2 + X[5]^2 + X[6]^2 )^-1.5 ,
> X[6] * ( X[4]^2 + X[5]^2 + X[6]^2 )^-1.5 ,
> ]
> 
> ```

And then there’s the manual gradient:

```julia
function ∇G(X)
    x = 1/(X[4]^2 + X[5]^2 + X[6]^2)
    y = x * sqrt(x)
    @SVector [
    X[1] ,
    X[2] ,
    X[3] ,
    X[4] * y, 
    X[5] * y,
    X[6] * y]
end

```

Benchmark:

```julia
1.7.2> @btime ∇H($x0);
  34.884 ns (0 allocations: 0 bytes)
        
1.7.2> @btime ∇G($x0);
  3.300 ns (0 allocations: 0 bytes)

```

---

<div class="post-metadata">

**Author:** ![Elrod](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/elrod/32/22461_2.png) [@Elrod](https://discourse.julialang.org/u/Elrod)\
**Post date:** [April 21, 2022, 5:05am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/10 "2022-04-21T05:05:20Z")

</div>

The `1.7.2>` prompts are really annoying, as they prevent me from copy/pasting code.  
The default `julia>` prompts would automatically be stripped upon pasting into a REPL.

Also, adding `@inline` to `G` helps the gradient time.

Also, starting Julia with `--math-mode=fast` also helps a lot.  
Starting Julia this way is not recommended, except in helping to identify optimization opportunities, e.g. that it’d be worth adding `@fastmath` support to `ForwardDiff`.  
Currently, `@fastmath` with `ForwardDiff` does not actually set any fast flags. All it does is prevent code from inlining, pessimizing it.

---

<div class="post-metadata">

**Author:** ![DNF](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dnf/32/10191_2.png) [@DNF](https://discourse.julialang.org/u/DNF)\
**Post date:** [April 21, 2022, 6:47am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/11 "2022-04-21T06:47:50Z")

</div>

> [@Elrod](#):
>
> The `1.7.2>` prompts are really annoying, as they prevent me from copy/pasting code.  
> The default `julia>` prompts would automatically be stripped upon pasting into a REPL.

Not sure what to do about that. I’m not planning to use `julia>`, it’s way too long without carrying useful information. I use either `version>` or `jl>`.

Prompt pasting never worked for me anyway (on any platform), so I didn’t really consider it.

Isn’t it easy enough to not copy the prompt? You wouldn’t want to copy the output anyway.

---

<div class="post-metadata">

**Author:** ![Gabrielmlando](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gabrielmlando/32/20628_2.png) [@Gabrielmlando](https://discourse.julialang.org/u/Gabrielmlando)\
**Post date:** [April 21, 2022, 8:42am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/12 "2022-04-21T08:42:30Z")

</div>

As far I can see from the answers, I can summarize what I learned as:

- My function can be made immediately more efficient (thanks @DNF)
- I shouldn’t be using Zygote (thanks @marius311, @ToucheSir, @goerch)

I’ll keep passing gradients, apparently. I still wonder, however, how people like @ChrisRackauckas were able to use ForwardDiff in their super efficient libraries. In DifferentialEquations, I know that for symplectic integrators one can provide the solver with the Hamiltonian only, and gradients are calculated by automatic differentiation. This is a mystery to me, since even for the gradient function provided by DNF, ForwardDiff is still 8x slower than writing the gradient manually. Maybe people behind DiferentialEquations chose to sacrifice efficiency in order to spare the user from writing the gradients, but the efficiency loss when one is propagating an ensemble of particles seems to be relevant (although, of course, the bottleneck of ODE solving is definitely happening somewhere else).

---

<div class="post-metadata">

**Author:** ![skleinbo](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/skleinbo/32/36080_2.png) [@skleinbo](https://discourse.julialang.org/u/skleinbo)\
**Post date:** [April 21, 2022, 9:29am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/13 "2022-04-21T09:29:25Z")

</div>

```julia
julia> G(X) = (X[1]^2 + X[2]^2 + X[3]^2)/2 - 1/sqrt(X[4]^2 + X[5]^2 + X[6]^2)
G (generic function with 1 method)

julia> @btime Zygote.gradient(G, $x0)
  3.874 μs (57 allocations: 2.27 KiB)
([0.8312830657406003, 0.42241553817070676, 0.029793466603427188, 0.685730593635655, 0.23572364881902477, 0.540322706597109],)

julia> @views K(X) = sum(abs2, X[1:3])/2 - 1/norm(X[4:6])
K (generic function with 1 method)

julia> @btime Zygote.gradient(K, $x0)
  475.292 ns (10 allocations: 624 bytes)
([0.8312830657406003, 0.42241553817070676, 0.029793466603427188, 0.685730593635655, 0.23572364881902477, 0.540322706597109],)

```

```julia

```

---

<div class="post-metadata">

**Author:** ![maxfreu](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/maxfreu/32/17468_2.png) [@maxfreu](https://discourse.julialang.org/u/maxfreu)\
**Post date:** [April 21, 2022, 10:20am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/14 "2022-04-21T10:20:38Z")

</div>

I think for many physical functions it’s best to derive them by hand, if possible. Then optimize them for the computer (like DNF did, avoid divisions, repeated calculations and non-integer exponents) and write a `ChainRules` rrule. That way Zygote will be able to work with it _and_ it will be relatively fast.

I’m not sure the following is 100% correct. Improves from 4.3us to 735ns (gradient(K, $x) gives 920ns).

```julia
# define before first call to gradient
function ChainRulesCore.rrule(::typeof(G), x) 
    pullback(Δy) = (NoTangent(), ∇G(x) * Δy)
    return G(x), pullback
end

```

Maybe someone more familiar with Zygote + ChainRules can optimize this even more.

Oh and while marius311’s answer above involving forwarddiff is super fast, it changes the definition of H. If that is not an option and you still want to use forwarddiff explicitly, you can use the freshly announced [`ForwardDiffPullbacks`](https://discourse.julialang.org/t/ann-forwarddiffpullbacks-jl-forwarddiff-based-chainrulescore-pullbacks/78737/1) package:

```julia
using ForwardDiffPullbacks
gradient(fwddiff(G), x) # 133ns vs 730ns on my machine, 0 allocs

```

---

<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 21, 2022, 10:49am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/15 "2022-04-21T10:49:00Z")

</div>

> [@Gabrielmlando](#):
>
> I’ll keep passing gradients, apparently. I still wonder, however, how people like @ChrisRackauckas were able to use ForwardDiff in their super efficient libraries. In DifferentialEquations, I know that for symplectic integrators one can provide the solver with the Hamiltonian only, and gradients are calculated by automatic differentiation.

We allow for direct derivation by ModelingToolkit.jl which is symbolic.

[https://diffeq.sciml.ai/stable/tutorials/faster\_ode\_example/#Automatic-Derivation-of-Jacobian-Functions](https://diffeq.sciml.ai/stable/tutorials/faster_ode_example/#Automatic-Derivation-of-Jacobian-Functions)

[https://mtk.sciml.ai/dev/mtkitize\_tutorials/modelingtoolkitize/](https://mtk.sciml.ai/dev/mtkitize_tutorials/modelingtoolkitize/)

You cannot automatically apply it, but it will outperform AD in a lot of cases when you can. So we document the ModelingToolkit/Symbolics tools necessary to achieve the top-notch performance, and leave it to the user to choose the right path.

---

<div class="post-metadata">

**Author:** ![Elrod](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/elrod/32/22461_2.png) [@Elrod](https://discourse.julialang.org/u/Elrod)\
**Post date:** [April 21, 2022, 11:40am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/16 "2022-04-21T11:40:18Z")

</div>

> [@DNF](#):
>
> Isn’t it easy enough to not copy the prompt? You wouldn’t want to copy the output anyway.

Julia strips the prompt and output so that you can copy/ past multiple lines.  
So, no, it’s not easy, as the point is to copy multiple lines. E.g., if these were `julua>`s, I could use a single copy/paste for all the lines:

```julia
1.7.2> x0 = @SVector rand(6);

1.7.2> @btime H($x0)
  30.071 ns (0 allocations: 0 bytes)
-0.2515642893020753

1.7.2> @btime G($x0)
  3.100 ns (0 allocations: 0 bytes)
-0.2515642893020753

1.7.2> @btime ∇2H($x0);
  91.802 ns (0 allocations: 0 bytes)

1.7.2> @btime ∇2G($x0);
  22.211 ns (0 allocations: 0 bytes)

```

This works on Linux and Mac.

---

<div class="post-metadata">

**Author:** ![DNF](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dnf/32/10191_2.png) [@DNF](https://discourse.julialang.org/u/DNF)\
**Post date:** [April 21, 2022, 11:43am UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/17 "2022-04-21T11:43:03Z")

</div>

That never worked for me, so I guess I didn’t consider it.

I guess I’ll have to manually edit the prompts each time then 🫤

---

<div class="post-metadata">

**Author:** ![oschulz](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oschulz/32/2998_2.png) [@oschulz](https://discourse.julialang.org/u/oschulz)\
**Post date:** [April 21, 2022, 9:25pm UTC](https://discourse.julialang.org/t/zygote-dozens-of-times-slower-than-manually-written-function/79748/18 "2022-04-21T21:25:48Z")

</div>

> [@maxfreu](#):
>
> you can use the freshly announced [`ForwardDiffPullbacks`](https://discourse.julialang.org/t/ann-forwarddiffpullbacks-jl-forwarddiff-based-chainrulescore-pullbacks/78737/1) package: […] 133ns vs 730ns on my machine

There’s still some room for improvement there, since ForwardDiffPullbacks currently calculates the primal result and the pullback for every argument separately. I’ve been thinking about adding a mode that does everything in one go, for use cases that will evaluate all Thunks anyway.

But as @ChrisRackauckas noted, analytical derivatives via ModelingToolkit/Symbolics (if applicable for your problem) will typically outperform any AD.
