# Using gradients from struct with Zygote

**URL:** https://discourse.julialang.org/t/using-gradients-from-struct-with-zygote/50766
**Category:** Machine Learning
**Tags:** zygote
**Created:** [November 25, 2020, 6:22pm UTC](https://discourse.julialang.org/t/using-gradients-from-struct-with-zygote/50766 "2020-11-25T18:22:09Z")
**Posts on this page:** 3
**Page:** 1

<div class="post-metadata">

### Author: ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)
#### Post date: [November 25, 2020, 6:22pm UTC](https://discourse.julialang.org/t/using-gradients-from-struct-with-zygote/50766/1 "2020-11-25T18:22:09Z")

</div>

Hi!

Let’s say I have some nested struct (in my case the nesting is arbitrarily complicated)

```julia
struct Foo
a::Vector
end

struct Bar
b::Foo
c::Float64
end

c = Bar(Foo([2.0]), 1.0)

```

When using

```julia
g = Zygote.gradient(c) do x
   do_something(x)...
end

```

`g[1]` will be a tuple of the form `(c=nothing, b=(a=[...]))`.

My problem is that I have no idea how to apply this gradient automatically on my object.

I know there is `Flux.params` to compute the gradients implicitly but it has big disadvantages like getting saving unwanted gradients and also it just fails in a lot of cases for me.

What should I do?

---

<div class="post-metadata">

### Author: ![DrChainsaw](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/drchainsaw/32/8497_2.png) [@DrChainsaw](https://discourse.julialang.org/u/DrChainsaw)
#### Post date: [November 26, 2020, 6:39am UTC](https://discourse.julialang.org/t/using-gradients-from-struct-with-zygote/50766/2 "2020-11-26T06:39:20Z")

</div>

> [@theogf](#):
>
> I know there is `Flux.params` to compute the gradients implicitly but it has big disadvantages like getting saving unwanted gradients and also it just fails in a lot of cases for me.

I guess you are referring to the mechanism for [implicit gradients](https://fluxml.ai/Zygote.jl/dev/#Gradients-of-ML-models-1) which is a mechanism for when you in advance know exactly which parts of that nested struct you want the gradients for. It should in other words not have the problem of getting unwanted gradients. `Flux.params` is just `Flux`s way of conveniently returning all `AbstractArray`s found in the nested struct, but afaik it is not tied to that mechanism.

Anyways, the docs (in the same section I linked) recommend to not use that approach so it might be better to work with the output you have got there.

I don’t know what is the best way, but I think you should be able to use `getfield` to traverse the struct. I have found that Julias multiple dispatch makes it relatively painless to recurse into nested structs. Here is an untested skeleton implementation:

```julia
apply_gradient(g::NamedTuple, s) = foreach(pairs(g)) do (fieldname, subgradient)
              apply_gradient(subgradient, getfield(s, fieldname))
end

function apply_gradient(::Nothing, x) end # No gradient -> do nothing

apply_gradient(g::AbstractArray, p::AbstractArray) = g .- p #might want to propagate some policy (e.g. a learning rate) as a third argument

```

You might need to add a few methods there depending on what one might find in your structure (e.g. if there are arrays or tuples of structs in there).

---

<div class="post-metadata">

### Author: ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)
#### Post date: [November 26, 2020, 12:35pm UTC](https://discourse.julialang.org/t/using-gradients-from-struct-with-zygote/50766/3 "2020-11-26T12:35:15Z")

</div>

Thanks I was exactly looking for something like this
