# Spike-and-Slab Implementation using Turing.jl?

**URL:** <https://discourse.julialang.org/t/spike-and-slab-implementation-using-turing-jl/132816>\
**Category:** New to Julia\
**Tags:** question, turing\
**Created:** [October 2, 2025, 5:24am UTC](https://discourse.julialang.org/t/spike-and-slab-implementation-using-turing-jl/132816 "2025-10-02T05:24:52Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![Debartha\_Paul](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/debartha_paul/32/202192_2.png) [@Debartha\_Paul](https://discourse.julialang.org/u/Debartha_Paul)\
**Post date:** [October 2, 2025, 5:24am UTC](https://discourse.julialang.org/t/spike-and-slab-implementation-using-turing-jl/132816/1 "2025-10-02T05:24:52Z")

</div>

So I was looking for an implementation of spike-and-slab prior using the Turing.jl package. Say I have a model as follows:

y\_{ig} \sim N(\theta\_g, \sigma^2) where g = 1, \cdots, G and i = 1, \cdots, n\_g independently  
And I would like the priors on \theta\_g to be independently a spike-and-slab as \pi\delta\_0 + (1-\pi)t\_v, which is a mixture of a degenerate mass at 0 and a central t distribution.

Can somebody give me some guidance as to how I can do this?

---

<div class="post-metadata">

**Author:** ![eteppo](https://avatars.discourse-cdn.com/v4/letter/e/90db22/32.png) [@eteppo](https://discourse.julialang.org/u/eteppo)\
**Post date:** [October 2, 2025, 9:31am UTC](https://discourse.julialang.org/t/spike-and-slab-implementation-using-turing-jl/132816/2 "2025-10-02T09:31:54Z")

</div>

Distributions.jl has a Dirac distribution but it didn’t seem to work here. One option could be something like this?

```julia-auto
using Turing
@model function M₁(g; G = maximum(g))
	ν ~ Exponential(10)
	π ~ Beta(1, 1)
	θ ~ filldist(MixtureModel([Normal(0, 0.001), TDist(ν)], [π, 1-π]), G)
	σ² ~ Exponential(1)
	y ~ MvNormal(θ[g], σ²*I)
end

```

---

<div class="post-metadata">

**Author:** ![bertschi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/bertschi/32/33462_2.png) [@bertschi](https://discourse.julialang.org/u/bertschi)\
**Post date:** [October 2, 2025, 6:22pm UTC](https://discourse.julialang.org/t/spike-and-slab-implementation-using-turing-jl/132816/3 "2025-10-02T18:22:00Z")

</div>

While you can probably define a spike-and-slab prior somehow, it will be hard to sample from – due to its non-differentiable density at the location of the point measure. There are some modern differentiable alternatives/approximation available, such as the horseshoe prior. [Betancourt](https://betanalpha.github.io/assets/case_studies/modeling_sparsity.html) has a very detailed discussion of sparsity priors and their properties.

---

<div class="post-metadata">

**Author:** ![simonsteiger](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/simonsteiger/32/209862_2.png) [@simonsteiger](https://discourse.julialang.org/u/simonsteiger)\
**Post date:** [May 29, 2026, 9:47am UTC](https://discourse.julialang.org/t/spike-and-slab-implementation-using-turing-jl/132816/4 "2026-05-29T09:47:17Z")

</div>

Sorry for reviving this old topic. I tried to write a spike-and-slab prior for a homework assignment and the model below returns results that agree with a horseshoe implementation of the same model. I thought I’d add it for posterity:

```julia
using Turing, Statistics, LinearAlgebra
using Downloads, CSV, DataFrames

standardize(x) = (x .- mean(x)) ./ std(x)

data = let
	path = "https://hastie.su.domains/ElemStatLearn/datasets/prostate.data"
	file = Downloads.download(path)
	tmp = CSV.read(file, DataFrame; delim='\t', header=1, drop=[1])
	covariates = names(tmp)[1:8]
	transform(tmp, covariates .=> standardize => identity)
end

@model function turing_spikeslab(y, X)
	n, p = size(X)
	
	ω ~ Beta(1, 1)
	β₀ ~ Normal(0, 100)
	β ~ filldist(Normal(0, 1), p) 
	z ~ filldist(Bernoulli(ω), p)
	σ² ~ InverseGamma(3, 2)
	
	μ = β₀ .+ X * (β .* z)
	y ~ MvNormal(μ, σ² * I)
end

model = turing_spikeslab(data.lpsa, Matrix(data[:, 1:8]))
sampler = Gibbs(:z => PG(20), (:ω, :β₀, :β, :σ²) => NUTS())
chain = sample(model, sampler, MCMCThreads(), 1500, 4; discard_initial=500)

```
