From ada8e64e0640fe939c2857316ccb84637622e2e3 Mon Sep 17 00:00:00 2001 From: Simon Dirmeier Date: Wed, 1 Jul 2026 13:34:18 +0200 Subject: [PATCH] refactor: extract shared run_blackjax MCMC driver --- sbijax/_src/mcmc/irmh.py | 32 ++++++---------------- sbijax/_src/mcmc/mala.py | 32 ++++++---------------- sbijax/_src/mcmc/nuts.py | 31 ++++++---------------- sbijax/_src/mcmc/rmh.py | 32 ++++++---------------- sbijax/_src/mcmc/util.py | 50 ++++++++++++++++++++++++++++------- sbijax/_src/mcmc/util_test.py | 30 +++++++++++++++++++++ 6 files changed, 102 insertions(+), 105 deletions(-) create mode 100644 sbijax/_src/mcmc/util_test.py diff --git a/sbijax/_src/mcmc/irmh.py b/sbijax/_src/mcmc/irmh.py index 613081d..a29d255 100644 --- a/sbijax/_src/mcmc/irmh.py +++ b/sbijax/_src/mcmc/irmh.py @@ -2,7 +2,7 @@ import jax from jax import random as jr -from sbijax._src.mcmc.util import sample_and_post_process_from_blackjax_samples +from sbijax._src.mcmc.util import run_blackjax # ruff: noqa: PLR0913, D417 @@ -39,31 +39,15 @@ def sample_with_imh( a JAX pytree with keys corresponding to the variables names and tensor values of dimension `n_chains x n_samples x dim_variable` """ - - def _inference_loop(rng_key, kernel, initial_state, n_samples): - @jax.jit - def _step(states, rng_key): - keys = jax.random.split(rng_key, n_chains) - states, _ = jax.vmap(kernel)(keys, states) - return states, states - - sampling_keys = jax.random.split(rng_key, n_samples) - _, states = jax.lax.scan(_step, initial_state, sampling_keys) - return states - - init_key, rng_key = jr.split(rng_key) - initial_states, kernel = _mh_init(init_key, n_chains, prior, lp) - - samples = sample_and_post_process_from_blackjax_samples( + return run_blackjax( rng_key, - _inference_loop, - kernel, - initial_states, - n_chains, - n_samples, - n_warmup, + _mh_init, + prior, + lp, + n_chains=n_chains, + n_samples=n_samples, + n_warmup=n_warmup, ) - return samples # ruff: noqa: E731 diff --git a/sbijax/_src/mcmc/mala.py b/sbijax/_src/mcmc/mala.py index 0637da7..174cc0a 100644 --- a/sbijax/_src/mcmc/mala.py +++ b/sbijax/_src/mcmc/mala.py @@ -2,7 +2,7 @@ import jax from jax import random as jr -from sbijax._src.mcmc.util import sample_and_post_process_from_blackjax_samples +from sbijax._src.mcmc.util import run_blackjax # ruff: noqa: PLR0913, D417 @@ -39,31 +39,15 @@ def sample_with_mala( a JAX pytree with keys corresponding to the variables names and tensor values of dimension `n_chains x n_samples x dim_variable` """ - - def _inference_loop(rng_key, kernel, initial_state, n_samples): - @jax.jit - def _step(states, rng_key): - keys = jax.random.split(rng_key, n_chains) - states, _ = jax.vmap(kernel)(keys, states) - return states, states - - sampling_keys = jax.random.split(rng_key, n_samples) - _, states = jax.lax.scan(_step, initial_state, sampling_keys) - return states - - init_key, rng_key = jr.split(rng_key) - initial_states, kernel = _mala_init(init_key, n_chains, prior, lp) - - samples = sample_and_post_process_from_blackjax_samples( + return run_blackjax( rng_key, - _inference_loop, - kernel, - initial_states, - n_chains, - n_samples, - n_warmup, + _mala_init, + prior, + lp, + n_chains=n_chains, + n_samples=n_samples, + n_warmup=n_warmup, ) - return samples # pylint: disable=missing-function-docstring,no-member diff --git a/sbijax/_src/mcmc/nuts.py b/sbijax/_src/mcmc/nuts.py index f2d084a..441fbfd 100644 --- a/sbijax/_src/mcmc/nuts.py +++ b/sbijax/_src/mcmc/nuts.py @@ -2,7 +2,7 @@ import jax from jax import random as jr -from sbijax._src.mcmc.util import sample_and_post_process_from_blackjax_samples +from sbijax._src.mcmc.util import run_blackjax # ruff: noqa: PLR0913, D417 @@ -39,30 +39,15 @@ def sample_with_nuts( a JAX pytree with keys corresponding to the variables names and tensor values of dimension `n_chains x n_samples x dim_variable` """ - - def _inference_loop(rng_key, kernel, initial_state, n_samples): - @jax.jit - def _step(states, rng_key): - keys = jax.random.split(rng_key, n_chains) - states, _ = jax.vmap(kernel)(keys, states) - return states, states - - sampling_keys = jax.random.split(rng_key, n_samples) - _, states = jax.lax.scan(_step, initial_state, sampling_keys) - return states - - init_key, rng_key = jr.split(rng_key) - initial_states, kernel = _nuts_init(init_key, n_chains, prior, lp) - samples = sample_and_post_process_from_blackjax_samples( + return run_blackjax( rng_key, - _inference_loop, - kernel, - initial_states, - n_chains, - n_samples, - n_warmup, + _nuts_init, + prior, + lp, + n_chains=n_chains, + n_samples=n_samples, + n_warmup=n_warmup, ) - return samples # pylint: disable=missing-function-docstring diff --git a/sbijax/_src/mcmc/rmh.py b/sbijax/_src/mcmc/rmh.py index a2c51a0..e04f3fc 100644 --- a/sbijax/_src/mcmc/rmh.py +++ b/sbijax/_src/mcmc/rmh.py @@ -4,7 +4,7 @@ from jax import random as jr from jax._src.flatten_util import ravel_pytree -from sbijax._src.mcmc.util import sample_and_post_process_from_blackjax_samples +from sbijax._src.mcmc.util import run_blackjax # ruff: noqa: PLR0913, D417 @@ -41,31 +41,15 @@ def sample_with_rmh( a JAX pytree with keys corresponding to the variables names and tensor values of dimension `n_chains x n_samples x dim_variable` """ - - def _inference_loop(rng_key, kernel, initial_state, n_samples): - @jax.jit - def _step(states, rng_key): - keys = jax.random.split(rng_key, n_chains) - states, _ = jax.vmap(kernel)(keys, states) - return states, states - - sampling_keys = jax.random.split(rng_key, n_samples) - _, states = jax.lax.scan(_step, initial_state, sampling_keys) - return states - - init_key, rng_key = jr.split(rng_key) - initial_states, kernel = _mh_init(init_key, n_chains, prior, lp) - - samples = sample_and_post_process_from_blackjax_samples( + return run_blackjax( rng_key, - _inference_loop, - kernel, - initial_states, - n_chains, - n_samples, - n_warmup, + _mh_init, + prior, + lp, + n_chains=n_chains, + n_samples=n_samples, + n_warmup=n_warmup, ) - return samples # pylint: disable=missing-function-docstring,no-member diff --git a/sbijax/_src/mcmc/util.py b/sbijax/_src/mcmc/util.py index ca6fe99..489363f 100644 --- a/sbijax/_src/mcmc/util.py +++ b/sbijax/_src/mcmc/util.py @@ -3,6 +3,7 @@ import arviz as az import jax import xarray +from jax import random as jr def mcmc_diagnostics(samples: xarray.DataTree): @@ -10,18 +11,47 @@ def mcmc_diagnostics(samples: xarray.DataTree): return MCMCDiagnostics(az.rhat(samples), az.ess(samples)) +def _inference_loop(rng_key, kernel, initial_state, n_chains, n_samples): + @jax.jit + def _step(states, rng_key): + keys = jr.split(rng_key, n_chains) + states, _ = jax.vmap(kernel)(keys, states) + return states, states + + sampling_keys = jr.split(rng_key, n_samples) + _, states = jax.lax.scan(_step, initial_state, sampling_keys) + return states + + # ruff: noqa: PLR0913 -def sample_and_post_process_from_blackjax_samples( - rng_key, - inf_fn, - kernel, - initial_states, - n_chains, - n_samples, - n_warmup, -): +def run_blackjax(rng_key, init_fn, prior, lp, *, n_chains, n_samples, n_warmup): + """Draw samples from a distribution using a BlackJAX kernel. + + Constructs the initial chain states and kernel via ``init_fn``, runs a + vectorised (over chains) sampling loop, discards the warmup draws and + reshapes the result to ``n_chains x (n_samples - n_warmup) x dim``. + + Args: + rng_key: a jax random key + init_fn: a callable ``(rng_key, n_chains, prior, lp) -> + (initial_states, kernel_step)`` constructing the initial BlackJAX + chain states and the kernel step function + prior: a distribution to sample the initial chain positions from + lp: the logdensity to sample from + n_chains: number of chains to sample + n_samples: number of samples per chain + n_warmup: number of samples to discard + + Returns: + a JAX pytree with keys corresponding to the variable names and tensor + values of dimension ``n_chains x (n_samples - n_warmup) x dim_variable`` + """ + init_key, sample_key = jr.split(rng_key) + initial_states, kernel = init_fn(init_key, n_chains, prior, lp) first_key = list(initial_states.position.keys())[0] - states = inf_fn(rng_key, kernel, initial_states, n_samples) + states = _inference_loop( + sample_key, kernel, initial_states, n_chains, n_samples + ) _ = states.position[first_key].block_until_ready() thetas = jax.tree_util.tree_map( lambda x: x[n_warmup:, ...].reshape(n_chains, n_samples - n_warmup, -1), diff --git a/sbijax/_src/mcmc/util_test.py b/sbijax/_src/mcmc/util_test.py new file mode 100644 index 0000000..3205043 --- /dev/null +++ b/sbijax/_src/mcmc/util_test.py @@ -0,0 +1,30 @@ +# pylint: skip-file + +import blackjax as bj +import chex +import jax +from jax import random as jr + +from sbijax._src.mcmc.util import run_blackjax + + +def _mala_init(rng_key, n_chains, prior, lp): + initial_positions = prior.sample(seed=rng_key, sample_shape=(n_chains,)) + kernel = bj.mala(lp, 0.1) + initial_state = jax.vmap(kernel.init)(initial_positions) + return initial_state, kernel.step + + +def test_run_blackjax_returns_chain_shaped_samples(prior_log_prob_tuple): + prior_fn, lp = prior_log_prob_tuple + samples = run_blackjax( + jr.PRNGKey(0), + _mala_init, + prior_fn(), + lp, + n_chains=8, + n_samples=200, + n_warmup=100, + ) + chex.assert_shape(samples["mean"], (8, 100, 2)) + chex.assert_shape(samples["std"], (8, 100, 1))