# Efficient automatic differentation for Julia version \`jax.scan\`?

**URL:** <https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853>\
**Category:** Performance\
**Tags:** ad, jax\
**Created:** [October 3, 2025, 2:38pm UTC](https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853 "2025-10-03T14:38:26Z")\
**Posts on this page:** 5\
**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:** [October 3, 2025, 2:38pm UTC](https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853/1 "2025-10-03T14:38:26Z")

</div>

Hi,

I am using in some JAX array code `jax.lax.scan` which implements something like this in my example:

```julia-auto
function jax_lax_scan(f, x; accumulator_init)
    acc = accumulator_init 
    for i in axes(x, 1)
           acc = f(acc, selectdim(x, 1, i))
     end 
     return acc
end

julia> A = rand(2,2,3,1)
2×2×3×1 Array{Float64, 4}:
[:, :, 1, 1] =
 0.983886 0.301163
 0.142536 0.898853

[:, :, 2, 1] =
 0.502673 0.460876
 0.118309 0.675352

[:, :, 3, 1] =
 0.142753 0.386542
 0.858919 0.898126

julia> jax_lax_scan((acc, x) -> acc .+ x.^2, A; accumulator_init=zeros((2,3,1)))
2×3×1 Array{Float64, 3}:
[:, :, 1] =
 0.466399 0.598419 0.670064
 0.638399 0.307632 1.4969

```

What’s the Julia equivalent here? And is this efficiently supported in any automatic differentation package (such as in [JAX](https://github.com/jax-ml/jax/discussions/3850#discussioncomment-44785))?

Yes, I could differentiate this with Zygote but afaik it would keep copies of each iteration in memory. My examples are arrays with `(100, 512, 512, 100)` so impossible to store.

I guess there is some Enzyme.jl + Reactant.jl way of doing that?

Thanks,

Felix

---

<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:** [October 3, 2025, 2:52pm UTC](https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853/2 "2025-10-03T14:52:13Z")

</div>

Ok apparently:

```julia-auto
julia> f = (acc, x) -> acc .+ x.^2
 
julia> result = foldl(f, eachslice(A, dims=1); init=zeros((2,3,1)))
2×3×1 Array{Float64, 3}:
[:, :, 1] =
 0.988349 0.266677 0.758121
 0.898635 0.668507 0.956045

```

Does e.g. Enzyme efficiently with `foldl`, `reduce` or `mapreduce`?

---

<div class="post-metadata">

**Author:** ![ForceBru](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/forcebru/32/21389_2.png) [@ForceBru](https://discourse.julialang.org/u/ForceBru)\
**Post date:** [October 3, 2025, 4:00pm UTC](https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853/3 "2025-10-03T16:00:00Z")

</div>

I’m interested in this as well. Do Julia’s autodiff packages have specific performant rules for `scan`-like operations? Or perhaps such a rule wouldn’t increase performance?

---

<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:** [October 3, 2025, 4:08pm UTC](https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853/4 "2025-10-03T16:08:21Z")

</div>

If you add a `@trace` before the for loop and use Reactant, it will use some really nice AD tricks for differentiating the loop (and handles more broad set of possibilities compared to lax.scan).

---

<div class="post-metadata">

**Author:** ![yolhan\_mannes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yolhan_mannes/32/220485_2.png) [@yolhan\_mannes](https://discourse.julialang.org/u/yolhan_mannes)\
**Post date:** [October 3, 2025, 4:14pm UTC](https://discourse.julialang.org/t/efficient-automatic-differentation-for-julia-version-jax-scan/132853/5 "2025-10-03T16:14:03Z")

</div>

quick answer is yes

```julia-auto
julia> using Enzyme

julia> function jax_lax_scan(f, x, accumulator_init)
           acc = accumulator_init
               for i in axes(x, 1)
                  acc = f(acc, selectdim(x, 1, i))
            end
            return acc
       end
jax_lax_scan (generic function with 1 method)

julia> A = rand(2,2,3,1);

julia> accumulator_init=zeros((2,3,1));

julia> f(acc, x) = acc .+ x.^2;

julia> Enzyme.jacobian(Enzyme.set_runtime_activity(Reverse),jax_lax_scan,Const(f),A,Const(accumulator_init))
(nothing, [1.6257113308613 0.0 0.0; 0.0 0.0 0.0;;;; 1.080987298281143 0.0 0.0; 0.0 0.0 0.0;;;;; 0.0 0.0 0.0; 0.9727596195934325 0.0 0.0;;;; 0.0 0.0 0.0; 0.2399608038078349 0.0 0.0;;;;;; 0.0 0.41852879546292643 0.0; 0.0 0.0 0.0;;;; 0.0 0.8543467156284372 0.0; 0.0 0.0 0.0;;;;; 0.0 0.0 0.0; 0.0 0.9046578942453445 0.0;;;; 0.0 0.0 0.0; 0.0 0.5680239161728668 0.0;;;;;; 0.0 0.0 1.2498992746934297; 0.0 0.0 0.0;;;; 0.0 0.0 0.7830888594771388; 0.0 0.0 0.0;;;;; 0.0 0.0 0.0; 0.0 0.0 1.5680193443731774;;;; 0.0 0.0 0.0; 0.0 0.0 1.8554432985549147;;;;;;;], nothing)

```

however, if you want to have a really efficient one, you may need Reactant.jl indeed, for instance it will get rid of all the buffers you use in the loop, if you stay with so little cases, then its fine

```julia-auto
julia> @btime Enzyme.jacobian(Enzyme.set_runtime_activity(Reverse),jax_lax_scan,Const($f),$A,Const($accumulator_init))
  9.500 μs (216 allocations: 15.39 KiB)

```

the other way is to make a better function from julia side

```julia-auto
function jax_lax_scan2(f!, x, accumulator_init)
    acc = copy(accumulator_init)
    for i in axes(x, 1)
        f!(acc, selectdim(x, 1, i))
    end
    return acc
end
function f!(acc, x)  
    acc .+= x.^2
    return nothing
end
g2 = Enzyme.jacobian(Enzyme.set_runtime_activity(Reverse),jax_lax_scan2,Const(f!),A,accumulator_init)[2]

```

leading to

```julia-auto
@btime Enzyme.jacobian(Enzyme.set_runtime_activity(Reverse),jax_lax_scan2,Const($f!),$A,$accumulator_init)[2]

```

`6.560 μs (123 allocations: 13.86 KiB)`
