diff --git a/src/pyrecest/filters/dirichlet_process_birth_tracker.py b/src/pyrecest/filters/dirichlet_process_birth_tracker.py index 0a812478c..e10755c0c 100644 --- a/src/pyrecest/filters/dirichlet_process_birth_tracker.py +++ b/src/pyrecest/filters/dirichlet_process_birth_tracker.py @@ -238,29 +238,44 @@ def _create_birth_component_from_measurement( birth_mean, birth_covariance, ) - self.last_birth_diagnostics.append(decision) if decision["action"] == "clutter": + self.last_birth_diagnostics.append(decision) return None + pruning_threshold = self._normalize_birth_atom_pruning_threshold() + maximum_number_of_birth_atoms = self._normalize_maximum_number_of_birth_atoms() + + atom_index = decision["atom_index"] if decision["action"] == "new_atom": atom = DirichletProcessBirthAtom(birth_mean, birth_covariance) - self.birth_atoms.append(atom) - decision["atom_index"] = len(self.birth_atoms) - 1 + atom_index = len(self.birth_atoms) + decision["atom_index"] = atom_index else: - atom = self.birth_atoms[decision["atom_index"]] + atom = self.birth_atoms[atom_index].copy() atom.update_from_measurement( measurement, measurement_matrix, measurement_covariance, ) - self._prune_and_cap_birth_atoms() - return BernoulliComponent( + birth_component = BernoulliComponent( decision["existence_probability"], GaussianDistribution(atom.mean, atom.covariance), label=label, ) + if decision["action"] == "new_atom": + self.birth_atoms.append(atom) + else: + self.birth_atoms[atom_index] = atom + + self._prune_and_cap_birth_atoms( + pruning_threshold=pruning_threshold, + maximum_number_of_birth_atoms=maximum_number_of_birth_atoms, + ) + self.last_birth_diagnostics.append(decision) + return birth_component + def _resolve_birth_covariance( self, measurement, diff --git a/tests/filters/test_dirichlet_process_birth_atomicity.py b/tests/filters/test_dirichlet_process_birth_atomicity.py new file mode 100644 index 000000000..c8e8c975e --- /dev/null +++ b/tests/filters/test_dirichlet_process_birth_atomicity.py @@ -0,0 +1,106 @@ +import numpy as np +import pyrecest.backend +import pytest +from pyrecest.filters.dirichlet_process_birth_tracker import ( + DirichletProcessBirthMultiBernoulliTracker, +) + +pytestmark = pytest.mark.skipif( + pyrecest.backend.__backend_name__ != "numpy", + reason="DP birth multi-Bernoulli tracker is NumPy-only.", +) + + +def _tracker(**overrides): + birth_atoms = overrides.pop("birth_atoms", None) + tracker_param = { + "birth_covariance": np.diag([1.0, 1.0, 4.0, 4.0]), + "birth_existence_probability": 0.8, + "clutter_intensity": 1e-6, + "dp_concentration": 0.05, + "dp_birth_threshold": 1.0, + "measurement_to_state_matrix": np.array( + [ + [1.0, 0.0], + [0.0, 1.0], + [0.0, 0.0], + [0.0, 0.0], + ] + ), + } + tracker_param.update(overrides) + measurement_matrix = np.array( + [ + [1.0, 0.0, 0.0, 0.0], + [0.0, 1.0, 0.0, 0.0], + ] + ) + measurement_covariance = np.eye(2) * 0.2 + return ( + DirichletProcessBirthMultiBernoulliTracker( + tracker_param=tracker_param, + birth_atoms=birth_atoms, + ), + measurement_matrix, + measurement_covariance, + ) + + +def _assert_atom_unchanged(atom, reference): + np.testing.assert_allclose(atom.mean, reference.mean) + np.testing.assert_allclose(atom.covariance, reference.covariance) + assert atom.count == reference.count + + +def test_new_birth_rejects_invalid_atom_cap_without_partial_state(): + tracker, measurement_matrix, measurement_covariance = _tracker( + maximum_number_of_birth_atoms=1.5 + ) + + with pytest.raises(ValueError, match="maximum_number_of_birth_atoms"): + tracker._create_birth_component_from_measurement( + np.array([2.0, 3.0]), + measurement_matrix, + measurement_covariance, + ) + + assert tracker.birth_atoms == [] + assert tracker.last_birth_diagnostics == [] + + +def test_existing_birth_rejects_invalid_pruning_without_mutating_atom(): + tracker, measurement_matrix, measurement_covariance = _tracker( + birth_atoms=[(np.zeros(4), np.eye(4), 2.0)], + dp_birth_atom_pruning_threshold=np.nan, + ) + atom_before = tracker.get_birth_atoms()[0] + + with pytest.raises(ValueError, match="dp_birth_atom_pruning_threshold"): + tracker._create_birth_component_from_measurement( + np.array([0.1, -0.1]), + measurement_matrix, + measurement_covariance, + ) + + assert len(tracker.birth_atoms) == 1 + _assert_atom_unchanged(tracker.birth_atoms[0], atom_before) + assert tracker.last_birth_diagnostics == [] + + +def test_failed_birth_component_construction_does_not_mutate_existing_atom(): + tracker, measurement_matrix, measurement_covariance = _tracker( + birth_atoms=[(np.zeros(4), np.eye(4), 2.0)], + birth_existence_probability=np.nan, + ) + atom_before = tracker.get_birth_atoms()[0] + + with pytest.raises(ValueError, match="existence_probability"): + tracker._create_birth_component_from_measurement( + np.array([0.1, -0.1]), + measurement_matrix, + measurement_covariance, + ) + + assert len(tracker.birth_atoms) == 1 + _assert_atom_unchanged(tracker.birth_atoms[0], atom_before) + assert tracker.last_birth_diagnostics == []