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
103 changes: 64 additions & 39 deletions src/pyrecest/filters/vbrm_tracker.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
# pylint: disable=no-name-in-module,no-member
# pylint: disable=too-many-instance-attributes,too-many-arguments
# pylint: disable=too-many-positional-arguments,too-many-locals
from pyrecest.backend import all as backend_all
from pyrecest.backend import (
array,
concatenate,
Expand Down Expand Up @@ -133,15 +134,18 @@ def __init__(
)

self.num_iterations = _as_positive_integer(num_iterations, "num_iterations")
self.forgetting_factor = float(forgetting_factor)
if self.forgetting_factor <= 0.0:
raise ValueError("forgetting_factor must be positive")
self.extent_scale = float(extent_scale)
if self.extent_scale <= 0.0:
raise ValueError("extent_scale must be positive")
self.covariance_regularization = float(covariance_regularization)
if self.covariance_regularization < 0.0:
raise ValueError("covariance_regularization must be non-negative")
self.forgetting_factor = float(
self._as_positive_scalar(forgetting_factor, "forgetting_factor")
)
self.extent_scale = float(
self._as_positive_scalar(extent_scale, "extent_scale")
)
self.covariance_regularization = float(
self._as_nonnegative_scalar(
covariance_regularization,
"covariance_regularization",
)
)

@staticmethod
def _symmetrize(matrix):
Expand Down Expand Up @@ -171,29 +175,34 @@ def _as_covariance_matrix(cls, value, dim, name, require_positive_definite=True)
matrix = diag(matrix)
if matrix.shape != (dim, dim):
raise ValueError(f"{name} must have shape ({dim}, {dim})")
if not bool(backend_all(isfinite(matrix))):
raise ValueError(f"{name} must contain only finite values")
matrix = cls._symmetrize(matrix)
if require_positive_definite:
cls._validate_positive_definite(matrix, name)
return matrix

@staticmethod
def _as_positive_scalar(value, name):
def _as_finite_scalar(value, name):
scalar = array(value)
if scalar.ndim == 1 and scalar.shape == (1,):
scalar = scalar[0]
if scalar.shape != ():
raise ValueError(f"{name} must be scalar")
if not bool(isfinite(scalar)):
raise ValueError(f"{name} must be finite")
return scalar

@classmethod
def _as_positive_scalar(cls, value, name):
scalar = cls._as_finite_scalar(value, name)
if float(scalar) <= 0.0:
raise ValueError(f"{name} must be positive")
return scalar

@staticmethod
def _as_nonnegative_scalar(value, name):
scalar = array(value)
if scalar.ndim == 1 and scalar.shape == (1,):
scalar = scalar[0]
if scalar.shape != ():
raise ValueError(f"{name} must be scalar")
@classmethod
def _as_nonnegative_scalar(cls, value, name):
scalar = cls._as_finite_scalar(value, name)
if float(scalar) < 0.0:
raise ValueError(f"{name} must be non-negative")
return scalar
Expand All @@ -205,6 +214,8 @@ def _as_positive_vector(value, dim, name):
vector = vector * array([1.0] * dim)
if vector.shape != (dim,):
raise ValueError(f"{name} must be scalar or have shape ({dim},)")
if not bool(backend_all(isfinite(vector))):
raise ValueError(f"{name} entries must be finite")
for index in range(dim):
if float(vector[index]) <= 0.0:
raise ValueError(f"{name} entries must be positive")
Expand All @@ -214,6 +225,8 @@ def _as_positive_vector(value, dim, name):
def _validate_shape_state(shape_state):
if shape_state.shape != (3,):
raise ValueError("shape_state must have shape (3,)")
if not bool(backend_all(isfinite(shape_state))):
raise ValueError("shape_state must contain only finite values")
if float(shape_state[1]) <= 0.0 or float(shape_state[2]) <= 0.0:
raise ValueError("shape semi-axis lengths must be positive")

Expand Down Expand Up @@ -457,40 +470,52 @@ def predict_linear(
require_positive_definite=False,
)

self.kinematic_state = system_matrix @ self.kinematic_state
if inputs is not None:
self.kinematic_state = self.kinematic_state + array(inputs)
self.covariance = self._symmetrize(
system_matrix @ self.covariance @ system_matrix.T + sys_noise
orientation_system_matrix = float(
self._as_finite_scalar(
orientation_system_matrix,
"orientation_system_matrix",
)
)

orientation_system_matrix = float(orientation_system_matrix)
orientation_sys_noise = self._as_nonnegative_scalar(
orientation_sys_noise,
"orientation_sys_noise",
)
self.orientation = orientation_system_matrix * self.orientation
self.orientation_variance = (
orientation_system_matrix**2 * self.orientation_variance
+ orientation_sys_noise
)

gamma = (
self.forgetting_factor
if forgetting_factor is None
else float(forgetting_factor)
else float(
self._as_positive_scalar(forgetting_factor, "forgetting_factor")
)
)
if gamma <= 0.0:
raise ValueError("forgetting_factor must be positive")
self.alpha = gamma * self.alpha
self.beta = gamma * self.beta
if float(self.alpha[0]) <= 1.0 or float(self.alpha[1]) <= 1.0:

next_alpha = gamma * self.alpha
next_beta = gamma * self.beta
if float(next_alpha[0]) <= 1.0 or float(next_alpha[1]) <= 1.0:
raise ValueError(
"The prediction made inverse-gamma alpha <= 1, so the extent "
"mean is undefined. Increase inverse_gamma_shape or use a larger "
"forgetting_factor."
)

next_kinematic_state = system_matrix @ self.kinematic_state
if inputs is not None:
next_kinematic_state = next_kinematic_state + array(inputs)
next_covariance = self._symmetrize(
system_matrix @ self.covariance @ system_matrix.T + sys_noise
)
next_orientation = orientation_system_matrix * self.orientation
next_orientation_variance = (
orientation_system_matrix**2 * self.orientation_variance
+ orientation_sys_noise
)

self.kinematic_state = next_kinematic_state
self.covariance = next_covariance
self.orientation = next_orientation
self.orientation_variance = next_orientation_variance
self.alpha = next_alpha
self.beta = next_beta

if self.log_prior_estimates:
self.store_prior_estimates()
if self.log_prior_extents:
Expand Down Expand Up @@ -540,10 +565,10 @@ def update(
measurement_matrix = self._get_measurement_matrix(meas_mat)
meas_noise_cov = self._get_measurement_noise(meas_noise_cov)
num_iterations = (
self.num_iterations if num_iterations is None else int(num_iterations)
self.num_iterations
if num_iterations is None
else _as_positive_integer(num_iterations, "num_iterations")
)
if num_iterations <= 0:
raise ValueError("num_iterations must be positive")

self._update_vbrm(
measurements.T,
Expand Down
123 changes: 123 additions & 0 deletions tests/filters/test_vbrm_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,123 @@
import unittest

import numpy as np
import numpy.testing as npt

import pyrecest.backend
from pyrecest.backend import array, diag, eye
from pyrecest.filters.vbrm_tracker import VBRMTracker


@unittest.skipIf(
pyrecest.backend.__backend_name__ != "numpy",
reason="VBRM validation tests currently use numpy.testing assertions",
)
class TestVBRMValidation(unittest.TestCase):
def setUp(self):
self.kinematic_state = array([0.0, 0.0, 1.0, -1.0])
self.covariance = diag(array([0.1, 0.1, 0.01, 0.01]))
self.shape_state = array([0.0, 2.0, 1.0])
self.measurement_matrix = array(
[
[1.0, 0.0, 0.0, 0.0],
[0.0, 1.0, 0.0, 0.0],
]
)
self.measurement_noise_cov = 0.01 * eye(2)

def _make_tracker(self, **overrides):
arguments = {
"kinematic_state": self.kinematic_state,
"covariance": self.covariance,
"shape_state": self.shape_state,
"orientation_variance": 0.1,
"inverse_gamma_shape": 10.0,
"measurement_noise_cov": self.measurement_noise_cov,
"measurement_matrix": self.measurement_matrix,
}
arguments.update(overrides)
return VBRMTracker(**arguments)

@staticmethod
def _snapshot(tracker):
return {
"kinematic_state": np.asarray(tracker.kinematic_state).copy(),
"covariance": np.asarray(tracker.covariance).copy(),
"orientation": float(tracker.orientation),
"orientation_variance": float(tracker.orientation_variance),
"alpha": np.asarray(tracker.alpha).copy(),
"beta": np.asarray(tracker.beta).copy(),
}

def _assert_snapshot_equal(self, tracker, snapshot):
npt.assert_allclose(tracker.kinematic_state, snapshot["kinematic_state"])
npt.assert_allclose(tracker.covariance, snapshot["covariance"])
self.assertEqual(float(tracker.orientation), snapshot["orientation"])
self.assertEqual(
float(tracker.orientation_variance), snapshot["orientation_variance"]
)
npt.assert_allclose(tracker.alpha, snapshot["alpha"])
npt.assert_allclose(tracker.beta, snapshot["beta"])

def test_constructor_rejects_nonfinite_hyperparameters(self):
cases = (
("orientation_variance", np.nan),
("orientation_variance", np.inf),
("inverse_gamma_shape", np.nan),
("inverse_gamma_shape", np.inf),
("forgetting_factor", np.nan),
("forgetting_factor", np.inf),
("extent_scale", np.nan),
("extent_scale", np.inf),
("covariance_regularization", np.nan),
("covariance_regularization", np.inf),
)
for name, value in cases:
with self.subTest(name=name, value=value), self.assertRaises(ValueError):
self._make_tracker(**{name: value})

def test_constructor_rejects_nonfinite_shape_state(self):
invalid_shapes = (
array([np.nan, 2.0, 1.0]),
array([np.inf, 2.0, 1.0]),
array([0.0, np.nan, 1.0]),
array([0.0, 2.0, np.inf]),
)
for shape_state in invalid_shapes:
with self.subTest(shape_state=shape_state), self.assertRaises(ValueError):
self._make_tracker(shape_state=shape_state)

def test_predict_validation_failures_are_atomic(self):
invalid_controls = (
{"orientation_system_matrix": np.nan},
{"orientation_sys_noise": np.inf},
{"forgetting_factor": np.nan},
{"forgetting_factor": 0.05},
)
for controls in invalid_controls:
with self.subTest(controls=controls):
tracker = self._make_tracker()
snapshot = self._snapshot(tracker)
with self.assertRaises(ValueError):
tracker.predict_linear(
2.0 * eye(4),
sys_noise=0.01 * eye(4),
**controls,
)
self._assert_snapshot_equal(tracker, snapshot)

def test_update_rejects_noninteger_iteration_overrides(self):
for num_iterations in (True, 1.5, "2", array([2])):
with self.subTest(num_iterations=num_iterations):
tracker = self._make_tracker()
snapshot = self._snapshot(tracker)
with self.assertRaises(ValueError):
tracker.update(
array([[0.1, 0.0]]),
num_iterations=num_iterations,
)
self._assert_snapshot_equal(tracker, snapshot)


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