diff --git a/src/pyrecest/filters/ggiw_tracker.py b/src/pyrecest/filters/ggiw_tracker.py index a76cb566f..88f32c963 100644 --- a/src/pyrecest/filters/ggiw_tracker.py +++ b/src/pyrecest/filters/ggiw_tracker.py @@ -18,6 +18,7 @@ sin, zeros, ) +from pyrecest.numerics import assert_covariance_matrix from .abstract_extended_object_tracker import AbstractExtendedObjectTracker @@ -29,6 +30,16 @@ def _validate_bool_flag(value, name: str) -> bool: return bool(value_array.item()) +def _validate_finite_scalar(value, name: str) -> float: + try: + scalar = float(value) + except (TypeError, ValueError, OverflowError) as exc: + raise ValueError(f"{name} must be a finite scalar") from exc + if not np.isfinite(scalar): + raise ValueError(f"{name} must be finite") + return scalar + + class GGIWTracker( AbstractExtendedObjectTracker ): # pylint: disable=too-many-instance-attributes @@ -81,11 +92,18 @@ def __init__( ) extent = array(extent) + if extent.ndim != 2 or extent.shape[0] != extent.shape[1]: + raise ValueError("extent must be a square matrix") self.measurement_dim = extent.shape[0] - self.extent_degrees_of_freedom = float(extent_degrees_of_freedom) - self.gamma_shape = float(gamma_shape) - self.gamma_rate = float(gamma_rate) - self.extent_innovation_weight = float(extent_innovation_weight) + self.extent_degrees_of_freedom = _validate_finite_scalar( + extent_degrees_of_freedom, "extent_degrees_of_freedom" + ) + self.gamma_shape = _validate_finite_scalar(gamma_shape, "gamma_shape") + self.gamma_rate = _validate_finite_scalar(gamma_rate, "gamma_rate") + self.extent_innovation_weight = _validate_finite_scalar( + extent_innovation_weight, "extent_innovation_weight" + ) + extent_is_scale = _validate_bool_flag(extent_is_scale, "extent_is_scale") self.subtract_measurement_noise_from_scatter = _validate_bool_flag( subtract_measurement_noise_from_scatter, "subtract_measurement_noise_from_scatter", @@ -102,7 +120,9 @@ def __init__( self.measurement_matrix = array(measurement_matrix) self._validate_measurement_matrix(self.measurement_matrix) - self.extent_scale = self._symmetrize(extent) + self.extent_scale = assert_covariance_matrix( + extent, name="extent", dim=self.measurement_dim + ) if not extent_is_scale: self.extent_scale = self.extent_scale * self._extent_mean_denominator() self._validate_positive_definite(self.extent_scale, "extent_scale") @@ -124,14 +144,17 @@ def _as_covariance_matrix(cls, value, dim, name): matrix = matrix * eye(dim) if matrix.shape != (dim, dim): raise ValueError(f"{name} must have shape ({dim}, {dim})") - matrix = cls._symmetrize(matrix) + matrix = assert_covariance_matrix(matrix, name=name, dim=dim) cls._validate_positive_definite(matrix, name) return matrix def _extent_mean_denominator(self, degrees_of_freedom=None): if degrees_of_freedom is None: degrees_of_freedom = self.extent_degrees_of_freedom - denominator = float(degrees_of_freedom) - 2.0 * self.measurement_dim - 2.0 + degrees_of_freedom = _validate_finite_scalar( + degrees_of_freedom, "extent_degrees_of_freedom" + ) + denominator = degrees_of_freedom - 2.0 * self.measurement_dim - 2.0 if denominator <= 0.0: raise ValueError( "extent_degrees_of_freedom must be larger than 2 * measurement_dim + 2" @@ -331,6 +354,9 @@ def update( meas_noise_cov = self._get_measurement_noise(meas_noise_cov) if extent_innovation_weight is None: extent_innovation_weight = self.extent_innovation_weight + extent_innovation_weight = _validate_finite_scalar( + extent_innovation_weight, "extent_innovation_weight" + ) if extent_innovation_weight < 0.0: raise ValueError("extent_innovation_weight must be non-negative") diff --git a/tests/filters/test_ggiw_input_validation.py b/tests/filters/test_ggiw_input_validation.py new file mode 100644 index 000000000..e5e89bb5b --- /dev/null +++ b/tests/filters/test_ggiw_input_validation.py @@ -0,0 +1,88 @@ +import numpy as np +import pytest +import pyrecest.backend +from pyrecest.backend import array, diag, eye +from pyrecest.filters import GGIWTracker + + +pytestmark = pytest.mark.skipif( + pyrecest.backend.__backend_name__ != "numpy", + reason="GGIW validation regressions use the NumPy-backed tracker", +) + + +def _make_tracker(**kwargs): + parameters = { + "kinematic_state": array([0.0, 0.0, 1.0, -1.0]), + "covariance": diag(array([1.0, 1.0, 0.25, 0.25])), + "extent": diag(array([4.0, 1.0])), + "extent_degrees_of_freedom": 12.0, + "gamma_shape": 4.0, + "gamma_rate": 2.0, + } + parameters.update(kwargs) + return GGIWTracker(**parameters) + + +@pytest.mark.parametrize( + ("name", "value"), + [ + ("gamma_shape", np.nan), + ("gamma_shape", np.inf), + ("gamma_rate", np.nan), + ("gamma_rate", np.inf), + ("extent_innovation_weight", np.nan), + ("extent_innovation_weight", np.inf), + ("extent_degrees_of_freedom", np.nan), + ("extent_degrees_of_freedom", np.inf), + ], +) +def test_constructor_rejects_nonfinite_scalar_hyperparameters(name, value): + with pytest.raises(ValueError, match=name): + _make_tracker(**{name: value}) + + +def test_constructor_rejects_non_boolean_extent_is_scale(): + with pytest.raises(TypeError, match="extent_is_scale"): + _make_tracker(extent_is_scale="False") + + +def test_constructor_rejects_asymmetric_covariance_and_extent(): + asymmetric_covariance = eye(4) + asymmetric_covariance[0, 1] = 0.5 + with pytest.raises(ValueError, match="covariance"): + _make_tracker(covariance=asymmetric_covariance) + + with pytest.raises(ValueError, match="extent"): + _make_tracker(extent=array([[4.0, 1.0], [0.0, 1.0]])) + + +def test_constructor_rejects_nonfinite_covariance_and_extent(): + with pytest.raises(ValueError, match="covariance"): + _make_tracker(covariance=diag(array([1.0, 1.0, np.nan, 1.0]))) + + with pytest.raises(ValueError, match="extent"): + _make_tracker(extent=array([[4.0, 0.0], [0.0, np.nan]])) + + +def test_update_rejects_nonfinite_extent_innovation_weight(): + tracker = _make_tracker() + + with pytest.raises(ValueError, match="extent_innovation_weight"): + tracker.update( + array([[0.0], [0.0]]), + extent_innovation_weight=np.nan, + ) + + +def test_finite_valid_inputs_remain_supported(): + tracker = _make_tracker( + extent_is_scale=np.bool_(False), + extent_innovation_weight=0.0, + ) + + tracker.predict_linear(eye(4), 0.01 * eye(4)) + tracker.update(array([[0.0], [0.0]]), meas_noise_cov=0.1 * eye(2)) + + assert np.isfinite(float(tracker.get_measurement_rate_estimate())) + assert np.all(np.isfinite(np.asarray(tracker.get_point_estimate_extent())))