# How to make Zygote avoid differentiating with respect to some fields in struct

**URL:** <https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520>\
**Category:** Machine Learning\
**Tags:** question, machine-learning, zygote, struct, autodiff\
**Created:** [July 29, 2021, 8:05pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520 "2021-07-29T20:05:57Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![Nikos\_Gianniotis](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nikos_gianniotis/32/11487_2.png) [@Nikos\_Gianniotis](https://discourse.julialang.org/u/Nikos_Gianniotis)\
**Post date:** [July 29, 2021, 8:05pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/1 "2021-07-29T20:05:57Z")

</div>

Hello everyone,

I have been always used ForwardDiff for my automatic differentiation needs, but recently I thought I should try out Zygote instead. My main motivation for doing so is the ability of Zygote to work with callable structs that can be used for implementing ML models.

The first obstacle I have encountered is the following: besides model parameters, a Struct may also hold additional information, e.g. random seed, various weights, some text description, that should not be included in the automatic differentiation. To make my question more concrete, I post below some code (basically a modified version of an example from the Zygote documentation):

```julia
using Zygote

struct Linear
         W
         b
         C
       end

(l::Linear)(x) = l.W * x .+ l.b

model = Linear(rand(2, 5), rand(2), rand(3))

N = 3

X = [randn(5) for i in 1:N]
Y = [randn(2) for i in 1:N]

function loss(model)
  loss = 0.0
  for n in 1:N
    loss += sum(model.C[n]*(model(X[n]) .- Y[n]).^2)
  end
  loss
end

```

Basically, we define a linear model, N=3 data item pairs of inputs X and targets Y and a loss function.  
In addition to the original example, the struct here also holds some weighing coefficients C which are constant and are not free model parameters that need to be optimised.

However, if I call `dmodel = gradient(loss, model)[1]`, Zygote will also offer the gradient wrt C:

```julia
(W = [4.109999051379186 12.539463966870912 … 5.154617937809002 3.818367921164797; 2.7677477951659277 8.485557353981928 … 0.33008326109979613 5.871635193000793], b = [10.270604809342977, 6.932605807555453], C = [5.114710562186542, 8.589047944559583, 15.909538432009729])

```

Is there perhaps a way of telling Zygote not to differentiate with respect to C? Many thanks.

---

<div class="post-metadata">

**Author:** ![jondeuce](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jondeuce/32/16378_2.png) [@jondeuce](https://discourse.julialang.org/u/jondeuce)\
**Post date:** [July 30, 2021, 3:40pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/2 "2021-07-30T15:40:07Z")

</div>

In Flux, this is typically accomplished using [@functor](https://fluxml.ai/Functors.jl/stable/#Basic-Usage-and-Implementation). From the documentation:

> To include only certain fields of a struct, one can pass a tuple of field names to [`@functor`](https://fluxml.ai/Functors.jl/stable/@ref):

```julia
julia> struct Baz
         x
         y
       end

julia> @functor Baz (x,)

julia> model = Baz(1, 2)
Baz(1, 2)

julia> fmap(float, model)
Baz(1.0, 2)

```

Apparently, this does not interact perfectly with Zygote, though; see the relevant issue [here](https://github.com/FluxML/Zygote.jl/issues/1042).

Per the discussion there, it seems that it is sometimes the case that correctly computing the gradient of one field will implicitly depend on the gradient of other fields, and therefore ignoring the other gradients would silently give incorrect results.

---

<div class="post-metadata">

**Author:** ![CarloLucibello](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/carlolucibello/32/3278_2.png) [@CarloLucibello](https://discourse.julialang.org/u/CarloLucibello)\
**Post date:** [July 30, 2021, 7:05pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/3 "2021-07-30T19:05:17Z")

</div>

A workaround is to define an auxiliary function blocking the gradient:

```julia
julia> nograd(x) = x
nograd (generic function with 1 method)

julia> Zygote.@nograd nograd

julia> function loss(model)
         loss = 0.0
         for n in 1:N
           c = nograd(model.C[n]) 
           loss += sum(c*(model(X[n]) .- Y[n]).^2)
         end
         loss
       end
loss (generic function with 1 method)

julia> dmodel = gradient(loss, model)[1]
(W = [-2.6711034296853375 -0.2254551144087391 … 1.4033779589003235 3.1219721905976305; 2.5271537498909424 -0.8264030561248071 … -1.663924873993493 -1.3721983219005442], b = [1.8739815151086066, -0.8248439951554849], C = nothing)

```

---

<div class="post-metadata">

**Author:** ![cortner](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cortner/32/204_2.png) [@cortner](https://discourse.julialang.org/u/cortner)\
**Post date:** [July 30, 2021, 10:00pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/4 "2021-07-30T22:00:15Z")

</div>

Nice

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [July 30, 2021, 10:40pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/5 "2021-07-30T22:40:31Z")

</div>

You can also use the built-in `Zygote.dropgrad` in place of `nograd`: [Utilities · Zygote](https://fluxml.ai/Zygote.jl/latest/utils/#Zygote.dropgrad). For a higher-level API, the linked Flux issue is definitely the one to follow.

---

<div class="post-metadata">

**Author:** ![Nikos\_Gianniotis](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nikos_gianniotis/32/11487_2.png) [@Nikos\_Gianniotis](https://discourse.julialang.org/u/Nikos_Gianniotis)\
**Post date:** [August 2, 2021, 8:07pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/6 "2021-08-02T20:07:41Z")

</div>

I had never come across functors, thanks very much for informing me about it. I thought I had looked thoroughly in the Zygote documentation, but apparently I didn’t. Looks like this could be useful to me in other contexts too…

---

<div class="post-metadata">

**Author:** ![Nikos\_Gianniotis](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nikos_gianniotis/32/11487_2.png) [@Nikos\_Gianniotis](https://discourse.julialang.org/u/Nikos_Gianniotis)\
**Post date:** [August 2, 2021, 8:08pm UTC](https://discourse.julialang.org/t/how-to-make-zygote-avoid-differentiating-with-respect-to-some-fields-in-struct/65520/7 "2021-08-02T20:08:51Z")

</div>

Thanks for your reply. This is easier for me to understand than the functor solution.
