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
60 changes: 53 additions & 7 deletions src/pyrecest/filters/velocity_aided_mem_qkf_tracker.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,40 @@
from __future__ import annotations

from math import isfinite

import numpy as np

# pylint: disable=no-name-in-module,no-member,too-many-arguments
# pylint: disable=too-many-positional-arguments,too-many-locals
from pyrecest.backend import arctan2, array, cos, maximum, sin

from .mem_qkf_tracker import MEMQKFTracker


def _as_finite_float(value, name: str) -> float:
"""Return a finite scalar float without accepting boolean controls."""
value_array = np.asarray(value)
if value_array.shape != () or value_array.dtype == np.bool_:
raise ValueError(f"{name} must be a finite scalar")
scalar = value_array.item()
if isinstance(scalar, (bool, np.bool_)):
raise ValueError(f"{name} must be a finite scalar")
try:
parsed = float(scalar)
except (TypeError, ValueError, OverflowError) as exc:
raise ValueError(f"{name} must be a finite scalar") from exc
if not isfinite(parsed):
raise ValueError(f"{name} must be a finite scalar")
return parsed


def _as_bool_flag(value, name: str) -> bool:
"""Return a strict Python boolean for public configuration flags."""
if isinstance(value, (bool, np.bool_)):
return bool(value)
raise ValueError(f"{name} must be a boolean")


class VelocityAidedMEMQKFTracker(MEMQKFTracker):
"""MEM-QKF with a soft axial velocity-heading pseudo-measurement.

Expand Down Expand Up @@ -66,21 +94,36 @@ def __init__(
super().__init__(*args, **kwargs)

self.velocity_indices = self._normalize_velocity_indices(velocity_indices)
self.speed_threshold = float(speed_threshold)
self.speed_threshold = _as_finite_float(speed_threshold, "speed_threshold")
if self.speed_threshold < 0.0:
raise ValueError("speed_threshold must be non-negative")

self.orientation_offset = float(orientation_offset)
self.heading_noise_variance = float(heading_noise_variance)
self.orientation_offset = _as_finite_float(
orientation_offset,
"orientation_offset",
)
self.heading_noise_variance = _as_finite_float(
heading_noise_variance,
"heading_noise_variance",
)
if self.heading_noise_variance < 0.0:
raise ValueError("heading_noise_variance must be non-negative")

self.minimum_heading_variance = float(minimum_heading_variance)
self.minimum_heading_variance = _as_finite_float(
minimum_heading_variance,
"minimum_heading_variance",
)
if self.minimum_heading_variance <= 0.0:
raise ValueError("minimum_heading_variance must be positive")

self.apply_heading_on_prediction = bool(apply_heading_on_prediction)
self.use_heading_constraint = bool(use_heading_constraint)
self.apply_heading_on_prediction = _as_bool_flag(
apply_heading_on_prediction,
"apply_heading_on_prediction",
)
self.use_heading_constraint = _as_bool_flag(
use_heading_constraint,
"use_heading_constraint",
)
self._heading_update_pending = False

def _normalize_velocity_indices(self, velocity_indices):
Expand Down Expand Up @@ -211,7 +254,10 @@ def update(
"""Update and fuse the velocity-heading pseudo-measurement once per scan."""
old_use_heading_constraint = self.use_heading_constraint
if use_heading_constraint is not None:
self.use_heading_constraint = bool(use_heading_constraint)
self.use_heading_constraint = _as_bool_flag(
use_heading_constraint,
"use_heading_constraint",
)
self._heading_update_pending = True
try:
super().update(
Expand Down
66 changes: 66 additions & 0 deletions tests/filters/test_velocity_aided_mem_qkf_validation.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,66 @@
# pylint: disable=protected-access
import numpy as np
import pytest
from pyrecest.backend import array, diag, eye
from pyrecest.filters.velocity_aided_mem_qkf_tracker import VelocityAidedMEMQKFTracker


def _make_tracker(**kwargs):
config = {
"kinematic_state": array([0.0, 0.0, 3.0, 4.0]),
"covariance": diag(array([1.0, 1.0, 0.04, 0.04])),
"shape_state": array([0.0, 2.0, 1.0]),
"shape_covariance": diag(array([0.5, 0.2, 0.2])),
"default_meas_noise_cov": 0.05 * eye(2),
"heading_noise_variance": 0.25,
}
config.update(kwargs)
return VelocityAidedMEMQKFTracker(**config)


@pytest.mark.parametrize(
"parameter",
(
"speed_threshold",
"orientation_offset",
"heading_noise_variance",
"minimum_heading_variance",
),
)
@pytest.mark.parametrize("value", (np.nan, np.inf, -np.inf))
def test_heading_scalar_controls_reject_nonfinite_values(parameter, value):
with pytest.raises(ValueError, match=parameter):
_make_tracker(**{parameter: value})


@pytest.mark.parametrize(
"parameter",
("apply_heading_on_prediction", "use_heading_constraint"),
)
@pytest.mark.parametrize("value", ("False", 0, 1, None, np.array(False)))
def test_heading_boolean_controls_reject_non_boolean_values(parameter, value):
with pytest.raises(ValueError, match=parameter):
_make_tracker(**{parameter: value})


def test_heading_boolean_controls_accept_numpy_booleans():
tracker = _make_tracker(
apply_heading_on_prediction=np.bool_(False),
use_heading_constraint=np.bool_(True),
)

assert tracker.apply_heading_on_prediction is False
assert tracker.use_heading_constraint is True


def test_per_update_heading_override_rejects_non_boolean_without_mutating_state():
tracker = _make_tracker()

with pytest.raises(ValueError, match="use_heading_constraint"):
tracker.update(
array([[1.0, 0.2]]),
use_heading_constraint="False",
)

assert tracker.use_heading_constraint is True
assert tracker._heading_update_pending is False
Loading