# NUTS speed is very slow for high dimension parameter inference in Turing.jl

**URL:** <https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349>\
**Category:** Probabilistic Programming\
**Tags:** turing\
**Created:** [May 2, 2022, 2:03am UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349 "2022-05-02T02:03:20Z")\
**Posts on this page:** 9\
**Page:** 1

<div class="post-metadata">

**Author:** ![zqn](https://avatars.discourse-cdn.com/v4/letter/z/e480ec/32.png) [@zqn](https://discourse.julialang.org/u/zqn)\
**Post date:** [May 2, 2022, 2:03am UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/1 "2022-05-02T02:03:20Z")

</div>

Hi, I’m trying Turing.jl for a high dimension hyperparameter inference problem, and here’s a toy example. To estimate tureE using fake testx and testy, I’m using a NUTS which costs no less than 5 hours according to the progress meter. I’ve changed the backed to reversediff and used a multivariate normal distribution. Was wondering if I made any mistakes to build the model? Many thanks.

```julia
testx = rand(1000,10000) 
trueE = rand(10000)
testy = testx*trueE
@model toy(x,y) = begin
    sig ~ InverseGamma(1,1)
    M ~ MvNormal(10000,sig)
    y ~ MvNormal(x*M,sig)
end
model = toy(testx,testy)
@time chain = sample(model, NUTS(0.65), 1000);

```

---

<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:** [May 2, 2022, 8:04am UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/2 "2022-05-02T08:04:46Z")

</div>

No obvious mistakes. In general though, you should not expect large models to sample quickly. 5 hours to sample 10,000 parameters does not sound unrealistic. For models like this though, where the compute time will be dominated by computing the gradient of the matrix-vector multiplication, I expect Zygote will perform better than ReverseDiff.

---

<div class="post-metadata">

**Author:** ![danielw2904](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/danielw2904/32/10890_2.png) [@danielw2904](https://discourse.julialang.org/u/danielw2904)\
**Post date:** [May 2, 2022, 8:11am UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/3 "2022-05-02T08:11:01Z")

</div>

What about `Turing.setrdcache(true)` in this case?

This worked well for me in a different situation:

> [@Turing.jl for Causal Inference model](https://discourse.julialang.org/t/turing-jl-for-causal-inference-model/75650/19):
>
> Here is my benchmark for two runs of each: Turing.setadbackend(:forwarddiff) Iterations = 1001:1:2000 Number of chains = 4 Samples per chain = 1000 Wall duration = 102.14 seconds Compute duration = 361.99 seconds Iterations = 1001:1:2000 Number of chains = 4 Samples per chain = 1000 Wall duration = 95.9 seconds Compute duration = 333.32 seconds Turing.setadbackend(:tracker) Iterations = 1001:1:2000 Number of chains = 4 Samples per chain = 1000 Wall duration…

---

<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:** [May 2, 2022, 1:10pm UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/4 "2022-05-02T13:10:19Z")

</div>

Two things.

**AD Backend** First, I think the default backend for Turing is still `forwarddiff`, which should be very inefficient for a high-dimensional model like yours. I just tried your model on my laptop with `reversediff` as the AD backend and progressmeter shows 2:30:16,

**Hierarchical Prior** Second, `InverseGamma(1,1)` is too weak and should result in hard-to-navigate tails. This will result in NUTS choosing long integration trajectories, and the length of the trajectory is roughly proportional to the iteration complexity. The motivation behind people using the inverse gamma is simply because it is a conjugate prior to the normal, which is irrelevant to MCMC, so we’re free to use alternative priors. A better choice is to use a more informative prior with lighter tails like the `truncated(Normal(0, 10), 0, Inf)`. This quickly reduced the projected sampling time to 1:20:45. See [the prior choice wiki](https://github.com/stan-dev/stan/wiki/Prior-Choice-Recommendations) for an up-to-date recommendation list by the Stan people.

**Forcing Short Trajectories** The last measure would be to reduce the max tree depth as `NUTS(0.65, max_depth=8)`. Combined with the informative prior above, progressmeter shows 0:40:18. The default parameter is 10 which results in a maximum of 2^10 leapfrog steps while 8 will result in 2^8. This measure, however, will negatively affect the statistical efficiency of the sampler, it’s recommended to instead fix the model so that the sampler does not hit the maximum limit. (But problems that are fundamentally hard to infer do exist, like sparse regression, stochastic volatility, etc… These are an open challenge to modern inference algorithms. So as an end-user, there is not much we can do about these…)

One of the difficulties with the current Bayesian workflow is that model design is not entirely independent from inference. It actually strongly affects the sampler’s performance both statistically and computationally. So you should tweak the model so that NUTS can do it’s job quickly and efficiently. Mike Betancourt’s blog have a lot of good guidelines on these aspects. See for example: [Identity Crisis](https://betanalpha.github.io/assets/case_studies/identifiability.html)

---

<div class="post-metadata">

**Author:** ![zqn](https://avatars.discourse-cdn.com/v4/letter/z/e480ec/32.png) [@zqn](https://discourse.julialang.org/u/zqn)\
**Post date:** [May 3, 2022, 2:16pm UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/6 "2022-05-03T14:16:26Z")

</div>

Many thanks for your detailed reply and it indeed reminds me of the pain points of current common Bayesian methods which I should spend more time to improve the model itself. I guess there are not many things I could do currently with Turing’s samplers.

---

<div class="post-metadata">

**Author:** ![zqn](https://avatars.discourse-cdn.com/v4/letter/z/e480ec/32.png) [@zqn](https://discourse.julialang.org/u/zqn)\
**Post date:** [May 3, 2022, 2:25pm UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/7 "2022-05-03T14:25:57Z")

</div>

Indeed after I changed to Zygote, I got a worse speed and I’m not clear how Turing choose to calculate the log joint and seems it tried to calculate the fully conditional posterior given any specific prior which is not an easy job for a high dimensional data.

---

<div class="post-metadata">

**Author:** ![zqn](https://avatars.discourse-cdn.com/v4/letter/z/e480ec/32.png) [@zqn](https://discourse.julialang.org/u/zqn)\
**Post date:** [May 3, 2022, 2:27pm UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/8 "2022-05-03T14:27:06Z")

</div>

It helps indeed and many thanks for your suggestion but the bottleneck here is probably the model itself as discussed in other posts here.

---

<div class="post-metadata">

**Author:** ![marius311](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/marius311/32/3953_2.png) [@marius311](https://discourse.julialang.org/u/marius311)\
**Post date:** [May 4, 2022, 8:02pm UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/9 "2022-05-04T20:02:30Z")

</div>

Do you care about the posterior on `M` or do you just need to marginalize it and only care about `sig`? If the latter, this is a perfect example problem for [https://cosmicmar.com/MuseInference.jl](https://cosmicmar.com/MuseInference.jl). You can beat the NUTS runtime by 100X or more.

The MUSE [algorithm](https://arxiv.org/abs/2112.09354) gives a Gaussian approximation to the marginal `sig` posterior, but is very accurate for high-dimensional latent spaces (due to central limit theorem) and is exact up to MC error if the likelihood is Gaussian (your case happens to be both).

Here’s an example code where I reduced the latent dimensionality to 1000 to give something easier to run. I’m getting about 100X faster than NUTS, and the relative improvement will get even more as you increase the dimensionality to your original 10000.

```julia
using MuseInference, Turing, Zygote, PyPlot
Turing.setadbackend(:zygote)

x = rand(1000,1000)
trueE = rand(1000)
testy = x * trueE
# you'll have to define the model without the observations
# in the arguments and instead use newer-style conditioning via `|`
@model function toy()
    sig ~ InverseGamma(1,1)
    M ~ MvNormal(1000,sig)
    y ~ MvNormal(x*M,sig)
end
model = toy() | (y=testy,)

# NUTS (~1000sec)
@time chain = sample(
    model, NUTS(100, 0.65), 100, progress=true
)

# MUSE (~10sec)
result = muse(
    model, (sig=0.5,), get_covariance=true,
    nsims=30, θ_rtol=1e-1, ∇z_logLike_atol=1,
)

# comparison plot
hist(chain[:sig], density=true)
sigs = range(xlim()...,length=1000)
plot(sigs, pdf.(result.dist, sigs))

```

![plot_86](https://global.discourse-cdn.com/julialang/original/3X/8/f/8fcbaaa7eac14590410034decf0bc21e256b46b4.png)

MUSE takes a starting guess for the `sig` value, which you can refine as you do more runs. The main parameters are the number of sims and tolerances. The error on the estimated mean relative to the uncertianty goes like `1/sqrt(nsims)`, so the like-to-like comparison sets this to the ESS of the chain, which I did above. `θ_rtol` is a solver tolerance on `sig` relative to its uncertainty and `∇z_logLike_atol` is the absolute tolerance on the gradient w.r.t. `M` which appears in an internal maximization over `M` that happens.

Let me know if you try it out if you run into any issues!

---

<div class="post-metadata">

**Author:** ![zqn](https://avatars.discourse-cdn.com/v4/letter/z/e480ec/32.png) [@zqn](https://discourse.julialang.org/u/zqn)\
**Post date:** [May 13, 2022, 4:38am UTC](https://discourse.julialang.org/t/nuts-speed-is-very-slow-for-high-dimension-parameter-inference-in-turing-jl/80349/11 "2022-05-13T04:38:46Z")

</div>

So sorry for the late reply due to the final week. I appreciate a lot for this interesting approximation methods but I’m afraid the posterior of M are what I’m interested in and seems I couldn’t run your sample code in Julia 1.6 environment.
