# Autodiff with Zygote: issues with setting seeds

**URL:** <https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957>\
**Category:** Statistics\
**Tags:** zygote, autodiff\
**Created:** [October 29, 2024, 5:35pm UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957 "2024-10-29T17:35:17Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![weltenbummler](https://avatars.discourse-cdn.com/v4/letter/w/82dd89/32.png) [@weltenbummler](https://discourse.julialang.org/u/weltenbummler)\
**Post date:** [October 29, 2024, 5:35pm UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/1 "2024-10-29T17:35:17Z")

</div>

Hello all,

I would like to differentiate a function with fixed seed but obtain a “can’t differentiate foreign call” error when using `Zygote`. Any advise would be appreciated.

The following is a minimal working example and one possible, but for me somewhat limiting, work around.

```julia
using Zygote
using Random

# Random.seed!( ) does not work with Zygote. It produces a "can't differentiate foreight call" error
function simulator(x, id::Int64)
    Random.seed!(id)
    return simulator(x)
end

"This is a work around. The simulator needs to take a rng as input."
function simulator(x, rng::AbstractRNG)
    noise1 = randn(rng)
    noise2 = randn(rng)
    @show noise1
    @show noise2
    return x+noise1+noise2
end

function simulator(x)
    noise1 = randn()
    noise2 = randn()
    @show noise1
    @show noise2
    return x+noise1+noise2
end

function distance(sim, obs)
    return sum((sim-obs).^2)
end

"This will work"
function loss(x, obsdata, id::Int64)
    rng = Xoshiro(id) 
    sim = simulator(x, rng)
    return distance(sim, obsdata)
end

"This won't work"
function loss_with_issue(x, obsdata, id::Int64)
    sim = simulator(x, id)
    return distance(sim, obsdata)
end

# data
myobs = 2.0;

# to fix the seed
id = 123

# test point
xtest = 3.0

# This works
Zygote.gradient(x->loss(x, myobs, id), xtest)
2*(simulator(xtest, Xoshiro(id))-myobs)

# This throws an error: can't differentiate foreigncall expression"
Zygote.gradient(x->loss_with_issue(x, myobs, id), xtest)

```

Arguably `simulator(x, rng::AbstractRNG)` is cleaner code and may be preferred anyway, but I needed to be able to differentiate my loss also for simulators such as `simulator(x)` that do not work with an explicit RNG instance.

Would someone know how to make `Zygote` work without having to pass around a RNG instance, i.e. for the `loss_with_issue` case?

Many thanks!

---

<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:** [October 29, 2024, 5:38pm UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/2 "2024-10-29T17:38:59Z")

</div>

I think you can just tell Zygote not to look inside that function, like so:

```julia
julia> Zygote.gradient(x->loss_with_issue(x, myobs, id), xtest)
noise1 = -0.6457306721039767
noise2 = -1.4632513788889214
ERROR: Can't differentiate foreigncall expression $(Expr(:foreigncall, :(:jl_get_current_task), Ref{Task}, svec(), 0, :(:ccall))).
Stacktrace:
...
  [4] setstate!
    @ /Applications/Julia-1.10.app/Contents/Resources/julia/share/julia/stdlib/v1.10/Random/src/Xoshiro.jl:132 [inlined]

julia> function simulator(x, id::Int64)
           Zygote.@ignore Random.seed!(id)
           return simulator(x)
       end
simulator (generic function with 3 methods)

julia> Zygote.gradient(x->loss_with_issue(x, myobs, id), xtest)
noise1 = -0.6457306721039767
noise2 = -1.4632513788889214
(-2.2179641019857965,)

```

I believe that could be made permanent by a one-line PR [here](https://github.com/JuliaDiff/ChainRules.jl/blob/main/src/rulesets/Random/random.jl#L38).

---

<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:** [October 29, 2024, 8:33pm UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/3 "2024-10-29T20:33:22Z")

</div>

Tangential remark: why do you want to differentate a function that returns random values? Autodiff engines are not designed to deal with such situations by default, so you might obtain unexpected (and backend-dependent) results

---

<div class="post-metadata">

**Author:** ![weltenbummler](https://avatars.discourse-cdn.com/v4/letter/w/82dd89/32.png) [@weltenbummler](https://discourse.julialang.org/u/weltenbummler)\
**Post date:** [October 29, 2024, 8:50pm UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/4 "2024-10-29T20:50:54Z")

</div>

Thank you very much. That indeed resolves the issue.

---

<div class="post-metadata">

**Author:** ![weltenbummler](https://avatars.discourse-cdn.com/v4/letter/w/82dd89/32.png) [@weltenbummler](https://discourse.julialang.org/u/weltenbummler)\
**Post date:** [October 29, 2024, 8:54pm UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/5 "2024-10-29T20:54:03Z")

</div>

The motivation for this is the implementation of a statistical inference procedure that works by fixing the seed of the stochastic generative model (the simulator). Details about the method would be [here](http://proceedings.mlr.press/v108/ikonomov20a.html).

---

<div class="post-metadata">

**Author:** ![Red-Portal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/red-portal/32/9102_2.png) [@Red-Portal](https://discourse.julialang.org/u/Red-Portal)\
**Post date:** [November 9, 2024, 3:45am UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/6 "2024-11-09T03:45:51Z")

</div>

@gdall Isn’t that basically what all stochastic gradient descent-based methods do? At least we do that all over the place in `AdvancedVI`

---

<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:** [November 9, 2024, 7:06am UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/7 "2024-11-09T07:06:55Z")

</div>

Not really. Stochastic gradient descent approximates the gradient of a deterministic function f(x) = \sum\_{i \in \mathcal{I}} f\_i(x) with a random subset \mathcal{S} \subset \mathcal{I} of the component’s gradients. Here, we’re talking about a function f which itself involves randomness. The right way to think about it is as a stochastic computational graph, see [https://arxiv.org/abs/1506.05254](https://arxiv.org/abs/1506.05254).

---

<div class="post-metadata">

**Author:** ![Red-Portal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/red-portal/32/9102_2.png) [@Red-Portal](https://discourse.julialang.org/u/Red-Portal)\
**Post date:** [November 9, 2024, 7:22am UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/8 "2024-11-09T07:22:51Z")

</div>

I would say that is only one specific type of SGD, where stochasticity is discrete due to subsampling. In variational inference, for example, we deal with more general type of stochastic gradient descent, where the gradient is defined as

\nabla\_{x} \; \mathbb{E}\_{\epsilon} f\left(x, \epsilon\right),

where \epsilon is general (often continuous) noise.  
To me, this is the same as differentiating a random function if we think as the randomness \epsilon being implicitly generated inside the function. In fact, abstractly speaking, `AdvancedVI` operates exactly as the snippet shown in the original post 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:** [November 9, 2024, 7:56am UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/9 "2024-11-09T07:56:01Z")

</div>

You’re right, but even in the general SGD you’re differentiating an expectation, which is a deterministic function, and you’re only replacing its gradient with a stochastic approximation. While the implementation may be similar, conceptually it is very different from differentiating a function with random outputs (the best review on Monte-Carlo gradients is [https://jmlr.org/papers/volume21/19-346/19-346.pdf](https://jmlr.org/papers/volume21/19-346/19-346.pdf)). I’m just highlighting that users should be aware of which function they’re considering, and whether it is inherently stochastic or not.

---

<div class="post-metadata">

**Author:** ![weltenbummler](https://avatars.discourse-cdn.com/v4/letter/w/82dd89/32.png) [@weltenbummler](https://discourse.julialang.org/u/weltenbummler)\
**Post date:** [November 11, 2024, 9:08am UTC](https://discourse.julialang.org/t/autodiff-with-zygote-issues-with-setting-seeds/121957/10 "2024-11-11T09:08:34Z")

</div>

Yes, it’s a rather different situation. In our case, by setting a random seed (or setting \epsilon to \epsilon\_0), we are working on a realisation of a random process, i.e. a single instance f(x, \epsilon\_0) of the random function. After changing x, we keep the seed \epsilon fixed. In mini-batch approaches to stochastic optimisation, one would take a new random sample (i.e. change \epsilon after updating x.
