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
5 changes: 3 additions & 2 deletions src/pyrecest/filters/abstract_filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import copy
from abc import ABC, abstractmethod

from pyrecest.backend import atleast_1d
from pyrecest.utils.history_recorder import HistoryRecorder


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