# Need help with example; Mixture Density Networks from Site

**URL:** <https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514>\
**Category:** Performance\
**Tags:** flux\
**Created:** [May 23, 2022, 1:49pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514 "2022-05-23T13:49:06Z")\
**Posts on this page:** 20\
**Page:** 1

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 23, 2022, 1:49pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/1 "2022-05-23T13:49:06Z")

</div>

Tried to reproduce the example that you can see on the site:  
[Mixture Density Networks](https://www.janisklaise.com/post/mdn_julia/)  
But I don’t know how to update the following lines of code:

```julia
# lowest-level?
data = [(y, x)]

for epoch = 1:n_epochs
    
    # forward
    pi_out = pi(y)
    sigma_out = sigma(y)
    mu_out = mu(y)
    l = mdn_loss(pi_out, sigma_out, mu_out, x)
    
    # backward
    Tracker.back!(l)
    for p in pars
        Tracker.update!(opt, p, Tracker.grad(p))
    end

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l)
    end
end

```

In order to achieve the final result:  
 ![Flux_ example_to_update](https://global.discourse-cdn.com/julialang/original/3X/7/b/7b515da40efbf1d8470a8e40f6d34ea60bfc0e78.jpeg)

I work with

```julia
(@v1.7) pkg> st Flux
      Status `C:\Users\Hermesr\.julia\environments\v1.7\Project.toml`
  [587475ba] Flux v0.13.0

```

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 2:38pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/2 "2022-05-23T14:38:23Z")

</div>

Hi @HerAdri

I guess you can change the backward pass  
from →

```julia
    l = mdn_loss(pi_out, sigma_out, mu_out, x)
    
    # backward
    Tracker.back!(l)
    for p in pars
        Tracker.update!(opt, p, Tracker.grad(p))
    end

```

to →

```julia
    # backward
    gs = gradient(pars) do
        l = mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

```

This should work-out in `Flux : v0.13.0`, i never checked-it myself though…

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 23, 2022, 3:10pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/3 "2022-05-23T15:10:41Z")

</div>

The lines of code are as follows (according to the proposal, ok?):

```julia
for epoch = 1:n_epochs
    
    # forward
    pi_out = pi(y)
    sigma_out = sigma(y)
    mu_out = mu(y)
    
    
    # backward
    gs = gradient(pars) do
        l = mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l)
    end
end

```

and we get the following alert

```julia
ERROR: UndefVarError: l not defined
Stacktrace:
 [1] top-level scope
   @ c:\projects\Julia Flux\Mixture Density Networks with Julia.jl:162

```

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 3:18pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/4 "2022-05-23T15:18:46Z")

</div>

ooh okay,  
Maybe we don’t need to assign loss value to variable `l`, can you try removing `l` like follows →

```julia
    # backward
    gs = gradient(pars) do
        mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

```

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 23, 2022, 3:27pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/5 "2022-05-23T15:27:07Z")

</div>

It didn’t work, I get the same error

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 3:30pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/6 "2022-05-23T15:30:33Z")

</div>

The problem is not with this line sorry →

```julia
    l = mdn_loss(pi_out, sigma_out, mu_out, x)

```

After `do` block completes its execution, we are loosing the var `l`, so it is undefined at line →

```julia
println("Epoch: ", epoch, " loss: ", l)

```

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 3:34pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/7 "2022-05-23T15:34:03Z")

</div>

You can make use of `Flux.withgradient` function here to access both loss value and gradients  
below code fit’s our need →

```julia
    # backward
    l, gs = Flux.withgradient(pars) do
        mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

```

And also defining `l = 0f0` outside the do block works fine with our initial solution (function `gradient(....)`)

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 23, 2022, 3:39pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/8 "2022-05-23T15:39:41Z")

</div>

No more error inside the loop.  
But the result is not as expected  
 ![Flux_ eerror](https://global.discourse-cdn.com/julialang/original/3X/2/2/225a3b1b81ff45829e5dcf0de273d684992d4d94.jpeg)

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 23, 2022, 3:44pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/9 "2022-05-23T15:44:53Z")

</div>

you illustrate me:  
“And also defining l = 0f0 outside the do block works fine with our initial solution (function gradient(…))”

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 3:45pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/10 "2022-05-23T15:45:50Z")

</div>

Are the loss values what you are getting comparable to what is shown in the blog ?  
Unfortunately i don’t know anything about `Mixture Density Networks` …

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 3:47pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/11 "2022-05-23T15:47:38Z")

</div>

For example like below →

```julia
for epoch = 1:n_epochs
    
    # forward
    pi_out = pi(y)
    sigma_out = sigma(y)
    mu_out = mu(y)
    
    # backward
    l = 0f0 # define l before entering into do block to access later
    gs = gradient(pars) do
        l = mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l) # we can access l here without any undefined error
    end
end

```

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 23, 2022, 3:49pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/12 "2022-05-23T15:49:51Z")

</div>

```julia
ERROR: UndefVarError: gradient not defined
Stacktrace:
 [1] top-level scope
   @ c:\projects\Julia Flux\Mixture Density Networks with Julia.jl:159

Epoch: 1000 loss: 0.0
Epoch: 2000 loss: 0.0
Epoch: 3000 loss: 0.0
Epoch: 4000 loss: 0.0
Epoch: 5000 loss: 0.0
Epoch: 6000 loss: 0.0
Epoch: 7000 loss: 0.0
Epoch: 8000 loss: 0.0

```

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 3:55pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/13 "2022-05-23T15:55:41Z")

</div>

is this the new error ? 🧐  
can you also paste the code, so that we can see what exactly we have @line-no-159

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 23, 2022, 4:00pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/14 "2022-05-23T16:00:54Z")

</div>

It would be better to summarize our 2 solutions so that it is not confusing …

1. Defining `l` outside of do block and using `gradient` / `Flux.gradient` function → 

```julia
for epoch = 1:n_epochs
    
    # forward
    pi_out = pi(y)
    sigma_out = sigma(y)
    mu_out = mu(y)
    
    # backward
    l = 0f0 # define l before entering into do block to access later
    gs = gradient(pars) do # or gs = Flux.gradient(pars) do
        l = mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l) # we can access l here without any undefined error
    end
end

```

1. using `Flux.withgradient` function → 

```julia
for epoch = 1:n_epochs
    
    # forward
    pi_out = pi(y)
    sigma_out = sigma(y)
    mu_out = mu(y)
    
    # backward
    l, gs = Flux.withgradient(pars) do
        mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    update!(opt, pars, gs)

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l) 
    end
end

```

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 24, 2022, 4:15am UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/15 "2022-05-24T04:15:44Z")

</div>

this is the code:

```julia
#https://www.janisklaise.com/post/mdn_julia/
using Distributions
using Flux
using Plots
using Random
Random.seed!(12345); # for reproducibility

function generate_data(n_samples)
    ϵ = rand(Normal(), 1, n_samples)
    x = rand(Uniform(-10.5, 10.5), 1, n_samples)
    y = 7sin.(0.75x) + 0.5x + ϵ
    return x, y
end

n_samples = 1000
x, y = generate_data(n_samples); # semicolon to suppress output as in MATLAB
scatter(transpose(y), transpose(x), alpha=0.2)

n_gaussians = 5
n_hidden = 20;

z_h = Dense(1, n_hidden, tanh)
z_π = Dense(n_hidden, n_gaussians)
z_σ = Dense(n_hidden, n_gaussians, exp)
z_μ = Dense(n_hidden, n_gaussians);

pi = Chain(z_h, z_π, softmax)
sigma = Chain(z_h, z_σ)
mu = Chain(z_h, z_μ);

#We need to implement this loss function ourselves:

function gaussian_distribution(y, μ, σ)
    # periods are used for element-wise operations
    result = 1 ./ ((sqrt(2π).*σ)).*exp.(-0.5((y .- μ)./σ).^2)
end;
function mdn_loss(π, σ, μ, y)
    result = π.*gaussian_distribution(y, μ, σ)
    result = sum(result, dims=1)
    result = -log.(result)
    return mean(result)
end;

pars = Flux.params(pi, sigma, mu)
opt = ADAM()
n_epochs = 8000;

#Finally we write the training loop.
# lowest-level?
data = [(y, x)]

#1 Defining l outside of do block and using gradient / Flux.gradient function 
for epoch = 1:n_epochs
    
    # forward
    pi_out = pi(y)
    sigma_out = sigma(y)
    mu_out = mu(y)
    
    
    # backward
    l = 0f0 # define l before entering into do block to access later
    gs = Flux.gradient(pars) do
        l = mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    Flux.update!(opt, pars, gs)

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l) # we can access l here without any undefined error
    end
end

x_test = range(-15, stop=15, length=n_samples)
pi_data = pi(transpose(collect(x_test)))
sigma_data = sigma(transpose(collect(x_test)))
mu_data = mu(transpose(collect(x_test)));
plot(collect(x_test), transpose(pi_data), lw=2)

plot(collect(x_test), transpose(sigma_data), lw=2)

plot(collect(x_test), transpose(mu_data), lw=2)

#We can also plot the mean μₖ(x) for each Gaussian together with the range μₖ(x) ± σₖ(x):
plot(collect(x_test), transpose(mu_data), lw=2, ribbon=transpose(sigma_data),
    ylim=(-12,12), fillalpha=0.3)
scatter!(transpose(y), transpose(x), alpha=0.05)

function gumbel_sample(x)
    z = rand(Gumbel(), size(x))
    return argmax(log.(x) + z, dims=1)
end

k = gumbel_sample(pi_data);

sampled = rand(Normal(),1, 1000).*sigma_data[k] + mu_data[k];

scatter(transpose(y), transpose(x), alpha=0.2)
scatter!(collect(x_test), transpose(sampled), alpha=0.5)

```

and with it the loss function does not evolve.:

```julia
Epoch: 1000 loss: 8.61979609988576
Epoch: 2000 loss: 8.61979609988576
Epoch: 3000 loss: 8.61979609988576
Epoch: 4000 loss: 8.61979609988576
Epoch: 5000 loss: 8.61979609988576
Epoch: 6000 loss: 8.61979609988576
Epoch: 7000 loss: 8.61979609988576
Epoch: 8000 loss: 8.61979609988576

```

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 24, 2022, 4:32am UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/16 "2022-05-24T04:32:25Z")

</div>

```julia
#with version:
#using Flux.withgradient function →
for epoch = 1:n_epochs
    
     # forward
     pi_out = pi(y)
     sigma_out = sigma(y)
     mu_out = mu(y)
    
     # backward
     l, gs = Flux.withgradient(pars) do
         mdn_loss(pi_out, sigma_out, mu_out, x)
     end
     update!(opt, pars, gs)

     if epoch % 1000 == 0
         println("Epoch: ", epoch, " loss: ", l)
     end
end

```

The same thing happens here, the loss function does not evolve.

```julia
Epoch: 1000 loss: 4.8080576028606465
Epoch: 2000 loss: 4.8080576028606465
Epoch: 3000 loss: 4.8080576028606465
Epoch: 4000 loss: 4.8080576028606465
Epoch: 5000 loss: 4.8080576028606465
Epoch: 6000 loss: 4.8080576028606465
Epoch: 7000 loss: 4.8080576028606465
Epoch: 8000 loss: 4.8080576028606465

```

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 24, 2022, 4:29pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/17 "2022-05-24T16:29:23Z")

</div>

It’s weird 😇  
I will try to play with the code and let you know if i find any success with it !!

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 27, 2022, 5:44am UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/18 "2022-05-27T05:44:06Z")

</div>

Sorry, i couldn’t able to figure out what the issue is, loss seems to be constant (w/o any improvement)  
maybe you could open an issue on GitHub [jklaise / personal\_website](https://github.com/jklaise/personal_website/blob/master/notebooks/mdn_julia.ipynb) so that the author of that blog might help in correcting the code !!  
or I hope someone here who knows about `flux` better would come up with the solution…!!

---

<div class="post-metadata">

**Author:** ![Karthik-d-k](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/karthik-d-k/32/35438_2.png) [@Karthik-d-k](https://discourse.julialang.org/u/Karthik-d-k)\
**Post date:** [May 30, 2022, 4:14pm UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/19 "2022-05-30T16:14:44Z")

</div>

Hey @HerAdri

I finally could able to solve something !!, thanks to @ToucheSir 😀where he mentioned in other post about how gradients are calculated in `Zygote` [here](https://discourse.julialang.org/t/methoderror-objects-of-type-float64-are-not-callable/81876/3).  
Taking that comment into consideration, i moved every calculations that needs to be tracked for backward pass into `gradient do block`  
so i changed the code as follows and it looks like for me it does the job (acc to me 😉) you should confirm if otherwise →

```julia
# lowest-level?
data = [(y, x)]

for epoch in 1:n_epochs
    
    # forward
    l = 0f0
    
    # backward
    gs = gradient(pars) do
        pi_out = model[:pi](y)
        sigma_out = model[:sigma](y)
        mu_out = model[:mu](y)
        l = mdn_loss(pi_out, sigma_out, mu_out, x)
    end
    Flux.update!(opt, pars, gs)

    if epoch % 1000 == 0
        println("Epoch: ", epoch, " loss: ", l)
    end
end

```

---

<div class="post-metadata">

**Author:** ![HerAdri](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/heradri/32/5816_2.png) [@HerAdri](https://discourse.julialang.org/u/HerAdri)\
**Post date:** [May 31, 2022, 4:51am UTC](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514/20 "2022-05-31T04:51:19Z")

</div>

> [@Karthik-d-k](#):
>
> ```julia
> # lowest-level?
> data = [(y, x)]
> 
> for epoch in 1:n_epochs
>     
> # forward
> l = 0f0
>     
> # backward
> gs = gradient(pars) do
> pi_out = model[:pi](y)
> sigma_out = model[:sigma](y)
> mu_out = model[:mu](y)
> l = mdn_loss(pi_out, sigma_out, mu_out, x)
> end
> Flux.update!(opt, pars, gs)
> 
> if epoch % 1000 == 0
> println("Epoch: ", epoch, " loss: ", l)
> end
> end
> 
> ```

```julia
ERROR: UndefVarError: model not defined
Stacktrace:
 [1] _pullback(::Zygote.Context, ::var"#9#10")
   @ Zygote C:\Users\Hermesr\.julia\packages\Zygote\DkIUK\src\compiler\interface2.jl:9
 [2] pullback(f::Function, ps::Zygote.Params{Zygote.Buffer{Any, Vector{Any}}})
   @ Zygote C:\Users\Hermesr\.julia\packages\Zygote\DkIUK\src\compiler\interface.jl:352
 [3] gradient(f::Function, args::Zygote.Params{Zygote.Buffer{Any, Vector{Any}}})
   @ Zygote C:\Users\Hermesr\.julia\packages\Zygote\DkIUK\src\compiler\interface.jl:75
 [4] top-level scope
   @ c:\projects\Julia Flux\Mixture Density Networks.jl:59

```

 ![model_not_defined](https://global.discourse-cdn.com/julialang/original/3X/f/f/ffcfa2c0e611db7f2e075d568b1aff9895434a93.jpeg)

[Next page](https://discourse.julialang.org/t/need-help-with-example-mixture-density-networks-from-site/81514.md?page=2)
