# Zygote.jl: How to get the gradient of sparse matrix

**URL:** <https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067>\
**Category:** General Usage\
**Tags:** question, package, differentiation, zygote, reversediff\
**Created:** [April 12, 2021, 3:03am UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067 "2021-04-12T03:03:39Z")\
**Posts on this page:** 12\
**Page:** 1

<div class="post-metadata">

**Author:** ![Richard-Li](https://avatars.discourse-cdn.com/v4/letter/r/85f322/32.png) [@Richard-Li](https://discourse.julialang.org/u/Richard-Li)\
**Post date:** [April 12, 2021, 3:03am UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/1 "2021-04-12T03:03:40Z")

</div>

Hi there, I am trying to differentiate through sparse matrix, and here is the MWE:

```julia
using LinearAlgebra
using SparseArrays
using ChainRulesCore
using Zygote

Zygote.@adjoint function SparseMatrixCSC{T,N}(arr) where {T,N}
    SparseMatrixCSC{T,N}(arr), Δ -> (collect(Δ),)
  end

function test1(a)
    A = sparse([1, 1, 2, 2],[1, 2, 1, 2],a, 2, 2)
    return sum(A)
end 

a = [1.0,2.0,3.0,4.0]
gradient(test1, a)

```

Then I got error:

```julia
ERROR: LoadError: Mutating arrays is not supported
Stacktrace:
 [1] error(::String) at ./error.jl:33
 [2] (::Zygote.var"#399#400")(::Nothing) at /root/.julia/packages/Zygote/CgsVi/src/lib/array.jl:58
 [3] (::Zygote.var"#2265#back#401"{Zygote.var"#399#400"})(::Nothing) at /root/.julia/packages/ZygoteRules/OjfTt/src/adjoint.jl:59
 [4] sparse! at /buildworker/worker/package_linux64/build/usr/share/julia/stdlib/v1.5/SparseArrays/src/sparsematrix.jl:862 [inlined]
 [5] (::typeof(∂(sparse!)))(::FillArrays.Fill{Float64,2,Tuple{Base.OneTo{Int64},Base.OneTo{Int64}}}) at /root/.julia/packages/Zygote/CgsVi/src/compiler/interface2.jl:0
 [6] sparse at /buildworker/worker/package_linux64/build/usr/share/julia/stdlib/v1.5/SparseArrays/src/sparsematrix.jl:703 [inlined]
 [7] (::typeof(∂(sparse)))(::FillArrays.Fill{Float64,2,Tuple{Base.OneTo{Int64},Base.OneTo{Int64}}}) at /root/.julia/packages/Zygote/CgsVi/src/compiler/interface2.jl:0
 [8] sparse at /buildworker/worker/package_linux64/build/usr/share/julia/stdlib/v1.5/SparseArrays/src/sparsematrix.jl:892 [inlined]
 [9] (::typeof(∂(sparse)))(::FillArrays.Fill{Float64,2,Tuple{Base.OneTo{Int64},Base.OneTo{Int64}}}) at /root/.julia/packages/Zygote/CgsVi/src/compiler/interface2.jl:0
 [10] test1 at /root/codes/test_zygote/test_sparse.jl:11 [inlined]
 [11] (::typeof(∂(test1)))(::Float64) at /root/.julia/packages/Zygote/CgsVi/src/compiler/interface2.jl:0
 [12] (::Zygote.var"#41#42"{typeof(∂(test1))})(::Float64) at /root/.julia/packages/Zygote/CgsVi/src/compiler/interface.jl:41
 [13] gradient(::Function, ::Array{Float64,1}) at /root/.julia/packages/Zygote/CgsVi/src/compiler/interface.jl:59
 [14] top-level scope at /root/codes/test_zygote/test_sparse.jl:16
in expression starting at /root/codes/test_zygote/test_sparse.jl:16

```

Then I changed `sparse` to `SparseMatrixCSC`， it works:

```julia
using LinearAlgebra
using SparseArrays
using ChainRulesCore
using Zygote

Zygote.@adjoint function SparseMatrixCSC{T,N}(arr) where {T,N}
    SparseMatrixCSC{T,N}(arr), Δ -> (collect(Δ),)
  end

function test2(a)
    A = SparseMatrixCSC(2,2,[1, 3, 5], [1, 2, 1, 2],a)
    return sum(A)
end 

a = [1.0,2.0,3.0,4.0]
gradient(test2, a)

```

output

```julia
([1.0, 1.0, 1.0, 1.0],)

```

As far as I am concerned, calling `sparse` will also create `SparseMatrixCSC`, so my questions are:

1. I know when two indices are the same in the input index array, calling `sparse` will **add** the values of duplicated entries, is this the root cause of “Mutating arrays is not supported” when calling `sparse`?
2. **How** can I directly differentiate `sparse`? Because I know calling `SparseMatrixCSC` rather than the exposed API `sparse` is not encouraged.

Thx for any reply.

---

<div class="post-metadata">

**Author:** ![sbhasan](https://avatars.discourse-cdn.com/v4/letter/s/ed655f/32.png) [@sbhasan](https://discourse.julialang.org/u/sbhasan)\
**Post date:** [June 11, 2023, 2:18pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/2 "2023-06-11T14:18:31Z")

</div>

I am facing similar problem. Did you ever manage to solve this problem?

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [June 11, 2023, 5:08pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/3 "2023-06-11T17:08:47Z")

</div>

Zygote can’t differentiate through sparse-matrix constructors AFAIK. You need to write a custom `rrule` for that part using ChainRulesCore.jl.

(Even with AD, at some point you need to learn to take derivatives yourself, at least for some pieces of your calculation.)

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [June 11, 2023, 5:37pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/4 "2023-06-11T17:37:37Z")

</div>

While ChainRulesCore is the recommended way to do things now, [custom adjoints](https://fluxml.ai/Zygote.jl/stable/adjoints/) can also be defined in Zygote directly like @Richard-Li did.

The `SparseMatricCSC` constructor is part of the API, so don’t feel bad about using it 🙂 Unfortunately, the function `sparse` mutates a lot of things before calling said constructor (see the [source](https://github.com/JuliaSparse/SparseArrays.jl/blob/2c8b8b150319f96beccd9f85cd432358fb562663/src/sparsematrix.jl#L1039-L1078)).  
So I suggest you assess if you really need `sparse` or if you can use the constructor directly. This will tell you which adjoint(s) to write. The one you made _looks_ correct at first glance but I didn’t think about it for long.

---

<div class="post-metadata">

**Author:** ![mohamed82008](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mohamed82008/32/18171_2.png) [@mohamed82008](https://discourse.julialang.org/u/mohamed82008)\
**Post date:** [June 11, 2023, 7:49pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/5 "2023-06-11T19:49:40Z")

</div>

I suggest writing an rrule for `sparse` and contributing it to ChainRules.jl. It’s doable and pretty straightforward.

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [June 13, 2023, 1:25pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/6 "2023-06-13T13:25:24Z")

</div>

> [@mohamed82008](#):
>
> I suggest writing an rrule for `sparse` and contributing it to ChainRules.jl. It’s doable and pretty straightforward.

Even if you do this, I’m worried that Zygote (or Enzyme) will still try to construct a dense matrix for the primal-tangent input to the `rrule`’s pullback?

For example, consider something as simple as the scalar-valued function f(p) = x^T A(p) y, where A(p) constructs an \ell \times m sparse matrix from some parameters p \in \mathbb{R}^n, while x \in \mathbb{R}^\ell and y \in \mathbb{R}^m are (dense) constant vectors. The partial derivatives are \frac{\partial f}{\partial p\_k} = x^T \frac{\partial A}{\partial p\_k} y = \mathrm{tr}[yx^T \frac{\partial A}{\partial p\_k}] = (xy^T) \cdot \frac{\partial A}{\partial p\_k}, where \cdot is the [Frobenius inner product](https://en.wikipedia.org/wiki/Frobenius_inner_product). In a sparse situation where \frac{\partial A}{\partial p\_k} has only O(1) nonzero entries, then \frac{\partial f}{\partial p\_k} can be computed in O(1) operations and the whole \nabla\_p f can be computed in O(n) operations with O(n) storage (like the calculation of f(p) itself) 😃 — this could easily be implemented in an `rrule` for f(p). However, if you instead define an `rrule` for A(p) (or for the `sparse` constructor), then the input tangent vector to the A(p) pullback is the rank-1 matrix xy^T, and if Zygote stores this as a dense matrix it will require O(\ell m) storage and time ☹.

Can Zygote (or Enzyme) be easily taught to store a low-rank tangent like xy^T implicitly?

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [June 13, 2023, 1:36pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/7 "2023-06-13T13:36:50Z")

</div>

What if we define an `rrule` for f(p) that calls back into the `rrule` for A(p)? Best of both worlds?

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [June 13, 2023, 1:40pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/8 "2023-06-13T13:40:07Z")

</div>

> [@gdalle](#):
>
> What if we define an `rrule` for f(p) that calls back into the `rrule` for A(p)? Best of both worlds?

You then lose the generic benefit of defining an `rrule` for the `sparse` constructor itself — you still need a manual `rrule` for any function that _uses_ sparse matrices, and to get the full benefit of reverse-mode AD you need to manually implement the chain rule connecting the sparse-matrix constructor(s) all the way to the first low-dimensional (e.g. scalar) outputs.

That is, it’s basically the same as the current situation.

---

<div class="post-metadata">

**Author:** ![mcabbott](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mcabbott/32/6603_2.png) [@mcabbott](https://discourse.julialang.org/u/mcabbott)\
**Post date:** [June 13, 2023, 2:08pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/9 "2023-06-13T14:08:25Z")

</div>

Yes. This starts to sound like the wrong level to solve the problem. Instead of writing special sparse rules for `x' * A * y` etc, to complement the sparse forward evaluation, perhaps the right level is to opt out of the existing rule for dense `x' * A * y` and instead differentiate the sparse forward implementation.

This is obviously what ForwardDiff does. It’s not impossible to make Zygote do this, although it’s going to involve a lot of indexing (which without thunks is expensive). It’s possible that Enzyme is already efficient at this, or could be made so.

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [June 13, 2023, 2:14pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/10 "2023-06-13T14:14:37Z")

</div>

> [@mcabbott](#):
>
> instead differentiate the sparse forward implementation.

I hope you’re not suggesting forward-mode AD here? That doesn’t scale to a large number of input parameters.

---

<div class="post-metadata">

**Author:** ![mohamed82008](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mohamed82008/32/18171_2.png) [@mohamed82008](https://discourse.julialang.org/u/mohamed82008)\
**Post date:** [June 13, 2023, 4:35pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/11 "2023-06-13T16:35:30Z")

</div>

> [@stevengj](#):
>
> Can Zygote (or Enzyme) be easily taught to store a low-rank tangent like xy^TxyTxy^T implicitly?

I think the answer is yes but I will need some work to back that up with a full example. Loss of sparsity and structure in the co-tangent is a problem that hasn’t received enough attention as far as I can tell, but not because it is technically impossible. Zygote passes special arrays and types as co-tangents all the time. And ChainRulesCore has a whole mechanism to project the input’s co-tangent onto the structure of the primal input. In theory, one can define a lazy low rank matrix like the following and then propagate that backward.

```julia
julia> using LazyArrays

julia> x = rand(1000);

julia> y = @~ x .* x';

julia> Base.summarysize(x)
8040

julia> Base.summarysize(y)
8096

```

Lazy arrays are not used as much in ChainRules but I think they should where possible.

---

<div class="post-metadata">

**Author:** ![mohamed82008](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mohamed82008/32/18171_2.png) [@mohamed82008](https://discourse.julialang.org/u/mohamed82008)\
**Post date:** [June 13, 2023, 5:17pm UTC](https://discourse.julialang.org/t/zygote-jl-how-to-get-the-gradient-of-sparse-matrix/59067/12 "2023-06-13T17:17:20Z")

</div>

> [@mohamed82008](#):
>
> I think the answer is yes but I will need some work to back that up with a full example.

Here is an example adapted from ChainRules.

```julia
using LazyArrays, ChainRulesCore, LinearAlgebra, Zygote

mydot(x, A, y) = dot(x, A, y)

function ChainRulesCore.rrule(::typeof(mydot), x::AbstractVector{<:Number}, A::AbstractMatrix{<:Number}, y::AbstractVector{<:Number})
    z = dot(x, A, y)
    function dot_pullback(Ω̄)
        Ay = @~ A * y
        ΔΩ = unthunk(Ω̄)
        cΔΩ = conj(ΔΩ)
        dx = @~(cΔΩ .* Ay)
        ay = adjoint(y)
        dA = @~(ΔΩ .* x .* ay)
        aA = adjoint(A)
        dy = @~(ΔΩ .* (aA * x))
        return (NoTangent(), dx, dA, dy)
    end
    dot_pullback(::ZeroTangent) = (NoTangent(), ZeroTangent(), ZeroTangent(), ZeroTangent())
    return z, dot_pullback
end

```

```julia
julia> x = rand(200); A = rand(200, 300); y = rand(300);

julia> Zygote.pullback(mydot, x, A, y)[2](1.0)[2] |> Base.summarysize
4184

julia> Base.summarysize(x) + Base.summarysize(y)
4080

```
