# Custom ChainRulesCore rrule with ReverseDiff

**URL:** <https://discourse.julialang.org/t/custom-chainrulescore-rrule-with-reversediff/77389>\
**Category:** Probabilistic Programming\
**Tags:** reversediff, chainrulescore\
**Created:** [March 4, 2022, 3:36am UTC](https://discourse.julialang.org/t/custom-chainrulescore-rrule-with-reversediff/77389 "2022-03-04T03:36:57Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![smharwood](https://avatars.discourse-cdn.com/v4/letter/s/bc8723/32.png) [@smharwood](https://discourse.julialang.org/u/smharwood)\
**Post date:** [March 4, 2022, 3:36am UTC](https://discourse.julialang.org/t/custom-chainrulescore-rrule-with-reversediff/77389/1 "2022-03-04T03:36:57Z")

</div>

I’m having trouble finding documentation to get ReverseDiff to use a custom adjoint defined using ChainRulesCore.rrule. MWE:

```julia
import Zygote, ReverseDiff
import ChainRulesCore
function f(x)
  x'*x
end
function ChainRulesCore.rrule(::typeof(f), x)
  y = f(x)
  function f_pullback(y_bar)
    @show "Calling custom pullback"
    return ChainRulesCore.NoTangent(), y_bar*2*x
  end
  return y, f_pullback
end

```

Running `Zygote.gradient(f, ones(3))` gives the expected `Calling custom pullback` output, while `ReverseDiff.gradient(f, ones(3))` does not.

I must be missing something; where can I find an example?

---

<div class="post-metadata">

**Author:** ![sethaxen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sethaxen/32/35604_2.png) [@sethaxen](https://discourse.julialang.org/u/sethaxen)\
**Post date:** [March 4, 2022, 12:47pm UTC](https://discourse.julialang.org/t/custom-chainrulescore-rrule-with-reversediff/77389/2 "2022-03-04T12:47:37Z")

</div>

By default, ReverseDiff ignores all `rrule`s, and one must opt into each rule one wants. Unfortunately this isn’t documented, but there’s an [open PR](https://github.com/JuliaDiff/ReverseDiff.jl/pull/196) to do so.

Here’s how to opt in:

```julia
julia> ReverseDiff.gradient(f, ones(3))
3-element Vector{Float64}:
 2.0
 2.0
 2.0

julia> ReverseDiff.@grad_from_chainrules f(x::TrackedArray)

julia> ReverseDiff.gradient(f, ones(3))
"Calling custom pullback" = "Calling custom pullback"
3-element Vector{Float64}:
 2.0
 2.0
 2.0

```

---

<div class="post-metadata">

**Author:** ![smharwood](https://avatars.discourse-cdn.com/v4/letter/s/bc8723/32.png) [@smharwood](https://discourse.julialang.org/u/smharwood)\
**Post date:** [March 4, 2022, 2:50pm UTC](https://discourse.julialang.org/t/custom-chainrulescore-rrule-with-reversediff/77389/3 "2022-03-04T14:50:43Z")

</div>

Brilliant. Many thanks!
