Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 8 additions & 24 deletions sbijax/_src/mcmc/irmh.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
32 changes: 8 additions & 24 deletions sbijax/_src/mcmc/mala.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
31 changes: 8 additions & 23 deletions sbijax/_src/mcmc/nuts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
32 changes: 8 additions & 24 deletions sbijax/_src/mcmc/rmh.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
50 changes: 40 additions & 10 deletions sbijax/_src/mcmc/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,25 +3,55 @@
import arviz as az
import jax
import xarray
from jax import random as jr


def mcmc_diagnostics(samples: xarray.DataTree):
MCMCDiagnostics = namedtuple("MCMCDiagnostics", "rhat ess")
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),
Expand Down
30 changes: 30 additions & 0 deletions sbijax/_src/mcmc/util_test.py
Original file line number Diff line number Diff line change
@@ -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))
Loading