# How to avoid ForwardDiff.jl generating a second-order derivative that wastes flops by eventually multiplying by zero

**URL:** <https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845>\
**Category:** Performance\
**Tags:** llvm, differentiation, forwarddiff\
**Created:** [January 23, 2021, 10:06pm UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845 "2021-01-23T22:06:02Z")\
**Posts on this page:** 11\
**Page:** 1

<div class="post-metadata">

**Author:** ![schneiderfelipe](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/schneiderfelipe/32/18004_2.png) [@schneiderfelipe](https://discourse.julialang.org/u/schneiderfelipe)\
**Post date:** [January 23, 2021, 10:06pm UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/1 "2021-01-23T22:06:02Z")

</div>

I played around with second derivatives and ForwardDiff.jl and found something unexpected, so I would like to ask about it. Basically, the generated LLVM representation is multiplying a whole code branch by zero before returning. Here’s a simple example:

```julia
julia> using ForwardDiff: derivative

julia> f(x) = x^3
f (generic function with 1 method)

julia> ∂f(x) = derivative(f, x)
∂f (generic function with 1 method)

julia> ∂²f(x) = derivative(∂f, x)
∂²f (generic function with 1 method)

julia> f(1.0), ∂f(1.0), ∂²f(1.0) # everything OK so far
(1.0, 3.0, 6.0)

```

And here are the LLVM representations for both the first- and second-order derivatives (with some annotations made by me - please tell me if I’m reading correctly):

```julia
julia> @code_llvm debuginfo=:none ∂f(1.0)

define double @"julia_\E2\88\82f_17305"(double) {
top:
  %1 = fmul double %0, %0 # %1 = x * x
  %2 = fmul double %1, 3.000000e+00 # %2 = 3 * x^2
  ret double %2 # very nice, great success
}

julia> @code_llvm debuginfo=:none ∂²f(1.0)

define double @"julia_\E2\88\82\C2\B2f_17314"(double) {
top:
  %1 = fmul double %0, %0 # %1 = x * x
  %2 = fmul double %0, 2.000000e+00 # %2 = 2 * x
  %3 = fmul double %1, 3.000000e+00 # %3 = 3 * x^2
  %4 = fmul double %2, 3.000000e+00 # %4 = 3 * 2 * x
  %5 = fmul double %3, 0.000000e+00 # %5 = 0 * 3 * x^2 # <= 🤔
  %6 = fadd double %4, %5 # %6 = 6 * x
  ret double %6 # %1, %3 and %5 are not required!
}

```

Is this expected? Can this be optimized? I think so since the above does not happen if `x` is an `Int64`:

```julia
julia> @code_llvm debuginfo=:none ∂²f(1)

define i64 @"julia_\E2\88\82\C2\B2f_17357"(i64) {
top:
  %1 = mul i64 %0, 6
  ret i64 %1
}

```

---

<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:** [January 23, 2021, 10:31pm UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/2 "2021-01-23T22:31:00Z")

</div>

Unfortunately IEEE floating points rules are not kind to making code optimiseable.  
In particular the there exists a value `x` such that `0.0 * x != 0.0`.  
Namely `x = NaN`.  
Thus the compiler is not generally allowed to simply remove that chain of expressions.

* * *

In theory ForwardDiff2.jl might handle thus better since it uses ChainRule’s types.  
ChainRules has a strong `Zero()` object which does have `Zero() * x = Zero()` for all `x`.  
I don’t know if it will show up in this circumstance.  
Also ForwardDiff2 is an early experiment and is on hiatus last I heard.

---

<div class="post-metadata">

**Author:** ![schneiderfelipe](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/schneiderfelipe/32/18004_2.png) [@schneiderfelipe](https://discourse.julialang.org/u/schneiderfelipe)\
**Post date:** [January 24, 2021, 1:59am UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/3 "2021-01-24T01:59:31Z")

</div>

Hi, @oxinabox thanks for the reply and for the suggestion. I’ll try to use ForwardDiff2.jl. Are there any other AD libraries that currently support ChainRules.jl?

---

<div class="post-metadata">

**Author:** ![patrick-kidger](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/patrick-kidger/32/20378_2.png) [@patrick-kidger](https://discourse.julialang.org/u/patrick-kidger)\
**Post date:** [January 24, 2021, 2:46am UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/4 "2021-01-24T02:46:51Z")

</div>

Strangely, I find that annotating with `@fastmath`:

```julia
julia> ∂²f(x) = @fastmath derivative(x->derivative(f, x), x)

```

doesn’t seem to produce the optimised version:

```julia
julia> @code_llvm debuginfo=:none ∂²f(1.0)

; Function Attrs: uwtable
define double @"julia_\E2\88\82\C2\B2f_475"(double) #0 {
top:
  %1 = fmul double %0, %0
  %2 = fmul double %0, 2.000000e+00
  %3 = fmul double %1, 3.000000e+00
  %4 = fmul double %2, 3.000000e+00
  %5 = fmul double %3, 0.000000e+00
  %6 = fadd double %4, %5
  ret double %6
}

```

Is this still expected?

---

<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:** [January 24, 2021, 8:41am UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/5 "2021-01-24T08:41:17Z")

</div>

> [@patrick-kidger](#):
>
> `@fastmath derivative(x->derivative(f, x), x)`

The reason this doesn’t help is that it doesn’t do anything, as the macro can’t see any operations it knows about.

But adding fastmath to the `x^3` alone was my first guess, too, and that still doesn’t improve what’s generated. This time because it gets thrown away at the first dispatch, DiffRules defines no methods for FastMath functions, and the fallback is to call the un-fast variant:

```julia
julia> @macroexpand @fastmath derivative(x->derivative(f, x), x)
:(derivative((x->begin
              #= REPL[6]:1 =#
              derivative(f, x)
          end), x))

julia> @macroexpand @fastmath x^3
:(Base.FastMath.pow_fast(x, Val{3}()))

julia> @less Base.FastMath.pow_fast(ForwardDiff.Dual(1.0, 1.0), Val{3}()) 
# @inline pow_fast(x, v::Val) = Base.literal_pow(^, x, v)

```

Edit: the obvious attempt does _not_ lead to an improvement:

```julia
julia> using ForwardDiff: Dual, partials, value, derivative

julia> import Base.FastMath: pow_fast, mul_fast

julia> function pow_fast(x::Dual{Z}, ::Val{p}) where {Z,p}
         y = pow_fast(value(x), Val(p))
         dys = map(partials(x).values) do dx
           mul_fast(p, dx, pow_fast(value(x), Val(p-1)))
         end
         Dual{Z}(y, dys)
       end;

julia> g(x) = @fastmath x^3

julia> ∂²g(x) = derivative(x -> derivative(g, x), x)

```

---

<div class="post-metadata">

**Author:** ![patrick-kidger](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/patrick-kidger/32/20378_2.png) [@patrick-kidger](https://discourse.julialang.org/u/patrick-kidger)\
**Post date:** [January 24, 2021, 11:43am UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/6 "2021-01-24T11:43:47Z")

</div>

> [@mcabbott](#):
>
> The reason this doesn’t help is that it doesn’t do anything, as the macro can’t see any operations it knows about.

Ah of course, fastmath just sees the symbols.

I’m not sure how to solve this either.

---

<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:** [January 24, 2021, 1:55pm UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/8 "2021-01-24T13:55:28Z")

</div>

> [@schneiderfelipe](#):
>
> I’ll try to use ForwardDiff2.jl.

I did not recommend that.  
Yingbo does not recommend that; because it is an experiment that is on hiatus.

That info was much more on medium term outlook, than a suggestion.

> Are there any other AD libraries that currently support ChainRules.jl?

Reverse mode: Zygote, and there is a branch of Nabla.  
Forward and Reverse: ReversePropagation.jl  
Reverse mode for some types, but not for rules: Yota.jl

but it doesn’t mean on any of this that this is solved.  
Just that there is an option to be

---

<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:** [January 24, 2021, 1:56pm UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/9 "2021-01-24T13:56:52Z")

</div>

> [@oxinabox](#):
>
> In particular the there exists a value `x` such that `0.0 * x != 0.0` .  
> Namely `x = NaN` .

It’s worse than that, because `Inf * 0.0` is also `NaN`. Because of that, the code will give an incorrect result when `3x^2` overflows:

```julia
julia> ∂²f(1e200)
NaN

```

---

<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:** [January 24, 2021, 2:02pm UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/10 "2021-01-24T14:02:07Z")

</div>

> [@mcabbott](#):
>
> DiffRules defines no methods for FastMath functions, and the fallback is to call the un-fast variant:

For interest:

ChainRules.jl does define fast math rules for FastMath functions.  
And the way the code works is really cool.  
Even if i do say so myself.  
The fast math rules are a code tranform (basically applying `@fastmath`) of the original rules

> <https://github.com/JuliaDiff/ChainRules.jl/blob/v0.7.49/src/rulesets/Base/fastmath_able.jl>

Even having those though is not enough always since `@fastmath` does not apply recursively.  
(ideally `@fastmath` would act on dynamic scope like a Cassette pass, not lexical scope)

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 25, 2021, 1:27am UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/11 "2021-01-25T01:27:07Z")

</div>

Obligatory posting of NaN+Inf computing:

[![](https://global.discourse-cdn.com/julialang/original/3X/a/3/a3c2157d9bb57b941a0db5726e0f62ec378bc4c2.jpeg "NaN Gates and Flip FLOPS") ](https://www.youtube.com/watch?v=5TFDG-y-EHs)

> [@oxinabox](#):
>
> Forward and Reverse: ReversePropagation.jl

Technically ReversePropogation.jl is only reverse mode. We could (and should) easily make it emit the forward mode. Then when SymbolicUtils.jl supports array symbols, it’ll be a fairly complete system and essentially do this with a switch to apply simplify rules.

---

<div class="post-metadata">

**Author:** ![dpsanders](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dpsanders/32/3573_2.png) [@dpsanders](https://discourse.julialang.org/u/dpsanders)\
**Post date:** [January 25, 2021, 2:28am UTC](https://discourse.julialang.org/t/how-to-avoid-forwarddiff-jl-generating-a-second-order-derivative-that-wastes-flops-by-eventually-multiplying-by-zero/53845/12 "2021-01-25T02:28:23Z")

</div>

It can already emit forward mode, it’s just a bit hidden.
