# ChainRulesCore and ForwardDiff

**URL:** https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705
**Category:** Numerics
**Tags:** question
**Created:** [May 24, 2021, 6:55am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705 "2021-05-24T06:55:28Z")
**Posts on this page:** 16
**Page:** 1

<div class="post-metadata">

### Author: ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)
#### Post date: [May 24, 2021, 6:55am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/1 "2021-05-24T06:55:28Z")

</div>

If I want to define derivatives for a function to work with ForwardDiff, is it sufficient to define a `ChainRulesCore.frule`, or is there something else?

The rule does not seem to be used, instead the original function is called with `ForwardDiff.Dual` arguments.

---

<div class="post-metadata">

### Author: ![simeonschaub](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/simeonschaub/32/216566_2.png) [@simeonschaub](https://discourse.julialang.org/u/simeonschaub)
#### Post date: [May 24, 2021, 8:57am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/2 "2021-05-24T08:57:13Z")

</div>

No, ForwardDiff does not currently use ChainRules, so you will have to add your rules as done in [ForwardDiff.jl/dual.jl at 909976d719fdbd5fec91a159c6e2d808c45a770f · JuliaDiff/ForwardDiff.jl · GitHub](https://github.com/JuliaDiff/ForwardDiff.jl/blob/909976d719fdbd5fec91a159c6e2d808c45a770f/src/dual.jl#L251).

There is [GitHub - YingboMa/ForwardDiff.jl: Forward Mode Automatic Differentiation for Julia](https://github.com/YingboMa/ForwardDiff.jl) which does use ChainRules, but I don’t think that is currently still being actively developed.

---

<div class="post-metadata">

### Author: ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)
#### Post date: [May 24, 2021, 2:15pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/3 "2021-05-24T14:15:49Z")

</div>

Thanks. I managed to adapt the docs example to use the `frule` with `ForwardDiff`.

This MWE uses bisection (to keep it simple) to solve

x^{\varepsilon\_1} + x^{\varepsilon\_2} = \theta \qquad \varepsilon\_1, \varepsilon\_2, \theta \> 0

If it sees `ForwardDiff.Dual`, it just invokes the `frule`. This works and I don’t see any obvious problems with inference, but suggestions to improve how I hook into `ForwardDiff` are welcome (the rest is really just an MWE).

```julia
using ForwardDiff, ChainRulesCore, FiniteDifferences
using ForwardDiff: value, partials, Dual

# A primitive bisection method, just to make the MWE self-contained.
function bisection(f, a::T, b::T;
                   xtol = abs(b - a) * √eps(max(a, b))) where {T <: AbstractFloat}
    fa = f(a)
    fb = f(b)
    fa * fb < 0 || error("not bracketed")
    for _ in 1:100
        m = (a + b) / 2
        fm = f(m)
        abs(a - b) ≤ xtol && return m, fm
        if fm * fa > 0
            a, fa = m, fm
        else
            b, fb = m, fm
        end
    end
    error("too many iterations")
end

bisection(f, a, b; kwargs...) = bisection(f, promote(float(a), float(b))...; kwargs...)

# analytically derived partial derivatives
∂x∂θ(ϵ1, ϵ2, θ, x) = 1 / (ϵ1 * x^(ϵ1 - 1) + ϵ2 * x^(ϵ2 - 1))
∂x∂ϵ1(ϵ1, ϵ2, θ, x) = -log(x)*x^ϵ1 * ∂x∂θ(ϵ1, ϵ2, θ, x)
∂x∂ϵ2(ϵ1, ϵ2, θ, x) = ∂x∂ϵ1(ϵ2, ϵ1, θ, x)

# solve using bisection
function _solve(ϵ1, ϵ2, θ)
    (ϵ1 + ϵ2 + θ) isa ForwardDiff.Dual && error("sanity check: wrong code path")
    (θ > 0 && ϵ1 > 0 && ϵ2 > 0) || throw(DomainError((; θ, ϵ1, ϵ2)))
    b = max(θ^(1/ϵ1), θ^(1/ϵ2))
    x, _ = bisection(x -> x^ϵ1 + x^ϵ2 - θ, 0, b)
    x
end

function ChainRulesCore.frule((Δself, Δϵ1, Δϵ2, Δθ),
                              ::typeof(_solve), ϵ1, ϵ2, θ)
    x = solve(ϵ1, ϵ2, θ)
    Δx = (∂x∂θ(ϵ1, ϵ2, θ, x) * Δθ + # this works with ForwardDiff.Partials …
          ∂x∂ϵ1(ϵ1, ϵ2, θ, x) * Δϵ1 + # … because they have + and * defined.
          ∂x∂ϵ2(ϵ1, ϵ2, θ, x) * Δϵ2)
    return x, Δx
end

# adapted from
# https://juliadiff.org/ChainRulesCore.jl/dev/autodiff/operator_overloading.html
function _solve_dual(::Type{T}, dual_args...) where {T<:Dual}
    ȧrgs = (NO_FIELDS, partials.(dual_args)...)
    args = (_solve, value.(dual_args)...)
    y, ẏ = frule(ȧrgs, args...)
    T(y, ẏ)
end

function solve(ϵ1::T1, ϵ2::T2, θ::T3) where {T1,T2,T3}
    T = promote_type(T1, T2, T3)
    if T <: Dual
        _solve_dual(T, ϵ1, ϵ2, θ)
    else
        _solve(ϵ1, ϵ2, θ)
    end
end

# rudimentary checks
f(x) = solve(0.3 + x, 0.2 + x, 1.2 + x)
d1 = central_fdm(5, 1)(f, 0.0)
d2 = ForwardDiff.derivative(f, 0.0)
@show isapprox(d1, d2; atol = 1e-4)

```

---

<div class="post-metadata">

### Author: ![oxinabox](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oxinabox/32/206603_2.png) [@oxinabox](https://discourse.julialang.org/u/oxinabox)
#### Post date: [May 24, 2021, 2:32pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/4 "2021-05-24T14:32:44Z")

</div>

Nice!.  
Do you think we could make a macro to do that.  
Like a `@import_frule foo` or something?

---

<div class="post-metadata">

### Author: ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)
#### Post date: [May 28, 2021, 2:31pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/5 "2021-05-28T14:31:03Z")

</div>

I haven’t found a clean way to _add_ an AD-method for `ForwardDiff.Dual`s, because in order to dispatch on the AD method I need a trait like the promoted type above.

Maybe some helper functions could make this easier and cleaner than a macro. Also, obtaining partials for non-duals is OK but would be wasteful for arrays and similar, where specializing on them not being duals would work better.

---

<div class="post-metadata">

### Author: ![oxinabox](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oxinabox/32/206603_2.png) [@oxinabox](https://discourse.julialang.org/u/oxinabox)
#### Post date: [May 28, 2021, 4:12pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/6 "2021-05-28T16:12:17Z")

</div>

> [@Tamas\_Papp](#):
>
> Also, obtaining partials for non-duals is OK but would be wasteful for arrays and similar, where specializing on them not being duals would work better.

Yeah, I am pretty convinced one wants to always keep the dual on the outside  
A Dual of Arrays, not an Array of Duals.  
Because it makes this easy, and also means you can use BLAS etc.

ForwardDiff2 did that

---

<div class="post-metadata">

### Author: ![cgeoga](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cgeoga/32/216186_2.png) [@cgeoga](https://discourse.julialang.org/u/cgeoga)
#### Post date: [May 28, 2021, 4:32pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/7 "2021-05-28T16:32:07Z")

</div>

If ForwardDiff2 is maybe not looking like it’s going to become a full replacement, is there anything else in the pipelines for forward-mode AD beyond ForwardDiff? I am in general a very happy use of ForwardDiff, but I sometimes have to do some hacky things that are less elegant than @Tamas_Papp’s solution here. I have a pattern I’ve ended up using a few times of writing methods like

```julia
function myfunction(x::Dual{T,V,N}) where{T,V<:AbstractFloat,N}
    # manual first derivative of my function here...
end

```

which are definitely kind of scary because I don’t really know what I’m doing and have gotten partials wrong in non-obvious ways before. I should probably start using Tamas’ method here instead, but in general I’d be curious to hear about forward-mode stuff. A lot of my problems are, say, 10 dimensional, and so the extra overhead of reverse-mode makes it less appealing.

---

<div class="post-metadata">

### Author: ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)
#### Post date: [May 29, 2021, 11:36am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/8 "2021-05-29T11:36:27Z")

</div>

> [@cgeoga](#):
>
> gotten partials wrong in non-obvious ways before

Note that my solution above is making `ForwardDiff` use _existing_ partials defined by `ChainRulesCore.frule` — in other words, you still need to get them right, but hopefully in one place only without duplicating code.

---

<div class="post-metadata">

### Author: ![oxinabox](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/oxinabox/32/206603_2.png) [@oxinabox](https://discourse.julialang.org/u/oxinabox)
#### Post date: [June 3, 2021, 5:56pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/9 "2021-06-03T17:56:38Z")

</div>

> [@cgeoga](#):
>
> If ForwardDiff2 is maybe not looking like it’s going to become a full replacement, is there anything else in the pipelines for forward-mode AD beyond ForwardDiff?

Diffractor (@keno) has both forwards and reverse mode.  
And it’s forward mode does use `ChainRulesCore.frule` natively.  
(and it’s reverse does use `rrule`)  
DIffractor has not been released yet, but I get the impression it is getting close

---

<div class="post-metadata">

### Author: ![ThummeTo](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/thummeto/32/26105_2.png) [@ThummeTo](https://discourse.julialang.org/u/ThummeTo)
#### Post date: [November 30, 2021, 3:25pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/10 "2021-11-30T15:25:04Z")

</div>

Hi all,

I am seeking for a “default procedure” to get ForwardDiff working together with already defined frules. The code example from @Tamas_Papp looks pretty well, but I need a more arbitrary way to solve this for array arguments. The use of array arguements results in `ForwardDiff.Dual` becoming `Vector{ForwardDiff.Dual{...}}` and `T = promote_type(T1, T2, T3)` results in `T=any` if I mix scalar and vector arguments.

I am playing around the whole day, but I can’t figure out a way to solve the problem for arbitrary functions, meaning for functions with mixed scalar and vector arguments.

Everyone with suggestions would make my day 🙂  
Thanks in advance!

PS: If there is a good tutorial on how to build custom AD rules for ForwardDiff for functions with multiple and/or vector arguments, it would be a nice hint, too.

---

<div class="post-metadata">

### Author: ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)
#### Post date: [February 3, 2022, 9:43am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/11 "2022-02-03T09:43:22Z")

</div>

> [@ThummeTo](#):
>
> The use of array arguements results in `ForwardDiff.Dual` becoming `Vector{ForwardDiff.Dual{...}}` and `T = promote_type(T1, T2, T3)` results in `T=any` if I mix scalar and vector arguments.

This is a rather late reply, hope it is still useful. For something similar to the code above to work, you need to extract the `eltype` of array arguments and promote those.

---

<div class="post-metadata">

### Author: ![ThummeTo](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/thummeto/32/26105_2.png) [@ThummeTo](https://discourse.julialang.org/u/ThummeTo)
#### Post date: [February 4, 2022, 6:47am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/12 "2022-02-04T06:47:57Z")

</div>

Thanks for the reply.  
Yes, we (prototypically) solved it in a similar way (cheching the array type and converting every element). As soon as I spend more time on this again, I will try to build and post a generic solution.

---

<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: [May 24, 2022, 1:08pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/13 "2022-05-24T13:08:07Z")

</div>

Just stumbled upon this thread, and I seem to recall that @mohamed82008 had an answer in the works

---

<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: [May 28, 2022, 3:10pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/14 "2022-05-28T15:10:37Z")

</div>

Check [NonconvexUtils.jl/runtests.jl at main · JuliaNonconvex/NonconvexUtils.jl · GitHub](https://github.com/JuliaNonconvex/NonconvexUtils.jl/blob/main/test/runtests.jl#L281)

---

<div class="post-metadata">

### Author: ![ThummeTo](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/thummeto/32/26105_2.png) [@ThummeTo](https://discourse.julialang.org/u/ThummeTo)
#### Post date: [October 7, 2022, 10:19am UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/15 "2022-10-07T10:19:12Z")

</div>

This works excellent!  
Should be added to ForwardDiff, there is already an open issue adressing this: [Automatic ChainRules compatibility · Issue #579 · JuliaDiff/ForwardDiff.jl · GitHub](https://github.com/JuliaDiff/ForwardDiff.jl/issues/579)

---

<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: [October 21, 2022, 4:14pm UTC](https://discourse.julialang.org/t/chainrulescore-and-forwarddiff/61705/16 "2022-10-21T16:14:20Z")

</div>

Update: The macro from NonconvexUtils was moved to the light weight package [GitHub - ThummeTo/ForwardDiffChainRules.jl](https://github.com/ThummeTo/ForwardDiffChainRules.jl) that’s now also registered.
