diff --git a/src/pyrecest/filters/_ukf.py b/src/pyrecest/filters/_ukf.py index f09c393b9..b5225367c 100644 --- a/src/pyrecest/filters/_ukf.py +++ b/src/pyrecest/filters/_ukf.py @@ -116,11 +116,19 @@ def predict(self, fx=None, dt=None, **fx_args): x_pred = einsum("i,ij->j", Wm, sigmas_f) - P_pred = zeros((self._model.dim_x, self._model.dim_x)) + process_covariance = asarray(self.Q, dtype=float64) + expected_process_shape = (self._model.dim_x, self._model.dim_x) + if process_covariance.shape != expected_process_shape: + raise ValueError( + "process noise covariance Q has shape " + f"{process_covariance.shape}, expected {expected_process_shape}" + ) + + P_pred = zeros(expected_process_shape) for i in range(n_sigmas): 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 = P_pred + process_covariance P_pred = 0.5 * (P_pred + transpose(P_pred)) self.x = x_pred diff --git a/tests/filters/test_ukf_process_noise_shape.py b/tests/filters/test_ukf_process_noise_shape.py new file mode 100644 index 000000000..4ec7273c1 --- /dev/null +++ b/tests/filters/test_ukf_process_noise_shape.py @@ -0,0 +1,35 @@ +import unittest + +import numpy.testing as npt + +# pylint: disable=no-name-in-module,no-member +import pyrecest.backend +from pyrecest.backend import array, diag +from pyrecest.distributions import GaussianDistribution +from pyrecest.filters.unscented_kalman_filter import UnscentedKalmanFilter + + +class UnscentedKalmanFilterProcessNoiseShapeTest(unittest.TestCase): + @unittest.skipIf( + pyrecest.backend.__backend_name__ in ("pytorch", "jax"), + reason="Not supported on this backend", + ) + def test_vector_process_covariance_is_rejected_without_state_mutation(self): + initial_mean = array([0.5, -0.25]) + initial_covariance = diag(array([1.2, 0.8])) + ukf = UnscentedKalmanFilter( + GaussianDistribution(initial_mean, initial_covariance) + ) + + with self.assertRaisesRegex( + ValueError, + r"process noise covariance Q has shape .* expected \(2, 2\)", + ): + ukf.predict_identity(array([0.4, 0.2])) + + npt.assert_allclose(ukf.get_point_estimate(), initial_mean) + npt.assert_allclose(ukf.filter_state.covariance(), initial_covariance) + + +if __name__ == "__main__": + unittest.main()