# Lux, ComponentArrays and flat parameters : computing the gradient works with Zygote but not with Enzyme

**URL:** <https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533>\
**Category:** New to Julia\
**Tags:** enzyme\
**Created:** [August 6, 2023, 1:29pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533 "2023-08-06T13:29:24Z")\
**Posts on this page:** 17\
**Page:** 1

<div class="post-metadata">

**Author:** ![mothmani](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mothmani/32/52088_2.png) [@mothmani](https://discourse.julialang.org/u/mothmani)\
**Post date:** [August 6, 2023, 1:29pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/1 "2023-08-06T13:29:24Z")

</div>

Hi all, I’m quite new (on and off) to Julia and I’m trying to play more consistently with it.  
I have an issue with AD : I can compute the gradient with respect to the weights of a very simple neural network with Zygote but it fails with Enzyme. I do not have the same problem when I compute the gradient with respect to the input of the neural network. I guess this is related to ComponentArrays.  
Would you have any clue about that ? Thanks.

Running the code below, I get the following error :

MethodError: no method matching EnzymeCore.Duplicated(::Vector{Float32}, ::Vector{Float64})  
Closest candidates are:  
EnzymeCore.Duplicated(::T, !Matched::T) where T at ~/data/.julia/packages/EnzymeCore/cwKkU/src/EnzymeCore.jl:64

Here is the code :

using Lux, ComponentArrays, Random, Enzyme  
import Zygote

rng = Random.MersenneTwister(1234)

# Define a basic neural network structure

NN = Lux.Chain( Lux.Dense(5 =\> 5, tanh),  
Lux.Dense(5 =\> 1) )

# Setup the network

ps, st = Lux.setup(rng, NN)

# Test the intialized network with some input values

xtest = [0.1, 0.2, 0.3, 0.4, 0.5]  
NN(xtest, ps, st)[1][1]

# No problem to compute the gradient with respect to the network input with Enzyme…

dx = zeros(size(xtest)[1])  
autodiff(Reverse, xt → NN(xt, ps, st)[1][1], Active, Duplicated(xtest, dx))  
dx

#…or Zygote  
Zygote.gradient(xt → NN(xt, ps, st)[1][1], xtest)

############################################################

# Now let’s try to differentiate wrt the weights of the neural network

ax\_test = getaxes( ComponentArray(ps) )  
θ\_test = getdata( ComponentArray(ps) )

# That works fine with Zygote…

Zygote.gradient( θ → NN(xtest, ComponentArray(θ, ax\_test), st)[1][1], θ\_test )

# …but not with Enzyme

dθ = zeros(size(θ\_test))  
autodiff(Reverse, θ → NN(xtest, ComponentArray(θ, ax\_init), st)[1][1], Active, Duplicated(θ\_test, dθ))

---

<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:** [August 6, 2023, 2:15pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/2 "2023-08-06T14:15:37Z")

</div>

> [@mothmani](#):
>
> MethodError: no method matching EnzymeCore.Duplicated(::Vector{Float32}, ::Vector{Float64})

From the error message `MethodError: no method matching EnzymeCore.Duplicated(::Vector{Float32}, ::Vector{Float64}) ` it looks like you’re trying to construct a duplicated with the primal of type Vector{Float32} and the shadow of type Vector{Float64}. The primal and shadow must be the same type.

What are the types of `θ_test` and `dθ`?

Also you should not create a capturing closure into the autodiff function, but instead pass the additional args. E.g.

`autodiff(Reverse, (NN, xtest, θ, ax_init, st) → NN(xtest, ComponentArray(θ, ax_init), st)[1][1], Active, Const(NN), Const(xtest), Duplicated(θ_test, dθ), Const(ax_init), Const(st))`

---

<div class="post-metadata">

**Author:** ![mothmani](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mothmani/32/52088_2.png) [@mothmani](https://discourse.julialang.org/u/mothmani)\
**Post date:** [August 6, 2023, 8:55pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/3 "2023-08-06T20:55:51Z")

</div>

You were right θ\_test is a Vector{Float32} (I guess it’s the default type for weights in Lux) while I define dθ with the zeros function which returns by default a Vector{Float64}.  
However defining  
dθ = zeros(Float32, size(θ\_test))

and using the code you suggest :  
autodiff(Reverse, (NN, xtest, θ, ax\_init, st) → NN(xtest, ComponentArray(θ, ax\_init), st)[1][1], Active, Const(NN), Const(xtest), Duplicated(θ\_test, dθ), Const(ax\_init), Const(st))

gives now the following error :  
UndefVarError: θ not defined

That’s quite cryptic to me…

---

<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:** [August 6, 2023, 11:11pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/4 "2023-08-06T23:11:42Z")

</div>

The theta Unicode symbol seems to have two different glyfps in the snippet (sorry I made my response on my phone). Can you just write theta in ascii instead of Unicode just to confirm.

---

<div class="post-metadata">

**Author:** ![mothmani](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mothmani/32/52088_2.png) [@mothmani](https://discourse.julialang.org/u/mothmani)\
**Post date:** [August 7, 2023, 9:36am UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/5 "2023-08-07T09:36:25Z")

</div>

Thank you for your reply. And sorry to all for the wrong code formatting in my previous messages.

1. There was an error in my code, apologies : the `ax_init` variable doesn’t exist, it’s `ax_test`. So I don’t understand the returned error message related to `θ`

2. My initial code `autodiff(Reverse, θ -> NN(xtest, ComponentArray(θ, ax_test), st)[1][1], Active, Duplicated(θ_test, dθ))` runs and gives me the same value as with Zygote. However you advised no to calculate the gradient in such a way and, more problematic, every time I change `xtest` the line has to compile again and it takes ages…

3. The code you suggested (corrected with `ax_test` instead of the wrong variable name `ax_init`) : `autodiff(Reverse, (NN, xtest, θ, ax_test, st) → NN(xtest, ComponentArray(θ, ax_test), st)[1][1], Active, Const(NN), Const(xtest), Duplicated(θ_test, dθ), Const(ax_test), Const(st))` still gives me the following error : `UndefVarError: θ not defined`

If you have any clue about that, thanks a lot !

---

<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:** [August 7, 2023, 5:48pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/6 "2023-08-07T17:48:04Z")

</div>

Mind posting an entire test.jl runnable file? Also did you try the unicode renaming (e.g. replacing all θ with theta, I’m still wondering if its a unicode copy paste bug that creates two different but similar looking symbols)?

Re 2, that’s great! The latter code is better because it means things will probably not be (type unstable) captured, which can be significant for both performance, and also because we’re not feature complete on Julia’s type unstable code (you’d get a very loud error if encountering something unsupported). The constant recompile time is odd, since it should be cached, but looking at the whole file will let me provide some better insights to it.

---

<div class="post-metadata">

**Author:** ![mothmani](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mothmani/32/52088_2.png) [@mothmani](https://discourse.julialang.org/u/mothmani)\
**Post date:** [August 9, 2023, 11:02am UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/7 "2023-08-09T11:02:00Z")

</div>

Sure, here is the code :

```julia
using Lux, ComponentArrays, Random, Enzyme
import Zygote

rng = Random.MersenneTwister(1234)

# Define a basic neural network structure
NN = Lux.Chain( Lux.Dense(5 => 5, tanh),
				Lux.Dense(5 => 1) )

# Setup the network
ps, st = Lux.setup(rng, NN)

# Test the intialized network with some input values
xtest = [0.1, 0.2, 0.3, 0.4, 0.5]
NN(xtest, ps, st)[1][1]

# No problem to compute the gradient with respect to the network input with Enzyme...
dx = zeros(size(xtest)[1])
autodiff(Reverse, xt -> NN(xt, ps, st)[1][1], Active, Duplicated(xtest, dx))
dx

#...or Zygote
Zygote.gradient(xt -> NN(xt, ps, st)[1][1], xtest)

############################################################
# Now let's try to differentiate wrt the weights of the neural network

ax_test = getaxes( ComponentArray(ps) )
θ_test = getdata( ComponentArray(ps) )
 
# No problem to differentiate with Zygote
Zygote.gradient( θ -> NN(xtest, ComponentArray(θ, ax_test), st)[1][1], θ_test )

# Works with Enzyme.autodiff with the following code
dθ = zeros(Float32 , size(θ_test))
autodiff(Reverse, θ -> NN(xtest .+ 0.01, ComponentArray(θ, ax_test), st)[1][1], Active, Duplicated(θ_test, dθ))
dθ

# But returns an error with this one
theta_test = copy(θ_test)
dtheta = copy(dθ) 
	
autodiff(Reverse, (NN, xtest, theta, ax_test, st) → NN(xtest, ComponentArray(theta, ax_test), st)[1][1], Active, Const(NN), Const(xtest), Duplicated(theta_test, dtheta), Const(ax_test), Const(st))

```

The last instruction returns :

```julia
UndefVarError: theta not defined

```

A few remarks :

- as you can see, the problem doesn’t seem to be the character `θ`
- I’m using Pluto.jl on Juliahub (very convenient by the way)

Hope this will help.

---

<div class="post-metadata">

**Author:** ![Sevi](https://avatars.discourse-cdn.com/v4/letter/s/c67d28/32.png) [@Sevi](https://discourse.julialang.org/u/Sevi)\
**Post date:** [August 10, 2023, 6:43am UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/8 "2023-08-10T06:43:12Z")

</div>

> [@mothmani](#):
>
> as you can see, the problem doesn’t seem to be the character `θ`

That’s right, the problem ist the arrow character for the anonymous function: it should be `->` (like in the above lines), but it is `→` in the last line. Pretty sneaky 😅

I do get an error `Enzyme execution failed.`, however, with this line (on Julia 1.8.5 and Enzyme 0.11.6)

```julia
autodiff(Reverse, θ -> NN(xtest .+ 0.01, ComponentArray(θ, ax_test), st)[1][1], Active, Duplicated(θ_test, dθ))

```

---

<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:** [August 12, 2023, 9:25pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/9 "2023-08-12T21:25:41Z")

</div>

What’s the full error?

Enzyme execution failed just means that it caught an exception (and then should print the actual error type and message below).

---

<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:** [August 12, 2023, 9:28pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/10 "2023-08-12T21:28:57Z")

</div>

Doing a quick local test, your code (with only one parameter passed to the closure), type unstable captures the variables xtets, ax\_test, st, etc – which leads to the currently unsupported error.

If you pass them in as additional args rather than capturing in the closure, like in my example above (however now fixing the arrow), it seems to work?

```julia
julia> autodiff(Reverse, (NN, xtest, theta, ax_test, st) -> NN(xtest, ComponentArray(theta, ax_test), st)[1][1], Active, Const(NN), Const(xtest), Duplicated(theta_test, dtheta), Const(ax_test), Const(st))

((nothing, nothing, nothing, nothing, nothing),)

```

---

<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:** [September 2, 2023, 1:43am UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/11 "2023-09-02T01:43:53Z")

</div>

I encounter this discussion as I was trying to efficiently implement the gradient of a NN using `Lux.jl` with respect to the input parameters. I noticed that all the examples here mentioned work, both with `Zygote.jl` and `Enzyme.jl`, but the running time to compute the gradient seems to see very slow if we consider that we are just trying to differentiate a small neural network.

Here the example from @mothmani

```julia
using Lux, ComponentArrays, Random, Enzyme
import Zygote
using ReverseDiff
using Base: @time

rng = Random.MersenneTwister(1234)

# Define a basic neural network structure
NN = Lux.Chain( Lux.Dense(5 => 5, tanh),
				Lux.Dense(5 => 1) )

# Setup the network
ps, st = Lux.setup(rng, NN)

# Test the intialized network with some input values
xtest = [0.1, 0.2, 0.3, 0.4, 0.5]
NN(xtest, ps, st)[1][1]

# No problem to compute the gradient with respect to the network input with Enzyme...
dx = zeros(size(xtest)[1])
@time autodiff(Reverse, xt -> NN(xt, ps, st)[1][1], Active, Duplicated(xtest, dx))
dx

#...or Zygote
@time Zygote.gradient(xt -> NN(xt, ps, st)[1][1], xtest)

```

The example with Zygote takes ~0.20s while Enzyme does it in ~0.10s.

I have an use case where I need to compute the gradient of the NN with respect to the input layer multiple times in one training epoch, which leads to very large running times. **Is there any way to make this code to run faster and/or compute the gradient with respect to multiple inputs at the same time?** I wonder if one of the tricks that @wsmoses mentioned here may work.

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [September 6, 2023, 1:50am UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/12 "2023-09-06T01:50:26Z")

</div>

The way you are timing the code globally, triggers recompilation for a new `xt -> NN(xt, ps, st)` every time. Your timing result should mention that most time is compilation.

```julia
@time Zygote.gradient(sum ∘ first ∘ NN, xtest, ps, st)
# 0.000220 seconds (107 allocations: 8.172 KiB)

```

---

<div class="post-metadata">

**Author:** ![pwingender](https://avatars.discourse-cdn.com/v4/letter/p/ba9def/32.png) [@pwingender](https://discourse.julialang.org/u/pwingender)\
**Post date:** [May 14, 2024, 7:03pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/13 "2024-05-14T19:03:13Z")

</div>

Hi! New to Julia here. I was using this example to train parameters on a NN. After updating Enzyme this morning, it doesn’t seem to be working anymore. Thanks!

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [May 14, 2024, 7:20pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/14 "2024-05-14T19:20:43Z")

</div>

what is the error you are getting? If anything I would expect enzyme support to be better as of this weekend

---

<div class="post-metadata">

**Author:** ![pwingender](https://avatars.discourse-cdn.com/v4/letter/p/ba9def/32.png) [@pwingender](https://discourse.julialang.org/u/pwingender)\
**Post date:** [May 14, 2024, 8:24pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/15 "2024-05-14T20:24:40Z")

</div>

Thanks for the reply. Here’s the error message:

julia\> autodiff(Reverse, θ → NN(xtest .+ 0.01, ComponentArray(θ, ax\_test), st)[1][1], Active, Duplicated(θ\_test, dθ))  
ERROR: Enzyme execution failed.  
Mismatched activity for: ret { { {} addrspace(10)_, double } } %7, !dbg !27 const val: %7 = insertvalue { { {} addrspace(10)_, double } } zeroinitializer, { {} addrspace(10)\*, double } %6, 0, !dbg !23  
Type tree: {}  
You may be using a constant variable as temporary storage for active memory ([FAQ · Enzyme.jl](https://enzyme.mit.edu/julia/stable/faq/#Activity-of-temporary-storage)). If not, please open an issue, and either rewrite this variable to no  
t be conditionally active or use Enzyme.API.runtimeActivity!(true) as a workaround for now

Stacktrace:  
[1] broadcasted  
@ .\broadcast.jl:0

Stacktrace:  
[1] throwerr(cstr::Cstring)  
@ Enzyme.Compiler E:\Profiles\PWingender.julia\packages\Enzyme\srACB\src\compiler.jl:1325

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [May 14, 2024, 8:30pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/16 "2024-05-14T20:30:59Z")

</div>

with runtime activity it works. I will open an issue to investigate it

```julia
Enzyme.API.runtimeActivity!(true)

function test_function(NN, x, ps, st)
    y, _ = NN(x, ps, st)
    return sum(y)
end

@time autodiff(
    Reverse, test_function, Active, Const(NN), Duplicated(xtest, dx), Const(ps), Const(st))

```

This is also to do with broadcasting which is being looked into [Error for JVP by Enzyme · Issue #628 · LuxDL/Lux.jl · GitHub](https://github.com/LuxDL/Lux.jl/issues/628#issuecomment-2108572561)

---

<div class="post-metadata">

**Author:** ![pwingender](https://avatars.discourse-cdn.com/v4/letter/p/ba9def/32.png) [@pwingender](https://discourse.julialang.org/u/pwingender)\
**Post date:** [May 14, 2024, 8:35pm UTC](https://discourse.julialang.org/t/lux-componentarrays-and-flat-parameters-computing-the-gradient-works-with-zygote-but-not-with-enzyme/102533/17 "2024-05-14T20:35:38Z")

</div>

Thanks for looking into it. Actually, using Enzyme.API.runtimeActivity!(true) doesn’t solve the issue when the gradient is taken wrt the parameters:

using Lux, ComponentArrays, Random, Enzyme  
Enzyme.API.runtimeActivity!(true)

rng = Random.default\_rng()

# Define a basic neural network structure

NN = Lux.Chain( Lux.Dense(5 =\> 5, tanh),  
Lux.Dense(5 =\> 1) )

# Setup the network

ps, st = Lux.setup(rng, NN)

# Test the intialized network with some input values

x\_test = [0.1, 0.2, 0.3, 0.4, 0.5]  
NN(x\_test, ps, st)[1][1]

dx\_test = zeros(size(x\_test)[1])

ax\_test = getaxes( ComponentArray(ps) )  
theta\_test = getdata( ComponentArray(ps) ) |\> f64  
dtheta\_test = zeros(size(theta\_test)[1])

function test\_function(NN, x, theta, ax, st)  
y, \_ = NN(x, ComponentArray(theta, ax), st)  
return sum(y)  
end

autodiff(Reverse, test\_function, Active, Const(NN), Duplicated(x\_test, dx\_test), Const(theta\_test), Const(ax\_test), Const(st))

autodiff(Reverse, test\_function, Active, Const(NN), Duplicated(x\_test, dx\_test), Duplicated(theta\_test, dtheta\_test), Const(ax\_test), Const(st))
