Skip to content

Commit c9b0c65

Browse files
authored
Make MEM-RBPF prediction failures atomic (#5446)
* Make MEM-RBPF prediction updates atomic * Add MEM-RBPF prediction atomicity regressions * Format MEM-RBPF atomicity regressions
1 parent b03c732 commit c9b0c65

2 files changed

Lines changed: 125 additions & 16 deletions

File tree

src/pyrecest/filters/mem_rbpf_tracker.py

Lines changed: 36 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -325,19 +325,17 @@ def predict_linear(
325325
raise ValueError(
326326
f"inputs must have shape {expected_shape}, got {inputs.shape}"
327327
)
328+
329+
next_system_matrix = self.system_matrix
328330
if system_matrix is not None:
329-
self.system_matrix = array(system_matrix)
330-
self._validate_system_matrix(self.system_matrix)
331+
next_system_matrix = array(system_matrix)
332+
self._validate_system_matrix(next_system_matrix)
333+
334+
next_sys_noise = self.sys_noise
331335
if sys_noise is not None:
332-
self.sys_noise = self._as_covariance(
336+
next_sys_noise = self._as_covariance(
333337
sys_noise, self.state_dim, "sys_noise", require_pd=False
334338
)
335-
self.kinematic_state = self.system_matrix @ self.kinematic_state
336-
if inputs is not None:
337-
self.kinematic_state = self.kinematic_state + inputs
338-
self.covariance = self._symmetrize(
339-
self.system_matrix @ self.covariance @ self.system_matrix.T + self.sys_noise
340-
)
341339

342340
axis_matrix = eye(2)
343341
if shape_system_matrix is not None:
@@ -348,26 +346,48 @@ def predict_linear(
348346
axis_matrix = shape_system_matrix
349347
else:
350348
raise ValueError("shape_system_matrix must be 3x3 or 2x2")
349+
350+
next_orientation_process_variance = self.orientation_process_variance
351+
next_axis_sys_noise = self.axis_sys_noise
351352
if shape_sys_noise is not None:
352353
shape_sys_noise = array(shape_sys_noise)
353354
if shape_sys_noise.shape == (3, 3):
354355
shape_sys_noise = self._as_covariance(
355356
shape_sys_noise, 3, "shape_sys_noise", require_pd=False
356357
)
357-
self.orientation_process_variance = float(shape_sys_noise[0, 0])
358-
self.axis_sys_noise = shape_sys_noise[1:, 1:]
358+
next_orientation_process_variance = float(shape_sys_noise[0, 0])
359+
next_axis_sys_noise = shape_sys_noise[1:, 1:]
359360
elif shape_sys_noise.shape == (2, 2):
360-
self.axis_sys_noise = self._as_covariance(
361+
next_axis_sys_noise = self._as_covariance(
361362
shape_sys_noise, 2, "shape_sys_noise", require_pd=False
362363
)
363364
else:
364365
raise ValueError("shape_sys_noise must be 3x3 or 2x2")
365-
self.axis = self.axis @ axis_matrix.T
366-
self.axis_covariances = self._symmetrize_stack(
366+
367+
next_kinematic_state = next_system_matrix @ self.kinematic_state
368+
if inputs is not None:
369+
next_kinematic_state = next_kinematic_state + inputs
370+
next_covariance = self._symmetrize(
371+
next_system_matrix @ self.covariance @ next_system_matrix.T
372+
+ next_sys_noise
373+
)
374+
next_axis = self.axis @ axis_matrix.T
375+
next_axis_covariances = self._symmetrize_stack(
367376
axis_matrix @ self.axis_covariances @ axis_matrix.T
368-
+ self.axis_sys_noise.reshape((1, 2, 2))
377+
+ next_axis_sys_noise.reshape((1, 2, 2))
369378
)
370-
self._apply_axis_floor()
379+
if self.axis_floor is not None:
380+
next_axis = maximum(next_axis, float(self.axis_floor))
381+
382+
self.system_matrix = next_system_matrix
383+
self.sys_noise = next_sys_noise
384+
self.orientation_process_variance = next_orientation_process_variance
385+
self.axis_sys_noise = next_axis_sys_noise
386+
self.kinematic_state = next_kinematic_state
387+
self.covariance = next_covariance
388+
self.axis = next_axis
389+
self.axis_covariances = next_axis_covariances
390+
371391
if self.log_prior_estimates:
372392
self.store_prior_estimates()
373393
if self.log_prior_extents:
Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,89 @@
1+
import numpy as np
2+
import numpy.testing as npt
3+
import pytest
4+
from pyrecest import backend
5+
from pyrecest.backend import array, diag, eye
6+
from pyrecest.filters.mem_rbpf_tracker import MEMRBPFTracker
7+
8+
9+
pytestmark = pytest.mark.skipif(
10+
backend.__backend_name__ == "jax",
11+
reason="MEMRBPFTracker is unsupported on JAX.",
12+
)
13+
14+
15+
def _make_tracker():
16+
return MEMRBPFTracker(
17+
kinematic_state=array([0.0, 0.0, 1.0, -0.5]),
18+
covariance=eye(4),
19+
shape_state=array([0.2, 2.0, 1.0]),
20+
shape_covariance=diag(array([0.05, 0.1, 0.1])),
21+
meas_noise_cov=0.05 * eye(2),
22+
sys_noise=0.01 * eye(4),
23+
shape_sys_noise=diag(array([0.01, 0.01, 0.01])),
24+
n_particles=8,
25+
resampling_threshold=0,
26+
rng=7,
27+
)
28+
29+
30+
def _to_numpy_copy(value):
31+
return np.asarray(backend.to_numpy(value)).copy()
32+
33+
34+
def _snapshot(tracker):
35+
return {
36+
"kinematic_state": _to_numpy_copy(tracker.kinematic_state),
37+
"covariance": _to_numpy_copy(tracker.covariance),
38+
"system_matrix": _to_numpy_copy(tracker.system_matrix),
39+
"sys_noise": _to_numpy_copy(tracker.sys_noise),
40+
"axis": _to_numpy_copy(tracker.axis),
41+
"axis_covariances": _to_numpy_copy(tracker.axis_covariances),
42+
"axis_sys_noise": _to_numpy_copy(tracker.axis_sys_noise),
43+
"orientation_process_variance": tracker.orientation_process_variance,
44+
}
45+
46+
47+
def _assert_snapshot_equal(tracker, snapshot):
48+
npt.assert_array_equal(
49+
_to_numpy_copy(tracker.kinematic_state), snapshot["kinematic_state"]
50+
)
51+
npt.assert_array_equal(
52+
_to_numpy_copy(tracker.covariance), snapshot["covariance"]
53+
)
54+
npt.assert_array_equal(
55+
_to_numpy_copy(tracker.system_matrix), snapshot["system_matrix"]
56+
)
57+
npt.assert_array_equal(_to_numpy_copy(tracker.sys_noise), snapshot["sys_noise"])
58+
npt.assert_array_equal(_to_numpy_copy(tracker.axis), snapshot["axis"])
59+
npt.assert_array_equal(
60+
_to_numpy_copy(tracker.axis_covariances), snapshot["axis_covariances"]
61+
)
62+
npt.assert_array_equal(
63+
_to_numpy_copy(tracker.axis_sys_noise), snapshot["axis_sys_noise"]
64+
)
65+
assert (
66+
tracker.orientation_process_variance
67+
== snapshot["orientation_process_variance"]
68+
)
69+
70+
71+
@pytest.mark.parametrize(
72+
("override", "message"),
73+
[
74+
({"system_matrix": eye(3)}, "system_matrix"),
75+
({"shape_system_matrix": eye(4)}, "shape_system_matrix"),
76+
(
77+
{"shape_sys_noise": array([[1.0, 2.0], [2.0, 1.0]])},
78+
"shape_sys_noise",
79+
),
80+
],
81+
)
82+
def test_predict_validation_failures_leave_tracker_unchanged(override, message):
83+
tracker = _make_tracker()
84+
snapshot = _snapshot(tracker)
85+
86+
with pytest.raises(ValueError, match=message):
87+
tracker.predict_linear(**override)
88+
89+
_assert_snapshot_equal(tracker, snapshot)

0 commit comments

Comments
 (0)