# Help with Zygote and parameters

**URL:** <https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329>\
**Category:** New to Julia\
**Tags:** zygote\
**Created:** [July 1, 2020, 1:22am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329 "2020-07-01T01:22:18Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![amrods](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/amrods/32/2543_2.png) [@amrods](https://discourse.julialang.org/u/amrods)\
**Post date:** [July 1, 2020, 1:22am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/1 "2020-07-01T01:22:19Z")

</div>

I’m having a hard time figuring out why this works:

```julia
using Zygote
W, b = rand(2, 3), rand(2)
predict(x) = W*x .+ b
g = gradient(() -> sum(predict([1,2,3])), Params([W, b]))
g[W], g[b]

```

but this doesn’t:

```julia
using Zygote
a = 2
x = 2
f(x) = x^a
gp = gradient(() -> f(x), Params(a))
gp[a]

```

I get the error:

```julia
ERROR: Only reference types can be differentiated with `Params`.

```

Can anyone help?

---

<div class="post-metadata">

**Author:** ![amrods](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/amrods/32/2543_2.png) [@amrods](https://discourse.julialang.org/u/amrods)\
**Post date:** [July 1, 2020, 2:18am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/2 "2020-07-01T02:18:46Z")

</div>

I think I figured it out. Everything has to be an array except the output of the function. So this now works:

```julia
using Zygote
a = [2]
x = [2]
f(x) = x.^a[1]
gp = gradient(() -> sum(f(x)), Params([a]))
gp[a]

```

---

<div class="post-metadata">

**Author:** ![ettersi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ettersi/32/6829_2.png) [@ettersi](https://discourse.julialang.org/u/ettersi)\
**Post date:** [July 1, 2020, 3:08am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/3 "2020-07-01T03:08:11Z")

</div>

Your second code works, but using arrays is quite a drag on performance:

```julia
# Integer version
julia> @btime $(Ref(2))[]^$(Ref(2))[]
  3.369 ns (0 allocations: 0 bytes)
4

# Array version
julia> @btime [2].^[2][1];
  82.958 ns (3 allocations: 288 bytes)

```

The easiest way to get what you want is obviously `gradient(a->2^a, 2)`, but I am assuming there are other considerations which lead you to the approach you proposed above. We might be able to help further if you share more details.

---

<div class="post-metadata">

**Author:** ![amrods](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/amrods/32/2543_2.png) [@amrods](https://discourse.julialang.org/u/amrods)\
**Post date:** [July 1, 2020, 3:26am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/4 "2020-07-01T03:26:00Z")

</div>

Thanks for offering. This is what I am trying to accomplish:  
I have a function that I would normally write as

```julia
F(L1, L2; a1=1,a2=1) = a1*(L1^a2 + L2^a2)^(1/a2)

```

`a1` and `a2` are parameters. I’m trying to get the derivatives with respect to `a1` and `a2`. I could write it using arrays:

```julia
F(L; a=ones(1,2)) = a[1]*(L[1]^a[2] + L[2]^a[2])^(1/a[2])

```

I was trying to find out why the following does not output an answer:

```julia
grads = gradient(() -> F(L), Params([a]))
grads[a]

```

Edit: Let me provide a little more context. Consider `L` as “data”, and `a` as parameters to be estimated later. The derivatives with respect to `L` have theoretical importance. The derivatives with respect to `a` are to be used in an optimization procedure later.

---

<div class="post-metadata">

**Author:** ![ettersi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ettersi/32/6829_2.png) [@ettersi](https://discourse.julialang.org/u/ettersi)\
**Post date:** [July 1, 2020, 8:42am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/5 "2020-07-01T08:42:45Z")

</div>

```julia
gradient((a1,a2)->F(L1,L2; a1=a1,a2=a2), a1,a2)

```

should do the trick, no?

---

<div class="post-metadata">

**Author:** ![amrods](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/amrods/32/2543_2.png) [@amrods](https://discourse.julialang.org/u/amrods)\
**Post date:** [July 1, 2020, 9:05am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/6 "2020-07-01T09:05:32Z")

</div>

That works. Then I’m unsure when it is necessary to employ `Params`. Could you explain a little about that?

---

<div class="post-metadata">

**Author:** ![ettersi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ettersi/32/6829_2.png) [@ettersi](https://discourse.julialang.org/u/ettersi)\
**Post date:** [July 1, 2020, 9:39am UTC](https://discourse.julialang.org/t/help-with-zygote-and-parameters/42329/7 "2020-07-01T09:39:38Z")

</div>

I’m no expert, but I would say the main occasion when `Params` comes in handy is if the function to differentiate consists of many nested, parametrised functions (e.g. a deep neural network). In this case, it can be annoying to explicitly pass the parameters through the callstack, and `Params` provides a means to avoid that. However, note that the [Zygote documentation](https://fluxml.ai/Zygote.jl/dev/) lists at least two other ways to achieve this and actually recommend against using `Params`:

> However, implicit parameters exist mainly for compatibility with Flux’s current AD; it’s recommended to use the other approaches unless you need this.
