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, first-order 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 Kalman filtering with explicit, replaceable Jacobian 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. Linear-Gaussian models are dense arrays. Nonlinear Gaussian and particle models use ordinary JAX callables, so model functions, Jacobians, proposals, and other algorithmic pieces can be replaced independently. Filtering and smoothing remain separate functions joined by typed posterior containers, allowing research code to replace one stage without subclassing or rerunning the other.

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 was informed by the functional state/information protocol in BlackJAX 1.6.2 and 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 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 nonlinear fixture.

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.9.0.tar.gz (59.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.9.0-py3-none-any.whl (69.9 kB view details)

Uploaded Python 3

File details

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

File metadata

  • Download URL: smcx-1.9.0.tar.gz
  • Upload date:
  • Size: 59.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.9.0.tar.gz
Algorithm Hash digest
SHA256 225e4a4ece27fc6686d5a637e01cbe352c226de2f3547eff25749dda5dd2de54
MD5 5600ef9b9c43b8c029eda17370c5d692
BLAKE2b-256 8212b1d064580f7f58c1a670da9fcb74912835bb81ee6b45596284429c0b556a

See more details on using hashes here.

Provenance

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

File metadata

  • Download URL: smcx-1.9.0-py3-none-any.whl
  • Upload date:
  • Size: 69.9 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.9.0-py3-none-any.whl
Algorithm Hash digest
SHA256 b04fcdeecfea73d86d4f2e792096c8066c82b57c0ab0bc248aaa69a843b78118
MD5 1711230e4776dbf864a3f8dab83974d5
BLAKE2b-256 e90bfec9606c8675c1774a22835813f7e7c27facb12a0de54ac74f4499b85720

See more details on using hashes here.

Provenance

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