# Is it possible to speed up PyMC sampling?

**URL:** <https://discourse.pymc.io/t/is-it-possible-to-speed-up-pymc-sampling/9453>\
**Category:** version agnostic\
**Created:** [May 20, 2022, 10:35am UTC](https://discourse.pymc.io/t/is-it-possible-to-speed-up-pymc-sampling/9453 "2022-05-20T10:35:21Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![jordan.howell2](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/jordan.howell2/32/1742_2.png) [@jordan.howell2](https://discourse.pymc.io/u/jordan.howell2)\
**Post date:** [May 20, 2022, 10:35am UTC](https://discourse.pymc.io/t/is-it-possible-to-speed-up-pymc-sampling/9453/1 "2022-05-20T10:35:21Z")

</div>

Hello.

I’m running a model on a fairly large dataset. I’ve just taken a job with unlimited use of the google cloud platform and thought, if I chose a higher performing CPU and higher memory, the sampling would noticeably speed up.

I haven’t seen that. The data I’m sampling has around 500,000 observations. It’s a time series based on daily data for five years with multiple items forecasted per year.

I’m currently using a 56core, 112GB RAM setup. No GPU as I’ve actually never ran PyMC with numpyro/jax on a GPU. Will the GPU help?

Do any of you have suggestions to try?

Thanks.

For Reference, here is my model and versions of programs:

```auto
with pm.Model(coords = coords) as model:
    
    # item_idx = pm.Data('item_idx', items, dims = "obs_id", mutable = False)
    
    k = pm.Normal('k', 0, 1)
    m = pm.Normal('m', 0, 5)
    delta = pm.Laplace('delta', 0, 0.1, shape = n_changepoints)
    
    growth = k + at.dot(A, delta)
    offset = m + at.dot(A, -s * delta)
    trend = growth * t + offset
    
# beta_weekly = pm.Normal('beta_weekly_seasonality', 0, 1, shape = weekly_n_components * 2)
# seasonality_weekly = at.dot(fourier(t, p = 7), beta_weekly)
    
# beta_monthly = pm.Normal('beta_monthly_seasonality', 0, 1, shape = monthly_n_components * 2)
# seasonality_monthly = at.dot(fourier(t, p = 30.5), beta_monthly)
    
# beta_yearly = pm.Normal('beta_yearly_seasonality', 0, 1, shape = yearly_n_components * 2)
# seasonality_yearly = at.dot(fourier(t, p = 365.25), beta_yearly)
    
    error = pm.HalfCauchy('sigma', .5)
    pm.Normal("predicted_sales", 
              trend, 
              error, 
              observed = train_y)
    
    trace = pymc.sampling_jax.sample_numpyro_nuts(tune=1000, chains = 4)

```

- PyMC/PyMC3 Version: 4.0.0b6
- Aesara/Theano Version: 2.5.1
- Python Version: 3.7.12
- Operating system: Debian via Google Cloud Platform
- How did you install PyMC/PyMC3: pip install pymc --pre ([Installation Guide (Linux) · pymc-devs/pymc Wiki · GitHub](https://github.com/pymc-devs/pymc/wiki/Installation-Guide-(Linux)))

---

<div class="post-metadata">

**Author:** ![cluhmann](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/cluhmann/32/3083_2.png) [@cluhmann](https://discourse.pymc.io/u/cluhmann)\
**Post date:** [May 20, 2022, 4:56pm UTC](https://discourse.pymc.io/t/is-it-possible-to-speed-up-pymc-sampling/9453/2 "2022-05-20T16:56:47Z")

</div>

In general, MCMC will use a CPU/core per chain, so cranking up the number of cores won’t improve speed. The results of sampling may be held in memory during sampling, so memory is necessary (and memory demands will be somewhat greater for large models) but not directly related to speed. Core speed (CPU or GPU) will impact sampling speed as will the geometry of your model (i.e., using difficult-to-sample priors, etc.). GPU can definitely help. See [here](https://martiningram.github.io/mcmc-comparison/) for an example.

---

<div class="post-metadata">

**Author:** ![jordan.howell2](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/jordan.howell2/32/1742_2.png) [@jordan.howell2](https://discourse.pymc.io/u/jordan.howell2)\
**Post date:** [May 20, 2022, 5:42pm UTC](https://discourse.pymc.io/t/is-it-possible-to-speed-up-pymc-sampling/9453/3 "2022-05-20T17:42:09Z")

</div>

So if I’m reading the git link correctly, all I have to do is the following and PyMC/JAX will use the GPU?

```auto
assert platform in ["cpu", "gpu"]

if platform == "cpu":
    # Disable GPU
    os.environ["CUDA_VISIBLE_DEVICES"] = ""

```

---

<div class="post-metadata">

**Author:** ![cluhmann](https://yyz2.discourse-cdn.com/flex036/user_avatar/discourse.pymc.io/cluhmann/32/3083_2.png) [@cluhmann](https://discourse.pymc.io/u/cluhmann)\
**Post date:** [May 20, 2022, 9:22pm UTC](https://discourse.pymc.io/t/is-it-possible-to-speed-up-pymc-sampling/9453/4 "2022-05-20T21:22:02Z")

</div>

I will let someone with more jax/gpu experience (@twiecki @ricardoV94 ?) weigh in on the implementation details. As [we have discussed](https://discourse.pymc.io/t/clear-pymc-workflow-for-beginners/9061), the user guide for v4/jax/gpu is not yet ready. But hopefully we can get you sorted out.
