| name | PyMC Samplers |
| description | Expert on PyMC MCMC sampling methods including NUTS, HMC, Metropolis variants, and pm.sample() API. Use for sampling errors, convergence issues, sampler configuration, or trace-related problems. |
PyMC Samplers Skill
You are an expert in PyMC sampling methods, helping migrate notebook code to work with the current stable PyMC version.
Core Sampling Functions
pm.sample() - Main Sampling Function
The primary function for MCMC sampling:
import pymc as pm
with pm.Model() as model:
trace = pm.sample(
draws=1000,
tune=1000,
chains=4,
cores=4,
random_seed=None
)
Full pm.sample() Parameters
trace = pm.sample(
draws=1000,
tune=1000,
chains=4,
cores=None,
random_seed=None,
step=None,
initvals=None,
init="auto",
n_init=200000,
progressbar=True,
return_inferencedata=True,
idata_kwargs=None,
compute_convergence_checks=True,
discard_tuned_samples=True,
target_accept=0.8,
callback=None,
)
Sampling Step Methods
NUTS (No-U-Turn Sampler)
Default and recommended for most models with continuous variables:
trace = pm.sample(1000)
trace = pm.sample(
1000,
target_accept=0.95,
max_treedepth=10
)
from pymc.step_methods import NUTS
with model:
step = NUTS()
trace = pm.sample(1000, step=step)
Metropolis Samplers
For simpler models or discrete variables:
from pymc.step_methods import Metropolis
with model:
step = Metropolis()
trace = pm.sample(5000, step=step)
step = Metropolis(proposal_dist=pm.NormalProposal)
BinaryMetropolis
Optimized for binary variables:
from pymc.step_methods import BinaryMetropolis
with model:
x = pm.Bernoulli("x", p=0.5)
step = BinaryMetropolis([x])
trace = pm.sample(5000, step=step)
Slice Sampler
For unimodal distributions:
from pymc.step_methods import Slice
with model:
step = Slice()
trace = pm.sample(5000, step=step)
HamiltonianMC
Standard Hamiltonian Monte Carlo:
from pymc.step_methods import HamiltonianMC
with model:
step = HamiltonianMC()
trace = pm.sample(1000, step=step)
CompoundStep - Multiple Samplers
Combine different samplers for different variables:
from pymc.step_methods import NUTS, BinaryMetropolis, CompoundStep
with model:
continuous_vars = [mu, sigma]
discrete_vars = [z]
step1 = NUTS(vars=continuous_vars)
step2 = BinaryMetropolis(vars=discrete_vars)
step = CompoundStep([step1, step2])
trace = pm.sample(1000, step=step)
Prior and Posterior Predictive Sampling
Prior Predictive Checks
Sample from the prior before seeing data:
with model:
prior_predictive = pm.sample_prior_predictive(
samples=500,
random_seed=None
)
import arviz as az
az.plot_ppc(prior_predictive, group="prior")
Posterior Predictive Checks
Sample predictions after fitting:
with model:
trace = pm.sample(1000)
posterior_predictive = pm.sample_posterior_predictive(
trace,
var_names=None,
random_seed=None
progressbar=True
)
az.plot_ppc(posterior_predictive)
Out-of-Sample Predictions
with model:
pm.set_data({"X": X_test})
predictions = pm.sample_posterior_predictive(
trace,
var_names=["y"]
)
Initialization Methods
Initialization Strategies
trace = pm.sample(1000, init="auto")
trace = pm.sample(1000, init="jitter+adapt_diag")
trace = pm.sample(1000, init="jitter+adapt_full")
trace = pm.sample(1000, init="adapt_diag")
trace = pm.sample(1000, init="adapt_diag_grad")
Custom Initial Values
initvals = {
"mu": 0.5,
"sigma": 1.0
}
trace = pm.sample(1000, initvals=initvals)
NUTS Initialization Helper
from pymc import init_nuts
with model:
init_vals, step = init_nuts(
init="auto",
chains=4,
random_seed=None
progressbar=True
)
trace = pm.sample(1000, step=step, initvals=init_vals)
Drawing Single Samples
with model:
sample = pm.draw(mu, draws=1)
samples = pm.draw(mu, draws=100)
sample_dict = pm.draw([mu, sigma], draws=100)
Computing Deterministics
After sampling, compute deterministic quantities:
with model:
trace = pm.sample(1000)
trace = pm.compute_deterministics(trace)
JAX-Based Samplers
BlackJAX NUTS
from pymc.sampling.jax import sample_blackjax_nuts
with model:
trace = sample_blackjax_nuts(
draws=1000,
tune=1000,
chains=4,
target_accept=0.8
)
NumPyro NUTS
from pymc.sampling.jax import sample_numpyro_nuts
with model:
trace = sample_numpyro_nuts(
draws=1000,
tune=1000,
chains=4,
target_accept=0.8
)
Common Migration Issues
PyMC3 → latest PyMC version
- Return type changes
trace = pm.sample(1000)
idata = pm.sample(1000, return_inferencedata=True)
idata.posterior
trace = pm.sample(1000, return_inferencedata=False)
- Step method imports
from pymc3.step_methods import NUTS, Metropolis
from pymc.step_methods import NUTS, Metropolis
- Tuning parameter
trace = pm.sample(draws=1000, n_tune=500)
trace = pm.sample(draws=1000, tune=500)
- Initialization
x = pm.Normal("x", mu=0, sigma=1, testval=0.5)
x = pm.Normal("x", mu=0, sigma=1)
trace = pm.sample(1000, initvals={"x": 0.5})
Sampling Best Practices
- Use NUTS for continuous models - It's efficient and robust
- Run multiple chains - At least 4 chains to check convergence
- Check convergence - Inspect Rhat (< 1.1) and ESS (> 200)
- Tune adequately - Default 1000 steps usually sufficient
- Adjust target_accept if needed - Increase if divergences occur
- Use random_seed - For reproducibility
- Prior predictive checks - Always check priors make sense
- Posterior predictive checks - Validate model fit
Troubleshooting Sampling Issues
Divergences
trace = pm.sample(1000, target_accept=0.95)
trace = pm.sample(1000, max_treedepth=15)
Slow Sampling
trace = pm.sample(1000, chains=2, cores=2)
trace = sample_numpyro_nuts(1000)
trace = pm.sample(1000, tune=500)
Poor Mixing
trace = pm.sample(1000, tune=2000)
trace = pm.sample(1000, init="adapt_diag")
High Rhat or Low ESS
trace = pm.sample(5000, tune=2000)
trace = pm.sample(1000, chains=8)
Example Usage
import pymc as pm
import numpy as np
import arviz as az
np.random.seed(42)
true_mu = 3.0
true_sigma = 1.5
data = np.random.normal(true_mu, true_sigma, size=100)
with pm.Model() as model:
mu = pm.Normal("mu", mu=0, sigma=10)
sigma = pm.HalfNormal("sigma", sigma=5)
y = pm.Normal("y", mu=mu, sigma=sigma, observed=data)
prior_pred = pm.sample_prior_predictive(samples=500)
trace = pm.sample(
draws=1000,
tune=1000,
chains=4,
cores=4,
target_accept=0.9,
random_seed=None
)
post_pred = pm.sample_posterior_predictive(trace)
print(az.summary(trace, var_names=["mu", "sigma"]))
az.plot_trace(trace, var_names=["mu", "sigma"])
az.plot_posterior(trace, var_names=["mu", "sigma"])
az.plot_ppc(post_pred)
When to Use This Skill
- Configuring sampling parameters
- Fixing divergences or sampling issues
- Choosing appropriate step methods
- Converting PyMC3 sampling code
- Implementing prior/posterior predictive checks
- Troubleshooting convergence problems
- Optimizing sampling performance
- Working with mixed discrete/continuous models