# Reparametrization trick in Flux.jl

**URL:** <https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489>\
**Category:** Machine Learning\
**Tags:** flux, machine-learning\
**Created:** [June 17, 2023, 1:18pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489 "2023-06-17T13:18:59Z")\
**Posts on this page:** 9\
**Page:** 1

<div class="post-metadata">

**Author:** ![josemanuel22](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/josemanuel22/32/20668_2.png) [@josemanuel22](https://discourse.julialang.org/u/josemanuel22)\
**Post date:** [June 17, 2023, 1:18pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/1 "2023-06-17T13:18:59Z")

</div>

Does `Flux.jl` have an equivalent to `rsample` in `PyTorch` that automatically implements these stochastic/policy gradients. That way the reparameterized sample becomes differentiable.

---

<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 17, 2023, 2:07pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/2 "2023-06-17T14:07:10Z")

</div>

I don’t think so, but @BatyLeo and I have been working on something like that. Maybe it’s time to open source it?

---

<div class="post-metadata">

**Author:** ![e3c6](https://avatars.discourse-cdn.com/v4/letter/e/e79b87/32.png) [@e3c6](https://discourse.julialang.org/u/e3c6)\
**Post date:** [June 17, 2023, 2:34pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/3 "2023-06-17T14:34:31Z")

</div>

+1 I would also find use for something like this. A package would be nice.

---

<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 17, 2023, 3:26pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/4 "2023-06-17T15:26:16Z")

</div>

In the meantime maybe [StochasticAD.jl](https://github.com/gaurav-arya/StochasticAD.jl) can be useful? What do you need this for?

---

<div class="post-metadata">

**Author:** ![josemanuel22](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/josemanuel22/32/20668_2.png) [@josemanuel22](https://discourse.julialang.org/u/josemanuel22)\
**Post date:** [June 17, 2023, 5:35pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/5 "2023-06-17T17:35:07Z")

</div>

I’m diving into Flux.jl for my research. My objective is to define a loss function that measures the convergence of the model towards a uniform data distribution, so to speak. I generate K random fictitious observations and compare how many of them are smaller than the true data in the training set. In other words, the model generates K simulations, and we determine the number of simulations where the generated data is smaller than the actual data. If the model is well trained, this distribution should converge to a uniform distribution. However, since the resulting histogram is not differentiable, I needed to approximate this idea using compositions of continuous differentiable functions (I did it).

The reason is that I believe this generates a random node, and therefore the need to apply the reparametrization trick arises from there.

---

<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 18, 2023, 3:57pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/6 "2023-06-18T15:57:18Z")

</div>

Is your model generating a discrete distribution or a continuous one? The trouble here is that not every distribution is easily amenable to reparametrization. There are extensions but they too have limits.  
What we are implementing with @BatyLeo is closer to the score function method, which is more generic but also suffers from high variance.  
See this paper for a great overview:

> **[Monte Carlo gradient estimation in machine learning | The Journal of Machine...](https://dl.acm.org/doi/abs/10.5555/3455716.3455848)**
>
> This paper is a broad and accessible survey of the methods we have at our disposal
> for Monte Carlo gradient estimation in machine learning and across the statistical sciences: the problem of computing
> the gradient of an expectation of a function with...

---

<div class="post-metadata">

**Author:** ![Tomas\_Pevny](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomas_pevny/32/25466_2.png) [@Tomas\_Pevny](https://discourse.julialang.org/u/Tomas_Pevny)\
**Post date:** [June 18, 2023, 7:48pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/7 "2023-06-18T19:48:28Z")

</div>

I think that the reason, why Julia lacks the `rsample` is that it is ridiculously simple to implement for distributions of interest and there is not a great need for it. To implement the classical Gaussian reparametrization, which covers most uses is effectively one line of code. I am adding a complete example, but it is effectively this `m.μ(x) .+ m.σ(x) .* r `, which is nicely similar to what is in papers.

```julia
using Flux
using Functors

struct Model{S,M}
	μ::M 
	σ::S
end

@functor Model

function (m::Model)(x)
	T = eltype(x)
	r = randn(T, 2, size(x,2))
	m.μ(x) .+ m.σ(x) .* r 
end

m = Model( 
	Chain(Dense(2,2,relu), Dense(2,2)),
	Chain(Dense(2,2,relu), Dense(2,2,softplus)),
	)

x = randn(Float32, 2, 11)

gradient(m -> sum(m(x)), m)

```

---

<div class="post-metadata">

**Author:** ![bertschi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bertschi/32/33462_2.png) [@bertschi](https://discourse.julialang.org/u/bertschi)\
**Post date:** [June 18, 2023, 10:14pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/8 "2023-06-18T22:14:41Z")

</div>

For several distributions this is true. Yet, torch also supports some distributions which are less trivial and even then, it would be a nice addition that I missed some times.

---

<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 10, 2024, 2:16pm UTC](https://discourse.julialang.org/t/reparametrization-trick-in-flux-jl/100489/9 "2024-06-10T14:16:38Z")

</div>

Update: I have put together a little package for differentiating through expectations. It includes both the REINFORCE and the reparametrization trick. Still very experimental, and currently being registered. I’d be excited to have your feedback 🙂

> **[GitHub - JuliaDecisionFocusedLearning/DifferentiableExpectations.jl: A Julia...](https://github.com/JuliaDecisionFocusedLearning/DifferentiableExpectations.jl)**
>
> A Julia package for differentiating through expectations with Monte-Carlo estimates - JuliaDecisionFocusedLearning/DifferentiableExpectations.jl
