Skip to content
Open
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
7 changes: 4 additions & 3 deletions src/pyrecest/filters/manifold_exponential_moving_average.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Exponential moving average for states on manifolds."""

import copy
from typing import Any, Callable

import numpy as np
Expand Down Expand Up @@ -107,16 +108,16 @@ def filter_state(self):

@filter_state.setter
def filter_state(self, new_state):
self._filter_state = new_state
self._filter_state = copy.deepcopy(new_state)

def update(self, sample):
"""Update the moving average with a new manifold-valued sample."""
if self._filter_state is None:
self._filter_state = sample
self.filter_state = sample
return

tangent_update = self.alpha * asarray(self.phi_inv(self._filter_state, sample))
self._filter_state = self.phi(self._filter_state, tangent_update)
self.filter_state = self.phi(self._filter_state, tangent_update)

def get_point_estimate(self):
"""Return the current manifold estimate."""
Expand Down
44 changes: 44 additions & 0 deletions tests/filters/test_manifold_exponential_moving_average_aliasing.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
"""Regression tests for manifold EMA state ownership."""

import numpy as np
import numpy.testing as npt

from pyrecest.filters import ManifoldExponentialMovingAverage


def _phi_euclidean(state, tangent):
return state + tangent


def _phi_inv_euclidean(state_ref, state):
return state - state_ref


def test_first_sample_is_copied_into_filter_state():
sample = np.array([2.0, 3.0])
ema = ManifoldExponentialMovingAverage(
initial_state=None,
alpha=0.5,
phi=_phi_euclidean,
phi_inv=_phi_inv_euclidean,
)

ema.update(sample)
sample[:] = -1.0

npt.assert_array_equal(ema.filter_state, np.array([2.0, 3.0]))


def test_explicit_state_assignment_does_not_alias_caller_array():
ema = ManifoldExponentialMovingAverage(
initial_state=np.array([0.0, 0.0]),
alpha=0.5,
phi=_phi_euclidean,
phi_inv=_phi_inv_euclidean,
)
replacement = np.array([4.0, 5.0])

ema.filter_state = replacement
replacement[:] = 0.0

npt.assert_array_equal(ema.filter_state, np.array([4.0, 5.0]))
Loading