# Using Zygote with ComponentArrays

**URL:** <https://discourse.julialang.org/t/using-zygote-with-componentarrays/109715>\
**Category:** Modelling & Simulations\
**Tags:** zygote, autodiff\
**Created:** [February 4, 2024, 10:29pm UTC](https://discourse.julialang.org/t/using-zygote-with-componentarrays/109715 "2024-02-04T22:29:43Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![ablaom](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ablaom/32/4889_2.png) [@ablaom](https://discourse.julialang.org/u/ablaom)\
**Post date:** [February 4, 2024, 10:29pm UTC](https://discourse.julialang.org/t/using-zygote-with-componentarrays/109715/1 "2024-02-04T22:29:43Z")

</div>

I am computing gradients of functions using Zygote. The argument of my function is a `ComponentArray`, because my functions depend on DifferentialEquation.jl solvers using adjoint method. I have been naively assuming Zygote.jl plays well with ComponentArrays.jl, and on that basis, the following feels like a bug to me:

```julia
using Zygote, ComponentArrays

x = (a=2.0, b=3.0) |> ComponentArray

g(x) = x.a + x. b

# two-step definition of `h`:
h(a, b) = a * b
h(x) = h(x...)

f(x) = g(x)*h(x)

# julia> f(x)
# 30.0

gradient(f, x)

# ERROR: MethodError: no method matching +(::Tuple{Float64, Float64}, ::ComponentVector{Float64, Vector{Float64}, Tuple{Axis{(a = 1, b = 2)}}})

```

I can remove the problem in three ways:

1. Leave `x` as a named tuple
2. Replace the two-step definition of `h` with the one-liner, `h(x) = x.a + x.b`.
3. Add the following hack (fix?) for the definition of `accum` in Zygote:

```julia
Zygote.accum(x::AbstractArray, y::Tuple) = accum(x, collect(y))
Zygote.accum(x::Tuple, y::AbstractArray) = accum(collect(x), y)

```

So, is this a bug, or am I expecting too much of Zygote?

---

<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:** [February 9, 2024, 8:14pm UTC](https://discourse.julialang.org/t/using-zygote-with-componentarrays/109715/2 "2024-02-09T20:14:04Z")

</div>

> [@ablaom](#):
>
> `h(x) = h(x...)`

Zygote thinks the gradient of a splat is always a Tuple, it’s a longstanding bug and I presume where this Tuple comes from. If that were fixed, perhaps it would make a gradient which is an array here, “natural” representation.

However, the gradient of `g(x) = x.a + x. b` is “structural”, a NamedTuple. These two gradient representations cannot by default be added. Although in this case, overloading `accum` may work.

ChainRules has a mechanism for standardising on one of these types. It’s used [here](https://github.com/jonniedie/ComponentArrays.jl/blob/94ab225f73b2837dda372d63c70e164ca0607358/src/compat/chainrulescore.jl#L35-L47) to turn both such representations into another ComponentArray. That would ideally allow them to be added.

tl;dr is that you should probably always avoid splats of arrays.
