# Zygote and StructArrays

**URL:** <https://discourse.julialang.org/t/zygote-and-structarrays/35537>\
**Category:** General Usage\
**Tags:** differentiation, flux, zygote, structarrays\
**Created:** [March 4, 2020, 8:54pm UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537 "2020-03-04T20:54:22Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [March 4, 2020, 8:54pm UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/1 "2020-03-04T20:54:22Z")

</div>

It seems `Flux.@functor` does not get along with `StructArray`s?

```julia
using StructArrays, Flux

struct A{SA}
	s::SA
end

struct Point
	x::Float64
	y::Float64
end

Flux.@functor A
Flux.@functor Point

a = A(StructArray((Point(i,2i) for i=0.0:10.0)))
Flux.params(a) # Params([])

```

I was expecting `Flux.params(a)` to detect the two parameter vectors `a.s.x` and `a.s.y` contained in the `StructArray`.

Any suggestions on how I can make this work?

---

<div class="post-metadata">

**Author:** ![findmyway](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/findmyway/32/4946_2.png) [@findmyway](https://discourse.julialang.org/u/findmyway)\
**Post date:** [March 5, 2020, 2:04am UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/2 "2020-03-05T02:04:00Z")

</div>

This is a very interesting question.

I just made some tests and found that the problem is the type constraint in the following line in Flux

[https://github.com/FluxML/Flux.jl/blob/5a4f1932a6f1d7e09aa0f70497ee689a428a2421/src/functor.jl#L74](https://github.com/FluxML/Flux.jl/blob/5a4f1932a6f1d7e09aa0f70497ee689a428a2421/src/functor.jl#L74)

If we remove the `<:Number` and define a function like this:

```julia
julia> Flux.params!(p::Flux.Params, x::AbstractArray, seen = IdSet()) = push!(p, x)

```

Then we can get:

```julia
julia> Flux.params(a)
Params([Point[Point(0.0, 0.0), Point(1.0, 2.0), Point(2.0, 4.0), Point(3.0, 6.0), Point(4.0, 8.0), Point(5.0, 10.0), Point(6.0, 12.0), Point(7.0, 14.0), Point(8.0, 16.0), Point(9.0, 18.0), Point(10.0, 20.0)]])

```

If this is what you want, maybe create an issue in Flux.jl?

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [March 5, 2020, 6:44am UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/3 "2020-03-05T06:44:59Z")

</div>

It doesn’t seem to work very well.

```julia
using StructArrays, Flux, Zygote

Flux.params!(p::Flux.Params, x::AbstractArray, seen = Zygote.IdSet()) = push!(p, x)

struct A{SA}
	s::SA
end

struct Point
	x::Float64
	y::Float64
end

Flux.@functor A
Flux.@functor Point

a = A(StructArray((Point(i,2i) for i=0.0:10.0)))

function fun(a::A)
	sum(a.s.x) + 2sum(a.s.y)
end

ps = Flux.params(a)
gs = gradient(ps) do
	fun(a)
end

julia> gs[a.s]
(fieldarrays = nothing, y = [2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0])

```

For some reason `a.s.x` is being ignored.

---

<div class="post-metadata">

**Author:** ![findmyway](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/findmyway/32/4946_2.png) [@findmyway](https://discourse.julialang.org/u/findmyway)\
**Post date:** [March 5, 2020, 7:01am UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/4 "2020-03-05T07:01:26Z")

</div>

Now I see. Sorry that I’m not very familiar with `StructArray` before.

You have defined the following `functor`:

```julia
Flux.@functor A
Flux.@functor Point

```

But the one for `StructArray` is missing. You can define it like this:

```julia
StructArray((Point(i,2i) for i=0.0:10.0))
b = StructArray((Point(i,2i) for i=0.0:10.0))
Flux.@functor typeof(b) (x, y)
a = A(b)

```

Here, we should be able to define a general functor function for `SubArray`s. You may take it as an exercise. 😆

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [March 6, 2020, 4:47pm UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/5 "2020-03-06T16:47:52Z")

</div>

No way to define a generic functor for `StructArray`, containing any data struct?

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [April 14, 2020, 8:57pm UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/6 "2020-04-14T20:57:54Z")

</div>

I did the following

```julia
function Flux.trainable(s::StructArray)
    getproperty.(Ref(s), propertynames(s))
end

```

In this case I get to populate the params correctly, but I still get an error when trying to differentiate. Here is an example:

```julia
using Flux, StructArrays
struct Point
    x::Float64
    y::Float64
end
pnts = StructArray([Point(randn(), randn()) for i = 1:10, j = 1:10])
fun(p::Point, v::Real) = (p.x^2 + p.y^2) * v
V = randn(10,10,10)
ps = Flux.params(pnts)
gs = gradient(ps) do
    sum(fun.(pnts, V))
end

```

This results in an error.

```julia
ERROR: MethodError: no method matching reducedim_init(::typeof(identity), ::typeof(Zygote.accum), ::Array{NamedTuple{(:x, :y),Tuple{Float64,Float64}},3}, ::Tuple{Int64,Int64,Int64})
Closest candidates are:
  reducedim_init(::Any, ::Union{typeof(+), typeof(Base.add_sum)}, ::AbstractArray, ::Any) at reducedim.jl:109
  reducedim_init(::Any, ::Union{typeof(*), typeof(Base.mul_prod)}, ::AbstractArray, ::Any) at reducedim.jl:112
  reducedim_init(::Any, ::typeof(min), ::AbstractArray, ::Any) at reducedim.jl:132
  ...
Stacktrace:
 [1] _mapreduce_dim(::Function, ::Function, ::NamedTuple{(),Tuple{}}, ::Array{NamedTuple{(:x, :y),Tuple{Float64,Float64}},3}, ::Tuple{Int64,Int64,Int64}) at ./reducedim.jl:317
 [2] #mapreduce#580 at ./reducedim.jl:307 [inlined]
 [3] #reduce#582 at ./reducedim.jl:352 [inlined]
 [4] accum_sum(::Array{NamedTuple{(:x, :y),Tuple{Float64,Float64}},3}; dims::Tuple{Int64,Int64,Int64}) at /home/cossio/.julia/packages/Zygote/4tJp5/src/lib/broadcast.jl:36
 [5] unbroadcast at /home/cossio/.julia/packages/Zygote/4tJp5/src/lib/broadcast.jl:48 [inlined]
 [6] map(::typeof(Zygote.unbroadcast), ::Tuple{StructArray{Point,2,NamedTuple{(:x, :y),Tuple{Array{Float64,2},Array{Float64,2}}},Int64},Array{Float64,3}}, ::Tuple{Array{NamedTuple{(:x, :y),Tuple{Float64,Float64}},3},Array{Float64,3}}) at ./tuple.jl:177
 [7] (::Zygote.var"#1734#1741"{Tuple{StructArray{Point,2,NamedTuple{(:x, :y),Tuple{Array{Float64,2},Array{Float64,2}}},Int64},Array{Float64,3}},Val{3},Array{typeof(∂(fun)),3}})(::FillArrays.Fill{Float64,3,Tuple{Base.OneTo{Int64},Base.OneTo{Int64},Base.OneTo{Int64}}}) at /home/cossio/.julia/packages/Zygote/4tJp5/src/lib/broadcast.jl:139
 [8] #4367#back at /home/cossio/.julia/packages/ZygoteRules/6nssF/src/adjoint.jl:49 [inlined]
 [9] (::Zygote.var"#175#176"{Zygote.var"#4367#back#1745"{Zygote.var"#1734#1741"{Tuple{StructArray{Point,2,NamedTuple{(:x, :y),Tuple{Array{Float64,2},Array{Float64,2}}},Int64},Array{Float64,3}},Val{3},Array{typeof(∂(fun)),3}}},Tuple{NTuple{4,Nothing},Tuple{}}})(::FillArrays.Fill{Float64,3,Tuple{Base.OneTo{Int64},Base.OneTo{Int64},Base.OneTo{Int64}}}) at /home/cossio/.julia/packages/Zygote/4tJp5/src/lib/lib.jl:170
 [10] #344#back at /home/cossio/.julia/packages/ZygoteRules/6nssF/src/adjoint.jl:49 [inlined]
 [11] broadcasted at ./broadcast.jl:1238 [inlined]
 [12] #5 at ./REPL[9]:2 [inlined]
 [13] (::typeof(∂(#5)))(::Float64) at /home/cossio/.julia/packages/Zygote/4tJp5/src/compiler/interface2.jl:0
 [14] (::Zygote.var"#48#49"{Zygote.Params,Zygote.Context,typeof(∂(#5))})(::Float64) at /home/cossio/.julia/packages/Zygote/4tJp5/src/compiler/interface.jl:108
 [15] gradient(::Function, ::Zygote.Params) at /home/cossio/.julia/packages/Zygote/4tJp5/src/compiler/interface.jl:45
 [16] top-level scope at REPL[9]:1

```

Issue: [reducedim\_init MethodError with structs · Issue #594 · FluxML/Zygote.jl · GitHub](https://github.com/FluxML/Zygote.jl/issues/594)

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [June 7, 2020, 11:50pm UTC](https://discourse.julialang.org/t/zygote-and-structarrays/35537/7 "2020-06-07T23:50:27Z")

</div>

I have created a small package that helps solves some of these issues: [GitHub - cossio/ZygoteStructArrays.jl](https://github.com/cossio/ZygoteStructArrays.jl)
