diff --git a/src/rulesets/SparseArrays/sparsematrix.jl b/src/rulesets/SparseArrays/sparsematrix.jl index 12ae76f80..99f0b4ba1 100644 --- a/src/rulesets/SparseArrays/sparsematrix.jl +++ b/src/rulesets/SparseArrays/sparsematrix.jl @@ -30,6 +30,7 @@ function rrule(::typeof(findnz), A::AbstractSparseMatrix) function findnz_pullback(Δ) _, _, V̄ = unthunk(Δ) + V̄ = unthunk(V̄) V̄ isa AbstractZero && return (NoTangent(), V̄) return NoTangent(), sparse(I, J, V̄, m, n) end @@ -43,6 +44,7 @@ function rrule(::typeof(findnz), v::AbstractSparseVector) function findnz_pullback(Δ) _, V̄ = unthunk(Δ) + V̄ = unthunk(V̄) V̄ isa AbstractZero && return (NoTangent(), V̄) return NoTangent(), sparsevec(I, V̄, n) end diff --git a/test/rulesets/SparseArrays/sparsematrix.jl b/test/rulesets/SparseArrays/sparsematrix.jl index ea0cf5199..2a3600619 100644 --- a/test/rulesets/SparseArrays/sparsematrix.jl +++ b/test/rulesets/SparseArrays/sparsematrix.jl @@ -77,6 +77,14 @@ end I, V = findnz(v) V̄ = rand!(similar(V)) test_rrule(findnz, v ⊢ dv, output_tangent=(zeros(length(I)), V̄)) + + # cotangent components may be thunks + IA, JA, VA = findnz(A) + V̄A = rand!(similar(VA)) + _, pb = rrule(findnz, A) + @test pb((ZeroTangent(), ZeroTangent(), @thunk(V̄A)))[2] == sparse(IA, JA, V̄A, 5, 5) + _, pb = rrule(findnz, v) + @test pb((ZeroTangent(), @thunk(V̄)))[2] == sparsevec(I, V̄, 5) end if Base.USE_GPL_LIBS # these rrules don't work without CHOLMOD from SuiteSparse.jl