Sequential Monte Carlo in JAX: particle filters, tempered SMC, and SMC2 on CPU, CUDA, TPU, and Apple-silicon GPUs
Project description
smcx
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), andliu_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 contract, log-domain weights throughout, float32-safe query grids.
- 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
diagnosesummary. store_history=Falseon every filter drops memory from O(T·N) to O(N) with a bit-identical evidence estimate.
Every sampler is validated against exact references — Kalman oracles for the filters, conjugate evidence for tempering, grid-integrated posteriors for SMC² — with Monte-Carlo-calibrated gates, not loose tolerances.
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 # unbiased evidence estimate (log-domain)
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 test suite itself uses Dynamax models this way).
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. See CITATION.cff for formal references.
License
Apache-2.0
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file smcx-1.2.0.tar.gz.
File metadata
- Download URL: smcx-1.2.0.tar.gz
- Upload date:
- Size: 37.7 kB
- Tags: Source
- Uploaded using Trusted Publishing? Yes
- Uploaded via: twine/6.1.0 CPython/3.13.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
9e73764e28b7c668adcd561a343c4948048d09849ae63a17d0a1049c5acea81e
|
|
| MD5 |
ad77ba27726c87ed61f425ea9cee650d
|
|
| BLAKE2b-256 |
f45fdea41383e7afd63d2732b7ca9d3958407e16b7351bac0dee17f7483d0197
|
Provenance
The following attestation bundles were made for smcx-1.2.0.tar.gz:
Publisher:
release.yml on michaelellis003/smcx
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
smcx-1.2.0.tar.gz -
Subject digest:
9e73764e28b7c668adcd561a343c4948048d09849ae63a17d0a1049c5acea81e - Sigstore transparency entry: 2195628077
- Sigstore integration time:
-
Permalink:
michaelellis003/smcx@52f3b62d1ab31d8e1958de27188e9cdb0f05a826 -
Branch / Tag:
refs/heads/main - Owner: https://github.com/michaelellis003
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
release.yml@52f3b62d1ab31d8e1958de27188e9cdb0f05a826 -
Trigger Event:
workflow_run
-
Statement type:
File details
Details for the file smcx-1.2.0-py3-none-any.whl.
File metadata
- Download URL: smcx-1.2.0-py3-none-any.whl
- Upload date:
- Size: 47.1 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? Yes
- Uploaded via: twine/6.1.0 CPython/3.13.13
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
0e3a45fbd497c26cab42d0600aad98c126d6d8d4956de46b79a8a14416cbd87d
|
|
| MD5 |
c91e910ef7151726dd74d02cae1f68da
|
|
| BLAKE2b-256 |
0af9725820ec26f596c5c97d651f936e42dd7ae5573a3fa29d03ab2be28b0bcc
|
Provenance
The following attestation bundles were made for smcx-1.2.0-py3-none-any.whl:
Publisher:
release.yml on michaelellis003/smcx
-
Statement:
-
Statement type:
https://in-toto.io/Statement/v1 -
Predicate type:
https://docs.pypi.org/attestations/publish/v1 -
Subject name:
smcx-1.2.0-py3-none-any.whl -
Subject digest:
0e3a45fbd497c26cab42d0600aad98c126d6d8d4956de46b79a8a14416cbd87d - Sigstore transparency entry: 2195628081
- Sigstore integration time:
-
Permalink:
michaelellis003/smcx@52f3b62d1ab31d8e1958de27188e9cdb0f05a826 -
Branch / Tag:
refs/heads/main - Owner: https://github.com/michaelellis003
-
Access:
public
-
Token Issuer:
https://token.actions.githubusercontent.com -
Runner Environment:
github-hosted -
Publication workflow:
release.yml@52f3b62d1ab31d8e1958de27188e9cdb0f05a826 -
Trigger Event:
workflow_run
-
Statement type: