# Error due to Dirichlet prior in Turing?

**URL:** <https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244>\
**Category:** Probabilistic Programming\
**Tags:** turing\
**Created:** [June 12, 2023, 6:34pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244 "2023-06-12T18:34:01Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![JanG](https://avatars.discourse-cdn.com/v4/letter/j/ecb155/32.png) [@JanG](https://discourse.julialang.org/u/JanG)\
**Post date:** [June 12, 2023, 6:34pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/1 "2023-06-12T18:34:02Z")

</div>

Hi,

This model below crashes for me with the error message `LoadError: DomainError with Dual{ForwardDiff.Tag{Turing.TuringTag, Float64}}(NaN,NaN,NaN): Bernoulli: the condition zero(p) <= p <= one(p) is not satisfied.`

```julia
using Turing, Random

Random.seed!(123)
@model function model(N = 1000, K = 2, y = rand(Bernoulli(0.5),N), g1 = rand(Categorical(K),N), g2 = rand(Categorical(K),N))
    x ~ Dirichlet(ones(K))
    for n in 1:N
        y[n] ~ Bernoulli(x[g1[n]]/(x[g1[n]]+x[g2[n]]))
    end
end

m = model()
chn = sample(m, NUTS(), 100)

```

Doing something like

```julia
p = x[g1[n]]/(x[g1[n]]+x[g2[n]])
if p >= 1.0
    p = 1.0
elseif p <= 0.0
    p = 0.0
end
y[n] ~ Bernoulli(p)

```

doesn’t fix the issue. I suspect this is due to `x ~ Dirichlet(ones(K))` as the model works if I replace that line with `x ~ filldist(Uniform(0,1),K)`. Does anybody have an idea what could be causing this?

Thanks in advance!

---

<div class="post-metadata">

**Author:** ![dlakelan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dlakelan/32/8491_2.png) [@dlakelan](https://discourse.julialang.org/u/dlakelan)\
**Post date:** [June 12, 2023, 6:36pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/2 "2023-06-12T18:36:50Z")

</div>

Possibly an initialization issue? If you provide valid initial conditions does it go at all?

---

<div class="post-metadata">

**Author:** ![JanG](https://avatars.discourse-cdn.com/v4/letter/j/ecb155/32.png) [@JanG](https://discourse.julialang.org/u/JanG)\
**Post date:** [June 12, 2023, 6:53pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/3 "2023-06-12T18:53:40Z")

</div>

Setting the initial conditions did indeed help for this seed but I can still find examples that don’t work. I.e. here’s an example that still doesn’t work (I’ve now set K=9)

```julia
using Turing, Random

Random.seed!(1)
@model function model(N = 1000, K = 9, y = rand(Bernoulli(0.5),N), g1 = rand(Categorical(K),N), g2 = rand(Categorical(K),N))
    x ~ Dirichlet(ones(K))
    # x ~ filldist(Uniform(0,1),K)
    for n in 1:N
        y[n] ~ Bernoulli(x[g1[n]]/(x[g1[n]]+x[g2[n]]))
    end
end

m = model()
chn = sample(m, NUTS(), 100, init_params = ones(9)/9)

```

---

<div class="post-metadata">

**Author:** ![Dan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dan/32/42581_2.png) [@Dan](https://discourse.julialang.org/u/Dan)\
**Post date:** [June 12, 2023, 11:03pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/4 "2023-06-12T23:03:18Z")

</div>

Using a different sampler instead of `NUTS` it works without the error:

```julia
chn = sample(m, HMC(0.01, 10), 100)

```

I think the problem is indeed with leaving the domain of Dirichlet distribution by the sampler. The solution of clamping to 0.0 - 1.0 isn’t enough as Dirichlet at 0.0 / 1.0 is an edge case with problematic derivative. Perhaps clamping to epsilon and 1-epsilon would resolve the issue.

For example:

```julia
p = x[g1[n]]/(x[g1[n]]+x[g2[n]])
pc = max(min(p,1.0-0.0001),0.0001)
y[n] ~ Bernoulli(pc)

```

allow completion of `NUTS` sampler with:

```julia
chn = sample(m, NUTS(), 100)

```

as was crashing in the OP.

---

<div class="post-metadata">

**Author:** ![JanG](https://avatars.discourse-cdn.com/v4/letter/j/ecb155/32.png) [@JanG](https://discourse.julialang.org/u/JanG)\
**Post date:** [June 13, 2023, 3:55pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/5 "2023-06-13T15:55:27Z")

</div>

Thanks for pointing out that it works with other samplers, that’s helpful.

I’m not sure clamping helps with NUTS though. Even when I set epsilon = 0.1, it still crashes for the second example I posted:

```julia
using Turing, Random

Random.seed!(1)
@model function model(N = 1000, K = 9, y = rand(Bernoulli(0.5),N), g1 = rand(Categorical(K),N), g2 = rand(Categorical(K),N))
    x ~ Dirichlet(ones(K))
    # x ~ filldist(Uniform(0,1),K)
    for n in 1:N
        p = x[g1[n]]/(x[g1[n]]+x[g2[n]])
        pc = max(min(p,1.0-0.1),0.1)
        y[n] ~ Bernoulli(pc)
    end
end

m = model()
chn = sample(m, NUTS(), 100, init_params = ones(9)/9)

```

---

<div class="post-metadata">

**Author:** ![JanG](https://avatars.discourse-cdn.com/v4/letter/j/ecb155/32.png) [@JanG](https://discourse.julialang.org/u/JanG)\
**Post date:** [June 13, 2023, 4:14pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/6 "2023-06-13T16:14:10Z")

</div>

I think it might have something to do with p being NaN. The code

```julia
using Turing, Random

Random.seed!(1)
@model function model(N = 1000, K = 9, y = rand(Bernoulli(0.5),N), g1 = rand(Categorical(K),N), g2 = rand(Categorical(K),N))
    x ~ Dirichlet(ones(K))
    # x ~ filldist(Uniform(0,1),K)
    for n in 1:N
        p = x[g1[n]]/(x[g1[n]]+x[g2[n]])
        if isnan(p)
            p = 0.000001
        end
        y[n] ~ Bernoulli(p)
    end
end

m = model()
chn = sample(m, NUTS(), 100, init_params = ones(9)/9)

```

runs without issues. I’m not quite sure where the NaNs are coming from though. I thought dividing by 0 produces Inf not NaN? Perhaps a very small denominator somehow causes NaNs when using NUTS?

edit: 0/0 is NaN, so it must be related to that!

---

<div class="post-metadata">

**Author:** ![Dan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dan/32/42581_2.png) [@Dan](https://discourse.julialang.org/u/Dan)\
**Post date:** [June 13, 2023, 6:28pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/7 "2023-06-13T18:28:19Z")

</div>

I think `(-Inf) - (-Inf)` is `NaN` , and logpdf outside Dirichlet distribution support is `-Inf` and this is how the `NaN`s creep in. There is not enough finiteness checking when calculating the gradient of the logpdf.  
There is an issue and a fix somewhere here, probably in AdvancedMHC.jl but haven’t zeroed in on it.

---

<div class="post-metadata">

**Author:** ![JanG](https://avatars.discourse-cdn.com/v4/letter/j/ecb155/32.png) [@JanG](https://discourse.julialang.org/u/JanG)\
**Post date:** [June 14, 2023, 12:06am UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/8 "2023-06-14T00:06:21Z")

</div>

I still think it’s the 0/0 issue; when I examine the values of `x[g1[n]]` and `x[g2[n]]` for cases where p is NaN, then they are both 0.0. The way to clamp this model successfully is, I think, something like this:

```julia
using Turing, Random

Random.seed!(1)
@model function model(N = 1000, K = 9, y = rand(Bernoulli(0.5),N), g1 = rand(Categorical(K),N), g2 = rand(Categorical(K),N))
    x ~ Dirichlet(ones(K))
    # x ~ filldist(Uniform(0,1),K)
    for n in 1:N
        a = max(x[g1[n]],0.00001)
        b = max(x[g2[n]],0.00001)
        y[n] ~ Bernoulli(a/(a+b))
    end
end

m = model()
chn = sample(m, NUTS(), 100, init_params = ones(9)/9)

```

---

<div class="post-metadata">

**Author:** ![sethaxen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/sethaxen/32/35604_2.png) [@sethaxen](https://discourse.julialang.org/u/sethaxen)\
**Post date:** [June 15, 2023, 1:17pm UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/9 "2023-06-15T13:17:36Z")

</div>

Without looking into this in any detail, I’d guess the issue is the one solved by [Make SimplexBijector actually bijective by sethaxen · Pull Request #263 · TuringLang/Bijectors.jl · GitHub](https://github.com/TuringLang/Bijectors.jl/pull/263). In short, how Turing currently unconstrains the simplex for sampling with gradient-based samplers makes the posterior improper. Technically the posterior variance for a single sampled parameter is infinite, so ironically, if warm-up works really well, that parameter will eventually overflow. This should be fixed soon.

---

<div class="post-metadata">

**Author:** ![dlakelan](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/dlakelan/32/8491_2.png) [@dlakelan](https://discourse.julialang.org/u/dlakelan)\
**Post date:** [June 16, 2023, 11:56am UTC](https://discourse.julialang.org/t/error-due-to-dirichlet-prior-in-turing/100244/10 "2023-06-16T11:56:49Z")

</div>

> [@sethaxen](#):
>
> This should be fixed soon.

Yikes, and fabulous at the same time!
