`@grad`/`@grad_from_chainrules` fail for rules returning structured arrays (`Diagonal`, `UpperTriangular`, ...)
Maintainers usually reply within 1 day
Nobody has claimed this yet.
Assessment
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Newbie friendliness
- 56/100
Research direction
Start with the two tracking paths in src/macros.jl named in the issue: @grad and @grad_from_chainrules. Reproduce the Diagonal and UpperTriangular examples, then trace how output tracking and forward replay handle structured arrays. Done means both examples complete without the IndexStyle assertion and return the expected gradients.
Written by the indexing model from the issue text.
Description
If a rule's primal output is an array without linear indexing, such as Diagonal or UpperTriangular, ReverseDiff fails when it wraps the output in a TrackedArray. This affects both @grad_from_chainrules and @grad.
using ReverseDiff, ChainRulesCore, LinearAlgebra
f(x) = Diagonal(x)
ChainRulesCore.rrule(::typeof(f), x) = f(x), Δ -> (NoTangent(), diag(unthunk(Δ)))
ReverseDiff.@grad_from_chainrules f(x::ReverseDiff.TrackedArray)
ReverseDiff.gradient(x -> sum(f(x)), [1.0, 2.0])
ERROR: LoadError: AssertionError: IndexStyle(value) === IndexLinear()
Stacktrace:
[1] ReverseDiff.TrackedArray{Float64, Float64, 2, Diagonal{Float64, Vector{Float64}}, Diagonal{Float64, Vector{Float64}}}(value::Diagonal{Float64, Vector{Float64}}, deriv::Diagonal{Float64, Vector{Float64}}, tape::ReverseDiff.InstructionTape)
...
[5] track(::typeof(f), x::ReverseDiff.TrackedArray{Float64, Float64, 1, Vector{Float64}, Vector{Float64}})
@ Main ~/.julia/dev/ReverseDiff/src/macros.jl:339
@grad with an UpperTriangular output fails the same way:
g(x) = UpperTriangular(x)
g(x::ReverseDiff.TrackedArray) = ReverseDiff.track(g, x)
ReverseDiff.@grad function g(x)
return UpperTriangular(ReverseDiff.value(x)), Δ -> (triu(Δ),)
end
ReverseDiff.gradient(x -> sum(g(x)), [1.0 2.0; 3.0 4.0])
# ERROR: AssertionError: IndexStyle(value) === IndexLinear()
Expected: [1.0, 1.0] and [1.0 1.0; 0.0 1.0]. Many ChainRules rules return structured matrices (Diagonal, Symmetric, triangular factors, ...), so importing them fails as soon as the output is tracked.
Cause: the macros call track(output_value, tp) on the primal output as is (@grad, @grad_from_chainrules). TrackedArray requires IndexLinear() storage, and even with #216 a Diagonal/UpperTriangular deriv couldn't store the off-structure entries the reverse pass seeds. Materializing such outputs (e.g. collect) before tracking would avoid this. The forward replay value!(output, out_value) has to do the same.
ReverseDiff master (b796032, v1.18.4), Julia 1.13.1.
- Dominant language
- Julia
- Stars
- 396
- Forks
- 61
- Avg merge
- 23h 18m
- Merged PRs (30d)
- 16
Getting set up
This project ships no dev container, Dockerfile or contributing guide, so setting up is up to you: start from its README, and see our first-contribution guide for the general steps.
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
More from JuliaDiff/ReverseDiff.jl
-
Difficulty 2/5 1-3 hours Newbie friendliness 86/100
JuliaDiff/ReverseDiff.jl#318 ·
Maintainers usually reply within 1 day
-
`@grad`/`@grad_from_chainrules` silently reuse a tangent if a pullback returns too few tangentsOpen
Difficulty 2/5 1-3 hours Newbie friendliness 83/100
JuliaDiff/ReverseDiff.jl#315 ·
Maintainers usually reply within 1 day
-
Difficulty 2/5 1-3 hours Newbie friendliness 84/100
JuliaDiff/ReverseDiff.jl#314 ·
Maintainers usually reply within 1 day
-
Difficulty 3/5 1-2 days Newbie friendliness 72/100
JuliaDiff/ReverseDiff.jl#320 ·
Maintainers usually reply within 1 day
-
Difficulty 3/5 1-2 days Newbie friendliness 68/100
JuliaDiff/ReverseDiff.jl#319 ·
Maintainers usually reply within 1 day
All issues in JuliaDiff/ReverseDiff.jl
Similar issues
-
Difficulty 2/5 1-3 hours Newbie friendliness 70/100
SciML/DiffEqNoiseProcess.jl#342 ·
-
Difficulty 1/5 Under an hour Newbie friendliness 88/100
-
Difficulty 2/5 1-3 hours Newbie friendliness 62/100
Maintainers usually reply within 1 day
-
Difficulty 2/5 1-3 hours Newbie friendliness 72/100
oxfordcontrol/COSMO.jl#211 ·
-
documentation
Difficulty 2/5 Half a day Newbie friendliness 65/100
Maintainers usually reply within 6 days