# Is there a way to teach Zygote to derive Diagonal \* Vector more efficiently?

**URL:** <https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145>\
**Category:** General Usage\
**Created:** [August 29, 2019, 12:22am UTC](https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145 "2019-08-29T00:22:34Z")\
**Posts on this page:** 5\
**Page:** 1

<div class="post-metadata">

**Author:** ![marius311](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marius311/32/3953_2.png) [@marius311](https://discourse.julialang.org/u/marius311)\
**Post date:** [August 29, 2019, 12:22am UTC](https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145/1 "2019-08-29T00:22:35Z")

</div>

I’m still getting my head around reverse-mode diff, but is there a way to teach Zygote to do the derivative of Diagonal \* Vector without allocating N^2 memory?

```julia
using LinearAlgebra, BenchmarkTools, Zygote

v = rand(4096)
D = Diagonal(v)

@btime gradient(α -> norm((α * D) * v), 1)
# 53.308 ms (32915 allocations: 129.41 MiB)

```

The `129.29MiB` is basically the size of `v*v'` which appears to get computed into a dense matrix in one of the adjoints. However, if I rewrite the exact same operation slightly differently I can get:

```julia
@btime gradient(α -> norm((α * D).diag .* v), 1)
# 871.463 μs (32919 allocations: 1.41 MiB)

```

So in theory it appears possible. I tried adding something inspired by the [adjoint rule](https://github.com/FluxML/Zygote.jl/blob/d74f3cf5ed3c185969ff7787ee36d15c105e6653/src/lib/broadcast.jl#L64-L65) for `Vector .* Vector`,

```julia
@adjoint *(x::Diagonal, y::Vector) = x.diag .* y,
  z̄ -> (unbroadcast(x, z̄ .* conj.(y)), unbroadcast(y, z̄ .* conj.(x)))

```

but this does not work (yields wrong answer, and memory consumption is the same).

Does anyone have any suggestions on if this is possible (seems it must be?), and if so, how to do it? Many thanks.

---

<div class="post-metadata">

**Author:** ![tkf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tkf/32/17635_2.png) [@tkf](https://discourse.julialang.org/u/tkf)\
**Post date:** [August 29, 2019, 12:48am UTC](https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145/2 "2019-08-29T00:48:44Z")

</div>

That’s probably because `z̄ .* conj.(x)` creates a matrix? This seems to work:

```julia
@adjoint *(x::Diagonal, y::Vector) = x.diag .* y,
    z̄ -> (Diagonal(unbroadcast(x.diag, z̄ .* conj.(y))), unbroadcast(y, z̄ .* conj.(x.diag)))

```

---

<div class="post-metadata">

**Author:** ![marius311](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marius311/32/3953_2.png) [@marius311](https://discourse.julialang.org/u/marius311)\
**Post date:** [August 29, 2019, 1:04am UTC](https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145/4 "2019-08-29T01:04:07Z")

</div>

Ah, I had messed it up a bit, your solutions makes sense, thanks!

---

<div class="post-metadata">

**Author:** ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)\
**Post date:** [August 29, 2019, 6:02am UTC](https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145/5 "2019-08-29T06:02:51Z")

</div>

Please consider contributing this to Zygote.jl.

---

<div class="post-metadata">

**Author:** ![tkf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tkf/32/17635_2.png) [@tkf](https://discourse.julialang.org/u/tkf)\
**Post date:** [August 29, 2019, 6:07am UTC](https://discourse.julialang.org/t/is-there-a-way-to-teach-zygote-to-derive-diagonal-vector-more-efficiently/28145/6 "2019-08-29T06:07:35Z")

</div>

See:  
[https://github.com/FluxML/Zygote.jl/issues/316](https://github.com/FluxML/Zygote.jl/issues/316)
