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 14 or later on 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 uses the base metadata in CITATION.cff to provide BibTeX and APA entries. Include the release tag or version and release date in the final citation.

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.14.18.tar.gz (92.1 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.14.18-py3-none-any.whl (106.6 kB view details)

Uploaded Python 3

File details

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

File metadata

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

File hashes

Hashes for smcx-1.14.18.tar.gz
Algorithm Hash digest
SHA256 5ca0608df7fab14e998e8991f3746c92fc843ec38292a5315fe88983cf68648d
MD5 48f49c988fc8ebe126cd4b666d6246d3
BLAKE2b-256 1583dc507fd4d67055609595df38412e427ea8fb665d3a86dd3b7834213647a8

See more details on using hashes here.

Provenance

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

File metadata

  • Download URL: smcx-1.14.18-py3-none-any.whl
  • Upload date:
  • Size: 106.6 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.14.18-py3-none-any.whl
Algorithm Hash digest
SHA256 46a51cb5688ab93d0e5745fcd8f7b4222a6accd50ed079fcd13b122b596bccaa
MD5 316e237d0ce65e385c737b74ca6ebf85
BLAKE2b-256 6a52b7081d3609340ab540900030f14fe9471a6704eba46fd8fcd9c61aabcbdf

See more details on using hashes here.

Provenance

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