Metadata-Version: 2.4
Name: smcx
Version: 2.0.0
Summary: State-space inference in JAX: Kalman and particle filters, tempered SMC, and SMC2 across CPU and accelerators
Keywords: jax,particle-filter,kalman-filter,sequential-monte-carlo,bayesian-inference,apple-silicon,state-space-models,hidden-markov-models
Author: Michael Ellis
Author-email: Michael Ellis <michaelellis003@gmail.com>
License-Expression: Apache-2.0
License-File: LICENSE
Classifier: Development Status :: 4 - Beta
Classifier: Intended Audience :: Science/Research
Classifier: License :: OSI Approved :: Apache Software License
Classifier: Operating System :: MacOS :: MacOS X
Classifier: Operating System :: POSIX :: Linux
Classifier: Programming Language :: Python :: 3
Classifier: Programming Language :: Python :: 3.11
Classifier: Programming Language :: Python :: 3.12
Classifier: Programming Language :: Python :: 3.13
Classifier: Topic :: Scientific/Engineering :: Mathematics
Classifier: Typing :: Typed
Requires-Dist: jax>=0.10,<0.11
Requires-Dist: jaxlib>=0.10,<0.11
Requires-Dist: jaxtyping>=0.3.7
Requires-Dist: numpy>=2.2.6
Requires-Dist: arviz>=0.23.4 ; extra == 'arviz'
Requires-Dist: jax-mps>=0.10.10,<0.11 ; platform_machine == 'arm64' and sys_platform == 'darwin' and extra == 'metal'
Requires-Python: >=3.11
Project-URL: Homepage, https://github.com/michaelellis003/smcx
Project-URL: Repository, https://github.com/michaelellis003/smcx
Project-URL: Issues, https://github.com/michaelellis003/smcx/issues
Project-URL: Documentation, https://michaelellis003.github.io/smcx/
Project-URL: Changelog, https://github.com/michaelellis003/smcx/blob/main/CHANGELOG.md
Provides-Extra: arviz
Provides-Extra: metal
Description-Content-Type: text/markdown

# smcx

Sequential inference for state-space models (SSMs) in JAX:
Kalman-family filters, particle filters, and sequential Monte Carlo
(SMC) samplers. Following a theme of some modern probabilistic
programming languages like [NumPyro](https://num.pyro.ai/en/stable/)
and other projects in the JAX ecosystem like
[dynestyx](https://github.com/BasisResearch/dynestyx), I aimed to
decouple the inference code from the model code. smcx goes one step
further and has no modeling language of its own. Users define their
models as plain JAX functions. This also lets them wrap models built
in other SSM libraries like
[dynamax](https://github.com/probml/dynamax).

An introduction to the Kalman and SMC methods is developed in
[the documentation](https://michaelellis003.github.io/smcx/guides/sequential-inference/).
Below is a quick start and a map of the methods.

```bash
pip install smcx
```

## Quick start

A simple model to start with is a linear-Gaussian model

$$
\begin{aligned}
y_t &\sim \mathcal{N}(\theta_t,\ 0.3), \\
\theta_t &\sim \mathcal{N}(0.8\,\theta_{t-1},\ 0.2), \qquad
\theta_0 \sim \mathcal{N}(0, 1).
\end{aligned}
$$

In this case we can calculate the exact filtering distribution in
closed form using the Kalman filter. The Kalman filter assumes the
model is linear and Gaussian, so all you need to provide are the
model parameters.

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

import smcx

# fmt: off
y = jnp.array([
    -0.54, -1.09, -0.77, -0.03, 0.92, -0.45, 1.19, 0.24, 1.13,
    -0.42, 0.63, 1.18, 1.13, 0.64, 1.35, 2.25, 1.98, 1.65, 2.01,
    1.63, 0.80, 0.39, -0.68, -0.87, -0.96,
])[:, None]
# fmt: on

m0 = jnp.zeros(1)
C0 = jnp.eye(1)
G = 0.8 * jnp.eye(1)
W = 0.2 * jnp.eye(1)
F = jnp.eye(1)
V = 0.3 * jnp.eye(1)

kalman = smcx.kalman_filter(m0, C0, G, W, F, V, y)
print(kalman.marginal_loglik)  # -29.26, exact
```

If you want to define your own model, of any form, you instead
write functions for the initial distribution, the transition
distribution, and the observation distribution, and pass them to a
particle filter. Here is the same model through the bootstrap
particle filter:

```python
def sample_initial(key, num_particles):
    return jr.normal(key, (num_particles, 1))


def sample_transition(key, state):
    return 0.8 * state + jnp.sqrt(0.2) * jr.normal(key, state.shape)


def log_observation(obs, state):
    residual = obs[0] - state[0]
    return -0.5 * (jnp.log(2 * jnp.pi * 0.3) + residual**2 / 0.3)


particle = smcx.bootstrap_filter(
    jr.key(0),
    sample_initial,
    sample_transition,
    log_observation,
    y,
    num_particles=10_000,
)
print(particle.marginal_loglik)  # -29.16, N = 10,000
```

The two estimates agree because the model is the same. If our model
leaves the linear-Gaussian family, we can no longer use the Kalman
filter. We only change the three functions of the bootstrap call to
the new densities. The table below maps each model class to its
methods, and the
[introduction in the documentation](https://michaelellis003.github.io/smcx/guides/sequential-inference/)
develops the theory with four worked examples, relaxing one
assumption at a time.

## Methods

smcx implements the standard sequential inference methods:

| Setting | Methods | Functions |
| --- | --- | --- |
| Linear-Gaussian, fully known | Kalman filter and RTS smoother, exact | `kalman_filter`, `rts_smoother` |
| Known nonlinear functions | Extended and unscented Kalman filters, approximate; the linearization strategy is an argument | `extended_kalman_filter`, `unscented_kalman_filter`, `gaussian_filter` |
| Observation variance unknown, variance-scaled | Conjugate DLM, exact | `dlm_filter` |
| Count and binary observations | Conjugate/linear-Bayes DGLM, approximate; the observation family is an argument | `dglm_filter` with `poisson()`, `bernoulli()`, or `binomial(trials=n)` |
| General densities | Bootstrap, auxiliary, and guided particle filters | `bootstrap_filter`, `auxiliary_filter`, `guided_filter` |
| Custom particle algorithms | Feynman–Kac derivations over one generic loop | `StateSpaceModel`, `FeynmanKac`, `run_smc`, `run_particle_filter` |
| Static parameters | Adaptive tempered SMC, SMC², and the joint Liu-West filter | `temper`, `smc2`, `liu_west_filter` |
| Simulation and prediction | Model simulation and posterior predictive draws | `simulate`, `posterior_predictive_sample` |
| Resampling | Systematic, stratified, multinomial, residual | `systematic`, `stratified`, `multinomial`, `residual` |
| Diagnostics and reporting | ESS, scoring rules, trajectory reconstruction, ArviZ export | `diagnose`, `crps`, `reconstruct_trajectories`, `to_arviz` |

smcx runs on CPU, CUDA, and TPU through JAX, and on Apple-silicon
GPUs through the optional
[jax-mps](https://github.com/tillahoffmann/jax-mps) backend.

## Installation

smcx requires Python 3.11 or later.

```bash
pip install smcx
```

Optional extras add Apple-silicon GPU execution or ArviZ reporting:

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

## Documentation

Available at
[michaelellis003.github.io/smcx](https://michaelellis003.github.io/smcx/).

## Citation

If smcx contributes to academic work, please cite the release used.
The repository's **Cite this repository** menu uses
[`CITATION.cff`](https://github.com/michaelellis003/smcx/blob/main/CITATION.cff)
to provide BibTeX and APA entries; include the version and release
date in the final citation.

## See also

**State-space models and SMC**

- [dynamax](https://github.com/probml/dynamax): probabilistic state-space
models with learning via EM and SGD.
- [dynestyx](https://github.com/BasisResearch/dynestyx): NumPyro-based
inference for dynamical systems.
- [particles](https://github.com/nchopin/particles): the reference
Python companion to Chopin and Papaspiliopoulos (2020).
- [BlackJAX](https://github.com/blackjax-devs/blackjax): MCMC and SMC
samplers for JAX.

**The JAX ecosystem**

- [Equinox](https://github.com/patrick-kidger/equinox): neural networks
and PyTree modules.
- [Diffrax](https://github.com/patrick-kidger/diffrax): numerical
differential equation solvers.
- [jaxtyping](https://github.com/patrick-kidger/jaxtyping): shape and
dtype annotations for arrays.
- [ArviZ](https://github.com/arviz-devs/arviz): exploratory analysis of
Bayesian models.

## Sources and attribution

The broader Feynman–Kac architecture follows Chopin and
Papaspiliopoulos's
[*An Introduction to Sequential Monte Carlo*](https://doi.org/10.1007/978-3-030-47845-2).
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](https://github.com/blackjax-devs/blackjax/blob/a9ef478c69d730a2caa13ca4b2d735c580e0feec/blackjax/smc/from_mcmc.py),
and by the separation of orchestration from history in
[particles 0.4](https://github.com/nchopin/particles/releases/tag/v0.4).
These are design credits; no code was copied or translated.
The implemented methods draw on these primary sources:

- Exact linear-Gaussian state estimation:
  [Kalman (1960)](https://doi.org/10.1115/1.3662552) and
  [Rauch, Tung, and Striebel (1965)](https://doi.org/10.2514/3.3166).
- Conjugate dynamic models:
  [West and Harrison (1997)](https://doi.org/10.1007/b98971) and
  [West, Harrison, and Migon (1985)](https://doi.org/10.1080/01621459.1985.10477131).
- Nonlinear Gaussian filtering:
  [Schmidt (1966)](https://doi.org/10.1016/B978-1-4831-6716-9.50011-4) and
  [Julier (2002)](https://doi.org/10.1109/ACC.2002.1025369).
- Particle filters: [Gordon, Salmond, and Smith (1993)](https://doi.org/10.1049/ip-f-2.1993.0015),
  [Pitt and Shephard (1999)](https://doi.org/10.1080/01621459.1999.10474153),
  [Doucet, Godsill, and Andrieu (2000)](https://doi.org/10.1023/A:1008935410038),
  and [Liu and West (2001)](https://doi.org/10.1007/978-1-4757-3437-9_10).
- Static and parameter inference:
  [Del Moral, Doucet, and Jasra (2006)](https://doi.org/10.1111/j.1467-9868.2006.00553.x)
  and [Chopin, Jacob, and Papaspiliopoulos (2013)](https://doi.org/10.1111/j.1467-9868.2012.01046.x).
- Resampling and diagnostics:
  [Douc, Cappé, and Moulines (2005)](https://doi.org/10.1109/ISPA.2005.195385),
  [Lee and Whiteley (2018)](https://doi.org/10.1093/biomet/asy028),
  [Zhang and Stephens (2009)](https://doi.org/10.1198/TECH.2009.08017),
  and [Vehtari et al. (2024)](https://jmlr.org/papers/v25/19-556.html).
- Scoring rules:
  [Matheson and Winkler (1976)](https://doi.org/10.1287/mnsc.22.10.1087)
  and [Gneiting and Raftery (2007)](https://doi.org/10.1198/016214506000001437).
- Reporting: [ArviZ](https://doi.org/10.21105/joss.01143).

## Contributing

Contributions are welcome. See
[`CONTRIBUTING.md`](https://github.com/michaelellis003/smcx/blob/main/CONTRIBUTING.md)
for the development setup and pull-request conventions.

## License

smcx is distributed under the
[Apache License 2.0](https://github.com/michaelellis003/smcx/blob/main/LICENSE).
