# Zygote differentiation very slow

**URL:** <https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751>\
**Category:** Performance\
**Tags:** zygote\
**Created:** [April 10, 2024, 3:55am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751 "2024-04-10T03:55:54Z")\
**Posts on this page:** 14\
**Page:** 1

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [April 10, 2024, 3:55am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/1 "2024-04-10T03:55:54Z")

</div>

I have an inner function I apply like this (a simple replicable example pasted here) . The input are chosen to be (5,5) but they are usually much larger. I did not do benchmark as this is not even close to acceptable. It usually crashes for larger realistic inputs.  
One issue I think is causing this is having to differentiate with respect to all the inputs when I only need with respect to first input. I did not find a way to do so other than computing with respect to all and indexing it afterwards.

EDIT: Code fixed to reflect the question

```julia
using Zygote
function inner_func(p1, p2, p3)
    return sum(p1) + sum(p2) + sum(p3)
end

# Define the function with the sample inner function
function coefficients(x,y,z)
    num_vectors = size(x,1)
    K = [inner_func(x[I,:],
                        y[J,:],
                        z[J,:]) for J in 1:num_vectors, I in 1:num_vectors]
    return K
end

x = rand(5,5)
y = rand(5,5)
z = rand(5,5)

inp_vector = [x,y,z]
# Compute the Jacobian and value at once, using Zygote
Zygote.jacobian(coefficients, inp_vector...)[1]

```

I want to wrap the function with loop in it and automatically differentiate through it rather than differentiating the inner function and building the matrices with outer loop. What is recommended when working with Zygote? Building matrices by myself every time I want to get higher order derivatives is inconvenient as shown below.

```julia
dK_dx = [Zygote.jacobian(inner_func, [x[I,:],y[J,:],z[J,:]]...)[1] for J in 1:size(x,1), I in 1:size(x,1) ]

```

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 10, 2024, 5:08am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/2 "2024-04-10T05:08:08Z")

</div>

> [@KapilKhanal](#):
>
> I did not do benchmark as this is not even close to acceptable. It usually crashes for larger realistic inputs.

I suspect the issue is that Zygote doesn’t like scalar indexing like `inner_func(x[I], y[J], z[J], x[J])`, it is optimized to work well on vectorized code.

> [@KapilKhanal](#):
>
> Building matrices by myself every time I want to get higher order derivatives is inconvenient.

What do you mean by higher-order derivatives?

> [@KapilKhanal](#):
>
> One issue I think is causing this is having to differentiate with respect to all the inputs when I only need with respect to first input.

In this case you could try Enzyme.jl, which has an interface to specify the differentiable vs constant inputs. if I understand correctly, you need the Jacobian of K wrt x only?

---

<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:** [April 10, 2024, 6:25am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/3 "2024-04-10T06:25:50Z")

</div>

Might be worth noting that `inner_func` is called with four numbers, but it looks like you expect arrays. Here `x[1] isa Float64`, although only the first 50 of 2500 elements are ever used:

> [@KapilKhanal](#):
>
> ```julia
> num_vectors = size(x,1)
> K = [inner_func(x[I],
> ...
> for I in 1:num_vectors]
> 
> ```

Such scalar indexing in a loop is indeed Zygote’s worst nightmare, as Guillaume says. But you might intend to be indexing `eachcol(x)` instead, which (aside from giving completely different answers) won’t be as bad for Zygote:

```julia
julia> begin
       x = rand(3,3) # smaller example
       y = rand(3,3)
       z = rand(3,3)

       inp_vector = [x,y,z]
       end;

julia> Zygote.jacobian(coefficients, inp_vector...)[1] # only first 3 entries of x are used
9×9 Matrix{Float64}:
 2.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0
 1.0 1.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0
 1.0 0.0 1.0 0.0 0.0 0.0 0.0 0.0 0.0
 1.0 1.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0
 0.0 2.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0
 0.0 1.0 1.0 0.0 0.0 0.0 0.0 0.0 0.0
 1.0 0.0 1.0 0.0 0.0 0.0 0.0 0.0 0.0
 0.0 1.0 1.0 0.0 0.0 0.0 0.0 0.0 0.0
 0.0 0.0 2.0 0.0 0.0 0.0 0.0 0.0 0.0

julia> function coefficients_2(x,y,z)
           xcols = eachcol(x)
           ycols = eachcol(y)
           zcols = eachcol(z)
           inds = eachindex(xcols)
           K = [inner_func(xcol, # this calls inner_func with 4 vectors
                           ycols[j],
                           zcols[j],
                           xcols[j]) for j in inds, xcol in xcols]
       end
coefficients_2 (generic function with 1 method)

julia> jacobian(x -> coefficients_2(x, y, z), x)[1] # jacobian w.r.t. x alone
9×9 Matrix{Float64}:
 2.0 2.0 2.0 0.0 0.0 0.0 0.0 0.0 0.0
 1.0 1.0 1.0 1.0 1.0 1.0 0.0 0.0 0.0
 1.0 1.0 1.0 0.0 0.0 0.0 1.0 1.0 1.0
 1.0 1.0 1.0 1.0 1.0 1.0 0.0 0.0 0.0
 0.0 0.0 0.0 2.0 2.0 2.0 0.0 0.0 0.0
 0.0 0.0 0.0 1.0 1.0 1.0 1.0 1.0 1.0
 1.0 1.0 1.0 0.0 0.0 0.0 1.0 1.0 1.0
 0.0 0.0 0.0 1.0 1.0 1.0 1.0 1.0 1.0
 0.0 0.0 0.0 0.0 0.0 0.0 2.0 2.0 2.0

```

---

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [April 10, 2024, 2:22pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/4 "2024-04-10T14:22:38Z")

</div>

Thank you for prompt help!

right ,the indexing could be one of the issues. By higher order, I mean second and third order derivative of `K with respect to x`, the first input. I actually only need first element of `K with respect to the first element of x` vector. Right now I am probably doing first element of `K with respect to all of the elements of x` vector which should mostly be zero. I am not sure how to do so in Zygote in the setup like this.

is there a utility function in zygote for derivative of a element matrix `K\_{ij}’ with respect to elements of another matrix ‘x\_{ij}’ ?

I had other issues with Enzyme and wrote most of the code to use Zygote. I might have to revert back to Enzyme then.

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 10, 2024, 3:44pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/5 "2024-04-10T15:44:38Z")

</div>

> [@mcabbott](#):
>
> Might be worth noting that `inner_func` is called with four numbers, but it looks like you expect arrays. Here `x[1] isa Float64`, although only the first 50 of 2500 elements are ever used:

@KapilKhanal could you maybe fix the example so that it does exactly what you intend it to do? Some aspects of it are weird at the moment. Perhaps you expect `x[i]` to be the `i`-th row of the matrix? But it is not, it is actually the `i`-th coefficient of the flattened matrix.

---

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [April 10, 2024, 4:13pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/6 "2024-04-10T16:13:57Z")

</div>

Oh Gotcha, that was not indexing correctly like I thought but I think the issue is still same in this case also.

I have fixed the code. I will have to think how to avoid scalar indexing within Zygote so that change is not reflected in the code yet.

---

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [April 10, 2024, 4:44pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/7 "2024-04-10T16:44:26Z")

</div>

thank you for pointing out the error in indexing. I thought it was giving me the first row. Is there a way to vectorize this properly? It’s very weird that julia itself does not care about vectorization but zygote does

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 10, 2024, 7:15pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/8 "2024-04-10T19:15:24Z")

</div>

> [@KapilKhanal](#):
>
> I actually only need first element of `K with respect to the first element of x` vector.

Why do you compute all of `K` and use all of `x` then? You are actually looking for a scalar-to-scalar derivative?

---

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [April 10, 2024, 7:47pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/9 "2024-04-10T19:47:54Z")

</div>

Yes element by element derivative. I have added the code that does that in the code above but I want to just take the derivative of the function with loop in it instead of differentiating inner function and building matrix. I have another code that calculates the 2nd and 3rd order derivative and I want it to be able to just differentiate through the matrix building as well.

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 11, 2024, 11:59am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/10 "2024-04-11T11:59:00Z")

</div>

> [@gdalle](#):
>
> Why do you compute all of `K` and use all of `x` then?

I’m sorry I still don’t understand that part

---

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [April 13, 2024, 4:42am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/11 "2024-04-13T04:42:27Z")

</div>

I need all of K element’s gradients with respect to corresponding index element in x. I wanted to write it such that its easier to understand when user just specifies dK/dx or d^2K/d^2x and the function would do so instead of just providing the derivative of the inner function and asking them to build matrices themselves. Also, that way code would be end to end differentiable by itself. There will be no need to manually get the matrices and get it working for example, if I am going to use this in an adjoint equation then manual would make sense but I am hoping to do reverse diff automatically without writing out the equations and getting the required partial via AD and having to do the adjoint derivation.

I am confusing you maybe because there’s something that I am not understanding here. Hopefully that makes sense. let me know if that’s not the right way of thinking about this

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 13, 2024, 6:35am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/12 "2024-04-13T06:35:47Z")

</div>

Let’s ignore the y and z which are constant for differentiation purposes. If I understand correctly, you have a function

x \in \mathbb{R}^{n \times m} \longmapsto K \in \mathbb{R}^{n \times m}

and you are only interested in “diagonal” partial derivatives like

\frac{\partial K\_{ij}}{\partial x\_{ij}} \quad \text{and} \quad \frac{\partial^2 K\_{ij}}{\partial^2 x\_{ij}} \quad \text{and} \quad \frac{\partial^3 K\_{ij}}{\partial^3 x\_{ij}}

but not in “non-diagonal” partial derivatives like

\frac{\partial K\_{ij}}{\partial x\_{kl}}

for (i,j) \neq (k,l). Does that sound right?

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 13, 2024, 8:07am UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/13 "2024-04-13T08:07:25Z")

</div>

If the above is correct, then this code computes one of the quantities you are interested in:

```julia
julia> using ForwardDiff, FillArrays

julia> K(x) = x * x
K (generic function with 1 method)

julia> x = rand(3, 3);

julia> dx(x, i, j) = OneElement(one(eltype(x)), (i, j), axes(x))
dx (generic function with 1 method)

julia> function diagonal_derivative(K, x, i, j)
           step(t) = K(x + t * dx(x, i, j))
           full_derivative = ForwardDiff.derivative(step, 0)
           return full_derivative[i, j]
       end
diagonal_derivative (generic function with 1 method)

julia> diagonal_derivative(K, x, 1, 2)
0.9432869747177048

```

You don’t gain anything by computing those derivatives in reverse mode, in this case you will need n^2 function calls either way, and forward mode is usually more efficient.

If you need even more efficiency, you should define your function K in a non-allocating way, like `K!(x_dest, x)`

---

<div class="post-metadata">

**Author:** ![KapilKhanal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/kapilkhanal/32/43976_2.png) [@KapilKhanal](https://discourse.julialang.org/u/KapilKhanal)\
**Post date:** [May 12, 2024, 11:39pm UTC](https://discourse.julialang.org/t/zygote-differentiation-very-slow/112751/14 "2024-05-12T23:39:39Z")

</div>

Gotcha. Thank you for your explanation. Just got to work on this project and it helped a lot!
