Skip to content

Backends: NumPy and JAX

Every filter takes backend="numpy" (the default) or backend="jax" in filter, and smooth uses the backend that produced its input.

NumPy JAX
Install always available pip install kalman-py[jax]
Step-by-step API yes no (use NumPy)
Batch speed (2-D tracking) ~2.8 µs/step ~0.6 µs/step
Results NumPy arrays JAX arrays, left on the device
import jax
import numpy as np

from kalman_py import KalmanFilter

jax.config.update("jax_enable_x64", True)

dt = 1.0
F = np.array([[1.0, dt], [0.0, 1.0]])
kf = KalmanFilter(F, [[1.0, 0.0]], 0.01 * np.eye(2), [[1.0]], x0=[0.0, 1.0], P0=np.eye(2))
zs = np.arange(1, 1001, dtype=float)[:, None]

result = kf.filter(zs, backend="jax")
jax.block_until_ready(result.means)  # JAX runs asynchronously
print(type(result.means).__name__)

Precision

JAX computes in float32 unless 64-bit mode is on. With float64 inputs and 64-bit mode off, JAX truncates them and warns. Turn it on at startup for results that match the NumPy backend:

jax.config.update("jax_enable_x64", True)

float32 problems stay float32 on both backends either way.

Compilation

The first call for a given set of shapes compiles; later calls with the same shapes and dtypes reuse the compiled code, even from different filter objects. When timing, warm up first and wait for the result with jax.block_until_ready.

For the EKF and UKF, the model functions are part of what's compiled. Define them once at module level and pass the same function objects to every filter: a new lambda per filter is a new function and compiles again.

Model functions

On the JAX backend, f, h and residual_z are traced by JAX, so they must use jax.numpy (jnp.array, jnp.hypot, ...) rather than NumPy calls that convert to concrete arrays. Code written with jax.numpy also works on the NumPy backend, so one version serves both.