# "static" autodiff

**URL:** <https://discourse.julialang.org/t/static-autodiff/94392>\
**Category:** New to Julia\
**Created:** [February 10, 2023, 6:29am UTC](https://discourse.julialang.org/t/static-autodiff/94392 "2023-02-10T06:29:02Z")\
**Posts on this page:** 14\
**Page:** 1

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 6:29am UTC](https://discourse.julialang.org/t/static-autodiff/94392/1 "2023-02-10T06:29:02Z")

</div>

Hey,

Is there a way to have an autodiff system such as forwarddiff “write out” the dual version of a piece of code for further manual optimisation ?

The function to derivate is not that complicated, but derivating it by hand would be a pain. On the other hand, I am pretty sure the strict running of dual numbers through it is suboptimal as the function value and it’s gradients might all share pieces of code, and I would like to be able to optimize it by hand.

---

<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:** [February 10, 2023, 9:12am UTC](https://discourse.julialang.org/t/static-autodiff/94392/2 "2023-02-10T09:12:51Z")

</div>

Hey there!  
The logic behind the ChainRules.jl package sounds vaguely related: it allows you to write rules that fuse primal and gradient computations to save memory and CPU. But from what I understand you want to get the rule automatically? I don’t know what your function looks like, but if it’s not too loopy, maybe Symbolics.jl can help?

---

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 9:22am UTC](https://discourse.julialang.org/t/static-autodiff/94392/3 "2023-02-10T09:22:39Z")

</div>

Hahaha you nailed it. It **is** really loopy, and Symbolics.jl cannot handle it.

What I want is to exploit chain rules already included into AD systems to **write code** computing the gradient. Then, by hand, I could curate the code and reason about it. Give me a sec I’ll write an exemple for you.

---

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 9:30am UTC](https://discourse.julialang.org/t/static-autodiff/94392/4 "2023-02-10T09:30:14Z")

</div>

@gdalle This might give you a clearer view of what i am looking for. Consider the following function:

```julia
function f(x,t)
    r = zero(eltype(x))
    n = length(x)
    for i in eachindex(x)
        r += exp(x[i]*t)
    end
    r = log(r)
    return r
end

x = rand(10)
t = 3.0
f(x,t)

```

Say i want tthe gradient of f w.r.t. x:

```julia
using ForwardDiff
g(x,t) = ForwardDiff.gradient(y -> f(y,t), x)

g(x,t) # works correctly; 

```

What I want is an automatic mechanisme that will output the following code :

```julia
# Wanted output : 
function ∂f_∂x(x,t)
    r = zero(eltype(x))
    r_dual = zero(r)
    n = length(x)
    for i in eachindex(x)
        r += exp(x[i]*t)
        r_dual += something(x,t)
    end
    r = log(r)
    r_dual = something_else(r)
    return r, r_dual
end

```

So that i can then curate by hand the result. My bet is that, for very loopy code, there will be opportunities for hand curating and improving efficiency of the gradient computation.

Even if the obtained code is cluttered, de-cluttering by hand would still be worth it in my application case.

---

<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:** [February 10, 2023, 9:36am UTC](https://discourse.julialang.org/t/static-autodiff/94392/5 "2023-02-10T09:36:46Z")

</div>

Right, that’s what I thought, and for your use case I probably don’t know the answer. Re-interpreting the code in that way is basically what most autodiff libraries do under the hood, but I’m not aware of any that lets you retrieve the gradient computation explicitly

---

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 9:40am UTC](https://discourse.julialang.org/t/static-autodiff/94392/6 "2023-02-10T09:40:21Z")

</div>

Yes, this is what they do but they don’t output this code, they compile it and run it (which is fine). I somehow want to interrupt this pipeline in the middle to take a look myself. I am talking about forwarddiff because i think that forward mode will be easier to reason about for me: the goal is not to produce an _efficient_ code automatically, but to produce _as little cluttered_ code as possible, since I will then benchmark it and modify it by hand to make it _efficient_.

Edit : if it does not exist and I / someone ends up building it, a funny name would be `ManualDiff.jl`😉

---

<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:** [February 10, 2023, 12:15pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/7 "2023-02-10T12:15:01Z")

</div>

The hard bit is loops.  
You basically can’t generate loops from a tracking approach (including ForwardDiff.jl, and Jax)  
Because it records the operation that run. Not the code that ran it.

So you get statically unrolled list.  
Which is fine if you don’t have dynamic control flow

There used to be code for doing this with Zygote.  
But Mike gave up in the end.

I think Tapenade can do this for C++ and Fortran.

---

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 12:43pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/8 "2023-02-10T12:43:41Z")

</div>

Thanks @oxinabox for taking the time.

I did not know about Tapenade, this stuff is beautifull.

Maybe for the loop problem, a good way of doing it would be feed the “tool” a function without the loop :

```julia
function f(x,t)
    r = zero(eltype(x))
    n = length(x)
    # for i in eachindex(x)
        r += exp(x[i]*t)
    # end
    r = log(r)
    return r
end

```

have it produce :

```julia
function ∂f_∂x(x,t)
    r = zero(eltype(x))
    r_dual = zero(r)
    n = length(x)
    # for i in eachindex(x)
        r += exp(x[i]*t)
        r_dual += something(x,t)
    # end
    r = log(r)
    r_dual = something_else(r)
    return r, r_dual
end

```

and then re-add the loop manually. So loops might not be my biggest problem if I can get back code from the AD system.

Maybe there are edge cases where this approaches wont work (cant think of one right now), but in my case that would be enough.

---

<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:** [February 10, 2023, 12:58pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/9 "2023-02-10T12:58:16Z")

</div>

I don’t understand what you mean here.

In general using functional constructs like `map` and `fold` are better.

---

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 1:10pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/10 "2023-02-10T13:10:27Z")

</div>

Well, please allow me to try to reformulate. My goal is to obtain the code. If the loops are a problem, I can simply remove the loops by comminting them out, even if this completely changes the meaning of the function.

If the emplacement of the (now removed) loops starts and stops are still noted in the code, i could re-add them by hand after while curating the obtained result.

---

<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:** [February 10, 2023, 1:31pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/11 "2023-02-10T13:31:13Z")

</div>

It’s not the loops themselves.  
Its how that interacts with tracing  
Tracing doesn’t record the code, only the operations that are run.  
So if the operations that are run change depending on the values of the input then you have a problem since your trace will be meaning less.

If on the other hand you do a source code transformation AD, like Zygote or Tapenade you can do this.  
But it is **much** harder to write.  
though it can be done.

I did just remember about @dfdx 's [XGrad.jl](https://dfdx.github.io/XGrad.jl/stable/tutorial.html) which does a lot like what you want.  
and does work via source code transformation

---

<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:** [February 10, 2023, 3:20pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/12 "2023-02-10T15:20:14Z")

</div>

> [@lrnv](#):
>
> I am pretty sure the strict running of dual numbers through it is suboptimal as the function value and it’s gradients might all share pieces of code, and I would like to be able to optimize it by hand.

Yes, this is why people use reverse mode for gradients. When you have a single output and many (n) inputs, forward mode costs roughly n times the cost of computing the output, whereas reverse mode costs roughly O(1) times.

But by the same token, the forward-mode derivative calculation is not a good starting point for deriving the reverse mode / “adjoint” calculation. (See also our [matrix calculus](https://github.com/mitmath/matrixcalc) course notes.)

---

<div class="post-metadata">

**Author:** ![lrnv](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lrnv/32/19373_2.png) [@lrnv](https://discourse.julialang.org/u/lrnv)\
**Post date:** [February 10, 2023, 3:34pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/13 "2023-02-10T15:34:17Z")

</div>

@stevengj Yes what you are saying makes a lot of sense.

Anyway, I could probably do something correct using pen & paper in a few days work, so if implementing something to do it for me is really hard as the discussion with @oxinabox suggests, maybe it is not worth it.

I’ll still check Xgrad.jl as it seems really cool !

---

<div class="post-metadata">

**Author:** ![dfdx](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dfdx/32/120_2.png) [@dfdx](https://discourse.julialang.org/u/dfdx)\
**Post date:** [February 10, 2023, 11:06pm UTC](https://discourse.julialang.org/t/static-autodiff/94392/14 "2023-02-10T23:06:13Z")

</div>

If you like XGrad.jl, you probably want to take a look at Yota.jl - same ideas, but evolved. None of them support loops with non-static number of iterations though.
