# Which direction: DifferentiatonInterface, Enzyme, Zygote with CUDA and FFTs?

**URL:** <https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225>\
**Category:** General Usage\
**Tags:** ad\
**Created:** [September 9, 2025, 10:22am UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225 "2025-09-09T10:22:32Z")\
**Posts on this page:** 20\
**Page:** 1

<div class="post-metadata">

**Author:** ![roflmaostc](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roflmaostc/32/30123_2.png) [@roflmaostc](https://discourse.julialang.org/u/roflmaostc)\
**Post date:** [September 9, 2025, 10:22am UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/1 "2025-09-09T10:22:32Z")

</div>

Hi!

For the third time, I am trying to implement some wave optics algorithms in Julia but each time I start, I suffer a bit from the current automatic differentiation status.

In principle, my functions look like:

```julia
function propagate(field::AbstractArray{T, 6}, distance::NumberOrArray, 
                   wavelengths::NumberOrArray, pixel_size)
    H = # some kernel calculated with distances and wavelengths

    field_p = # some processing on the field with wavelengths and distance

    field_prop = ifft(fft(field_p) .* H) # roughly
    return field_prop
end

```

The adjoint is a bit nasty to write manually since `distance` or `wavelengths` can be both scalars or vectors. The output will be 6 dimensional (`(x,y,z, wavelengths, polarization, batch)`), but some dimensions are singleton depending on the types of `distance` and `wavelengths`.

Anyway, it seems like Enzyme.jl is not supporting FFTs at the moment (latest is [this issue](https://github.com/SciML/NonlinearSolve.jl/issues/597) or older [this](https://github.com/EnzymeAD/Enzyme.jl/issues/1717)). So in my [current code](https://github.com/JuliaPhysics/WaveOpticsPropagation.jl/blob/9aefad71a1d5d5c645d5432a8a30a1840e3bc28c/src/angular_spectrum.jl) I use a mix of Zygote and custom written ChainRules.jl.

I’d love to use DifferentiationInterface.jl but if FFTs do not work, that’s gonna be a hard time.

Has anyone an idea what I should do? I feel like using Zygote.jl is not quite future proof as everyone else seem to switch to Enzyme.jl.

Honestly, I seriously consider to implement it in another language because this pain has been ongoing for years (I started hitting those general issues in 2020). I know there is progress on this front and things have been improved but they haven’t reached my feature demands yet (CUDA + FFT + complex numbers). And implementing an Enzyme rule seems just very hard for me. I would love to help out a bit but right now I want to focus on my real research problems.

I see similar packages implemented in PyTorch and JAX and it works well for them, so I am really a bit hopeless and cannot recommend Julia to them right now.

Best,

Felix

---

<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:** [September 9, 2025, 10:57am UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/2 "2025-09-09T10:57:52Z")

</div>

Hi Felix!

> [@roflmaostc](#):
>
> I’d love to use [DifferentiationInterface.jl](https://juliaregistries.github.io/General/packages/redirect_to_repo/DifferentiationInterface) but if FFTs do not work, that’s gonna be a hard time.

I just want to point out that the question of using DifferentiationInterface is pretty much orthogonal to the choice of backend. At least if you fit inside the restrictions of DI (a single active argument), you can choose between Enzyme, Zygote or anything else you fancy. Whether FFTs work or not depends on the existence of rules for a given backend, not on the DI infrastructure, see [this docs page](https://juliadiff.org/DifferentiationInterface.jl/DifferentiationInterface/stable/faq/differentiability/) for details.

In fact, using DI can actually make your code more resilient and allow easy switches between autodiff systems down the road. One of the main motivations for developing DI was to encourage this kind of seamless transition, e.g. from Zygote to Enzyme. Of course it comes with a caveat that DI induces some overhead and/or bugs which can be avoided with a backend’s native API. If you run into something like that, please open an issue so we can work to fix it!

> [@roflmaostc](#):
>
> Has anyone an idea what I should do? I feel like using [Zygote.jl](https://juliaregistries.github.io/General/packages/redirect_to_repo/Zygote) is not quite future proof as everyone else seem to switch to [Enzyme.jl](https://juliaregistries.github.io/General/packages/redirect_to_repo/Enzyme).

Have you tried implementing the Enzyme FFT rule yourself? Maybe it’s not out of reach?

Alternately, have you tried Mooncake.jl? Even if it doesn’t have FFT rules either, its [rule system](https://chalk-lab.github.io/Mooncake.jl/stable/understanding_mooncake/rule_system/) is less complex than that of Enzyme, so it may be an easier starting point.  
EDIT: I don’t think Mooncake [fully supports CUDA](https://github.com/chalk-lab/Mooncake.jl/issues/648) at the moment, so it doesn’t answer your full question. Sorry.

> [@roflmaostc](#):
>
> Honestly, I seriously consider to implement it in another language because this pain has been ongoing for years (I started hitting those general issues in 2020).

> [@roflmaostc](#):
>
> I see similar packages implemented in PyTorch and JAX and it works well for them, so I am really a bit hopeless and cannot recommend Julia to them right now.

Honestly, that’s fair. If it ain’t broke, don’t fix it. It would be great to get this in Julia too, but if you need it for your research and Python has it, no one will blame you ^^

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [September 9, 2025, 12:58pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/3 "2025-09-09T12:58:51Z")

</div>

Reactant.jl supports FFT (and its differentiation), used quite extensively in NeuralOperators.jl (it also runs on CUDA, TPU, CPU and whatever other accelerator you might want).

Is your `field_p` a 6-dimensional tensor? (we added upto 3-dims for FFTs in reactant though it is easy to expand if that is what is needed)

---

<div class="post-metadata">

**Author:** ![roflmaostc](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roflmaostc/32/30123_2.png) [@roflmaostc](https://discourse.julialang.org/u/roflmaostc)\
**Post date:** [September 9, 2025, 1:01pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/4 "2025-09-09T13:01:24Z")

</div>

> [@avikpal](#):
>
> [Reactant.jl](https://juliaregistries.github.io/General/packages/redirect_to_repo/Reactant) supports FFT (and its differentiation), used quite extensively in [NeuralOperators.jl](https://juliaregistries.github.io/General/packages/redirect_to_repo/NeuralOperators) (it also runs on CUDA, TPU, CPU and whatever other accelerator you might want).

But it’s not supported by Enzyme? What would I to do to use it?

> [@avikpal](#):
>
> Is your `field_p` a 6-dimensional tensor? (we added upto 3-dims for FFTs in reactant though it is easy to expand if that is what is needed)

Yes, it’s a 6dim array. FFTs are only done along 1st and 2nd dimension though.

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [September 9, 2025, 1:02pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/5 "2025-09-09T13:02:42Z")

</div>

> [@roflmaostc](#):
>
> But it’s not supported by Enzyme? What would I to do to use it?

Reactant.jl uses Enzyme.jl for autodiff support. So you just need to `@compile Enzyme.gradient(....)`.

> [@roflmaostc](#):
>
> Yes, it’s a 6dim array. FFTs are only done long 1st and 2nd dimension though.

Then it should just work

---

<div class="post-metadata">

**Author:** ![roflmaostc](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roflmaostc/32/30123_2.png) [@roflmaostc](https://discourse.julialang.org/u/roflmaostc)\
**Post date:** [September 9, 2025, 1:05pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/6 "2025-09-09T13:05:00Z")

</div>

Is there any further link? Didn’t see anything specific in the docs.

So reactant introduces some FFT rules but they are not included in Enzyme itself?

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [September 9, 2025, 1:34pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/7 "2025-09-09T13:34:05Z")

</div>

> [@roflmaostc](#):
>
> Is there any further link? Didn’t see anything specific in the docs.

[Compiling Lux Models using Reactant.jl | Lux.jl Docs](https://lux.csail.mit.edu/stable/manual/compiling_lux_models) has an example of compiling enzyme.gradient. I just realized we don’t show Enzyme examples in Reactant (will add something today)

> [@roflmaostc](#):
>
> So reactant introduces some FFT rules but they are not included in Enzyme itself?

So Reactant uses EnzymeMLIR, in-contract to EnzymeLLVM both interfaced via Enzyme.jl (so no change is required on user side except moving data via Reactant.to\_rarray and doing a @compile). We have the FFT rules defined for EnzymeMLIR.

---

<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:** [September 9, 2025, 2:06pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/8 "2025-09-09T14:06:48Z")

</div>

Not quite correct, but an analogy is that Reactant is similar to jax.jit, and Enzyme is similar to jax.gradient. The difference here is that if you Reactant.@compile a function, it converts your julia code into a nice tensor program that can be automatically optimized/fused/etc and executed on whatever backend you want, including CPU, GPU, TPU, and distributed versions thereof. Essentially all programs in this tensor IR are differentiable by Enzyme(MLIR), so if it reactant.compile’s, you’re good!

As for Enzyme.jl without Reactant compilation, there’s no technical reason for it – I think someone just needs to write the rule and/or use Enzyme.import\_rrule to import the chainrule like described in the issue you linked above. Maybe you, @ptiede or someone are interested in getting that over the line?

That said I’d recommend just using Reactant.@compile, with Enzyme on the inside. It’ll give you all the things you say are missing, and much more 🙂 .

---

<div class="post-metadata">

**Author:** ![giordano](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/giordano/32/2166_2.png) [@giordano](https://discourse.julialang.org/u/giordano)\
**Post date:** [September 9, 2025, 2:10pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/9 "2025-09-09T14:10:14Z")

</div>

> [@Autodifferentiation with FFT, with Enzyme?](https://discourse.julialang.org/t/autodifferentiation-with-fft-with-enzyme/128674/2):
>
> It would be great to have some Enzyme rules in AbstractFFTs.jl (there’s a draft PR at [Add EnzymeRules by sethaxen · Pull Request #103 · JuliaMath/AbstractFFTs.jl · GitHub](https://github.com/JuliaMath/AbstractFFTs.jl/pull/103), but someone has to adopt the project and get it over the finish line). In the meantime, maybe you’ll find the following Enzyme reverse rule useful, extracted from a private project I’m working on. It’s not a complete solution (no forward rules, no batched rules, no rule for mul!), but I think it’s correct for the cases it cov…

---

<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:** [September 9, 2025, 2:11pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/10 "2025-09-09T14:11:26Z")

</div>

Wait so when you `Reactant.@compile` with `Enzyme.gradient` on the _inside_, Enzyme nonetheless differentiates the Reactant-generated code?

---

<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:** [September 9, 2025, 2:19pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/11 "2025-09-09T14:19:10Z")

</div>

yes

```julia
julia> using Reactant, Enzyme

julia> function foo(x)
           return sum(x)
       end

julia> x = Reactant.to_rarray(ones(10));

julia> @code_hlo optimize=false Enzyme.gradient(Reverse, foo, x)
module @reactant_gradient attributes {mhlo.num_partitions = 1 : i64, mhlo.num_replicas = 1 : i64} {
  func.func private @identity_broadcast_scalar(%arg0: tensor<f64>) -> tensor<f64> {
    return %arg0 : tensor<f64>
  }
  func.func private @"Const{typeof(foo)}(Main.foo)_autodiff"(%arg0: tensor<10xf64>) -> (tensor<f64>, tensor<10xf64>) {
    %0 = stablehlo.transpose %arg0, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %cst = stablehlo.constant dense<0.000000e+00> : tensor<f64>
    %1 = stablehlo.convert %cst : tensor<f64>
    %2 = enzyme.batch @identity_broadcast_scalar(%0) {batch_shape = array<i64: 10>} : (tensor<10xf64>) -> tensor<10xf64>
    %cst_0 = stablehlo.constant dense<0.000000e+00> : tensor<f64>
    %3 = stablehlo.convert %cst_0 : tensor<f64>
    %cst_1 = stablehlo.constant dense<0.000000e+00> : tensor<f64>
    %4 = stablehlo.convert %cst_1 : tensor<f64>
    %5 = stablehlo.reduce(%2 init: %1) applies stablehlo.add across dimensions = [0] : (tensor<10xf64>, tensor<f64>) -> tensor<f64>
    %6 = stablehlo.transpose %2, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    return %5, %6 : tensor<f64>, tensor<10xf64>
  }
  func.func @main(%arg0: tensor<10xf64> {tf.aliasing_output = 1 : i32}) -> (tensor<10xf64>, tensor<10xf64>) {
    %0 = stablehlo.transpose %arg0, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %cst = stablehlo.constant dense<0.000000e+00> : tensor<f64>
    %1 = stablehlo.convert %cst : tensor<f64>
    %cst_0 = stablehlo.constant dense<0.000000e+00> : tensor<10xf64>
    %cst_1 = stablehlo.constant dense<0.000000e+00> : tensor<f64>
    %2 = stablehlo.convert %cst_1 : tensor<f64>
    %3 = stablehlo.broadcast_in_dim %2, dims = [] : (tensor<f64>) -> tensor<10xf64>
    %cst_2 = stablehlo.constant dense<1.000000e+00> : tensor<f64>
    %4 = stablehlo.transpose %0, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %5 = stablehlo.transpose %3, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %6:2 = enzyme.autodiff @"Const{typeof(foo)}(Main.foo)_autodiff"(%4, %cst_2, %5) {activity = [#enzyme<activity enzyme_active>], ret_activity = [#enzyme<activity enzyme_activenoneed>, #enzyme<activity enzyme_active>]} : (tensor<10xf64>, tensor<f64>, tensor<10xf64>) -> (tensor<10xf64>, tensor<10xf64>)
    %7 = stablehlo.transpose %6#0, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %8 = stablehlo.transpose %6#1, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %9 = stablehlo.transpose %8, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    %10 = stablehlo.transpose %7, dims = [0] : (tensor<10xf64>) -> tensor<10xf64>
    return %9, %10 : tensor<10xf64>, tensor<10xf64>
  }
}

julia> @code_hlo Enzyme.gradient(Reverse, foo, x)
module @reactant_gradient attributes {mhlo.num_partitions = 1 : i64, mhlo.num_replicas = 1 : i64} {
  func.func @main(%arg0: tensor<10xf64>) -> tensor<10xf64> {
    %cst = stablehlo.constant dense<1.000000e+00> : tensor<10xf64>
    return %cst : tensor<10xf64>
  }
}

```

---

<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:** [September 9, 2025, 2:33pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/12 "2025-09-09T14:33:28Z")

</div>

> [@roflmaostc](#):
>
> But it’s not supported by Enzyme? What would I to do to use it?

Generally using Reactant compilation makes thing more likely to be efficiently differentiated by Enzyme. It removes type instabilities, preserves structure in IR, and lots of other good things that make differentiation (and also just the performance of the original code) better.

---

<div class="post-metadata">

**Author:** ![jtravs](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jtravs/32/2010_2.png) [@jtravs](https://discourse.julialang.org/u/jtravs)\
**Post date:** [September 9, 2025, 3:26pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/13 "2025-09-09T15:26:47Z")

</div>

The code snippets pulled together in this discourse thread:

> [@Autodifferentiation with FFT, with Enzyme?](https://discourse.julialang.org/t/autodifferentiation-with-fft-with-enzyme/128674/12):
>
> As I mentioned just now in the [Autodifferentiation with FFT and Enzyme? · Issue #597 · SciML/NonlinearSolve.jl · GitHub](https://github.com/SciML/NonlinearSolve.jl/issues/597) this actually works well in Julia 1.10, and I am happily using Enzyme with FFTW and it is efficient. Thanks again @danielwe The issues appears to be on Julia 1.11 (I tested this a few weeks back, so that might have changed). Having a general set of rules for Enzyme and FFTW somewhere would be nice though to make this a bit easier, and the above seems to work. For reference, I…

work fine on Julia 1.10. So far I cannot get them to work on 1.11 or 1.12. I also haven’t tested on GPU.

---

<div class="post-metadata">

**Author:** ![avikpal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/avikpal/32/6550_2.png) [@avikpal](https://discourse.julialang.org/u/avikpal)\
**Post date:** [September 9, 2025, 3:57pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/14 "2025-09-09T15:57:50Z")

</div>

A more detailed autodiff tutorial in Reactant [Automatic Differentiation | Reactant.jl](https://enzymead.github.io/Reactant.jl/dev/tutorials/automatic-differentiation)

---

<div class="post-metadata">

**Author:** ![ptiede](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ptiede/32/17579_2.png) [@ptiede](https://discourse.julialang.org/u/ptiede)\
**Post date:** [September 12, 2025, 11:58am UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/15 "2025-09-12T11:58:09Z")

</div>

I have some private code where I added Enzyme support for FFTs (non-reactant version). I didn’t consider it production ready but I can share it clean it up and we could make a PR.

---

<div class="post-metadata">

**Author:** ![roflmaostc](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roflmaostc/32/30123_2.png) [@roflmaostc](https://discourse.julialang.org/u/roflmaostc)\
**Post date:** [September 13, 2025, 8:54am UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/16 "2025-09-13T08:54:21Z")

</div>

Does it only work with `fft(x)` but not with FFT plans `p = plan_fft(x); p * x`?

```julia-auto
# ╔═╡ 3eadd964-9079-11f0-2f63-f7227db0c8f8
using Enzyme, Reactant, ImageShow, ImageIO, FFTW

# ╔═╡ 45f5d151-e496-4526-9333-5ee68d38c496
arr = rand(Float32, (100, 100));

# ╔═╡ 3d59315e-cc6a-4d19-aeb2-06cdb7372f8f
p = plan_fft(arr)

# ╔═╡ 597cddb8-cceb-4af7-afae-1f0958ce3d55
loss_function(x) = sum(abs2.((p * x)))

# ╔═╡ d11dba54-6d60-4a90-b783-774c14d8c1c6
x = Reactant.to_rarray(arr);

# Compute gradient using reverse mode

# ╔═╡ d06c79f7-502c-47c5-ae0a-21f48fa8601b
function f(x)
	return @jit Enzyme.gradient(Reverse, loss_function, x)
end

# ╔═╡ d43b834f-0a96-4caa-a9f2-82abc8f546c0
f_compiled = @compile f(x)

```

> **Error**
>
> Scalar indexing is disallowed.
> 
> Invocation of getindex(::TracedRArray, ::Union{Int, TracedRNumber{Int}}) resulted in scalar indexing of a GPU array.
> 
> This is typically caused by calling an iterating implementation of a method.
> 
> Such implementations _do not_ execute on the GPU, but very slowly on the CPU,
> 
> and therefore should be avoided.
> 
> If you want to allow scalar iteration, use `allowscalar` or `@allowscalar`
> 
> to enable scalar iteration globally or for the operations in question.
> 
> Stack trace
> 
> Here is what happened, the most recent locations are first:
> 
> 1. **error** 
> 
> from _error.jl:35_
> 
> 1. **(::Nothing)** (none::typeof(error), none::String)
> 
> from **Reactant** →
> 
> - **ErrorException** 
> 
> from _boot.jl:323_
> 
> - **error** 
> 
> from _error.jl:35_
> 
> - **call\_with\_reactant** (::Reactant.MustThrowError, ::typeof(error), ::String)
> 
> from **Reactant** → _utils.jl_
> 
> - **errorscalar** 
> 
> from _GPUArraysCore.jl:151_
> 
> - **(::Nothing)** (none::typeof(GPUArraysCore.errorscalar), none::String)
> 
> from **Reactant** →
> 
> - **string** 
> 
> from _substring.jl:236_
> 
> - **scalardesc** 
> 
> from _GPUArraysCore.jl:134_
> 
> - **errorscalar** 
> 
> from _GPUArraysCore.jl:150_
> 
> - **call\_with\_reactant** (::Reactant.MustThrowError, ::typeof(GPUArraysCore.errorscalar), ::String)
> 
> from **Reactant** → _utils.jl_
> 
> - **\_assertscalar** 
> 
> from _GPUArraysCore.jl:124_
> 
> - **(::Nothing)** (none::typeof(GPUArraysCore.\_assertscalar), none::String, none::GPUArraysCore.ScalarIndexing)
> 
> from **Reactant** →
> 
> - **\_assertscalar** 
> 
> from _GPUArraysCore.jl:123_
> 
> - **call\_with\_reactant** (::typeof(GPUArraysCore.\_assertscalar), ::String, ::GPUArraysCore.ScalarIndexing)
> 
> from **Reactant** → _utils.jl_
> 
> - **assertscalar** 
> 
> from _GPUArraysCore.jl:112_
> 
> - **(::Nothing)** (none::typeof(GPUArraysCore.assertscalar), none::String)
> 
> from **Reactant** →
> 
> - **current\_task** 
> 
> from _task.jl:152_
> 
> - **task\_local\_storage** 
> 
> from _task.jl:280_
> 
> - **assertscalar** 
> 
> from _GPUArraysCore.jl:97_
> 
> - **call\_with\_reactant** (::typeof(GPUArraysCore.assertscalar), ::String)
> 
> from **Reactant** → _utils.jl_
> 
> - **getindex** 
> 
> from _TracedRArray.jl:120_
> 
> - **opaque closure** (none::typeof(getindex), none::Reactant.TracedRArray{…}, none::Int64) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** →
> 
> - **getindex** 
> 
> from _TracedRArray.jl:120_
> 
> - **call\_with\_reactant** (::typeof(getindex), ::Reactant.TracedRArray{…}, ::Int64) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _utils.jl_
> 
> - **copyto\_unaliased!** 
> 
> from _abstractarray.jl:1081_
> 
> - **copyto!** 
> 
> from _abstractarray.jl:1061_
> 
> - **opaque closure** (none::typeof(copyto!), none::Matrix{…}, none::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** →
> 
> - **getproperty** 
> 
> from _Base.jl:49_
> 
> - **size** 
> 
> from _TracedRArray.jl:489_
> 
> - **length** 
> 
> from _abstractarray.jl:315_
> 
> - **isempty** 
> 
> from _abstractarray.jl:1212_
> 
> - **copyto!** 
> 
> from _abstractarray.jl:1055_
> 
> - **call\_with\_reactant** (::typeof(copyto!), ::Matrix{…}, ::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _utils.jl_
> 
> - **circcopy!** 
> 
> from _multidimensional.jl:1303_
> 
> - **copy1** 
> 
> from _definitions.jl:54_
> 
> - \*\*\*\*\* 
> 
> from _definitions.jl:224_
> 
> - **opaque closure** (none::typeof(\*), none::FFTW.cFFTWPlan{…}, none::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** →
> 
> - **getproperty** 
> 
> from _Base.jl:49_
> 
> - **size** 
> 
> from _TracedRArray.jl:489_
> 
> - **axes** 
> 
> from _abstractarray.jl:98_
> 
> - **copy1** 
> 
> from _definitions.jl:53_
> 
> - \*\*\*\*\* 
> 
> from _definitions.jl:224_
> 
> - **call\_with\_reactant** (::typeof(\*), ::FFTW.cFFTWPlan{…}, ::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _utils.jl_
> 
> - **loss\_function** 
> 
> from [Other cell: _line 1_](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#597cddb8-cceb-4af7-afae-1f0958ce3d55)
> 
> [
> 
> ```julia-auto
> loss_function(x) = sum(abs2.((p * x)))
> 
> ```
> 
> ]([http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#597cddb8-cceb-4af7-afae-1f0958ce3d55](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#597cddb8-cceb-4af7-afae-1f0958ce3d55))
> 
> - **opaque closure** (none::typeof(loss\_function), none::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** →
> 
> - **loss\_function** 
> 
> from [Other cell: _line 1_](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#597cddb8-cceb-4af7-afae-1f0958ce3d55)
> 
> [
> 
> ```julia-auto
> loss_function(x) = sum(abs2.((p * x)))
> 
> ```
> 
> ]([http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#597cddb8-cceb-4af7-afae-1f0958ce3d55](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#597cddb8-cceb-4af7-afae-1f0958ce3d55))
> 
> - **call\_with\_reactant** (::typeof(loss\_function), ::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _utils.jl_
> 
> - **#make\_mlir\_fn#6** (f::typeof(loss\_function), args::Tuple{…}, kwargs::Tuple{}, name::String, concretein::Bool; toscalar::Bool, return\_dialect::Symbol, args\_in\_result::Symbol, construct\_function\_without\_args::Bool, do\_transpose::Bool, input\_shardings::Nothing, output\_shardings::Nothing, runtime::Nothing, verify\_arg\_names::Nothing, argprefix::Symbol, resprefix::Symbol, resargprefix::Symbol, num\_replicas::Int64, optimize\_then\_pad::Bool) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _TracedUtils.jl:330_
> 
> - **make\_mlir\_fn** 
> 
> from _TracedUtils.jl:260_
> 
> - **overload\_autodiff** (::EnzymeCore.ReverseMode{…}, f::EnzymeCore.Const{…}, ::Type{…}, args::EnzymeCore.Duplicated{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _Enzyme.jl:303_
> 
> - **#autodiff** (rmode::EnzymeCore.ReverseMode{…}, f::EnzymeCore.Const{…}, rt::Type{…}, args::EnzymeCore.Duplicated{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _Overlay.jl:21_
> 
> - **autodiff** 
> 
> from _Enzyme.jl:538_
> 
> - **macro expansion** 
> 
> from _sugar.jl:324_
> 
> - **gradient** 
> 
> from _sugar.jl:262_
> 
> - **f** 
> 
> from [Other cell: _line 2509_](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#d06c79f7-502c-47c5-ae0a-21f48fa8601b)
> 
> - **opaque closure** (none::typeof(f), none::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** →
> 
> - **GenericMemory** 
> 
> from _boot.jl:516_
> 
> - **IdDict** 
> 
> from _iddict.jl:31_
> 
> - **IdDict** 
> 
> from _iddict.jl:49_
> 
> - **make\_zero** 
> 
> from _EnzymeCore.jl:587_
> 
> - **macro expansion** 
> 
> from _sugar.jl:321_
> 
> - **gradient** 
> 
> from _sugar.jl:262_
> 
> - **f** 
> 
> from [Other cell: _line 2509_](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#d06c79f7-502c-47c5-ae0a-21f48fa8601b)
> 
> - **call\_with\_reactant** (::typeof(f), ::Reactant.TracedRArray{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _utils.jl_
> 
> - **#make\_mlir\_fn#6** (f::typeof(f), args::Tuple{…}, kwargs::@NamedTuple{}, name::String, concretein::Bool; toscalar::Bool, return\_dialect::Symbol, args\_in\_result::Symbol, construct\_function\_without\_args::Bool, do\_transpose::Bool, input\_shardings::Nothing, output\_shardings::Nothing, runtime::Val{…}, verify\_arg\_names::Nothing, argprefix::Symbol, resprefix::Symbol, resargprefix::Symbol, num\_replicas::Int64, optimize\_then\_pad::Bool) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _TracedUtils.jl:330_
> 
> - **#compile\_mlir!#15** (mod::Reactant.MLIR.IR.Module, f::Function, args::Tuple{…}, compile\_options::Reactant.CompileOptions, callcache::Dict{…}, sdycache::Dict{…}; fn\_kwargs::@NamedTuple{}, backend::String, runtime::Val{…}, legalize\_stablehlo\_to\_mhlo::Bool, kwargs::@Kwargs{}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _Compiler.jl:1544_
> 
> - **compile\_mlir!** 
> 
> from _Compiler.jl:1511_
> 
> - **#compile\_xla#58** (f::Function, args::Tuple{…}; before\_xla\_optimizations::Bool, client::Nothing, serializable::Bool, kwargs::@Kwargs{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _Compiler.jl:3420_
> 
> - **compile\_xla** 
> 
> from _Compiler.jl:3393_
> 
> - **#compile#59** (f::Function, args::Tuple{…}; kwargs::@Kwargs{…}) […show types…](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#)
> 
> from **Reactant** → _Compiler.jl:3492_
> 
> - **compile** 
> 
> from _Compiler.jl:3489_
> 
> - **macro expansion** 
> 
> from _Compiler.jl:2573_
> 
> - from [This cell: _line 1_](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#d43b834f-0a96-4caa-a9f2-82abc8f546c0)
> 
> [
> 
> ```julia-auto
> f_compiled = @compile f(x)
> 
> ```
> 
> ]([http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#d43b834f-0a96-4caa-a9f2-82abc8f546c0](http://localhost:1234/edit?id=3eadd8bc-9079-11f0-2f23-a30cd90fb73b#d43b834f-0a96-4caa-a9f2-82abc8f546c0))

---

<div class="post-metadata">

**Author:** ![jgreener64](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jgreener64/32/2483_2.png) [@jgreener64](https://discourse.julialang.org/u/jgreener64)\
**Post date:** [September 15, 2025, 10:05pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/17 "2025-09-15T22:05:31Z")

</div>

I think an Enzyme rule that works with FFT plans is:

```julia
using Enzyme, FFTW

grad_safe_fft!(arr, fft_plan) = fft_plan * arr

function EnzymeRules.augmented_primal(config, ::Const{typeof(grad_safe_fft!)}, t,
                                      arr, fft_plan)
    fft_plan.val * arr.val
    return EnzymeRules.AugmentedReturn(nothing, nothing, nothing)
end

function EnzymeRules.reverse(config, ::Const{typeof(grad_safe_fft!)}, dret, tape,
                             arr, fft_plan)
    arr.dval .= bfft(arr.dval)
    return (nothing, nothing)
end

```

This is based on the [FFT rrule in AbstractFFTs.jl](https://github.com/JuliaMath/AbstractFFTs.jl/blob/04a14f5a3491697d9b5a6e17ffa994cf82ce2ab0/ext/AbstractFFTsChainRulesCoreExt.jl), the other variants are available there.

---

<div class="post-metadata">

**Author:** ![roflmaostc](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roflmaostc/32/30123_2.png) [@roflmaostc](https://discourse.julialang.org/u/roflmaostc)\
**Post date:** [September 16, 2025, 6:52am UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/18 "2025-09-16T06:52:29Z")

</div>

> [@jgreener64](#):
>
> ```julia-auto
> function EnzymeRules.reverse(config, ::Const{typeof(grad_safe_fft!)}, dret, tape,
> arr, fft_plan)
> arr.dval .= bfft(arr.dval)
> return (nothing, nothing)
> end
> 
> ```

Is the `bfft` correct here? Why don’t you use `inv(fft_plan)`?

---

<div class="post-metadata">

**Author:** ![jgreener64](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jgreener64/32/2483_2.png) [@jgreener64](https://discourse.julialang.org/u/jgreener64)\
**Post date:** [September 16, 2025, 5:22pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/19 "2025-09-16T17:22:17Z")

</div>

`bfft` matches [the rrule](https://github.com/JuliaMath/AbstractFFTs.jl/blob/04a14f5a3491697d9b5a6e17ffa994cf82ce2ab0/ext/AbstractFFTsChainRulesCoreExt.jl#L15) and passes tests for my use case. See some discussion on [the original PR](https://github.com/JuliaMath/AbstractFFTs.jl/pull/58). `inv(fft_plan)` may also work.

---

<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:** [September 16, 2025, 5:26pm UTC](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225/20 "2025-09-16T17:26:28Z")

</div>

> [@roflmaostc](#):
>
> Is the `bfft` correct here? Why don’t you use `inv(fft_plan)`?

`bfft` is the correct adjoint operator (it is literally the adjoint, the conjugate-transpose, of `fft`). If you used `inv(fft_plan)` or `ifft` then it would introduce an additional 1/n scale factor that you don’t want.

[Next page](https://discourse.julialang.org/t/which-direction-differentiatoninterface-enzyme-zygote-with-cuda-and-ffts/132225.md?page=2)
