# Turing: selective truncated distribution error

**URL:** <https://discourse.julialang.org/t/turing-selective-truncated-distribution-error/118598>\
**Category:** Probabilistic Programming\
**Tags:** turing, distributions, autodiff\
**Created:** [August 25, 2024, 8:51pm UTC](https://discourse.julialang.org/t/turing-selective-truncated-distribution-error/118598 "2024-08-25T20:51:59Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![Sam\_P](https://avatars.discourse-cdn.com/v4/letter/s/ed8c4c/32.png) [@Sam\_P](https://discourse.julialang.org/u/Sam_P)\
**Post date:** [August 25, 2024, 8:51pm UTC](https://discourse.julialang.org/t/turing-selective-truncated-distribution-error/118598/1 "2024-08-25T20:51:59Z")

</div>

Hi all,

I’ve been trying to fit a model to truncated discrete data using Turing, with the aim of recovering the non-truncated distribution’s parameter, and use the fitted model to predict new values based on new truncation points.

Below is a MWP example. Dummy data consists of a set of observations where the `data` is above or at the `truncation` threshold

```julia
Random.seed!(1234)

# Generate Poisson-distributed data
λ_true = 10.0 # True Poisson rate parameter
n_samples = 100 # Number of samples
poisson_data = rand(Poisson(λ_true), n_samples)

λ_threshold = 7.0
trunc = rand(Poisson(λ_threshold), n_samples)

poisson_data_truncated = DataFrame(obs_data = poisson_data, truncation = trunc)

poisson_data_truncated = @subset(poisson_data_truncated, :obs_data .>= :truncation)

```

The `T()` notation is from [this link](https://discourse.julialang.org/t/turing-error-with-truncated-distribution/53418), whilst the `lower` kwarg is from a more recent post [here](https://github.com/JuliaStats/Distributions.jl/issues/1889)

I then set up the `Turing` model:

```julia
@model function truncated_poisson_1(obs_data, trunc_thres)
    # Prior for λ
    λ ~ Exponential(5)
    T = typeof(λ)
    
    # Likelihood for each data point, truncated Poisson
    for i in 1:length(obs_data)
        obs_data[i] ~ truncated(Distributions.Poisson(λ); lower = T(trunc_thres[i]) , upper = T(100))
    end
end

model_1 = truncated_poisson_1(poisson_data_truncated.obs_data, poisson_data_truncated.truncation)
chain_1 = sample(model_1, NUTS(), 10)

```

and got the following error:

```julia
MethodError: no method matching _gammalogccdf(::ForwardDiff.Dual{ForwardDiff.Tag{DynamicPPL.DynamicPPLTag, Float64}, Float64, 1}, ::ForwardDiff.Dual{ForwardDiff.Tag{DynamicPPL.DynamicPPLTag, Float64}, Float64, 1}, ::ForwardDiff.Dual{ForwardDiff.Tag{DynamicPPL.DynamicPPLTag, Float64}, Float64, 1})

```

1. Is there anything wrong with my code above? I can switch to `Gibbs(MH())` for sampling and it works - but the convergence is quote poor that I would much prefer `NUTS()` if possible

2. I see from the earlier [link](https://discourse.julialang.org/t/turing-error-with-truncated-distribution/53418) that it worked with the `Exponential` distribution? I understand from various posts that `NUTS()` uses AD which causes some issues, but I was under the impression from [Distributions.jl](https://juliastats.org/Distributions.jl/stable/truncate/) that the `logccdf` is available for all univariate distributions?

3. Seems like others from `StatsFuns.jl` and `Distributions.jl` tried doing something about it but then stopped [here](https://github.com/JuliaStats/StatsFuns.jl/issues/161) and [here](https://github.com/JuliaStats/Distributions.jl/issues/745).

As a brute-force - if I am to combine the use of `NUTS()` and to be able to use `predict` down the line, then I’ve come up with the below as a workaround:

```julia
@model function truncated_poisson_2(obs_data, trunc_thres, response_data)
    
    λ ~ Exponential(1)
    
    for i in 1:length(obs_data)
        if obs_data[i] > trunc_thres[i]
            # Poisson log-likelihood
            logp = logpdf(Poisson(λ), obs_data[i])
            
            tmp_cdf = zero(Float64)
            for j in 0:trunc_thres[i]
                tmp_cdf += pdf(Poisson(λ), j)
            end
            tmp_cdf = min(one(eltype(tmp_cdf)), tmp_cdf) # just in case

            trunc_logp = log1p(-tmp_cdf) # errors out if log1p(-cdf(Poisson(λ), trunc_thres[i]))??
            
            Turing.@addlogprob! logp - trunc_logp
        end
        # for predicting data
        if response_data[i] === missing
            response_data[i] ~ Poisson(λ)
        end
    end
end

model_2 = truncated_poisson_2(poisson_data_truncated.obs_data, poisson_data_truncated.truncation, poisson_data_truncated.obs_data)

chain_2 = sample(model_2, NUTS(), MCMCThreads(), 100, 2)

```

Quite messy. I’ll need to add on Censoring later on too so really hoping for something along the lines of:

```julia
response_data[I] ~ censored(truncated(Poisson(λ); lower = trunc_thres[I]), upper = cen_thres[I])

```

Lastly, I also tried `Normal` distribution - same syntax - seems to have worked fine?

```julia
n_samples = 100  
m_true = 3.0
s_true = 0.5
normal_data = rand(Normal(m_true, s_true), n_samples)

μ_threshold = 2.5
trunc_2 = rand(Normal(μ_threshold, s_true), n_samples)

normal_data_truncated = DataFrame(obs_data = normal_data, truncation = trunc_2)
normal_data_truncated = @subset(normal_data_truncated, :obs_data .>= :truncation)

@model function truncated_norm_1(obs_data, trunc_thres)
    m ~ Exponential(1)
    s ~ Exponential(0.1)
    for i in 1:length(obs_data)
        obs_data[i] ~ truncated(Normal(m, s), trunc_thres[i], 100)
    end
end

model_3 = truncated_norm_1(normal_data_truncated.obs_data, normal_data_truncated.truncation)
chain_3 = sample(model_3, NUTS(), 10)

```

---

<div class="post-metadata">

**Author:** ![Red-Portal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/red-portal/32/9102_2.png) [@Red-Portal](https://discourse.julialang.org/u/Red-Portal)\
**Post date:** [September 12, 2024, 1:19am UTC](https://discourse.julialang.org/t/turing-selective-truncated-distribution-error/118598/2 "2024-09-12T01:19:11Z")

</div>

Yeah seems like an unncessary type constraint upstream that is not playing nice with ForwardDiff. Try changing the backend. The last resort would be to use gradient-free MCMC algorithms. If it boils down to that, try SliceSampling.jl it should work much better than Metropolis-Hastings.
