Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 14 additions & 2 deletions src/pyrecest/filters/_ukf.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,15 +8,18 @@

# pylint: disable=no-name-in-module,no-member
from pyrecest.backend import (
abs,
asarray,
einsum,
expand_dims,
eye,
float64,
linalg,
maximum,
reshape,
stack,
transpose,
where,
zeros,
)
from pyrecest.sampling.sigma_points import JulierSigmaPoints, MerweScaledSigmaPoints
Expand Down Expand Up @@ -47,6 +50,15 @@ def _as_vector(value, description):
return value


def _stable_symmetric_average(matrix):
"""Average a matrix with its transpose without finite overflow."""
transposed = transpose(matrix)
scale = maximum(abs(matrix), abs(transposed))
safe_scale = where(scale == 0.0, 1.0, scale)
normalized_average = 0.5 * (matrix / safe_scale + transposed / safe_scale)
return scale * normalized_average


# ---------------------------------------------------------------------------
# Unscented Kalman Filter
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -121,7 +133,7 @@ def predict(self, fx=None, dt=None, **fx_args):
d = expand_dims(sigmas_f[i] - x_pred, -1)
P_pred = P_pred + Wc[i] * (d @ transpose(d))
P_pred = P_pred + asarray(self.Q, dtype=float64)
P_pred = 0.5 * (P_pred + transpose(P_pred))
P_pred = _stable_symmetric_average(P_pred)

self.x = x_pred
self.P = P_pred
Expand Down Expand Up @@ -213,7 +225,7 @@ def update(self, z, R=None, hx=None, **hx_args):

self.x = self.x + K @ (z - z_pred)
self.P = self.P - K @ Pz @ transpose(K)
self.P = 0.5 * (self.P + transpose(self.P))
self.P = _stable_symmetric_average(self.P)

self._sigmas_f = None # clear cached sigma points

Expand Down
36 changes: 36 additions & 0 deletions tests/filters/test_ukf_covariance_stability.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
import unittest

import numpy as np
import numpy.testing as npt

# pylint: disable=no-name-in-module,no-member
import pyrecest.backend
from pyrecest.backend import array, float64
from pyrecest.distributions import GaussianDistribution
from pyrecest.filters.unscented_kalman_filter import UnscentedKalmanFilter


class UnscentedKalmanFilterCovarianceStabilityTest(unittest.TestCase):
@unittest.skipIf(
pyrecest.backend.__backend_name__ in ("pytorch", "jax"),
reason="Not supported on this backend",
)
def test_predict_preserves_extreme_finite_process_covariance(self):
largest = np.finfo(np.float64).max
ukf = UnscentedKalmanFilter(
GaussianDistribution(
array([0.0], dtype=float64),
array([[1.0]], dtype=float64),
)
)

with np.errstate(over="raise", invalid="raise"):
ukf.predict_identity(array([[largest]], dtype=float64))

covariance = ukf.filter_state.covariance()
self.assertTrue(np.isfinite(covariance).all())
npt.assert_array_equal(covariance, array([[largest]], dtype=float64))


if __name__ == "__main__":
unittest.main()
Loading