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
27 changes: 21 additions & 6 deletions src/pyrecest/filters/dirichlet_process_birth_tracker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
106 changes: 106 additions & 0 deletions tests/filters/test_dirichlet_process_birth_atomicity.py
Original file line number Diff line number Diff line change
@@ -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 == []
Loading