Skip to content

MCMC Inference

Internal API reference for filter-based MCMC/SGMCMC inference orchestration.

MCMCInference

Provides a high-level interface for MCMC inference, consistent between NumPyro and BlackJAX backends.

Models must take in obs_times, obs_values, ctrl_times, ctrl_values as arguments (and optionally, *model_args, **model_kwargs).

Attributes:

Name Type Description
mcmc_config

Sampler configuration dataclass (NUTSConfig, HMCConfig, AdaptiveMetropolisConfig, SGLDConfig, or MALAConfig).

model

Callable probabilistic model with signature model(obs_times=..., obs_values=..., ctrl_times=..., ctrl_values=..., *model_args, **model_kwargs).

get_diagnostics() -> dict[str, jax.Array]

Return compact diagnostics from the most recent successful run.

NUTS reports mean_acceptance_rate and num_divergences per chain. Adaptive Metropolis reports mean_acceptance_rate and final_proposal_scale per chain and unconstrained coordinate.

Raises:

Type Description
RuntimeError

If inference has not completed successfully.

run(rng_key: jnp.ndarray, obs_times: jnp.ndarray, obs_values: jnp.ndarray, ctrl_times: jnp.ndarray | None = None, ctrl_values: jnp.ndarray | None = None, *model_args, **model_kwargs) -> dict

Run inference and return posterior samples.

Parameters:

Name Type Description Default
rng_key ndarray

JAX PRNG key.

required
obs_times ndarray

Observation times.

required
obs_values ndarray

Observation values.

required
ctrl_times ndarray | None

Control times.

None
ctrl_values ndarray | None

Control values.

None
*model_args

Additional positional arguments passed to model.

()
**model_kwargs

Additional keyword arguments passed to model.

{}

Returns:

Type Description
dict

Dict-like pytree of posterior samples.

_blackjax_mcmc(mcmc_config: BaseMCMCConfig, rng_key: jnp.ndarray, model: Callable, obs_times: jnp.ndarray, obs_values: jnp.ndarray, ctrl_times: jnp.ndarray | None = None, ctrl_values: jnp.ndarray | None = None, *model_args, **model_kwargs) -> tuple[dict, dict[str, jax.Array]]

Run BlackJAX inference via the BlackJAX integration module.

_numpyro_mcmc(mcmc_config: BaseMCMCConfig, rng_key: jnp.ndarray, model: Callable, obs_times: jnp.ndarray, obs_values: jnp.ndarray, ctrl_times: jnp.ndarray | None = None, ctrl_values: jnp.ndarray | None = None, *model_args, **model_kwargs) -> tuple[dict, dict[str, jax.Array]]

Run NumPyro MCMC and return samples plus compact diagnostics.