diff --git a/src/pyrecest/filters/abstract_filter.py b/src/pyrecest/filters/abstract_filter.py index 0a049bd20..96eb05989 100644 --- a/src/pyrecest/filters/abstract_filter.py +++ b/src/pyrecest/filters/abstract_filter.py @@ -3,6 +3,7 @@ import copy from abc import ABC, abstractmethod +from pyrecest.backend import atleast_1d from pyrecest.utils.history_recorder import HistoryRecorder @@ -33,8 +34,8 @@ def filter_state(self, new_state): self._filter_state = copy.deepcopy(new_state) def get_point_estimate(self): - """Get a point estimate""" - return self.filter_state.mean() + """Get a point estimate as a one-dimensional state vector.""" + return atleast_1d(self.filter_state.mean()) @property def dim(self) -> int: diff --git a/tests/filters/test_track_manager_atomic_filter_state.py b/tests/filters/test_track_manager_atomic_filter_state.py index dc4efecfc..235459975 100644 --- a/tests/filters/test_track_manager_atomic_filter_state.py +++ b/tests/filters/test_track_manager_atomic_filter_state.py @@ -3,7 +3,9 @@ import unittest import numpy as np +from pyrecest.distributions import VonMisesDistribution from pyrecest.filters.track_manager import TrackManager +from pyrecest.filters.von_mises_filter import VonMisesFilter class TrackManagerAtomicFilterStateTest(unittest.TestCase): @@ -32,3 +34,19 @@ def test_invalid_replacement_preserves_existing_track_bank(self): manager.tracks[0].get_point_estimate(), np.array([1.0]), ) + + def test_scalar_filter_point_estimate_stacks_as_one_dimensional_state(self): + circular_filter = VonMisesFilter() + circular_filter.filter_state = VonMisesDistribution(0.25, 2.0) + manager = TrackManager( + extract_confirmed_only=False, + keep_history=False, + log_prior_estimates=False, + log_posterior_estimates=False, + ) + manager.initialize_from_states([circular_filter], confirmed=True) + + point_estimate = manager.get_point_estimate() + + self.assertEqual(point_estimate.shape, (1, 1)) + np.testing.assert_allclose(point_estimate, np.array([[0.25]]))