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
40 changes: 33 additions & 7 deletions src/pyrecest/filters/ggiw_tracker.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
sin,
zeros,
)
from pyrecest.numerics import assert_covariance_matrix

from .abstract_extended_object_tracker import AbstractExtendedObjectTracker

Expand All @@ -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
Expand Down Expand Up @@ -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",
Expand All @@ -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")
Expand All @@ -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"
Expand Down Expand Up @@ -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")

Expand Down
88 changes: 88 additions & 0 deletions tests/filters/test_ggiw_input_validation.py
Original file line number Diff line number Diff line change
@@ -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())))
Loading