# Help With Automatic Differentiation

**URL:** <https://discourse.julialang.org/t/help-with-automatic-differentiation/100844>\
**Category:** General Usage\
**Tags:** autodiff\
**Created:** [June 26, 2023, 1:59pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844 "2023-06-26T13:59:14Z")\
**Posts on this page:** 9\
**Page:** 1

<div class="post-metadata">

**Author:** ![Devetak](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/devetak/32/50611_2.png) [@Devetak](https://discourse.julialang.org/u/Devetak)\
**Post date:** [June 26, 2023, 1:59pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/1 "2023-06-26T13:59:14Z")

</div>

Good afternoon,

I wanted to ask some advice about Automatic Differentiation in Julia. I am aware of the fact that in AD mutating arrays are not supported, but I find it difficult to write my code in a non mutating fashion, therefore I wanted to ask if somebody has any help/advice on how to do that. For example:

```julia
N = 10
Y = rand(N)
W = rand(N)
A = rand(N)
for i in 1:N
    p = Y[i] * W[i]
    if p < 0.3
        A[i] = A[i] - p
    else
        A[i] = 0
    end
end

```

my naive approach was to:

```julia
newA = zeros(N)
p = Y .* W
pIdx = findall(p .< 0.3)
newA[pIdx] = A[pIdx] - p[pIdx]
notIdx = findall(p .>= 0.3)
newA[notIdx] = 0
A = newA

```

which of course does not work. Does anybody have any suggestions/advice or resources on the matter?

One of the sources I encountered suggest defining your own pullback for the mutating parts. Would anybody know where to start in these regard?

Thank you in advance!

---

<div class="post-metadata">

**Author:** ![JordiBolibar](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jordibolibar/32/24307_2.png) [@JordiBolibar](https://discourse.julialang.org/u/JordiBolibar)\
**Post date:** [June 26, 2023, 2:11pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/2 "2023-06-26T14:11:28Z")

</div>

This can perhaps be useful to you: [GitHub - rakeshvar/Zygote-Mutating-Arrays-WorkAround.jl: A tutorial on how to work around ‘Mutating arrays is not supported’ error while performing automatic differentiation (AD) using the Julia package Zygote.](https://github.com/rakeshvar/Zygote-Mutating-Arrays-WorkAround.jl)

Bear in mind that the mutation limitation is intrinsic to Zygote.jl. You could try using Enzyme.jl (not sure if it will work), or ReverseDiff.jl.

---

<div class="post-metadata">

**Author:** ![Devetak](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/devetak/32/50611_2.png) [@Devetak](https://discourse.julialang.org/u/Devetak)\
**Post date:** [June 26, 2023, 2:28pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/3 "2023-06-26T14:28:42Z")

</div>

Thank you. This was the resource I spoke about in the post! Say that I wanted to differentiate the above to the respect of Y, would then the appropiate thing to do be to make the for loop in a separa function and then set the derivative to be -W[i] if p\<0.3 and 0 otherwise?(when differentiating A?)

---

<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:** [June 26, 2023, 3:23pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/4 "2023-06-26T15:23:07Z")

</div>

> [@Devetak](#):
>
> One of the sources I encountered suggest defining your own pullback for the mutating parts. Would anybody know where to start in these regard?

The place to start is the (very thorough) documentation of ChainRules.jl, on which Zygote.jl relies

> **[GitHub - JuliaDiff/ChainRules.jl: forward and reverse mode automatic...](https://github.com/JuliaDiff/ChainRules.jl)**
>
> forward and reverse mode automatic differentiation primitives for Julia Base + StdLibs - GitHub - JuliaDiff/ChainRules.jl: forward and reverse mode automatic differentiation primitives for Julia Ba...

---

<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:** [June 26, 2023, 3:24pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/5 "2023-06-26T15:24:57Z")

</div>

> [@Devetak](#):
>
> I find it difficult to write my code in a non mutating fashion, therefore I wanted to ask if somebody has any help/advice on how to do that.

If your code relies on mutation for efficiency, a custom chain rule is what you want anyway

---

<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:** [June 26, 2023, 4:55pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/6 "2023-06-26T16:55:43Z")

</div>

Enzyme works fine on this code:

```julia
wmoses@beast:~/git/Enzyme.jl ((HEAD detached from 0a742b9)) $ ./julia-1.9.0/bin/julia --project test.jl 
(dY, dW, dA) = ([0.0, 0.0, -0.36481667164909504, 0.0, -0.40760801938330093, 0.0, -0.29238901363265246, -0.7601463773026405, -0.10528419303498171, -0.21089297714276434], [0.0, 0.0, -0.17407141133367543, 0.0, -0.17535508982666037, 0.0, -0.5764836125183987, -0.12627017591563605, -0.13579863783844937, -0.1190684947298164], [0.0, 0.0, 1.0, 0.0, 1.0, 0.0, 1.0, 1.0, 1.0, 1.0])
wmoses@beast:~/git/Enzyme.jl ((HEAD detached from 0a742b9)) $ cat test.jl 
using Enzyme

function f(N, Y, W, A)
    for i in 1:N
        p = Y[i] * W[i]
        if p < 0.3
            A[i] = A[i] - p
        else
            A[i] = 0
        end
    end
end

N = 10
Y = rand(N)
dY = zeros(N)
W = rand(N)
dW = zeros(N)

A = rand(N)
# Derivative we backpropagate
dA = ones(N)

Enzyme.autodiff(Reverse, f, Const(N), Duplicated(Y, dY), Duplicated(W, dW), Duplicated(A, dA))
@show dY, dW, dA

```

---

<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:** [June 26, 2023, 6:12pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/7 "2023-06-26T18:12:36Z")

</div>

Comprehensions (or `map`) and broadcasting are ways of avoiding mutation. E.g.

```julia
julia> function orig!(A, Y, W) 
         for i in eachindex(A)
           p = Y[i] * W[i]
           if p < 0.3
               A[i] = A[i] - p
           else
               A[i] = 0
           end
         end
         A
       end;

julia> nonmut(A, Y, W) = map(eachindex(A)) do i
           p = Y[i] * W[i]
           (p < 0.3) * (A[i] - p) # this makes consistent type
       end;

julia> nonmut(A, Y, W) ≈ orig!(copy(A), Y, W)
true

julia> noindex(A, Y, W) = @. ((Y * W) < 0.3) * (A - Y * W); # Zygote also dislikes indexing

julia> noindex(A, Y, W) ≈ orig!(copy(A), Y, W)
true

julia> using Zygote, BenchmarkTools

julia> @btime gradient(a -> sum(abs2, nonmut(a,$Y,$W)), $A)
  min 1.625 μs, mean 2.413 μs (51 allocations, 11.69 KiB)
([-0.4549349972052379, 1.471142478710321, 0.927932031171403, 0.0, 1.055719617078303, 0.3149281759965271, 0.07592684856551069, 0.0, 0.0, 1.3759356933258697],)

julia> @btime gradient(a -> sum(abs2, noindex(a,$Y,$W)), $A) # 1/4 the memory, probably saves more time at larger N
  min 1.271 μs, mean 1.415 μs (36 allocations, 3.16 KiB)
([-0.4549349972052379, 1.471142478710321, 0.927932031171403, -0.0, 1.055719617078303, 0.3149281759965271, 0.07592684856551069, 0.0, -0.0, 1.3759356933258697],)

```

---

<div class="post-metadata">

**Author:** ![Devetak](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/devetak/32/50611_2.png) [@Devetak](https://discourse.julialang.org/u/Devetak)\
**Post date:** [June 26, 2023, 6:33pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/8 "2023-06-26T18:33:09Z")

</div>

Thanks to all for their different perspectives all of them are super useful!

---

<div class="post-metadata">

**Author:** ![jling](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jling/32/212909_2.png) [@jling](https://discourse.julialang.org/u/jling)\
**Post date:** [June 26, 2023, 7:12pm UTC](https://discourse.julialang.org/t/help-with-automatic-differentiation/100844/9 "2023-06-26T19:12:03Z")

</div>

> [@Devetak](#):
>
> I am aware of the fact that in AD mutating arrays are not supported,

ForwardDiff.jl should work right?
