# What are \`autodiff\_deferred\` and \`autodiff\_thunk\` for in Enzyme?

**URL:** <https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065>\
**Category:** Machine Learning\
**Tags:** autodiff, enzyme, gradient\
**Created:** [March 25, 2024, 7:26am UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065 "2024-03-25T07:26:46Z")\
**Posts on this page:** 20\
**Page:** 1

<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:** [March 25, 2024, 7:26am UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/1 "2024-03-25T07:26:46Z")

</div>

I think I understand the basic `autodiff` function (see [this post](https://discourse.julialang.org/t/do-i-understand-enzyme-properly/97760)) but it has three variants:

- [`autodiff_deferred`](https://enzymead.github.io/Enzyme.jl/stable/api/#EnzymeCore.autodiff_deferred-Union%7BTuple%7BReturnPrimal%7D,%20Tuple%7BA%7D,%20Tuple%7BFA%7D,%20Tuple%7BReverseMode%7BReturnPrimal%7D,%20FA,%20Type%7BA%7D,%20Vararg%7BAny%7D%7D%7D%20where%20%7BFA%3C:Annotation,%20A%3C:Annotation,%20ReturnPrimal%7D)
- [`autodiff_thunk`](https://enzymead.github.io/Enzyme.jl/stable/api/#EnzymeCore.autodiff_thunk-Union%7BTuple%7BRABI%7D,%20Tuple%7BModifiedBetweenT%7D,%20Tuple%7BWidth%7D,%20Tuple%7BReturnShadow%7D,%20Tuple%7BReturnPrimal%7D,%20Tuple%7BA%7D,%20Tuple%7BFA%7D,%20Tuple%7BEnzymeCore.ReverseModeSplit%7BReturnPrimal,%20ReturnShadow,%20Width,%20ModifiedBetweenT,%20RABI%7D,%20Type%7BFA%7D,%20Type%7BA%7D,%20Vararg%7BAny%7D%7D%7D%20where%20%7BFA%3C:Annotation,%20A%3C:Annotation,%20ReturnPrimal,%20ReturnShadow,%20Width,%20ModifiedBetweenT,%20RABI%3C:EnzymeCore.ABI%7D) (which seems [broken on 1.11](https://github.com/EnzymeAD/Enzyme.jl/issues/1358))
- [`autodiff_deferred_thunk`](https://enzymead.github.io/Enzyme.jl/stable/api/#EnzymeCore.autodiff_deferred_thunk-Union%7BTuple%7BRABI%7D,%20Tuple%7BModifiedBetweenT%7D,%20Tuple%7BWidth%7D,%20Tuple%7BReturnShadow%7D,%20Tuple%7BReturnPrimal%7D,%20Tuple%7BA%7D,%20Tuple%7BFA%7D,%20Tuple%7BEnzymeCore.ReverseModeSplit%7BReturnPrimal,%20ReturnShadow,%20Width,%20ModifiedBetweenT,%20RABI%7D,%20Type%7BFA%7D,%20Type%7BA%7D,%20Vararg%7BAny%7D%7D%7D%20where%20%7BFA%3C:Annotation,%20A%3C:Annotation,%20ReturnPrimal,%20ReturnShadow,%20Width,%20ModifiedBetweenT,%20RABI%3C:EnzymeCore.ABI%7D)

Is there someone who can explain to me:

- why “deferring” computations is useful for GPU or higher-order?
- in which cases I should use the “thunk” version?
- whether there are performance differences?

---

<div class="post-metadata">

**Author:** ![vchuravy](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/vchuravy/32/8_2.png) [@vchuravy](https://discourse.julialang.org/u/vchuravy)\
**Post date:** [March 25, 2024, 12:16pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/2 "2024-03-25T12:16:20Z")

</div>

The `autodiff_defered` variants are for higher order or GPU use cases. They defer compilation until we actually call things, this allows us to detect them when doing compilation for the GPU or when seeing a call to within a compilation request for `autodiff`.

Long-term I think `autodiff_deferred` should just be the normal interface.

The `thunk` variants don’t immediately call the augmented primal or reverse function, but instead return callable objects (thunks) that you can use to perform the computation.

---

<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:** [March 25, 2024, 12:57pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/3 "2024-03-25T12:57:56Z")

</div>

Okay thanks! Could you help me understand how to make it work with a vector output? The example in the docs is a scalar output, and I don’t understand how I’m supposed to pass the cotangent:

```julia
julia> using Enzyme

julia> g(x) = abs2.(x)
g (generic function with 1 method)

julia> x = [2.0, 3.0];

julia> dx = zero(x);

julia> dy = [5.0, 7.0];

julia> forw, rev = autodiff_thunk(ReverseSplitWithPrimal, Const{typeof(g)}, Duplicated, Duplicated{typeof(x)})
(Enzyme.Compiler.AugmentedForwardThunk{Ptr{Nothing}, Const{typeof(g)}, Duplicated{Vector{Float64}}, Tuple{Duplicated{Vector{Float64}}}, Val{1}, Val{true}(), @NamedTuple{1, 2, 3, 4, 5::Bool, 6::Bool, 7::Core.LLVMPtr{Float64, 0}, 8::Core.LLVMPtr{Float64, 0}, 9::Core.LLVMPtr{Float64, 0}, 10::Core.LLVMPtr{Float64, 0}, 11::Core.LLVMPtr{Float64, 0}, 12::Core.LLVMPtr{Float64, 0}, 13::Core.LLVMPtr{Float64, 0}, 14::Core.LLVMPtr{Float64, 0}, 15::Core.LLVMPtr{Float64, 0}}}(Ptr{Nothing} @0x00007871c8880450), Enzyme.Compiler.AdjointThunk{Ptr{Nothing}, Const{typeof(g)}, Duplicated{Vector{Float64}}, Tuple{Duplicated{Vector{Float64}}}, Val{1}, @NamedTuple{1, 2, 3, 4, 5::Bool, 6::Bool, 7::Core.LLVMPtr{Float64, 0}, 8::Core.LLVMPtr{Float64, 0}, 9::Core.LLVMPtr{Float64, 0}, 10::Core.LLVMPtr{Float64, 0}, 11::Core.LLVMPtr{Float64, 0}, 12::Core.LLVMPtr{Float64, 0}, 13::Core.LLVMPtr{Float64, 0}, 14::Core.LLVMPtr{Float64, 0}, 15::Core.LLVMPtr{Float64, 0}}}(Ptr{Nothing} @0x00007871c8880a40))

julia> tape, y, shadow_y = forw(Const(g), Duplicated(x, dx))
(var"1" = @NamedTuple{1, 2, 3, 4, 5::Bool, 6::Bool, 7::Core.LLVMPtr{Float64, 0}, 8::Core.LLVMPtr{Float64, 0}, 9::Core.LLVMPtr{Float64, 0}, 10::Core.LLVMPtr{Float64, 0}, 11::Core.LLVMPtr{Float64, 0}, 12::Core.LLVMPtr{Float64, 0}, 13::Core.LLVMPtr{Float64, 0}, 14::Core.LLVMPtr{Float64, 0}, 15::Core.LLVMPtr{Float64, 0}}(([0.0, 0.0], [4.0, 9.0], nothing, nothing, false, false, Core.LLVMPtr{Float64, 0}(0x000078715e006610), Core.LLVMPtr{Float64, 0}(0x7ffffffffffffffe), Core.LLVMPtr{Float64, 0}(0x000078715e0065e0), Core.LLVMPtr{Float64, 0}(0x000078715e181730), Core.LLVMPtr{Float64, 0}(0x000078725cb10311), Core.LLVMPtr{Float64, 0}(0x000078724ca32f60), Core.LLVMPtr{Float64, 0}(0x000078725ca6bee0), Core.LLVMPtr{Float64, 0}(0x000078724ca32dc0), Core.LLVMPtr{Float64, 0}(0x0000000009900c00))), var"2" = [4.0, 9.0], var"3" = [0.0, 0.0])

julia> y
2-element Vector{Float64}:
 4.0
 9.0

julia> rev(Const(g), Duplicated(x, dx), Duplicated(y, dy), tape)
ERROR: AssertionError: length(argtypes) + needs_tape == length(argexprs)
Stacktrace:
     ⋮ internal @ Enzyme.Compiler, GPUCompiler, Core, Unknown
 [6] (::Enzyme.Compiler.AdjointThunk{Ptr{…}, Const{…}, Duplicated{…}, Tuple{…}, Val{…}, @NamedTuple{…}})(::Const{typeof(g)}, ::Duplicated{Vector{…}}, ::Vararg{Any})
   @ Enzyme.Compiler ~/.julia/packages/Enzyme/l4FS0/src/compiler.jl:5004
Use `err` to retrieve the full stack trace.
Some type information was truncated. Use `show(err)` to see complete types.

```

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 25, 2024, 1:20pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/4 "2024-03-25T13:20:09Z")

</div>

You should be passing the same args to the primal as the shadow, plus the shadow return if it is active, plus the tape.

In other words you should not pass y and dy.

As for Julia 1.11 it is not presently supported in Enzyme

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 25, 2024, 1:20pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/5 "2024-03-25T13:20:46Z")

</div>

As for your example set shadow\_y to your desired cotangent

---

<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:** [March 25, 2024, 2:44pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/6 "2024-03-25T14:44:17Z")

</div>

> [@wsmoses](#):
>
> In other words you should not pass y and dy.

I don’t understand, where do I plug in `dy` in order to compute the VJP \partial g(x)^\top (\mathrm{d}y)?

---

<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:** [March 25, 2024, 2:51pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/7 "2024-03-25T14:51:31Z")

</div>

> [@wsmoses](#):
>
> As for your example set shadow\_y to your desired cotangent

`shadow_y` is an output, not an input, right? I never plug it back in

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 25, 2024, 3:37pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/8 "2024-03-25T15:37:41Z")

</div>

If it is needed to compute the reverse pass it will be captured (by reference) in the tape

---

<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:** [March 25, 2024, 3:51pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/9 "2024-03-25T15:51:59Z")

</div>

Okay so where do I plug it in? Can you show me the code on this simple example? I’m really struggling to guess here ^^

---

<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:** [March 25, 2024, 4:49pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/10 "2024-03-25T16:49:20Z")

</div>

I got it, I need to do

```julia
shadow_y .= dy

```

---

<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:** [March 25, 2024, 4:57pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/11 "2024-03-25T16:57:32Z")

</div>

But what happens if `shadow_y` cannot be mutated in place?

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 25, 2024, 5:51pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/12 "2024-03-25T17:51:27Z")

</div>

So in reverse mode Reverse Mode should only use Duplicated for mutable memory. For immutable data you should use Active for the return (and then you pass that into the reverse pass function).

---

<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:** [March 25, 2024, 6:49pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/13 "2024-03-25T18:49:09Z")

</div>

What about immutable data that has mutable fields, like `Tuple{Vector,Vector}`?

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 25, 2024, 6:56pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/14 "2024-03-25T18:56:15Z")

</div>

You can do `shadow_xy0] .= dy` and `shadow_y[1] .= dy` on the outside just the original vector case. This would still use Duplicated (sorry I should have been more specific to say that if derivative data is immutable on the top level register you should use active. In this case the inner differentiable data is in a mutable construct so you should still use Duplicated (in this case a vector, even if behind a tuple)).

---

<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:** [March 26, 2024, 1:40pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/15 "2024-03-26T13:40:16Z")

</div>

I’m tempted to use the `autodiff_thunk` option everywhere and put the part

```julia
forw, rev = autodiff_thunk(ReverseSplitWithPrimal, Const{typeof(g)}, Duplicated, Duplicated{typeof(x)})

```

in a preparation step, so that it only runs once when I want to compute several gradients in a row.  
Is that reasonable? Does `autodiff` call `autodiff_thunk` like that under the hood? Or is there a conceptual difference?

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 26, 2024, 4:05pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/16 "2024-03-26T16:05:44Z")

</div>

No they are not equivalent, at least from a performance standpoint. The reason is that if you know you’re computing the forward and reverse pass together in one go, Enzyme can do a lot more performance optimizations.

Within autodiff there is a compilation cache so asking for the same function and activities multiple times won’t incur additional cost

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [March 26, 2024, 4:06pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/17 "2024-03-26T16:06:44Z")

</div>

We could also make a thunk for combined mode (currently thunk only supports split mode), if desired.

---

<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:** [March 26, 2024, 4:47pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/18 "2024-03-26T16:47:46Z")

</div>

Okay, I thought the thunk was specific to split mode, it is helpful to know that it could exist in combined mode too.

Does the output of `autodiff_thunk` (before even running the forward pass) save some time if we reuse it for different inputs, or do you think it would be negligible?  
Because for different inputs we cannot reuse the forward pass of course.

---

<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:** [April 5, 2024, 6:04pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/19 "2024-04-05T18:04:21Z")

</div>

Hey @wsmoses sorry for pinging again, I was just curious about the purpose of the `thunk` mechanism.  
One obvious application would be computing a Jacobian, where you only do one forward sweep and then as many pullbacks as there are output dimensions. But if I call the `thunk`-ed pullback more than once, will the answer still be correct? Or would one reverse sweep somehow alter the tape, so that the pullback closure is no longer working? I’m asking specifically because that’s the situation in Tapir.jl, where each reverse sweep must be directly preceded by a forward sweep

See also [Split reverse mode for Tapir · Issue #115 · withbayes/Tapir.jl · GitHub](https://github.com/withbayes/Tapir.jl/issues/115)

---

<div class="post-metadata">

**Author:** ![wsmoses](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wsmoses/32/26497_2.png) [@wsmoses](https://discourse.julialang.org/u/wsmoses)\
**Post date:** [April 5, 2024, 7:02pm UTC](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065/20 "2024-04-05T19:02:17Z")

</div>

The thunks are stateless and can be called as many times as you like.

However, the extra “tape” (really value cache) for the reverse pass may be different, depending on the function (e.g. if a shadow pointer is captured/overwritten/etc).

If certain properties hold (or you swap out all uses of the capured shadow with your new shadow) you can do that.

Of course use enzyme’s batchduplicated/etc to perform as many reverse passes in one reverse sweep as you want (thunk or no thunk).

[Next page](https://discourse.julialang.org/t/what-are-autodiff-deferred-and-autodiff-thunk-for-in-enzyme/112065.md?page=2)
