Skip to content

Nonlinear filters

For x_k = f(x_{k-1}, dt) + w and z_k = h(x_k) + v, kalman-py has an extended Kalman filter, which linearizes f and h with their Jacobians, and an unscented Kalman filter, which propagates sigma points through them. Both take f(x, dt) and h(x), a dt per step or one for all, and both smooth.

import jax
import jax.numpy as jnp
import numpy as np

from kalman_py import ExtendedKalmanFilter, UnscentedKalmanFilter

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


def f(x, dt):  # 2-D constant velocity: (px, py, vx, vy)
    return jnp.array([x[0] + dt * x[2], x[1] + dt * x[3], x[2], x[3]])


def h(x):  # range and bearing from a sensor at the origin
    return jnp.array([jnp.hypot(x[0], x[1]), jnp.arctan2(x[1], x[0])])


def wrap_bearing(z, z_pred):
    d = jnp.asarray(z) - jnp.asarray(z_pred)
    return d.at[1].set((d[1] + jnp.pi) % (2 * jnp.pi) - jnp.pi)


model = dict(Q=0.01 * np.eye(4), R=np.diag([0.25, 1e-4]), x0=[100.0, 5.0, -1.0, 0.0], P0=np.eye(4))
zs = np.asarray(jax.vmap(h)(jnp.array([[100.0 - k, 5.0, -1.0, 0.0] for k in range(1, 51)])))

ekf = ExtendedKalmanFilter(f, h, **model, residual_z=wrap_bearing)
ukf = UnscentedKalmanFilter(f, h, **model, residual_z=wrap_bearing)
ekf_result = ekf.filter(zs, dt=1.0, backend="jax")
ukf_result = ukf.filter(zs, dt=1.0, backend="jax")

Jacobians (EKF)

Pass jac_f(x, dt) and jac_h(x) if you have them. Otherwise they are derived with jax.jacfwd, on either backend, which needs JAX and jax.numpy model functions. With both Jacobians supplied, the EKF runs on NumPy without JAX.

Angles and other wrapped quantities

residual_z(z, z_pred) replaces z - z_pred. For angles, wrap the difference into [-π, π), as wrap_bearing does above. The UKF also uses it to average its sigma-point measurements (relative to the central point), so bearings near ±π need no separate mean function.

Sigma points (UKF)

The UKF uses Van der Merwe's scaled sigma points with alpha, beta, kappa. The defaults alpha=1, beta=2, kappa=0 keep every covariance weight nonnegative, so the predicted covariance can't become indefinite. Small alpha (e.g. 1e-3, common in the literature) concentrates the points but makes the central weight strongly negative.

Vectorized model functions (UKF)

On the NumPy backend the UKF calls f and h once per sigma point. With vectorized=True it passes the whole (2n + 1, n) stack in one call instead, which is much faster (about 118 → 49 µs per step on a range-bearing problem). Write such functions with [..., i] indexing so they work on single points too. The JAX backend batches with vmap either way.

def f_stack(x, dt):
    x = jnp.asarray(x)
    return jnp.stack(
        [x[..., 0] + dt * x[..., 2], x[..., 1] + dt * x[..., 3], x[..., 2], x[..., 3]], axis=-1
    )


def h_stack(x):
    x = jnp.asarray(x)
    return jnp.stack([jnp.hypot(x[..., 0], x[..., 1]), jnp.arctan2(x[..., 1], x[..., 0])], axis=-1)


def wrap_stack(z, z_pred):
    d = jnp.asarray(z) - jnp.asarray(z_pred)
    return d.at[..., 1].set((d[..., 1] + jnp.pi) % (2 * jnp.pi) - jnp.pi)


fast_ukf = UnscentedKalmanFilter(f_stack, h_stack, **model, residual_z=wrap_stack, vectorized=True)
print(np.abs(fast_ukf.filter(zs, dt=1.0).means - ukf.filter(zs, dt=1.0).means).max())