# Turing: indicator variables vs control flow

**URL:** <https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539>\
**Category:** Probabilistic Programming\
**Created:** [January 18, 2021, 1:22pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539 "2021-01-18T13:22:32Z")\
**Posts on this page:** 10\
**Page:** 1

<div class="post-metadata">

**Author:** ![drbenvincent](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/drbenvincent/32/3046_2.png) [@drbenvincent](https://discourse.julialang.org/u/drbenvincent)\
**Post date:** [January 18, 2021, 1:22pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/1 "2021-01-18T13:22:32Z")

</div>

I have a change point detection model which I’m converting over from JAGS to Turing, and this one is interesting as it raises the question of whether you should use indicator variables (common to other PPL’s) or to use control flow.

I’ve implemented both, and they both ‘work’ in that sampling does not through an error. However having tried a couple of different samplers, the sampling is extremely bad. Any tips on why this is the case, which would be the best sampler, or whether one should favour use if indicator variables over control flow?

## Indicator variable model

```julia
using Turing, StatsPlots

@model function model(c)
    tₘₐₓ = length(c) 
    t = [1:tₘₐₓ;]
    # priors
    μ = Vector(undef, 2)
    μ[1] ~ Normal(0, 100)
    μ[2] ~ Normal(0, 100)
    σ ~ Uniform(0, 100)
    τ ~ Uniform(1, tₘₐₓ)
    # indicator variable
    z = Vector(undef, tₘₐₓ)
    z[t.<τ] .= 1
    z[t.≥τ] .= 2
    # likelihood
    for i in 1:tₘₐₓ
        c[i] ~ Normal(μ[z[i]], σ)
    end
end

# generate data
τ_true, μ₁, μ₂, σ_true = 500, 45, 30, 4
c = vcat(rand(Normal(μ₁,σ_true), τ_true), 
         rand(Normal(μ₂,σ_true), 1000-τ_true))

chain = sample(model(c), MH(), 5000)

plot(chain)

# data
plot([1:length(c);], c, xlabel="Time", ylabel="Count", title="Data space", legend=false)
# mean posterior predictive
plot!([1, mean(chain[:τ])], [mean(chain["μ[1]"]), mean(chain["μ[1]"])], lw=6, color=:black)
plot!([mean(chain[:τ]), length(c)], [mean(chain["μ[2]"]), mean(chain["μ[2]"])], lw=6, color=:black)

```

results in terrible chains

 ![Screenshot 2021-01-18 at 13.13.31](https://global.discourse-cdn.com/julialang/original/3X/1/f/1f1c568bdbadbd6f9d7fefb80849bbab2dc063a1.png)  
and  
 ![Screenshot 2021-01-18 at 13.13.22](https://global.discourse-cdn.com/julialang/original/3X/5/e/5ec0875b8593b0f46bce20422431147efeff0826.png)

## If else model

```julia
using Turing, StatsPlots

@model function model(c)
    tₘₐₓ = length(c)
    # priors
    μ₁ ~ Normal(0, 100)
    μ₂ ~ Normal(0, 100)
    σ ~ Uniform(0, 100)
    τ ~ Uniform(1, tₘₐₓ)
    # likelihood
    for t in 1:tₘₐₓ
        if t < τ        
            c[t] ~ Normal(μ₁, σ)
        else
            c[t] ~ Normal(μ₂, σ)
        end
    end
end

# generate data
τ_true, μ₁, μ₂, σ_true = 500, 45, 30, 4
c = vcat(rand(Normal(μ₁,σ_true), τ_true), 
         rand(Normal(μ₂,σ_true), 1000-τ_true))

chain = sample(model(c), MH(), 5000)

plot(chain)

# data
plot([1:length(c);], c, xlabel="Time", ylabel="Count", title="Data space", legend=false)
# mean posterior predictive
plot!([1, mean(chain[:τ])], [mean(chain[:μ₁]), mean(chain[:μ₁])], lw=6, color=:black)
plot!([mean(chain[:τ]), length(c)], [mean(chain[:μ₂]), mean(chain[:μ₂])], lw=6, color=:black)

```

with similarly bad chains.

---

<div class="post-metadata">

**Author:** ![Tamas\_Papp](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tamas_papp/32/25949_2.png) [@Tamas\_Papp](https://discourse.julialang.org/u/Tamas_Papp)\
**Post date:** [January 19, 2021, 9:46am UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/2 "2021-01-19T09:46:18Z")

</div>

If the posterior is otherwise equivalent, go with what’s simpler/faster — possibly your second option.

Bad mixing may indicate a misfit of your model to the data, or just a model with multiple modes — MH is not particularly good at dealing with that.

---

<div class="post-metadata">

**Author:** ![BradGroff](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bradgroff/32/21156_2.png) [@BradGroff](https://discourse.julialang.org/u/BradGroff)\
**Post date:** [January 19, 2021, 2:46pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/3 "2021-01-19T14:46:33Z")

</div>

Hi Ben,

I have 2 thoughts.

1. There are some stan examples [here](https://mc-stan.org/docs/2_26/stan-users-guide/change-point-section.html) that might help with the coding details.
2. You might be better off with a smooth changepoint function (sigmoid). I don’t know the internals well enough to be sure but I suspect that might provide a better gradient for \tau. In particular, with incorrect values of \tau near 0, you’d expect \mu\_2 to be somewhere like an average over the whole data which is what you see, so changepoint sampling issues seem pretty explanatory. I remember a post somewhere about this that I can’t find but iirc it was something like:

```julia
cp ~ Uniform(0.0, length) # Real uniform
switch = sigmoid(cp, bandwidth)     
# ^ center = cp, scale = bandwidth, exercise for reader to code :)
data ~ Normal(switch * mu_1 + (1 - switch) * mu_2, 1)

```

Here bandwidth controls how “strict” your cutoff could be. This function gives you a gradient at all values of `switch`. I’d love to hear any thoughts on this approach from someone with more experience!

Take care,  
Brad

---

<div class="post-metadata">

**Author:** ![BradGroff](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bradgroff/32/21156_2.png) [@BradGroff](https://discourse.julialang.org/u/BradGroff)\
**Post date:** [January 19, 2021, 2:58pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/4 "2021-01-19T14:58:37Z")

</div>

[Here’s](https://statmodeling.stat.columbia.edu/2016/03/18/i-definitely-wouldnt-frame-it-as-to-determine-if-the-time-series-has-a-change-point-or-not-the-time-series-whatever-it-is-has-a-change-point-at-every-time-the-question/#comment-266539) a comment from Daniel Lakeland on Gelman’s blog that suggests a similar approach.

---

<div class="post-metadata">

**Author:** ![drbenvincent](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/drbenvincent/32/3046_2.png) [@drbenvincent](https://discourse.julialang.org/u/drbenvincent)\
**Post date:** [January 19, 2021, 3:45pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/5 "2021-01-19T15:45:21Z")

</div>

Thanks both. Doing more experimentation, it does seem that the chain for the switch point `τ` is all over the place. Even when tightening up the priors around the true values for the other parameters, it still does a poor job. So yes I suspect you’re right re. the switch point.

Although just visualising in my mind, I’d assume as the switch point gets closer to the true switch point then the likelihood would improve IF the sample for mu1 is higher than mu2, otherwise it would presumably be worse. So yes, possibly it’s just a hard problem to solve re the switch point even though it’s superficially trivial.

What does _not_ help, is if you constrain mu2 \> mu1 for example.

It also does _not_ help if you change the model and add some gradient in a linear discontinuity type model.

 ![Screenshot 2021-01-19 at 15.43.31](https://global.discourse-cdn.com/julialang/original/3X/9/8/981a2e594ef7371fea91694fec8f26cace7a1e28.png)

So yep, I’ll experiment with your suggestions with a sigmoid around the switchpoint.

---

<div class="post-metadata">

**Author:** ![BradGroff](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bradgroff/32/21156_2.png) [@BradGroff](https://discourse.julialang.org/u/BradGroff)\
**Post date:** [January 19, 2021, 9:45pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/6 "2021-01-19T21:45:30Z")

</div>

Hi Ben.

I took a stab myself and identified a few things:

1. Perhaps not surprising, standardizing the data really helps. Alternately, we could have centered in the model instead. Here’s what I did:

```julia
τ_true, μ₁, μ₂, σ_true = 500, 45, 30, 4
c = vcat(rand(Normal(μ₁,σ_true), τ_true), 
         rand(Normal(μ₂,σ_true), 1000-τ_true));

std_c = (c .- mean(c)) ./ std(c); 

```

1. The sigmoid approach also worked, here’s the code:

```julia
logit = bijector(Beta()) # bijection: (0, 1) → ℝ
inv_logit = inv(logit) # bijection: ℝ → (0, 1)

function sigm(μ, σ, x) # scaled, shifted sigmoid
    return inv_logit(σ*(x - μ))
end;

@model function changepoint(data)
    # priors
    spec = 0.01
    μ_1 ~ Normal(0, 2)
    μ_2 ~ Normal(0, 2)
    τ ~ Uniform(1, 1000)
    σ ~ truncated(Normal(1, 2), 0, 20)
    
    # likelihood
    for i in 1:length(data)
        switch = sigm(τ, spec, i)
        z = (1-switch) * μ_1 + switch * μ_2
        data[i] ~ Normal(z, σ)
    end
end;

# Settings of the Hamiltonian Monte Carlo (HMC) sampler.
iterations = 2000
ϵ = 0.005
τ = 10;

cp_chain = sample(
    changepoint(std_c), 
    HMC(ϵ, τ), iterations, 
    progress=true, drop_warmup=false);

StatsPlots.plot(cp_chain)

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/3/0/30372fe65643b9241cab5b61dfa8256e3c4cf589.png)

Note that the traces are against the standardized data, which looks like:

 ![image](https://global.discourse-cdn.com/julialang/original/3X/e/5/e539c0ef210178cddf8a21f2d8c68d26a6d209d2.png)

---

<div class="post-metadata">

**Author:** ![cpfiffer](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/cpfiffer/32/208747_2.png) [@cpfiffer](https://discourse.julialang.org/u/cpfiffer)\
**Post date:** [January 19, 2021, 10:16pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/7 "2021-01-19T22:16:06Z")

</div>

An excellent answer!

---

<div class="post-metadata">

**Author:** ![drbenvincent](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/drbenvincent/32/3046_2.png) [@drbenvincent](https://discourse.julialang.org/u/drbenvincent)\
**Post date:** [January 22, 2021, 10:59am UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/8 "2021-01-22T10:59:46Z")

</div>

Thanks for this excellent answer @BradGroff. It’s also spurred me to add bijectors to my “to learn” list.

I tried to delve into this a bit further. It’s still confusing to me why estimation of change points in JAGS works fine, but not here. I get why the sigmoid might help things, but I still can’t quite intuit why a change point model would fail here.

To experiment with that, I tried a change point model where all the parameters were known other than the change point. Inferring the change point alone does indeed work. Also works with MH.

```julia
using Turing, StatsPlots

@model function model(c, μ₁, μ₂, σ)
    tₘₐₓ = length(c)
    # prior
    τ ~ Uniform(1, tₘₐₓ) 
    # likelihood
    for t in 1:tₘₐₓ
        if t < τ        
            c[t] ~ Normal(μ₁, σ)
        else
            c[t] ~ Normal(μ₂, σ)
        end
    end
end

# generate data
τ, μ₁, μ₂, σ = 350, 45, 30, 4
c = vcat(rand(Normal(μ₁,σ), τ), 
         rand(Normal(μ₂,σ), 1000-τ))

chain = sample(model(c,μ₁, μ₂, σ), HMC(0.005, 10), 2000)

plot(chain)

```

 ![Screenshot 2021-01-22 at 10.08.52](https://global.discourse-cdn.com/julialang/original/3X/e/a/ea0782e5421e942f4c66ead0512608f595088e7b.png)

Where the problem seems to start is when inferring a change point _and_ a mean. If you have custom selected priors, it’s doable, but the convergence is very slow…

```julia
@model function model(c, μ₂, σ)
    tₘₐₓ = length(c)
    # prior
    τ ~ Uniform(1, tₘₐₓ) 
    μ₁ ~ Normal(45, σ)
    # likelihood
    for t in 1:tₘₐₓ
        if t < τ        
            c[t] ~ Normal(μ₁, σ)
        else
            c[t] ~ Normal(μ₂, σ)
        end
    end
end

# generate data
τ, μ₁, μ₂, σ = 350, 45, 30, 4
c = vcat(rand(Normal(μ₁,σ), τ), 
         rand(Normal(μ₂,σ), 1000-τ))

chain = sample(model(c, μ₂, σ), HMC(0.005, 10), 5000)
plot(chain)

```

 ![Screenshot 2021-01-22 at 10.18.52](https://global.discourse-cdn.com/julialang/original/3X/0/3/03ed31b174a620e58de69bf7cd56ef527f261ad4.png)

And the convergence gets unworkably slow with non hand-picked priors.

So I guess it looks like the problem is not the change point / step function alone, but estimating that in conjunction with one or both means. I still don’t quite get the intuition of the problem, but certainly get why the sigmoid solution would work. I’ll certainly be going forward with that… thanks again for the thorough reply 🙂

---

<div class="post-metadata">

**Author:** ![drbenvincent](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/drbenvincent/32/3046_2.png) [@drbenvincent](https://discourse.julialang.org/u/drbenvincent)\
**Post date:** [January 22, 2021, 12:08pm UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/9 "2021-01-22T12:08:33Z")

</div>

It’s also a bit odd… if you standardise the data as in @BradGroff 's example it works fine. But if you try to adjust the priors to the data then that does not work well, even with a ton of samples.

```julia
@model function changepoint(c)
    tₘₐₓ = length(c)
    spec = 0.01
    # priors
    μ_1 ~ Normal(mean(c), std(c)*2)
    μ_2 ~ Normal(mean(c), std(c)*2)
    σ ~ TruncatedNormal(0, std(c)*2, 0, std(c)*20)
    σ ~ Uniform(0, std(c)*20)
    τ ~ Uniform(1, tₘₐₓ)
    # likelihood
    for t in 1:length(c)
        switch = sigmoid(τ, spec, t)
        z = (1-switch) * μ_1 + switch * μ_2
        c[t] ~ Normal(z, σ)
    end
end

# generate data
τ_true, μ₁, μ₂, σ_true = 350, 45, 30, 4
c = vcat(rand(Normal(μ₁,σ_true), τ_true), 
         rand(Normal(μ₂,σ_true), 1000-τ_true))

chain = sample(changepoint(c), HMC(0.005, 10), 10000)
plot(chain)

```

 ![Screenshot 2021-01-22 at 12.07.04](https://global.discourse-cdn.com/julialang/original/3X/b/f/bf582e7db0fa2632fd085b6625b91c1af0d31349.png)

EDIT: Seems to work _way_ better using NUTS, compared to HMC.

 ![Screenshot 2021-01-22 at 12.14.45](https://global.discourse-cdn.com/julialang/original/3X/d/2/d248d9bb70c91b60b7234d4a67b1eb764f152d9f.png)

---

<div class="post-metadata">

**Author:** ![trappmartin](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/trappmartin/32/1165_2.png) [@trappmartin](https://discourse.julialang.org/u/trappmartin)\
**Post date:** [January 23, 2021, 10:44am UTC](https://discourse.julialang.org/t/turing-indicator-variables-vs-control-flow/53539/10 "2021-01-23T10:44:39Z")

</div>

> [@drbenvincent](#):
>
> It’s also a bit odd… if you standardise the data as in @BradGroff 's example it works fine. But if you try to adjust the priors to the data then that does not work well, even with a ton of samples.
> 
> ```julia
> 
> ```

This is to be expected as inference in the model with very wide priors is more difficult for HMC. You might want to switch to NUTS but in any case a model reparameterisation would be a good idea in this case.
