# Conditional Branching of Parameter in Turing.jl

**URL:** <https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019>\
**Category:** Modelling & Simulations\
**Tags:** question, package, diffeq\
**Created:** [November 12, 2020, 2:54am UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019 "2020-11-12T02:54:30Z")\
**Posts on this page:** 11\
**Page:** 1

<div class="post-metadata">

**Author:** ![nyubachi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nyubachi/32/18455_2.png) [@nyubachi](https://discourse.julialang.org/u/nyubachi)\
**Post date:** [November 12, 2020, 2:54am UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/1 "2020-11-12T02:54:30Z")

</div>

I’m new to the Julia language and to modeling, so I may ask a strange question.

When estimating parameters in Turing.jl, is it possible to conditionally branch for one parameter with another variable?

For example, in the code below, I want to branch the value of the parameter γ depending on the value of var1.

In such a case, I don’t know where and how to write the code.

Please let me know if there is a better way.

I would appreciate your advice.  
Thank you for your help.

```julia
using Turing, Distributions, DataFrames, DifferentialEquations, DiffEqSensitivity
using MCMCChains, Plots, StatsPlots
using Random
Random.seed!(12);

function lotka_volterra(du,u,p,t)
    x, y = u
    α, β, δ, γ = p
    du[1] = dx = (α - β*y)x
    du[2] = dy = (δ*x - γ)y
end
p = [1.5, 1.0, 3.0, 1.0]
u0 = [1.0,1.0]
prob = ODEProblem(lotka_volterra,u0,(0.0,10.0),p)
sol = solve(prob,Tsit5())
plot(sol)

odedata1 = Array(solve(prob,Tsit5(),saveat=0.1))
odedata2 = odedata1 .+ rand()
odedata3 = odedata1 .+ rand()
odedata = zeros(Float64, 2, 101, 3)
odedata[:,:,1] = odedata1
odedata[:,:,2] = odedata2
odedata[:,:,3] = odedata3

# I don’t know where and how to write the following code.

# var1 = [80, 40, 70]
# if var1 > 60
# γ = 4
# else
# γ = 1.5
# end

Turing.setadbackend(:forwarddiff)

@model function fitlv(data)
    σ ~ InverseGamma(2, 3)
    α ~ truncated(Normal(1.5,0.5),0.5,2.5)
    β ~ truncated(Normal(1.2,0.5),0,2)
    γ ~ truncated(Normal(3.0,0.5),1,4)
    δ ~ truncated(Normal(1.0,0.5),0,2)

    p = [α,β,γ,δ]
    prob = ODEProblem(lotka_volterra,u0,(0.0,10.0),p)
    predicted = solve(prob,Tsit5(),saveat=0.1)

    for k in 1:ndims(data)
        for i = 1:length(predicted)
            data[:,i,k] ~ MvNormal(predicted[i], σ)
        end
    end
end

model = fitlv(odedata)
chain = sample(model, NUTS(.65),1000)
plot(chain)

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [November 12, 2020, 3:20am UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/2 "2020-11-12T03:20:48Z")

</div>

> [@nyubachi](#):
>
> ```julia
> # var1 = [80, 40, 70]
> # if var1 > 60
> # γ = 4
> # else
> # γ = 1.5
> # end
> 
> ```

Just stick that into the model.

---

<div class="post-metadata">

**Author:** ![nyubachi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nyubachi/32/18455_2.png) [@nyubachi](https://discourse.julialang.org/u/nyubachi)\
**Post date:** [November 12, 2020, 2:36pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/3 "2020-11-12T14:36:17Z")

</div>

Thank you so much for your advice.  
Your activities have helped me a lot.

I’ve tried many things, but it didn’t work.

For example, in the following code, I removed γ from the parameter and wrote a conditional branch in the model, but it doesn’t seem to be working properly.

I think I’ve made a fundamental mistake.  
I would appreciate it if you could show me a concrete method.

```julia
using Turing, Distributions, DifferentialEquations
using MCMCChains, Plots, StatsPlots
using Random
Random.seed!(12);

function lotka_volterra(du,u,p,t)
    x, y = u
    α, β, δ = p
    γ = 3.0
    du[1] = dx = (α - β*y)x
    du[2] = dy = (δ*x - γ)y
end
p = [1.5, 1.1, 1.0]
u0 = [1.0,1.0]
prob = ODEProblem(lotka_volterra,u0,(0.0,10.0),p)
sol = solve(prob,Tsit5())
plot(sol)

odedata1 = Array(solve(prob,Tsit5(),saveat=0.1))
odedata2 = odedata1 .+ rand()
odedata3 = odedata1 .+ rand()
odedata = zeros(Float64, 2, 101, 3)
odedata[:,:,1] = odedata1
odedata[:,:,2] = odedata2
odedata[:,:,3] = odedata3

Turing.setadbackend(:forwarddiff)

@model function fitlv(data)
    σ ~ InverseGamma(2, 3)
    α ~ truncated(Normal(1.5,0.5),0.5,2.5)
    β ~ truncated(Normal(1.2,0.5),0,2)
    # γ ~ truncated(Normal(3.0,0.5),1,4)
    δ ~ truncated(Normal(1.0,0.5),0,2)

    var1 = [80, 40, 70]
    γ = zeros(Float64, length(var1))

    for l in 1:length(var1)
        if var1[l] > 60
            γ[l] = 5.0
        else
            γ[l] = 2.0
        end
    end

    p = [α,β,δ]
    prob = ODEProblem(lotka_volterra,u0,(0.0,10.0),p)
    predicted = solve(prob,Tsit5(),saveat=0.1)

    for k in 1:ndims(data)
        for i = 1:length(predicted)
            data[:,i,k] ~ MvNormal(predicted[i], σ)
        end
    end
end

model = fitlv(odedata)
chain = sample(model, NUTS(.65),1000)
plot(chain)

```

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [November 12, 2020, 2:43pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/4 "2020-11-12T14:43:29Z")

</div>

> [@nyubachi](#):
>
> ```julia
> for l in 1:length(var1)
> if var1[l] > 60
> γ[l] = 5.0
> else
> γ[l] = 2.0
> end
> end
> 
> p = [α,β,δ]
> prob = ODEProblem(lotka_volterra,u0,(0.0,10.0),p)
> 
> ```

I think you meant:

```julia
    for l in 1:length(var1)
        if var1[l] > 60
            γ = 5.0
        else
            γ = 2.0
        end
    end

    p = [α,β,δ,γ]
    prob = ODEProblem(lotka_volterra,u0,(0.0,10.0),p)

```

---

<div class="post-metadata">

**Author:** ![wulpuqu](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wulpuqu/32/17640_2.png) [@wulpuqu](https://discourse.julialang.org/u/wulpuqu)\
**Post date:** [November 12, 2020, 3:32pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/5 "2020-11-12T15:32:04Z")

</div>

> [@nyubachi](#):
>
> For example, in the code below, I want to branch the value of the parameter γ depending on the value of var1.
> 
> In such a case, I don’t know where and how to write the code.

Under the hood of `Turing.jl`, the parsing of model definitions is done by `DynamicPPL.jl`.

In practice, starting from the definition of your model, a compiler looks for any expressions like `LHS ~ RHS` and replace it with a function that returns the value of your random variable and update the metadata of your model, sampler, likelihood, etc…

So as long as your are writing proper code, you can write (almost\*) anything you want next to random assignements `LHS ~ RHS`.

(\* : Many definitions of a random variable with the same symbol is overwriting the same field in the metadata, definitions are not position-dependent in the code)

Example:

```julia
julia> using DynamicPPL

julia> @macroexpand @model function test()
           a ~ Normal()
           newfunc() = anyfunc()
           a = fct(a)+b
           if a > 0
               return "out"
           else
               b ~ Normal(a)
           end
       end
quote
    $(Expr(:meta, :doc))
    function test(; )
        var"##evaluator#271" = ((_rng::Random.AbstractRNG, _model::Model, _varinfo::AbstractVarInfo, _sampler::AbstractMCMC.AbstractSampler, _context::DynamicPPL.AbstractContext)->begin
                    begin
                        #= REPL[6]:2 =#
                        begin
                            var"##tmpright#263" = Normal()
                            var"##tmpright#263" isa Union{Distributions.Distribution, AbstractVector{<:Distributions.Distribution}} || throw(ArgumentError("Right-hand side of a ~ must be subtype of Distribution or a vector of Distributions."))
                            var"##vn#265" = a
                            var"##inds#266" = ()
                            a = (DynamicPPL.tilde_assume)(_rng, _context, _sampler, var"##tmpright#263", var"##vn#265", var"##inds#266", _varinfo)
                        end
                        #= REPL[6]:4 =#
                        newfunc() = begin
                                #= REPL[6]:4 =#
                                anyfunc()
                            end
                        #= REPL[6]:6 =#
                        a = fct(a) + b
                        #= REPL[6]:7 =#
                        if a > 0
                            #= REPL[6]:8 =#
                            return "out"
                        else
                            var"##tmpright#267" = Normal(a)
                            var"##tmpright#267" isa Union{Distributions.Distribution, AbstractVector{<:Distributions.Distribution}} || throw(ArgumentError("Right-hand side of a ~ must be subtype of Distribution or a vector of Distributions."))
                            var"##vn#269" = b
                            var"##inds#270" = ()
                            b = (DynamicPPL.tilde_assume)(_rng, _context, _sampler, var"##tmpright#267", var"##vn#269", var"##inds#270", _varinfo)
                        end
                    end
                end)
        return (Model)(:test, var"##evaluator#271", NamedTuple(), NamedTuple())
    end
end

```

So as you can see here, your model definition (model function definition) is transformed into a function that outputs a `Model` struct with as internal evaluator your initial model function where all random assignement have been replaced with a block of code that

- test if you are actually assigning a proper `Distribution` for your random variables
- looks for indices in case you are assigning value to an indexable variable
- use `DynamicPPL.tilde_assume` which outputs your RVs, update the sampler, the context and the VarInfo (the latter is the metadata of your model)

---

<div class="post-metadata">

**Author:** ![nyubachi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nyubachi/32/18455_2.png) [@nyubachi](https://discourse.julialang.org/u/nyubachi)\
**Post date:** [November 12, 2020, 3:44pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/6 "2020-11-12T15:44:38Z")

</div>

Thank you for your advice.

I have tried many things, but I may have to use MonteCarloProblem (EnsembleProblem ?) in order to reach my goal.

In this example, I want to map odedata to var1.  
(i.e., I want to change the value of γ for each odedata.)

I’m going to try to figure out how to write for a while.

Thank you so much.

---

<div class="post-metadata">

**Author:** ![nyubachi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nyubachi/32/18455_2.png) [@nyubachi](https://discourse.julialang.org/u/nyubachi)\
**Post date:** [November 12, 2020, 3:47pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/7 "2020-11-12T15:47:05Z")

</div>

Thank you for your thoughtful advice.

I’m a beginner, but I think I understand a lot better.

Thank you very much.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [November 12, 2020, 9:54pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/8 "2020-11-12T21:54:44Z")

</div>

I’m not sure what your question is then 🤷‍♂️

---

<div class="post-metadata">

**Author:** ![nyubachi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nyubachi/32/18455_2.png) [@nyubachi](https://discourse.julialang.org/u/nyubachi)\
**Post date:** [November 12, 2020, 10:06pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/9 "2020-11-12T22:06:13Z")

</div>

Sorry.

I guess I don’t fully understand my problem myself.

I’ll have to rethink it myself.

Thank you.

---

<div class="post-metadata">

**Author:** ![wulpuqu](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/wulpuqu/32/17640_2.png) [@wulpuqu](https://discourse.julialang.org/u/wulpuqu)\
**Post date:** [November 13, 2020, 1:18pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/10 "2020-11-13T13:18:18Z")

</div>

Another funny example which shows that `DynamicPPL.jl` takes the goal of replacing **ALL appearences** of `LHS ~ RHS` serious:

```julia
julia> @model function test()
           :(a~b)
       end
test (generic function with 1 method)

julia> vi=VarInfo();

julia> test()(vi)
quote
    var"##tmpright#268" = b
    var"##tmpright#268" isa Union{Distributions.Distribution, AbstractVector{<:Distributions.Distribution}} || throw(ArgumentError("Right-hand side of a ~ must be subtype of Distribution or a vector of Distributions."))
    var"##vn#270" = a
    var"##inds#271" = ()
    a = (DynamicPPL.tilde_assume)(_rng, _context, _sampler, var"##tmpright#268", var"##vn#270", var"##inds#271", _varinfo)
end

```

---

<div class="post-metadata">

**Author:** ![nyubachi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/nyubachi/32/18455_2.png) [@nyubachi](https://discourse.julialang.org/u/nyubachi)\
**Post date:** [November 13, 2020, 3:07pm UTC](https://discourse.julialang.org/t/conditional-branching-of-parameter-in-turing-jl/50019/11 "2020-11-13T15:07:16Z")

</div>

I learned a lot and it’s interesting.  
Thank you.
