# Flux.jl Restrict Gradients to Non-Zero values in sparse layer

**URL:** <https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124>\
**Category:** Machine Learning\
**Tags:** flux\
**Created:** [November 26, 2021, 4:16pm UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124 "2021-11-26T16:16:47Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![james-a-mcmanus](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/james-a-mcmanus/32/14368_2.png) [@james-a-mcmanus](https://discourse.julialang.org/u/james-a-mcmanus)\
**Post date:** [November 26, 2021, 4:16pm UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124/1 "2021-11-26T16:16:47Z")

</div>

Hi, I’m trying to make a custom network layer with sparse weights.  
I only want to train the non-zero values of the layer, however when I try to create a gradient from this the gradient returns nothing.

Here is a MWE.

```julia
using Flux, SparseArrays

struct sparsewrapper{T}
	weights::SparseMatrixCSC{T, Int64}
end

Flux.trainable(l::sparsewrapper) = [l.weights.nzval]
layer = sparsewrapper(SparseArrays.sprand(100,100,0.2))
model(x) = layer.weights * x
loss(x,y) = Flux.mse(model(x),y)
g = Flux.gradient(()->loss(rand(100),rand(100)), Flux.params(layer))
println(g[layer.weights.nzval])

```

How do I use Flux.trainable to pick out the non-zero values.

Thanks in advance!  
James

---

<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 27, 2021, 10:47am UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124/2 "2021-11-27T10:47:25Z")

</div>

I’m guessing here that the AD (Zygote) does not try to differentiate the internals of the SparseMatrix multiplication operation and therefore never sees that `layer.weights.nzval` plays any part.

If you just point to weights as the trainable it seems like it is smart enough to give you a sparse gradient so it seems like what you want should work without any extra effort:

```julia
julia> Flux.trainable(l::sparsewrapper) = (l.weights,)

julia> g = Flux.gradient(()->loss(rand(100),rand(100)), Flux.params(layer))
Grads(...)

julia> g[layer.weights]
100×100 SparseMatrixCSC{Float64, Int64} with 1922 stored entries:
⠑⢎⠸⠐⡬⠀⠌⠓⠢⠐⠒⠱⠊⢈⠀⠂⠠⠂⡂⠐⢢⠀⠀⡐⣠⠂⠆⢈⢢⠀⠴⢉⠀⡧⢐⠅⠒⠂⠲⠄⠡⢤⠤⠄⠂⢀⡀⡙⢈⠶
⠡⠊⢘⠀⠂⢀⠅⠐⠀⢘⡃⠤⠠⠁⡈⠀⠄⢐⡠⠤⠐⠂⢄⡑⠐⡆⠂⠡⠑⡑⠣⢄⠂⡐⡁⠡⡅⠁⠜⡠⠂⠀⠀⣀⢓⣄⠀⠂⠄⠄
⠀⡈⢑⢀⠀⡀⡃⠐⠀⠂⠢⢵⠔⢁⠙⠀⡀⠂⢓⣈⡂⠀⡪⣀⠉⣀⢀⠰⠀⢊⡀⡀⢄⠄⠀⢨⢁⠐⠨⠄⠵⠁⠊⠈⡂⠀⠀⣀⠩⠀
⢁⡀⢁⠈⡀⠂⠈⠀⠂⢀⠀⠀⠉⢌⠀⠈⠀⠀⠨⠁⠉⠡⠀⣂⠀⠀⣠⠑⠆⠀⠃⠃⢀⡐⠨⢀⡀⠈⢰⢄⠠⠌⠀⠀⡔⠡⢄⠀⠠⡁
⠐⠠⠀⢀⠐⠄⠄⢐⡄⠀⠁⢇⢀⠡⠀⠀⠒⠠⠀⠂⠀⠒⠠⠀⢈⠀⢂⠁⠀⡊⠀⠆⠄⠀⡁⢀⢁⢂⢀⠂⠂⢐⠀⡕⠁⠀⠀⠁⡀⠀
⢆⠂⠁⠂⠨⠀⠑⠁⠊⠀⠐⡐⡈⠀⠀⠌⠈⠀⠀⡄⡔⡌⠉⠈⢰⠌⢬⠐⠚⢐⠆⠆⠒⠄⢀⠀⠄⠃⢈⠬⠔⠆⠀⠞⠴⢈⡄⠈⠒⠀
⢀⠄⢀⠠⡀⡈⠂⠅⠠⠐⠃⠑⠠⠸⠔⠢⠀⡀⠀⠂⡈⢆⠀⢐⠰⡂⠀⠀⠠⠀⠀⠁⣂⡀⡈⡒⠄⣈⠢⡂⠢⡁⠀⢌⠄⡂⠁⢀⠀⠀
⠅⠄⡉⠀⡀⠕⠠⠐⠤⠄⠐⢹⡀⡉⠄⠅⠀⢐⠅⡄⠬⢠⠄⢈⠀⠔⢀⡤⠁⠨⠂⢤⠈⠀⠄⠀⠀⡀⠁⠐⠈⣀⠀⠑⢀⠁⠀⢀⠒⠄
⠬⢄⠁⠀⠐⡄⡁⠀⡑⢠⠀⣀⠈⢀⢀⠠⢂⠀⠈⠘⠁⠦⠠⠄⠠⡃⠢⠀⢀⠀⠋⡡⣐⢀⠂⠄⢈⢄⠀⠠⠉⠠⠐⠀⠌⠡⠄⡪⠐⠫
⠁⠐⠀⢀⡂⠄⠂⡀⡘⠃⠃⢀⠈⡡⠀⠄⠀⢀⢕⠠⡢⢃⠀⠀⠀⠄⢊⠁⢊⠀⢠⠅⢂⡁⢀⠕⠈⢀⠬⠠⠂⡁⠠⢀⠘⠁⡁⢁⡂⣑
⠈⢌⠠⡆⠤⠁⢂⠑⠸⢂⣐⠌⠀⠀⢐⢔⢐⠁⡀⢀⡆⢑⠁⠄⠀⢐⠀⠌⡀⡠⠀⠪⠐⡁⠀⠈⠀⠂⠈⡀⠐⠀⡐⠂⠆⠔⠒⠀⠘⠀
⠆⡴⠀⠒⠂⠀⠄⢀⠢⠂⠀⠂⠢⠀⠄⠀⡈⠦⡂⠋⢖⢠⡀⢙⠈⠁⠈⠅⠑⠉⠂⠆⡇⢈⠀⠈⡀⠁⢢⠁⠨⠠⢒⠈⠨⡀⢀⡥⠠⡃
⢠⠀⠆⢀⠄⢴⠢⠢⡂⢈⠂⠁⠈⠀⠂⠊⣢⢠⠀⠐⡉⠡⠨⠀⠎⠂⠰⢂⠱⠈⡰⠪⠑⢀⠀⡂⠐⠈⠂⠀⡌⡂⠈⢄⠈⠂⢀⠢⠄⠉
⢷⣂⠈⠈⡦⠊⠠⠴⢠⠀⠐⢄⢩⢶⠈⠁⠒⠀⠀⢠⡂⢚⠐⢱⡠⣅⢖⠠⠘⢑⠀⡑⡀⠱⢔⠁⠂⢀⠐⠤⠂⠀⠀⢘⢁⠆⠑⡀⠃⠀
⠀⠄⣄⡀⢌⡀⠤⡌⡔⠐⠂⠁⠨⠐⠑⠄⢉⠈⢚⠄⠆⠡⠢⡀⡣⢴⠀⠀⠠⢂⡄⠬⢀⢰⣀⠀⠁⢀⠀⣀⠁⣁⠈⠀⠠⠌⡀⠅⢁⡀
⠤⠀⠀⢁⠀⠀⠔⠀⠩⠀⠀⠐⡐⡈⠀⠁⡒⠀⠀⠄⠰⠉⠡⠀⡁⠅⠂⠁⢂⠀⡀⢁⠀⠴⣀⠀⠤⠰⠢⠍⡀⠤⠀⣄⡉⠠⠀⠄⡂⠄
⡌⠠⠀⡐⢈⠧⢂⠂⠂⡃⠀⠈⠠⠃⠂⠴⠘⡀⠀⠘⠠⠁⢊⡄⠊⣁⠁⠂⢂⠁⢰⠀⠄⠀⠈⠥⠠⢒⢈⠀⠀⠱⠤⠀⠠⡘⠀⠈⠐⠀
⠑⡀⠰⣀⠈⠈⠀⢈⡁⠐⢠⡊⡂⡀⠨⠂⠌⠀⣁⣑⠒⣁⠀⠪⠠⠂⣀⠢⡆⠀⢊⠆⠄⠠⢈⠈⠈⠂⡄⠸⠠⠂⠠⠊⠀⠀⠀⠼⠘⠐
⠠⠀⢀⠀⠆⡂⠀⡐⠈⠀⠂⠑⠀⠀⠻⠤⠥⠅⠢⢂⠀⠑⢡⠀⠤⢠⠁⠀⠠⠀⢤⠑⠄⢐⢀⠈⠑⠀⠰⠰⠃⠠⠁⠀⡈⠂⡒⠀⢠⠠
⠘⠓⠀⠀⣂⠁⠠⡐⠨⠂⠰⡀⠌⠰⡢⢁⠠⠠⡈⡄⢐⡈⠡⢀⢀⡈⡂⢀⡅⠐⠄⣄⢃⠑⠓⢁⡰⢊⡠⣌⢠⠀⠀⠰⡈⠘⢀⡒⠆⠊
⢁⠐⢀⠈⠝⠠⠂⢂⠂⠀⠀⠀⢨⠲⡀⢄⠂⣀⣄⡠⡤⠀⠀⠀⢁⠁⡀⠀⠂⢐⠄⠀⠌⠀⠡⠂⠄⡐⢀⣀⡒⢒⠀⡀⢠⠌⢀⡀⢂⡅
⠀⠙⠔⢐⢀⢙⠀⠡⠤⣆⡂⠄⢈⠀⠉⠂⢲⠐⢁⠀⠅⠐⠔⠪⠀⠐⡘⠊⠃⢍⠈⠒⠀⠤⠁⠀⡀⢀⢐⡄⡔⠄⠌⠀⠀⠆⠁⡄⠀⠄
⡀⠡⢒⠢⡀⠈⠀⢠⡢⠀⡀⠂⠈⢘⠁⢐⠀⢑⠀⠠⠂⠢⡙⠐⡑⡠⢈⠀⠉⠄⡅⠀⠁⠠⠸⡀⣉⠠⠁⠴⠁⡐⠲⢸⢔⠯⠂⠄⠠⡁
⠘⠂⠑⢐⢄⠀⠠⠀⠠⠁⠌⡈⠄⠈⠍⡸⠙⢀⠊⠘⡀⠁⠐⢅⠐⢒⠀⠡⠠⢀⡀⡠⠈⢀⠀⠠⠠⠌⢈⢀⡑⠂⠁⢀⠂⠁⠀⠄⠈⠂
⠘⡀⠁⠁⠅⠀⡂⠩⢁⠁⠀⠀⡠⢠⢀⠠⠤⢀⢀⠃⡠⢑⢀⠁⠒⠒⠢⠄⠦⠁⠀⠁⠠⠂⠁⡈⡨⠀⠘⠬⠀⠄⠀⠡⢈⡀⢄⠀⠂⣀

```

---

<div class="post-metadata">

**Author:** ![dhairyagandhi96](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dhairyagandhi96/32/7589_2.png) [@dhairyagandhi96](https://discourse.julialang.org/u/dhairyagandhi96)\
**Post date:** [December 1, 2021, 7:26pm UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124/3 "2021-12-01T19:26:42Z")

</div>

Also prefer to use `Flux.@functor` over `Flux.trainable`.

---

<div class="post-metadata">

**Author:** ![james-a-mcmanus](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/james-a-mcmanus/32/14368_2.png) [@james-a-mcmanus](https://discourse.julialang.org/u/james-a-mcmanus)\
**Post date:** [December 3, 2021, 4:38pm UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124/4 "2021-12-03T16:38:51Z")

</div>

Thanks for the reply, that’s odd, I try the same thing but  
`g[layer.weights]` returns a dense matrix. (This is with Flux v0.11.3)

---

<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:** [December 4, 2021, 1:55am UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124/5 "2021-12-04T01:55:33Z")

</div>

That’s a pretty old version of Flux. I believe only the more recent 0.12.x versions can return non-dense gradients

---

<div class="post-metadata">

**Author:** ![james-a-mcmanus](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/james-a-mcmanus/32/14368_2.png) [@james-a-mcmanus](https://discourse.julialang.org/u/james-a-mcmanus)\
**Post date:** [January 5, 2022, 11:07am UTC](https://discourse.julialang.org/t/flux-jl-restrict-gradients-to-non-zero-values-in-sparse-layer/72124/6 "2022-01-05T11:07:06Z")

</div>

> [@DrChainsaw](#):
>
> `julia> g = Flux.gradient(()->loss(rand(100),rand(100)), Flux.params(layer))`

Oh yes, upgrading the Flux version fixed this - thanks everyone!
