Skip to content
Merged
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
31 changes: 28 additions & 3 deletions src/pyrecest/filters/kalman_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,21 +46,29 @@ def _get_optional_model_attribute(model, *names):
return None


def _validate_kalman_state_covariance(covariance, dim):
def _validate_kalman_covariance(covariance, *, name, dim=None):
"""Validate a Kalman covariance without leaving the active backend."""
covariance = validate_covariance_matrix(
covariance,
name="state.covariance",
name=name,
dim=dim,
allow_scalar=True,
check_symmetric=True,
)
eigenvalues = linalg.eigvalsh(covariance)
if not bool(backend_all(eigenvalues >= -_STATE_COVARIANCE_EIGENVALUE_ATOL)):
raise ValueError("state.covariance must be positive semidefinite.")
raise ValueError(f"{name} must be positive semidefinite.")
return covariance


def _validate_kalman_state_covariance(covariance, dim):
return _validate_kalman_covariance(
covariance,
name="state.covariance",
dim=dim,
)


class KalmanFilter(AbstractFilter, EuclideanFilterMixin):
"""Kalman filter for linear Gaussian Euclidean state-space models.

Expand Down Expand Up @@ -168,6 +176,11 @@ def predict_linear(
sys_input : array-like, shape (n,), optional
Additive deterministic input ``u``.
"""
sys_noise_cov = _validate_kalman_covariance(
sys_noise_cov,
name="sys_noise_cov",
dim=self.dim,
)
new_mean, new_covariance = linear_gaussian_predict(
mean=self._filter_state.mu,
covariance=self._filter_state.C,
Expand Down Expand Up @@ -246,6 +259,10 @@ def update_identity(

def innovation_linear(self, measurement, measurement_matrix, meas_noise):
"""Return innovation and innovation covariance for a linear measurement."""
meas_noise = _validate_kalman_covariance(
meas_noise,
name="meas_noise",
)
return linear_gaussian_innovation(
self._filter_state.mu,
self._filter_state.C,
Expand Down Expand Up @@ -298,6 +315,10 @@ def update_linear(
action : str, optional
Caller-defined diagnostic label for the update action.
"""
meas_noise = _validate_kalman_covariance(
meas_noise,
name="meas_noise",
)
result = linear_gaussian_update(
mean=self._filter_state.mu,
covariance=self._filter_state.C,
Expand Down Expand Up @@ -341,6 +362,10 @@ def update_linear_robust(
``"none"``. With ``robust_update=None`` and ``gate_threshold`` set,
measurements above the gate are rejected and the prior state is kept.
"""
meas_noise = _validate_kalman_covariance(
meas_noise,
name="meas_noise",
)
result = linear_gaussian_update_robust(
mean=self._filter_state.mu,
covariance=self._filter_state.C,
Expand Down
74 changes: 74 additions & 0 deletions tests/filters/test_kalman_filter_noise_covariance_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
import numpy as np
import pytest

from pyrecest import backend
from pyrecest.filters import KalmanFilter


def _one_dimensional_filter():
return KalmanFilter((backend.array([0.0]), backend.array([[1.0]])))


def _assert_unit_state_unchanged(kalman_filter):
np.testing.assert_allclose(backend.to_numpy(kalman_filter.filter_state.mu), [0.0])
np.testing.assert_allclose(
backend.to_numpy(kalman_filter.filter_state.C),
[[1.0]],
)


def test_predict_rejects_non_psd_process_noise_atomically():
kalman_filter = _one_dimensional_filter()

with pytest.raises(ValueError, match="sys_noise_cov must be positive semidefinite"):
kalman_filter.predict_identity(backend.array([[-2.0]]))

_assert_unit_state_unchanged(kalman_filter)


def test_update_rejects_non_psd_measurement_noise_atomically():
kalman_filter = _one_dimensional_filter()

with pytest.raises(ValueError, match="meas_noise must be positive semidefinite"):
kalman_filter.update_identity(
backend.array([[-0.5]]),
backend.array([0.0]),
)

_assert_unit_state_unchanged(kalman_filter)


def test_innovation_rejects_non_psd_measurement_noise():
kalman_filter = _one_dimensional_filter()

with pytest.raises(ValueError, match="meas_noise must be positive semidefinite"):
kalman_filter.innovation_linear(
backend.array([0.0]),
backend.array([[1.0]]),
backend.array([[-0.5]]),
)


def test_robust_update_rejects_non_psd_measurement_noise_atomically():
kalman_filter = _one_dimensional_filter()

with pytest.raises(ValueError, match="meas_noise must be positive semidefinite"):
kalman_filter.update_linear_robust(
backend.array([0.0]),
backend.array([[1.0]]),
backend.array([[-0.5]]),
)

_assert_unit_state_unchanged(kalman_filter)


def test_singular_noise_covariances_remain_supported():
kalman_filter = _one_dimensional_filter()

kalman_filter.predict_identity(backend.array([[0.0]]))
kalman_filter.update_identity(
backend.array([[0.0]]),
backend.array([0.0]),
)

np.testing.assert_allclose(backend.to_numpy(kalman_filter.filter_state.C), [[0.0]])
Loading