Skip to main content

State-space inference in JAX: Kalman and particle filters, tempered SMC, and SMC2 across CPU and accelerators

Project description

smcx

smcx is a JAX library for state-space inference: exact linear-Gaussian filtering and smoothing, extended and unscented nonlinear Gaussian filtering, particle filters, adaptive tempered SMC, and SMC² with a small, function-oriented API. It runs on CPU, CUDA, and TPU through JAX, and on Apple-silicon GPUs through the optional jax-mps backend.

Features include:

  • exact linear-Gaussian Kalman filtering and RTS smoothing;
  • extended and scaled unscented Kalman filtering with shared mean callbacks;
  • bootstrap, auxiliary, guided, and Liu–West particle filters;
  • a public runner for caller-owned particle-filter kernels;
  • adaptive tempered SMC and nested SMC² parameter inference;
  • systematic, stratified, multinomial, and residual resampling;
  • filtering diagnostics, scoring rules, trajectory reconstruction, and ArviZ export; and
  • structured latent-state PyTrees and explicit time-varying inputs.

smcx supplies inference algorithms, not model or distribution classes. Its catalog covers core families rather than every named variant; additions need concrete user or research demand and an independent validator. Linear-Gaussian models use dense arrays. The EKF and UKF share mean callbacks; the EKF alone adds explicit Jacobians. Particle models use JAX callbacks, and run_particle_filter lets research code own the filter kernel. Typed posteriors keep filtering and smoothing composable without a class hierarchy.

Installation

smcx requires Python 3.11 or later.

pip install smcx

Install the optional extras for Apple-silicon GPU execution or ArviZ reporting with:

pip install "smcx[metal]"
pip install "smcx[arviz]"

The metal extra uses jax-mps and is available on macOS arm64. Metal is float32-only; releases are tested on a physical M-series GPU as well as on CPU.

Documentation

The documentation includes a quickstart, guides for custom models and custom particle filters and ArviZ reporting, and the complete API reference.

Quick example

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

import smcx

a, q, r = 0.9, 0.5, 0.3


def initial_sampler(key, num_particles):
    return jr.normal(key, (num_particles, 1))


def transition_sampler(key, state):
    return a * state + jnp.sqrt(q) * jr.normal(key, state.shape)


def log_observation(y, state):
    error = y[0] - state[0]
    return -0.5 * (jnp.log(2 * jnp.pi * r) + error**2 / r)


observations = jnp.array([0.2, -0.1, 0.4, 0.7, 0.3])[:, None]
posterior = smcx.bootstrap_filter(
    jr.key(0),
    initial_sampler,
    transition_sampler,
    log_observation,
    observations,
    num_particles=10_000,
)

posterior.marginal_loglik
smcx.weighted_mean(posterior)
smcx.diagnose(posterior)

Callbacks describe one particle; smcx vectorizes them over the cloud. Every stochastic operation takes an explicit PRNG key, and posterior containers are JAX PyTrees.

Citation

If smcx contributes to academic work, cite the release used. The repository's Cite this repository menu is generated from CITATION.cff and provides BibTeX and APA entries.

Sources and attribution

The broader Feynman–Kac architecture follows Chopin and Papaspiliopoulos's An Introduction to Sequential Monte Carlo. The caller-owned particle-filter runner and dependency-free tempering mutation boundary were informed by BlackJAX's functional state/information protocol and pinned SMC-from-MCMC split, and by the separation of orchestration from history in particles 0.4. These are design credits; no code was copied or translated. The implemented methods draw on these primary sources:

Numerical validation references

The linear Kalman and RTS outputs are independently validated against Dynamax 1.0.2 and statsmodels 0.14.6; the details are recorded with the frozen linear fixture.

The extended and unscented Kalman outputs are independently validated against Stone Soup 1.9.1, cross-checked with Dynamax 1.0.2, and checked against SciPy 1.18.0 innovation log densities. Exact commits, environments, licenses, and observed differences are recorded with the frozen extended and unscented fixtures.

The auxiliary-guided runner recipe is distributionally cross-checked against particles 0.4's pinned AuxiliaryPF composition and SMC bookkeeping.

These projects are numerical comparison implementations, not code lineage; no implementation code was copied or translated.

Contributing

Contributions are welcome. See CONTRIBUTING.md for the development setup and pull-request conventions.

License

smcx is distributed under the Apache License 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.13.4.tar.gz (64.8 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.13.4-py3-none-any.whl (76.9 kB view details)

Uploaded Python 3

File details

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

File metadata

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

File hashes

Hashes for smcx-1.13.4.tar.gz
Algorithm Hash digest
SHA256 825bfd01e6af29c22a1bbb907ac147aca742c93dd9348e5361f949ba4ac47d28
MD5 ab650f70eba20b3c2ef87008b058f387
BLAKE2b-256 193bd7ff11598b112072fd082b93ace5929e0a1b3b8a122013e649d4a1ca8f4c

See more details on using hashes here.

Provenance

The following attestation bundles were made for smcx-1.13.4.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.13.4-py3-none-any.whl.

File metadata

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

File hashes

Hashes for smcx-1.13.4-py3-none-any.whl
Algorithm Hash digest
SHA256 bdb52f0280a9d544c10f0a2f8f5ab009e81404c6242a6d6356cd366391671234
MD5 34bb2ba5864e7039ec9e76594572f55f
BLAKE2b-256 a30cc199c96f8e4193e52571836ccb384180a6bcaa7ef9a28e9d164ad87ec29e

See more details on using hashes here.

Provenance

The following attestation bundles were made for smcx-1.13.4-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