Skip to main content

Sequential Monte Carlo in JAX: particle filters, tempered SMC, and SMC2 on CPU, CUDA, TPU, and Apple-silicon GPUs

Project description

smcx

CI PyPI License

Sequential Monte Carlo in JAX: particle filters, adaptive tempered SMC, and SMC² with a small, flat API. Runs on CPU, CUDA, and TPU through stock JAX, and on Apple-silicon GPUs through the optional jax-mps backend.

Install

pip install smcx            # CPU / CUDA / TPU via your jax install
pip install "smcx[metal]"   # + jax-mps for Apple-silicon GPUs

What's in the box

  • Filters: bootstrap_filter, guided_filter (general g·f/q proposal weights), auxiliary_filter (twisted potentials), and liu_west_filter (joint state–parameter, labeled approximate).
  • Static targets: temper — adaptive tempered SMC with an ESS-bisection schedule and covariance-adapted random-walk moves.
  • Parameter inference: smc2 — nested SMC² with vmapped inner filters and PMMH rejuvenation.
  • Resampling: systematic, stratified, multinomial, and residual — one probability-space contract and float32-safe query grids; filters retain their weights in the log domain.
  • Diagnostics: ESS traces, quantile tail-ESS, Pareto-k reliability, single-run log-evidence variance from the genealogy (Lee & Whiteley 2018), trajectory reconstruction, CRPS, cumulative log score, Bayes factors, posterior-predictive sampling, and a one-call diagnose summary.
  • store_history=False on every filter drops memory from O(T·N) to O(N) with a bit-identical evidence estimate.

Permanent SMC/PF tests use dependency-free mathematical oracles and fixed, Monte-Carlo-calibrated gates. The outside SMC/PF implementations were run once in isolated environments; their pinned sources, licenses, settings, and numerical summaries are retained beside the corresponding unit tests, not as test dependencies. Diagnostics retain their separately documented ArviZ cross-validation dependency.

Component Permanent oracle One-time independent comparison
Resampling Exact offspring moments particles, BlackJAX
Bootstrap / APF / guided Kalman evidence and moments particles; TFP where applicable
Liu–West Conjugate posterior nimbleSMC/NIMBLE under matched semantics
Tempering Conjugate evidence and moments particles, BlackJAX
SMC² Grid-converged Kalman integral particles; limited TFP diagnostic

Quick start

import jax.numpy as jnp
import jax.random as jr

import smcx

# A 1-D linear-Gaussian state-space model.
A, Q, R = 0.9, 0.5, 0.3


def init(key, n):
    return jr.normal(key, (n, 1))


def transition(key, z):
    return A * z + jnp.sqrt(Q) * jr.normal(key, z.shape)


def log_observation(y, z):
    return -0.5 * (jnp.log(2 * jnp.pi * R) + (y[0] - z[0]) ** 2 / R)


def emission(key, z):
    return z + jnp.sqrt(R) * jr.normal(key, z.shape)


_, emissions = smcx.simulate(
    jr.key(1),
    lambda key: init(key, 1)[0],
    transition,
    emission,
    num_timesteps=100,
)

post = smcx.bootstrap_filter(
    jr.key(0),
    init,
    transition,
    log_observation,
    emissions,
    num_particles=10_000,
)
post.marginal_loglik  # log of an unbiased evidence estimate
smcx.diagnose(post)  # ESS / diversity / Pareto-k health summary

Callbacks are per-particle; smcx vmaps them internally. Everything takes an explicit PRNG key, and posteriors are NamedTuples — ordinary JAX pytrees. The four filters and simulate accept a keyword-only inputs sequence for controlled dynamics and covariate-driven observations; input-aware callbacks receive the aligned input_t as their final argument.

smcx is deliberately just the inference engine: it defines no model classes and no distributions. Models enter as JAX callables — your own closures, or thin wrappers around a model library such as Dynamax (the example notebook uses Dynamax models this way; unit tests use frozen dependency-free oracles).

Apple silicon

The [metal] extra runs the same code on M-series GPUs via jax-mps. Filter correctness on Metal is gate-verified in this repository's test suite (SMCX_TEST_PLATFORM=mps runs it on the GPU), and several of the performance fixes that make the backend fast for SMC-shaped workloads were contributed upstream from this project (jax-mps #215, #216, #220). Metal is float32-only; the suite runs float64 on CPU and float32 on Metal automatically.

Development

Requires Python 3.11+ and uv.

git clone https://github.com/michaelellis003/smcx.git
cd smcx
uv sync
uv run pre-commit install
uv run pre-commit install --hook-type commit-msg
uv run pre-commit install --hook-type pre-push

A Makefile covers common tasks:

make test        # lint + pytest
make lint        # ruff check, format check, license headers, ty
make format      # add license headers, ruff format, ruff fix
make docs        # build docs

Releases are automated: python-semantic-release reads conventional commits on merge to main, bumps the version, tags, and publishes.

Acknowledgments

smcx's design draws on the SMC ecosystem: particles and Chopin & Papaspiliopoulos's An Introduction to Sequential Monte Carlo (the Feynman-Kac architecture), BlackJAX (the resampling contract), Dynamax (container conventions), TensorFlow Probability (criterion/trace hooks), and design lessons from PyMC, FilterPy, pfilter, pyfilter, Stone Soup, pomp, nimbleSMC, and ArviZ. Algorithm docstrings and validation tests carry the formal references and immutable implementation pins.

License

Apache-2.0

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

smcx-1.2.1.tar.gz (39.0 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

smcx-1.2.1-py3-none-any.whl (48.6 kB view details)

Uploaded Python 3

File details

Details for the file smcx-1.2.1.tar.gz.

File metadata

  • Download URL: smcx-1.2.1.tar.gz
  • Upload date:
  • Size: 39.0 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.13

File hashes

Hashes for smcx-1.2.1.tar.gz
Algorithm Hash digest
SHA256 9fcd30014f4c5c782ed815c3c4677016e760b874adc80e2720fd015c92eaca2f
MD5 cf2264ee27cb54516e7ab44d96424d28
BLAKE2b-256 ed889a458c65799ddf884957c0453c9cc8e5136ec23d578823801244f47e2965

See more details on using hashes here.

Provenance

The following attestation bundles were made for smcx-1.2.1.tar.gz:

Publisher: release.yml on michaelellis003/smcx

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file smcx-1.2.1-py3-none-any.whl.

File metadata

  • Download URL: smcx-1.2.1-py3-none-any.whl
  • Upload date:
  • Size: 48.6 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.13

File hashes

Hashes for smcx-1.2.1-py3-none-any.whl
Algorithm Hash digest
SHA256 f8a38e28c4f945e19c945d84fb27ad80c38c67a968d95095ceb8444a5d61ba1c
MD5 4662bc606afb35470b5d7406b4c44def
BLAKE2b-256 8954886d3fe325b47014729cc93e32cf8bab6ab0968bf72360405e3fbe132e64

See more details on using hashes here.

Provenance

The following attestation bundles were made for smcx-1.2.1-py3-none-any.whl:

Publisher: release.yml on michaelellis003/smcx

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page