From 83bb4a798b82cc913d004fcc4d54958ad13ce477 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Tue, 15 Sep 2026 14:02:03 +0200 Subject: [PATCH 1/3] Add normalized CheapJPDAF implementation, regression tests, and guide --- .github/workflows/implement-cheap-jpdaf.yml | 58 ++++ docs/cheap-jpdaf.md | 136 +++++++++ scripts/_integrate_cheap_jpdaf.py | 95 ++++++ ...t_probabilistic_data_association_filter.py | 199 +++++++++++++ ...t_probabilistic_data_association_filter.py | 274 ++++++++++++++++++ 5 files changed, 762 insertions(+) create mode 100644 .github/workflows/implement-cheap-jpdaf.yml create mode 100644 docs/cheap-jpdaf.md create mode 100644 scripts/_integrate_cheap_jpdaf.py create mode 100644 src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py create mode 100644 tests/filters/test_cheap_joint_probabilistic_data_association_filter.py diff --git a/.github/workflows/implement-cheap-jpdaf.yml b/.github/workflows/implement-cheap-jpdaf.yml new file mode 100644 index 0000000000..a2b87687eb --- /dev/null +++ b/.github/workflows/implement-cheap-jpdaf.yml @@ -0,0 +1,58 @@ +name: Integrate and validate Cheap JPDAF + +on: + push: + branches: + - feature/cheap-jpdaf-20260915 + +permissions: + contents: write + +jobs: + integrate: + runs-on: ubuntu-latest + timeout-minutes: 20 + env: + PYRECEST_BACKEND: numpy + BRANCH_NAME: feature/cheap-jpdaf-20260915 + steps: + - uses: actions/checkout@v7 + with: + ref: feature/cheap-jpdaf-20260915 + - uses: actions/setup-python@v7 + with: + python-version: '3.12' + cache: pip + - name: Install package and validation tools + run: python -m pip install -e . pytest parameterized black isort pylint + - name: Integrate shared solver and public API metadata + run: python scripts/_integrate_cheap_jpdaf.py + - name: Format changed Python files + run: | + python -m isort src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py + python -m black src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py src/pyrecest/filters/joint_probabilistic_data_association_filter.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py + - name: Exact and cheap JPDAF regression tests + run: | + python -m pytest -q tests/filters/test_cheap_joint_probabilistic_data_association_filter.py tests/filters/test_joint_probabilistic_data_association_filter.py tests/filters/test_jpdaf_logdet_underflow.py + - name: Validate API documentation and imports + run: | + python scripts/render_backend_api_matrix.py --check docs/backend-api-matrix.md + python scripts/check_public_api_registry.py --check docs/public-api-registry.md + python scripts/check_minimal_imports.py + python -m pylint src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py src/pyrecest/filters/joint_probabilistic_data_association_filter.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py + python - <<'PY' + import re + from pathlib import Path + example = re.search(r'```python\n(.*?)```', Path('docs/cheap-jpdaf.md').read_text(), re.S) + exec(compile(example.group(1), 'docs/cheap-jpdaf.md', 'exec')) + print('Cheap JPDAF documentation example passed') + PY + - name: Commit validated integration and remove bootstrap files + run: | + git config user.name 'github-actions[bot]' + git config user.email '41898282+github-actions[bot]@users.noreply.github.com' + git rm .github/workflows/implement-cheap-jpdaf.yml scripts/_integrate_cheap_jpdaf.py + git add src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py src/pyrecest/filters/joint_probabilistic_data_association_filter.py src/pyrecest/filters/__init__.py src/pyrecest/_backend/capabilities.py src/pyrecest/api_registry.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py docs/cheap-jpdaf.md docs/api-overview.md docs/backend-compatibility.md docs/backend-api-matrix.md docs/public-api-registry.md + git diff --cached --stat + git commit -m 'Integrate CheapJPDAF with shared Gaussian updates and API metadata' + git push origin "HEAD:$BRANCH_NAME" diff --git a/docs/cheap-jpdaf.md b/docs/cheap-jpdaf.md new file mode 100644 index 0000000000..e1232d22a5 --- /dev/null +++ b/docs/cheap-jpdaf.md @@ -0,0 +1,136 @@ +# Cheap joint probabilistic data association + +`CheapJPDAF` (also `CJPDAF` and +`CheapJointProbabilisticDataAssociationFilter`) is a NumPy-only, +linear-Gaussian alternative to the exact `JPDAF`. It retains soft association +and Gaussian mixture moment matching without enumerating joint events. +It is a Fitzgerald-style approximation with an explicit normalization repair, +not a numerically equivalent fast implementation of exact JPDA. + +## Usage + +```python +import numpy as np +from pyrecest.distributions import GaussianDistribution +from pyrecest.filters import CheapJPDAF, KalmanFilter + +tracker = CheapJPDAF( + [ + KalmanFilter(GaussianDistribution(np.array([-1.0, 0.0]), np.eye(2))), + KalmanFilter(GaussianDistribution(np.array([1.0, 0.0]), np.eye(2))), + ], + association_param={ + "detection_probability": 0.95, + "clutter_intensity": 1e-3, + "gating_distance_threshold": 9.21, + }, +) +tracker.predict_linear(np.eye(2), 0.01 * np.eye(2)) +measurements = np.array([[-0.8, 0.9], [0.1, -0.1]]) # (measurement_dim, n_meas) +tracker.update_linear(measurements, np.eye(2), 0.1 * np.eye(2)) + +beta = tracker.latest_association_probabilities +assert np.allclose(beta.sum(axis=1), 1.0) +assert np.all(beta[:, 1:].sum(axis=0) <= 1.0 + 1e-12) +print(tracker.get_point_estimate()) +``` + +The constructor, `filter_state`, `predict_linear`, and `update_linear` follow +`JPDAF`. Measurement covariances may be shared `(d, d)` or per measurement +`(d, d, n_meas)`. Clutter intensity may be a positive finite scalar or one +positive finite value per measurement. Detection probability must satisfy +`0 < P_D < 1`. Gating compares **squared** Mahalanobis distance with the +threshold; the inherited default is the 99.9% chi-square quantile for the +measurement dimension. `max_enumerated_events` is ignored by this class; +it remains an enforced limit in exact `JPDAF`. + +## Association and normalization convention + +Let `L[i,j]` be the predicted Gaussian likelihood, set to zero outside the +gate, and let `kappa[j]` be the clutter intensity. Define dimensionless odds + +```text +w[i,j] = P_D * L[i,j] / ((1 - P_D) * kappa[j]) +r[i] = sum_j w[i,j] +c[j] = sum_i w[i,j] +beta[i,j+1] = w[i,j] / (1 + r[i] + c[j] - w[i,j]) +beta[i,0] = 1 - sum_j beta[i,j+1] +``` + +For scalar clutter, this is the usual likelihood-form denominator +`T[i] + M[j] - L[i,j] + B` with `B=(1-P_D)*kappa/P_D`. +The row/column competition formula is reproduced in equations (49)-(52) of +[US patent application 20160245949](https://patents.justia.com/patent/20160245949), +which attributes it to Fitzgerald. Its original separate miss expression, +`1/(1+r[i])`, generally does **not** normalize a track's weights when tracks +compete. PyRecEst deliberately uses the complementary mass instead, and does +not claim to reproduce that unnormalized original variant. + +The denominator is at least `1+r[i]` and at least `1+c[j]`. Thus each row's +detection mass is at most one, and each measurement's total allocation is at +most one. Adding the complementary missed-detection mass gives a valid +per-track mixture. These bounds do not make the marginals exact Bayesian +association probabilities. + +For numerical stability, the implementation uses log likelihood ratios and +prefix/suffix log sums for the other tracks' likelihoods. It avoids both a +single global likelihood rescaling and subtraction of a dominant entry from +a column sum. The miss is evaluated through the positive-term identity + +```text +competition[i,j] = c[j] - w[i,j] +beta[i,0] = 1/(1+r[i]) + + sum_j (w[i,j]/(1+r[i])) + * competition[i,j]/(1+r[i]+competition[i,j]) +``` + +rather than floating-point subtraction from one. This preserves tiny miss +probabilities. No optional dependency or iterative association solver is added. + +The approximation coincides with the existing exact JPDAF association model +for a single track, a single measurement, or disjoint single-track validation +components. General ambiguous multi-track/multi-measurement cases differ. +For example, with a 2-by-2 matrix of unit odds, cheap JPDA assigns each track +`[miss=1/2, measurement_1=1/4, measurement_2=1/4]`, whereas exact JPDA gives +`[3/7, 2/7, 2/7]`. + +## State update and diagnostics + +The update reuses the exact JPDAF's per-pair Kalman hypotheses and mixture +moment matching, including the between-hypothesis covariance term. All-gated +and empty-measurement scans leave the prior unchanged. Like exact JPDAF here, +the miss prior is `1-P_D`: the gate does not introduce an additional gate +probability `P_G`. No track-birth or deletion mechanism is added. + +`find_association_probabilities(...)` returns `(beta, greedy_assignment)`. +`beta[:,0]` contains misses, and the remaining columns follow measurement order. +The second value is a **track-order greedy likelihood-ratio diagnostic**, with +`-1` for misses and no measurement reused. It is not a MAP event and is not used +by `update_linear`. Ties with a miss are resolved as misses; equal detection +scores use measurement order. Consequently, the diagnostic may depend on track +order even though the soft marginals are permutation-equivariant. + +The diagnostic is stored in `latest_greedy_association` and returned by +`find_association(...)`. `latest_map_association` remains `None`, so callers +cannot accidentally interpret the diagnostic as an exact MAP result. Empty +banks retain the exact class's legacy probability-array shape `(0,1)`. + +## Cost and limitations + +After pairwise Gaussian likelihoods are available, marginals and the greedy +diagnostic require **O(n_targets * n_meas)** work and memory, not a combinatorial +number of events. The whole update has this target/measurement scaling when +state and measurement dimensions are fixed. Gaussian linear algebra still +contributes its usual dimension-dependent cost. + +Competition can move extra mass to the missed-detection component; the result +is intentionally conservative in ambiguous examples, but is not a calibrated +probability guarantee. It does not solve track coalescence or establish a +trajectory-level association model. Use exact `JPDAF` as a small-problem +reference rather than expecting identical output in dense scenes. + +Focused regressions are in +`tests/filters/test_cheap_joint_probabilistic_data_association_filter.py` and +cover reference probabilities, normalization and allocation bounds, exact +special cases, numerical extremes, diagnostic semantics, event-limit +independence, covariance moment matching, and input/backend contracts. diff --git a/scripts/_integrate_cheap_jpdaf.py b/scripts/_integrate_cheap_jpdaf.py new file mode 100644 index 0000000000..efd4b79b1b --- /dev/null +++ b/scripts/_integrate_cheap_jpdaf.py @@ -0,0 +1,95 @@ +from pathlib import Path +import subprocess +import sys + + +def replace_once(path, old, new): + path = Path(path) + text = path.read_text() + if text.count(old) != 1: + raise RuntimeError(f'Expected one integration anchor in {path}: {old!r}') + path.write_text(text.replace(old, new, 1)) + + +renderers = { + 'docs/backend-api-matrix.md': 'scripts/render_backend_api_matrix.py', + 'docs/public-api-registry.md': 'scripts/check_public_api_registry.py', +} +old_tables = {doc: subprocess.check_output([sys.executable, script], text=True) + for doc, script in renderers.items()} + +path = Path('src/pyrecest/filters/joint_probabilistic_data_association_filter.py') +text = path.read_text() +start = text.index(' track_order = sorted(') +end = text.index(' self.latest_association_probabilities = association_probabilities', start) +solver = text[start:end] +replacement = ''' association_probabilities, map_association = self._compute_association_probabilities( + log_likelihoods, + eligible_measurements, + detection_probability, + clutter_intensity, + ) + +''' +text = text[:start] + replacement + text[end:] +method = ''' def _compute_association_probabilities( + self, + log_likelihoods, + eligible_measurements, + detection_probability, + clutter_intensity, + ): + """Solve the gated association problem by exact joint-event enumeration.""" + n_targets, n_meas = log_likelihoods.shape +''' +method += solver + ' return association_probabilities, map_association\n\n' +anchor = ' def find_association(\n' +assert text.count(anchor) == 1 +text = text.replace(anchor, method + anchor, 1) +path.write_text(text) + +names = ('CheapJointProbabilisticDataAssociationFilter', 'CheapJPDAF', 'CJPDAF') +exports = ''.join(f' "{name}": ".cheap_joint_probabilistic_data_association_filter",\n' for name in names) +anchor = ' "JPDAF": ".joint_probabilistic_data_association_filter",\n' +replace_once('src/pyrecest/filters/__init__.py', anchor, anchor + exports) + +capabilities = ''.join( + f''' "{name}": {{ + "numpy": "supported", + "pytorch": "unsupported", + "jax": "unsupported", + "notes": "Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration.", + }}, +''' for name in names) +replace_once('src/pyrecest/_backend/capabilities.py', 'API_BACKEND_CAPABILITIES: Final = {\n', + 'API_BACKEND_CAPABILITIES: Final = {\n' + capabilities) +registry = ''.join( + f''' "{name}": {{ + "module": "pyrecest.filters", + "category": "experimental", + "backend_contract": "{name}", + "notes": "Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only.", + }}, +''' for name in names) +replace_once('src/pyrecest/api_registry.py', 'PUBLIC_API_REGISTRY: Final = {\n', + 'PUBLIC_API_REGISTRY: Final = {\n' + registry) +for doc, script in renderers.items(): + new_table = subprocess.check_output([sys.executable, script], text=True) + replace_once(doc, old_tables[doc], new_table) + +with Path('docs/api-overview.md').open('a') as stream: + stream.write('\n## Cheap joint probabilistic data association\n\n' + '`CheapJPDAF` / `CJPDAF` provides normalized soft association for\n' + 'linear-Gaussian tracking without joint-event enumeration. It reuses\n' + 'the exact `JPDAF` Gaussian update and is explicitly NumPy-only.\n' + 'See [Cheap JPDAF](cheap-jpdaf.md) for usage, normalization,\n' + 'complexity, and the non-MAP greedy diagnostic.\n') +with Path('docs/backend-compatibility.md').open('a') as stream: + stream.write('\n### Cheap JPDAF\n\n' + '`CheapJPDAF`, `CJPDAF`, and\n' + '`CheapJointProbabilisticDataAssociationFilter` are NumPy-only.\n' + 'Association and measurement updates reject other backends explicitly.\n' + 'See [Cheap JPDAF](cheap-jpdaf.md) for the approximation contract.\n') +path = Path('tests/filters/test_cheap_joint_probabilistic_data_association_filter.py') +text = path.read_text() +path.write_text('# pylint: disable=protected-access,no-name-in-module,no-member\n' + text) diff --git a/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py b/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py new file mode 100644 index 0000000000..70adea6173 --- /dev/null +++ b/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py @@ -0,0 +1,199 @@ +"""Fitzgerald-style cheap JPDA with normalized missed-detection weights.""" + +from math import log, log1p + +import numpy as np +from scipy.special import logsumexp + +from .joint_probabilistic_data_association_filter import ( + JointProbabilisticDataAssociationFilter, +) + + +class CheapJointProbabilisticDataAssociationFilter( + JointProbabilisticDataAssociationFilter +): + r"""Linear-Gaussian cheap JPDA, without joint-event enumeration. + + For gated measurement likelihoods ``L[i, j]`` and clutter intensities + ``kappa[j]``, define ``w[i, j] = P_D * L[i, j] / ((1-P_D) * kappa[j])``. + Writing ``r[i] = sum_j w[i, j]`` and ``c[j] = sum_i w[i, j]``, use + + .. math:: + + \beta_{ij} = \frac{w_{ij}}{1+r_i+c_j-w_{ij}},\qquad + \beta_{i0} = 1-\sum_j\beta_{ij}. + + This retains Fitzgerald's row/column competition approximation, but uses + the complementary missed-detection mass, rather than the generally + non-normalizing original choice ``1 / (1 + r[i])``. With scalar clutter, + this is the likelihood-form formula with ``B=(1-P_D)*kappa/P_D``. + It is an approximation, not exact JPDA or independent per-track PDA. + + Gating, Gaussian hypotheses, prediction, and moment-matched covariance + updates are inherited from :class:`JointProbabilisticDataAssociationFilter`. + All association sums are evaluated in log space. The association stage + takes O(n_targets * n_meas) time and memory; no enumeration limit is used. + Measurement/state dimensions are held fixed in this complexity statement. + + ``find_association_probabilities`` returns probabilities followed by a + feasible, track-order greedy likelihood-ratio assignment (not joint MAP). + This diagnostic is available as ``latest_greedy_association`` and does not + affect the soft update. ``latest_map_association`` remains ``None`` because + no joint MAP event is computed. Misses are encoded as ``-1``. + + Uses the same association parameters as JPDAF. ``max_enumerated_events`` + is accepted for configuration compatibility, but ignored. Only the NumPy + backend and linear-Gaussian measurement models are supported. As in JPDAF, + the miss prior is ``1-P_D``; gating does not introduce a separate ``P_G``. + See ``docs/cheap-jpdaf.md`` for the normalization convention and limitations. + """ + + def __init__( + self, + initial_prior=None, + association_param=None, + log_prior_estimates=True, + log_posterior_estimates=True, + ): + super().__init__( + initial_prior, + association_param, + log_prior_estimates, + log_posterior_estimates, + ) + self.latest_greedy_association = None + + @staticmethod + def _prepare_clutter_intensity(clutter_intensity, n_meas): + clutter_intensity = ( + JointProbabilisticDataAssociationFilter._prepare_clutter_intensity( + clutter_intensity, n_meas + ) + ) + if not np.all(np.isfinite(clutter_intensity)): + raise ValueError("clutter_intensity must be finite and strictly positive.") + return clutter_intensity + + @staticmethod + def _cheap_marginals(log_weights): + """Compute normalized marginals from gated log detection-to-miss odds. + + ``-inf`` denotes a gated-out pair. The caller supplies a two-dimensional + array containing only finite values or ``-inf``. Prefix/suffix log sums + avoid subtracting a dominant pair from a column total, and do not share + one global scaling factor across disconnected association components. + """ + n_targets, n_meas = log_weights.shape + probabilities = np.zeros((n_targets, n_meas + 1)) + if n_targets == 0: + return probabilities + if n_meas == 0: + probabilities[:, 0] = 1.0 + return probabilities + + before = np.full_like(log_weights, -np.inf) + after = np.full_like(log_weights, -np.inf) + if n_targets > 1: + before[1:] = np.logaddexp.accumulate(log_weights[:-1], axis=0) + after[:-1] = np.logaddexp.accumulate(log_weights[:0:-1], axis=0)[::-1] + log_competition = np.logaddexp(before, after) + log_row_denominator = np.logaddexp( + 0.0, logsumexp(log_weights, axis=1, keepdims=True) + ) + log_denominator = np.logaddexp(log_row_denominator, log_competition) + probabilities[:, 1:] = np.exp(log_weights - log_denominator) + + # Algebraically 1 - sum(beta), but expressed as a sum of positive terms: + # 1/(1+r) + sum_j [w_j/(1+r)] * [competition_j/(1+r+competition_j)]. + # This preserves small miss probabilities without cancellation. + probabilities[:, 0] = np.exp(-log_row_denominator[:, 0]) + np.sum( + np.exp( + (log_weights - log_row_denominator) + + (log_competition - log_denominator) + ), + axis=1, + ) + return probabilities + + @staticmethod + def _greedy_assignment(log_weights): + """Return a feasible O(n_targets * n_meas) diagnostic, not joint MAP.""" + n_targets, n_meas = log_weights.shape + assignment = np.full(n_targets, -1, dtype=int) + available = np.ones(n_meas, dtype=bool) + if n_meas == 0: + return assignment + for track_index, row in enumerate(log_weights): + scores = np.where(available, row, -np.inf) + measurement_index = int(np.argmax(scores)) + # A unit detection-to-miss likelihood ratio ties with a miss. + if scores[measurement_index] > 0.0: + assignment[track_index] = measurement_index + available[measurement_index] = False + return assignment + + def _compute_association_probabilities( + self, + log_likelihoods, + eligible_measurements, + detection_probability, + clutter_intensity, + ): + """Replace the exact event solver, leaving Gaussian updates unchanged.""" + del eligible_measurements # The gated log-likelihood matrix encodes these. + if np.any(np.isnan(log_likelihoods)) or np.any(np.isposinf(log_likelihoods)): + raise ValueError("Gated log likelihoods must be finite or negative infinity.") + log_weights = ( + log_likelihoods + + log(detection_probability) + - log1p(-detection_probability) + - np.log(clutter_intensity)[None, :] + ) + return self._cheap_marginals(log_weights), self._greedy_assignment(log_weights) + + def find_association_probabilities( + self, + measurements, + measurement_matrix, + cov_mats_meas, + warn_on_no_meas_for_track=True, + ): + """Return normalized marginals and a greedy diagnostic, not joint MAP. + + Column zero is the missed-detection mass; remaining columns correspond + to measurements. Empty banks retain JPDAF's legacy ``(0, 1)`` shape. + """ + probabilities, greedy_assignment = super().find_association_probabilities( + measurements, + measurement_matrix, + cov_mats_meas, + warn_on_no_meas_for_track=warn_on_no_meas_for_track, + ) + self.latest_association_probabilities = probabilities + self.latest_greedy_association = greedy_assignment + self.latest_map_association = None + if probabilities.shape[0] == 0: + self._latest_posterior_hypotheses = [] + return probabilities, greedy_assignment + + def find_association( + self, + measurements, + measurement_matrix, + cov_mats_meas, + warn_on_no_meas_for_track=True, + ): + """Return the track-order greedy diagnostic, not a MAP joint event.""" + return super().find_association( + measurements, + measurement_matrix, + cov_mats_meas, + warn_on_no_meas_for_track=warn_on_no_meas_for_track, + ) + + +CheapJPDAF = CheapJointProbabilisticDataAssociationFilter +CJPDAF = CheapJointProbabilisticDataAssociationFilter + +__all__ = ["CheapJointProbabilisticDataAssociationFilter", "CheapJPDAF", "CJPDAF"] diff --git a/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py b/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py new file mode 100644 index 0000000000..e358e306ea --- /dev/null +++ b/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py @@ -0,0 +1,274 @@ +"""Cheap JPDA reference cases, normalization, and Gaussian-update contracts.""" + +import numpy as np +import numpy.testing as npt +import pytest + +import pyrecest.backend +from pyrecest.distributions import GaussianDistribution +from pyrecest.filters import ( + CJPDAF, + JPDAF, + CheapJointProbabilisticDataAssociationFilter, + CheapJPDAF, + KalmanFilter, +) + +pytestmark = pytest.mark.skipif( + pyrecest.backend.__backend_name__ != "numpy", + reason="Cheap JPDAF is explicitly NumPy-only.", +) + + +def _log_weights(weights): + weights = np.asarray(weights, dtype=float) + result = np.full_like(weights, -np.inf) + np.log(weights, out=result, where=weights > 0) + return result + + +def _marginals(weights, tracker_class=CheapJPDAF): + log_weights = _log_weights(weights) + eligible = [np.flatnonzero(np.isfinite(row)).tolist() for row in log_weights] + tracker = tracker_class() + # P_D=0.5 and kappa=1 make likelihoods equal detection-to-miss odds. + return tracker._compute_association_probabilities( + log_weights, eligible, 0.5, np.ones(log_weights.shape[1]) + )[0] + + +def _tracker(tracker_class=CheapJPDAF, means=(-1.0, 1.0), **parameters): + return tracker_class( + [ + KalmanFilter(GaussianDistribution(np.array([mean, 0.0]), np.eye(2))) + for mean in means + ], + association_param={ + "detection_probability": 0.9, + "clutter_intensity": 0.01, + "gating_distance_threshold": 100.0, + **parameters, + }, + ) + + +def test_public_aliases(): + assert CheapJPDAF is CJPDAF is CheapJointProbabilisticDataAssociationFilter + assert issubclass(CheapJPDAF, JPDAF) + + +def test_hand_computed_competition_and_complementary_miss(): + weights = np.array([[2.0, 3.0], [5.0, 7.0]]) + beta = _marginals(weights) + expected = np.array([[2 / 11, 3 / 13], [5 / 15, 7 / 16]]) + npt.assert_allclose(beta[:, 1:], expected) + npt.assert_allclose(beta[:, 0], 1 - expected.sum(axis=1)) + # The uncorrected Fitzgerald miss weight does not normalize these rows. + assert np.all(1 / (1 + weights.sum(axis=1)) < beta[:, 0]) + + +@pytest.mark.parametrize( + "weights", + [ + [[2.0, 3.0, 0.0]], + [[2.0], [5.0], [0.0]], + [[2.0, 3.0, 0.0, 0.0], [0.0, 0.0, 5.0, 7.0]], + [[0.0, 0.0], [0.0, 0.0]], + ], +) +def test_agrees_with_exact_jpda_in_simple_cases(weights): + npt.assert_allclose(_marginals(weights), _marginals(weights, JPDAF), atol=1e-14) + + +def test_general_case_is_an_approximation_not_exact_enumeration(): + weights = np.ones((2, 2)) + npt.assert_allclose(_marginals(weights), [[0.5, 0.25, 0.25]] * 2) + npt.assert_allclose(_marginals(weights, JPDAF), [[3 / 7, 2 / 7, 2 / 7]] * 2) + + +def test_randomized_formula_rows_and_measurement_exclusivity(): + rng = np.random.default_rng(117) + for _ in range(100): + weights = rng.lognormal(0, 3, (11, 13)) + weights[rng.random(weights.shape) < 0.5] = 0.0 + beta = _marginals(weights) + denominator = ( + 1 + weights.sum(axis=1, keepdims=True) + + weights.sum(axis=0, keepdims=True) - weights + ) + npt.assert_allclose(beta[:, 1:], weights / denominator, atol=1e-14) + npt.assert_allclose(beta.sum(axis=1), 1.0, atol=1e-14) + assert np.all(np.isfinite(beta)) + assert np.all(beta >= 0.0) + assert np.all(beta <= 1.0 + 1e-14) + assert np.all(beta[:, 1:].sum(axis=0) <= 1.0 + 1e-14) + assert np.all(beta[:, 1:][weights == 0.0] == 0.0) + + +def test_permutation_equivariance_of_soft_marginals(): + weights = np.array([[2.0, 0.0, 3.0], [5.0, 7.0, 0.1], [0.0, 1.0, 4.0]]) + tracks, measurements = [2, 0, 1], [1, 2, 0] + expected = _marginals(weights)[tracks][:, [0, 2, 3, 1]] + npt.assert_allclose(_marginals(weights[tracks][:, measurements]), expected) + + +def test_log_space_does_not_underflow_disconnected_components(): + log_weights = np.array([[1000.0, -np.inf], [-np.inf, 0.0]]) + with np.errstate(over="raise", invalid="raise", divide="raise"): + beta = CheapJPDAF._cheap_marginals(log_weights) + npt.assert_allclose(beta, [[0.0, 1.0, 0.0], [0.5, 0.0, 0.5]]) + + +def test_tiny_missed_detection_probability_is_preserved(): + beta = CheapJPDAF._cheap_marginals(np.array([[700.0]])) + assert beta[0, 0] > 0.0 + npt.assert_allclose(beta[0, 0], np.exp(-700.0), rtol=1e-13, atol=0) + + +def test_extreme_competition_and_empty_rows_stay_normalized(): + log_weights = np.array([[1000.0, -1000.0], [1000.0, -np.inf], [-np.inf, -np.inf]]) + with np.errstate(over="raise", invalid="raise", divide="raise"): + beta = CheapJPDAF._cheap_marginals(log_weights) + npt.assert_allclose(beta, [[0.5, 0.5, 0.0], [0.5, 0.5, 0.0], [1.0, 0.0, 0.0]]) + + +def test_greedy_diagnostic_is_feasible_but_not_claimed_to_be_map(): + # Track-order greedy chooses (0, 1), while the joint MAP event is (1, 0). + log_weights = _log_weights([[10.0, 9.0], [100.0, 1.1]]) + npt.assert_array_equal(CheapJPDAF._greedy_assignment(log_weights), [0, 1]) + tracker = _tracker(means=(0.0, 0.0)) + beta, diagnostic = tracker.find_association_probabilities( + np.zeros((2, 1)), np.eye(2), np.eye(2) + ) + npt.assert_array_equal(diagnostic, [0, -1]) + npt.assert_array_equal(tracker.latest_greedy_association, diagnostic) + npt.assert_allclose(tracker.latest_association_probabilities, beta) + assert tracker.latest_map_association is None + npt.assert_array_equal(tracker.find_association(np.zeros((2, 1)), np.eye(2), np.eye(2)), diagnostic) + + +def test_exact_event_limit_is_not_used(monkeypatch): + measurements = np.zeros((2, 12)) + exact = _tracker(JPDAF, means=(0.0,) * 12, max_enumerated_events=2) + with pytest.raises(RuntimeError, match="max_enumerated_events"): + exact.find_association_probabilities(measurements, np.eye(2), np.eye(2)) + + def no_enumeration(*_args, **_kwargs): + raise AssertionError("Cheap JPDAF must not call the exact event solver") + + monkeypatch.setattr(JPDAF, "_compute_association_probabilities", no_enumeration) + cheap = _tracker(means=(0.0,) * 12, max_enumerated_events=1) + cheap.update_linear(measurements, np.eye(2), np.eye(2)) + npt.assert_allclose(cheap.latest_association_probabilities.sum(axis=1), 1.0) + assert np.isfinite(cheap.get_point_estimate()).all() + + +def test_no_measurements_leave_state_and_covariance_unchanged(): + tracker = _tracker() + means = tracker.get_point_estimate().copy() + covariances = [state.C.copy() for state in tracker.filter_state] + tracker.update_linear(np.empty((2, 0)), np.eye(2), np.eye(2)) + npt.assert_allclose(tracker.latest_association_probabilities, [[1.0], [1.0]]) + npt.assert_array_equal(tracker.latest_greedy_association, [-1, -1]) + npt.assert_allclose(tracker.get_point_estimate(), means) + for state, covariance in zip(tracker.filter_state, covariances): + npt.assert_allclose(state.C, covariance) + + +def test_empty_bank_clears_cached_diagnostics(): + tracker = _tracker() + tracker.find_association_probabilities(np.zeros((2, 1)), np.eye(2), np.eye(2)) + tracker.filter_state = [] + with pytest.warns(UserWarning, match="zero targets"): + beta, diagnostic = tracker.find_association_probabilities( + np.zeros((2, 1)), np.eye(2), np.eye(2) + ) + assert beta.shape == (0, 1) + assert diagnostic.shape == (0,) + assert tracker.latest_association_probabilities.shape == (0, 1) + assert tracker.latest_greedy_association.shape == (0,) + assert tracker.latest_map_association is None + assert tracker._latest_posterior_hypotheses == [] + + +def test_all_measurements_outside_gate_leave_priors_unchanged(): + tracker = _tracker(gating_distance_threshold=0.1) + before = tracker.get_point_estimate().copy() + with pytest.warns(UserWarning, match="gating threshold"): + tracker.update_linear(np.full((2, 3), 100.0), np.eye(2), np.eye(2)) + npt.assert_allclose(tracker.latest_association_probabilities, [[1, 0, 0, 0]] * 2) + npt.assert_allclose(tracker.get_point_estimate(), before) + for state in tracker.filter_state: + npt.assert_allclose(state.C, np.eye(2)) + + +def test_heterogeneous_covariances_and_vector_clutter_match_exact_single_track(): + parameters = {"clutter_intensity": np.array([0.01, 0.04])} + cheap, exact = (_tracker(cls, means=(0.0,), **parameters) for cls in (CheapJPDAF, JPDAF)) + measurements = np.array([[-0.5, 1.0], [0.1, -0.2]]) + covariances = np.stack([0.2 * np.eye(2), 2 * np.eye(2)], axis=2) + for tracker in (cheap, exact): + tracker.update_linear(measurements, np.eye(2), covariances) + npt.assert_allclose(cheap.latest_association_probabilities, exact.latest_association_probabilities) + npt.assert_allclose(cheap.get_point_estimate(), exact.get_point_estimate()) + npt.assert_allclose(cheap.filter_state[0].C, exact.filter_state[0].C) + + +def test_mixture_update_includes_between_hypothesis_covariance(): + tracker = _tracker(means=(0.0,)) + measurements = np.array([[-2.0, 2.0], [0.0, 0.0]]) + tracker.update_linear(measurements, np.eye(2), np.eye(2)) + beta = tracker.latest_association_probabilities[0] + means = [np.zeros(2), np.array([-1.0, 0.0]), np.array([1.0, 0.0])] + covariances = [np.eye(2), 0.5 * np.eye(2), 0.5 * np.eye(2)] + mean = sum(weight * value for weight, value in zip(beta, means)) + covariance = sum( + weight * (cov + np.outer(value - mean, value - mean)) + for weight, value, cov in zip(beta, means, covariances) + ) + npt.assert_allclose(tracker.filter_state[0].mu, mean, atol=1e-14) + npt.assert_allclose(tracker.filter_state[0].C, covariance) + assert covariance[0, 0] > 1.0 + assert np.linalg.eigvalsh(covariance).min() > 0.0 + + +def test_prediction_uses_existing_filter_bank(): + tracker = _tracker() + before = tracker.get_point_estimate().copy() + tracker.predict_linear(np.eye(2), 0.1 * np.eye(2)) + npt.assert_allclose(tracker.get_point_estimate(), before) + for state in tracker.filter_state: + npt.assert_allclose(state.C, 1.1 * np.eye(2)) + + +@pytest.mark.parametrize("backend", ["jax", "pytorch"]) +def test_unsupported_backends_fail_explicitly(monkeypatch, backend): + tracker = _tracker() + monkeypatch.setattr(pyrecest.backend, "__backend_name__", backend) + with pytest.raises(NotImplementedError, match="numpy backend"): + tracker.update_linear(np.zeros((2, 1)), np.eye(2), np.eye(2)) + with pytest.raises(NotImplementedError, match="numpy backend"): + tracker.find_association_probabilities(np.zeros((2, 1)), np.eye(2), np.eye(2)) + + +@pytest.mark.parametrize("clutter", [0.0, -1.0, np.nan, np.inf, [0.01, np.nan]]) +def test_invalid_clutter_rejected(clutter): + tracker = _tracker(clutter_intensity=clutter) + with pytest.raises(ValueError, match="clutter_intensity"): + tracker.find_association_probabilities(np.zeros((2, 2)), np.eye(2), np.eye(2)) + + +@pytest.mark.parametrize("probability", [0.0, 1.0, -0.1, np.nan, np.inf]) +def test_invalid_detection_probability_rejected(probability): + tracker = _tracker(detection_probability=probability) + with pytest.raises(ValueError, match="detection_probability"): + tracker.find_association_probabilities(np.zeros((2, 1)), np.eye(2), np.eye(2)) + + +def test_invalid_covariance_shape_and_pairwise_costs_rejected(): + tracker = _tracker() + measurements = np.zeros((2, 1)) + with pytest.raises(ValueError, match="cov_mats_meas must have shape"): + tracker.find_association_probabilities(measurements, np.eye(2), np.ones((1, 1))) + with pytest.raises(NotImplementedError, match="pairwise_cost_matrix"): + tracker.update_linear(measurements, np.eye(2), np.eye(2), pairwise_cost_matrix=np.zeros((2, 1))) From 438f116ef90f4c9fdaae5804f1458c4b2224156a Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Tue, 15 Sep 2026 14:06:06 +0200 Subject: [PATCH 2/3] Isolate empty-bank association regression from inherited history logging --- scripts/_integrate_cheap_jpdaf.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/scripts/_integrate_cheap_jpdaf.py b/scripts/_integrate_cheap_jpdaf.py index efd4b79b1b..b0e9174913 100644 --- a/scripts/_integrate_cheap_jpdaf.py +++ b/scripts/_integrate_cheap_jpdaf.py @@ -93,3 +93,10 @@ def replace_once(path, old, new): path = Path('tests/filters/test_cheap_joint_probabilistic_data_association_filter.py') text = path.read_text() path.write_text('# pylint: disable=protected-access,no-name-in-module,no-member\n' + text) +replace_once( + path, + ' tracker.filter_state = []\n', + ' # Isolate association caches from the inherited empty-bank history logger.\n' + ' tracker.log_prior_estimates = False\n' + ' tracker.filter_state = []\n', +) From cb1abef77903d79c54289ce3ab427c4ccf186f43 Mon Sep 17 00:00:00 2001 From: "github-actions[bot]" <41898282+github-actions[bot]@users.noreply.github.com> Date: Tue, 15 Sep 2026 12:07:09 +0000 Subject: [PATCH 3/3] Integrate CheapJPDAF with shared Gaussian updates and API metadata --- .github/workflows/implement-cheap-jpdaf.yml | 58 ---------- docs/api-overview.md | 8 ++ docs/backend-api-matrix.md | 33 +++--- docs/backend-compatibility.md | 7 ++ docs/public-api-registry.md | 33 +++--- scripts/_integrate_cheap_jpdaf.py | 102 ------------------ src/pyrecest/_backend/capabilities.py | 18 ++++ src/pyrecest/api_registry.py | 18 ++++ src/pyrecest/filters/__init__.py | 3 + ...t_probabilistic_data_association_filter.py | 4 +- ...t_probabilistic_data_association_filter.py | 28 ++++- ...t_probabilistic_data_association_filter.py | 25 +++-- 12 files changed, 136 insertions(+), 201 deletions(-) delete mode 100644 .github/workflows/implement-cheap-jpdaf.yml delete mode 100644 scripts/_integrate_cheap_jpdaf.py diff --git a/.github/workflows/implement-cheap-jpdaf.yml b/.github/workflows/implement-cheap-jpdaf.yml deleted file mode 100644 index a2b87687eb..0000000000 --- a/.github/workflows/implement-cheap-jpdaf.yml +++ /dev/null @@ -1,58 +0,0 @@ -name: Integrate and validate Cheap JPDAF - -on: - push: - branches: - - feature/cheap-jpdaf-20260915 - -permissions: - contents: write - -jobs: - integrate: - runs-on: ubuntu-latest - timeout-minutes: 20 - env: - PYRECEST_BACKEND: numpy - BRANCH_NAME: feature/cheap-jpdaf-20260915 - steps: - - uses: actions/checkout@v7 - with: - ref: feature/cheap-jpdaf-20260915 - - uses: actions/setup-python@v7 - with: - python-version: '3.12' - cache: pip - - name: Install package and validation tools - run: python -m pip install -e . pytest parameterized black isort pylint - - name: Integrate shared solver and public API metadata - run: python scripts/_integrate_cheap_jpdaf.py - - name: Format changed Python files - run: | - python -m isort src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py - python -m black src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py src/pyrecest/filters/joint_probabilistic_data_association_filter.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py - - name: Exact and cheap JPDAF regression tests - run: | - python -m pytest -q tests/filters/test_cheap_joint_probabilistic_data_association_filter.py tests/filters/test_joint_probabilistic_data_association_filter.py tests/filters/test_jpdaf_logdet_underflow.py - - name: Validate API documentation and imports - run: | - python scripts/render_backend_api_matrix.py --check docs/backend-api-matrix.md - python scripts/check_public_api_registry.py --check docs/public-api-registry.md - python scripts/check_minimal_imports.py - python -m pylint src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py src/pyrecest/filters/joint_probabilistic_data_association_filter.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py - python - <<'PY' - import re - from pathlib import Path - example = re.search(r'```python\n(.*?)```', Path('docs/cheap-jpdaf.md').read_text(), re.S) - exec(compile(example.group(1), 'docs/cheap-jpdaf.md', 'exec')) - print('Cheap JPDAF documentation example passed') - PY - - name: Commit validated integration and remove bootstrap files - run: | - git config user.name 'github-actions[bot]' - git config user.email '41898282+github-actions[bot]@users.noreply.github.com' - git rm .github/workflows/implement-cheap-jpdaf.yml scripts/_integrate_cheap_jpdaf.py - git add src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py src/pyrecest/filters/joint_probabilistic_data_association_filter.py src/pyrecest/filters/__init__.py src/pyrecest/_backend/capabilities.py src/pyrecest/api_registry.py tests/filters/test_cheap_joint_probabilistic_data_association_filter.py docs/cheap-jpdaf.md docs/api-overview.md docs/backend-compatibility.md docs/backend-api-matrix.md docs/public-api-registry.md - git diff --cached --stat - git commit -m 'Integrate CheapJPDAF with shared Gaussian updates and API metadata' - git push origin "HEAD:$BRANCH_NAME" diff --git a/docs/api-overview.md b/docs/api-overview.md index fc61bd94c5..87931237bc 100644 --- a/docs/api-overview.md +++ b/docs/api-overview.md @@ -205,3 +205,11 @@ Common starting points include: dedicated tutorials. - Use module docstrings and class docstrings for detailed mathematical notes where they are available. + +## Cheap joint probabilistic data association + +`CheapJPDAF` / `CJPDAF` provides normalized soft association for +linear-Gaussian tracking without joint-event enumeration. It reuses +the exact `JPDAF` Gaussian update and is explicitly NumPy-only. +See [Cheap JPDAF](cheap-jpdaf.md) for usage, normalization, +complexity, and the non-MAP greedy diagnostic. diff --git a/docs/backend-api-matrix.md b/docs/backend-api-matrix.md index 2cdf0f9277..673a2dadbe 100644 --- a/docs/backend-api-matrix.md +++ b/docs/backend-api-matrix.md @@ -31,21 +31,24 @@ in CI so the user-facing matrix cannot silently drift from the executable metada ## Public API Rows -| API | NumPy | PyTorch | JAX | Notes | -|--------------------------------|-----------|-------------|-------------|----------------------------------------------------------------------------------------------------------------------------------| -| `BackendFacade` | supported | partial | partial | Facade names are importable across backends, but some functions are bridged or explicitly unsupported. | -| `DiscreteStateUtilities` | supported | bridged | bridged | Finite-state HMM and IMM utilities operate on NumPy arrays and SciPy sparse matrices; non-NumPy inputs are coerced. | -| `DistributionConversion` | supported | partial | partial | Euclidean particle/Gaussian conversions are portable; grid, Fourier, and manifold routes are route-specific. | -| `EuclideanParticleFilter` | supported | partial | partial | Particle operations are portable where sampling and resampling helpers preserve backend semantics. | -| `EvaluationUtilities` | supported | bridged | bridged | Some plotting, assignment, and summary operations remain NumPy/SciPy oriented and may not preserve device or gradient semantics. | -| `GaussianDistribution` | supported | supported | supported | Basic construction, moment access, and portable operations should remain backend portable. | -| `KalmanFilter` | supported | supported | supported | Linear Gaussian operations are part of the portable baseline. | -| `LinearDiracDistribution` | supported | supported | supported | Used by representation conversion and particle-style workflows. | -| `MultiBernoulliTracker` | supported | partial | unsupported | Tracking workflows rely on assignment and measurement-set utilities that are currently NumPy-oriented. | -| `PointSetRegistration` | supported | partial | unsupported | Registration utilities may copy through NumPy/SciPy and should not be assumed differentiable. | -| `SphericalHarmonicsEOTTracker` | supported | unsupported | unsupported | Depends on spherical harmonics and SciPy-adjacent functionality. | -| `UKFOnManifolds` | supported | partial | unsupported | The current implementation documents explicit JAX exclusions for predict/update. | -| `UnscentedKalmanFilter` | supported | partial | partial | Portable for backend-compatible model functions; advanced paths may still bridge through NumPy/SciPy. | +| API | NumPy | PyTorch | JAX | Notes | +|------------------------------------------------|-----------|-------------|-------------|----------------------------------------------------------------------------------------------------------------------------------| +| `BackendFacade` | supported | partial | partial | Facade names are importable across backends, but some functions are bridged or explicitly unsupported. | +| `CJPDAF` | supported | unsupported | unsupported | Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration. | +| `CheapJPDAF` | supported | unsupported | unsupported | Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration. | +| `CheapJointProbabilisticDataAssociationFilter` | supported | unsupported | unsupported | Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration. | +| `DiscreteStateUtilities` | supported | bridged | bridged | Finite-state HMM and IMM utilities operate on NumPy arrays and SciPy sparse matrices; non-NumPy inputs are coerced. | +| `DistributionConversion` | supported | partial | partial | Euclidean particle/Gaussian conversions are portable; grid, Fourier, and manifold routes are route-specific. | +| `EuclideanParticleFilter` | supported | partial | partial | Particle operations are portable where sampling and resampling helpers preserve backend semantics. | +| `EvaluationUtilities` | supported | bridged | bridged | Some plotting, assignment, and summary operations remain NumPy/SciPy oriented and may not preserve device or gradient semantics. | +| `GaussianDistribution` | supported | supported | supported | Basic construction, moment access, and portable operations should remain backend portable. | +| `KalmanFilter` | supported | supported | supported | Linear Gaussian operations are part of the portable baseline. | +| `LinearDiracDistribution` | supported | supported | supported | Used by representation conversion and particle-style workflows. | +| `MultiBernoulliTracker` | supported | partial | unsupported | Tracking workflows rely on assignment and measurement-set utilities that are currently NumPy-oriented. | +| `PointSetRegistration` | supported | partial | unsupported | Registration utilities may copy through NumPy/SciPy and should not be assumed differentiable. | +| `SphericalHarmonicsEOTTracker` | supported | unsupported | unsupported | Depends on spherical harmonics and SciPy-adjacent functionality. | +| `UKFOnManifolds` | supported | partial | unsupported | The current implementation documents explicit JAX exclusions for predict/update. | +| `UnscentedKalmanFilter` | supported | partial | partial | Portable for backend-compatible model functions; advanced paths may still bridge through NumPy/SciPy. | When adding a new public API, add a row to the matrix, update docs if the row is diff --git a/docs/backend-compatibility.md b/docs/backend-compatibility.md index b71b2aae69..145dcfa17f 100644 --- a/docs/backend-compatibility.md +++ b/docs/backend-compatibility.md @@ -175,3 +175,10 @@ When adding or changing an API with backend-specific behavior: - mention the restriction in the relevant tutorial, example, or API notes; - prefer implementing missing backend facade functions over direct imports when the operation should be portable. + +### Cheap JPDAF + +`CheapJPDAF`, `CJPDAF`, and +`CheapJointProbabilisticDataAssociationFilter` are NumPy-only. +Association and measurement updates reject other backends explicitly. +See [Cheap JPDAF](cheap-jpdaf.md) for the approximation contract. diff --git a/docs/public-api-registry.md b/docs/public-api-registry.md index b29af69add..3dcf231947 100644 --- a/docs/public-api-registry.md +++ b/docs/public-api-registry.md @@ -34,19 +34,22 @@ mapped to the same implementation module as their canonical form. New examples and documentation should use the canonical spelling. -| API | Module | Category | Backend contract | Notes | -|--------------------------------|-------------------------------------|------------------|--------------------------------|----------------------------------------------------------------------------------------------------------------------| -| `BackendFacade` | `pyrecest.backend` | backend-specific | `BackendFacade` | Facade names are importable across backends, with bridged or unsupported functions documented in the backend matrix. | -| `DiscreteStateUtilities` | `pyrecest.filters` | backend-specific | `DiscreteStateUtilities` | Finite-state HMM and IMM utilities operate on NumPy arrays and SciPy sparse matrices. | -| `DistributionConversion` | `pyrecest.distributions.conversion` | backend-specific | `DistributionConversion` | Euclidean Gaussian/particle routes are portable; grid, Fourier, and manifold routes are route-specific. | -| `EuclideanParticleFilter` | `pyrecest.filters` | backend-specific | `EuclideanParticleFilter` | Particle behavior depends on sampler and resampling support in the active backend. | -| `EvaluationUtilities` | `pyrecest.evaluation` | backend-specific | `EvaluationUtilities` | Plotting, assignment, summaries, and result helpers are only partly backend-portable. | -| `GaussianDistribution` | `pyrecest.distributions` | stable | `GaussianDistribution` | Basic construction, moment access, and portable operations are part of the core distribution API. | -| `KalmanFilter` | `pyrecest.filters` | stable | `KalmanFilter` | Linear Gaussian filtering is part of the portable baseline. | -| `LinearDiracDistribution` | `pyrecest.distributions` | stable | `LinearDiracDistribution` | Core particle-style representation used by conversion and filtering workflows. | -| `MultiBernoulliTracker` | `pyrecest.filters` | backend-specific | `MultiBernoulliTracker` | Tracking workflows rely on assignment and measurement-set utilities with NumPy-oriented paths. | -| `PointSetRegistration` | `pyrecest.utils` | backend-specific | `PointSetRegistration` | Registration helpers may bridge through NumPy/SciPy and are not guaranteed differentiable. | -| `SphericalHarmonicsEOTTracker` | `pyrecest.filters` | backend-specific | `SphericalHarmonicsEOTTracker` | Depends on spherical-harmonics and SciPy-adjacent functionality. | -| `UKFOnManifolds` | `pyrecest.filters` | backend-specific | `UKFOnManifolds` | Current predict/update paths explicitly exclude JAX. | -| `UnscentedKalmanFilter` | `pyrecest.filters` | backend-specific | `UnscentedKalmanFilter` | Portable for backend-compatible model functions; advanced paths may bridge through NumPy/SciPy. | +| API | Module | Category | Backend contract | Notes | +|------------------------------------------------|-------------------------------------|------------------|------------------------------------------------|----------------------------------------------------------------------------------------------------------------------| +| `BackendFacade` | `pyrecest.backend` | backend-specific | `BackendFacade` | Facade names are importable across backends, with bridged or unsupported functions documented in the backend matrix. | +| `CJPDAF` | `pyrecest.filters` | experimental | `CJPDAF` | Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only. | +| `CheapJPDAF` | `pyrecest.filters` | experimental | `CheapJPDAF` | Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only. | +| `CheapJointProbabilisticDataAssociationFilter` | `pyrecest.filters` | experimental | `CheapJointProbabilisticDataAssociationFilter` | Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only. | +| `DiscreteStateUtilities` | `pyrecest.filters` | backend-specific | `DiscreteStateUtilities` | Finite-state HMM and IMM utilities operate on NumPy arrays and SciPy sparse matrices. | +| `DistributionConversion` | `pyrecest.distributions.conversion` | backend-specific | `DistributionConversion` | Euclidean Gaussian/particle routes are portable; grid, Fourier, and manifold routes are route-specific. | +| `EuclideanParticleFilter` | `pyrecest.filters` | backend-specific | `EuclideanParticleFilter` | Particle behavior depends on sampler and resampling support in the active backend. | +| `EvaluationUtilities` | `pyrecest.evaluation` | backend-specific | `EvaluationUtilities` | Plotting, assignment, summaries, and result helpers are only partly backend-portable. | +| `GaussianDistribution` | `pyrecest.distributions` | stable | `GaussianDistribution` | Basic construction, moment access, and portable operations are part of the core distribution API. | +| `KalmanFilter` | `pyrecest.filters` | stable | `KalmanFilter` | Linear Gaussian filtering is part of the portable baseline. | +| `LinearDiracDistribution` | `pyrecest.distributions` | stable | `LinearDiracDistribution` | Core particle-style representation used by conversion and filtering workflows. | +| `MultiBernoulliTracker` | `pyrecest.filters` | backend-specific | `MultiBernoulliTracker` | Tracking workflows rely on assignment and measurement-set utilities with NumPy-oriented paths. | +| `PointSetRegistration` | `pyrecest.utils` | backend-specific | `PointSetRegistration` | Registration helpers may bridge through NumPy/SciPy and are not guaranteed differentiable. | +| `SphericalHarmonicsEOTTracker` | `pyrecest.filters` | backend-specific | `SphericalHarmonicsEOTTracker` | Depends on spherical-harmonics and SciPy-adjacent functionality. | +| `UKFOnManifolds` | `pyrecest.filters` | backend-specific | `UKFOnManifolds` | Current predict/update paths explicitly exclude JAX. | +| `UnscentedKalmanFilter` | `pyrecest.filters` | backend-specific | `UnscentedKalmanFilter` | Portable for backend-compatible model functions; advanced paths may bridge through NumPy/SciPy. | diff --git a/scripts/_integrate_cheap_jpdaf.py b/scripts/_integrate_cheap_jpdaf.py deleted file mode 100644 index b0e9174913..0000000000 --- a/scripts/_integrate_cheap_jpdaf.py +++ /dev/null @@ -1,102 +0,0 @@ -from pathlib import Path -import subprocess -import sys - - -def replace_once(path, old, new): - path = Path(path) - text = path.read_text() - if text.count(old) != 1: - raise RuntimeError(f'Expected one integration anchor in {path}: {old!r}') - path.write_text(text.replace(old, new, 1)) - - -renderers = { - 'docs/backend-api-matrix.md': 'scripts/render_backend_api_matrix.py', - 'docs/public-api-registry.md': 'scripts/check_public_api_registry.py', -} -old_tables = {doc: subprocess.check_output([sys.executable, script], text=True) - for doc, script in renderers.items()} - -path = Path('src/pyrecest/filters/joint_probabilistic_data_association_filter.py') -text = path.read_text() -start = text.index(' track_order = sorted(') -end = text.index(' self.latest_association_probabilities = association_probabilities', start) -solver = text[start:end] -replacement = ''' association_probabilities, map_association = self._compute_association_probabilities( - log_likelihoods, - eligible_measurements, - detection_probability, - clutter_intensity, - ) - -''' -text = text[:start] + replacement + text[end:] -method = ''' def _compute_association_probabilities( - self, - log_likelihoods, - eligible_measurements, - detection_probability, - clutter_intensity, - ): - """Solve the gated association problem by exact joint-event enumeration.""" - n_targets, n_meas = log_likelihoods.shape -''' -method += solver + ' return association_probabilities, map_association\n\n' -anchor = ' def find_association(\n' -assert text.count(anchor) == 1 -text = text.replace(anchor, method + anchor, 1) -path.write_text(text) - -names = ('CheapJointProbabilisticDataAssociationFilter', 'CheapJPDAF', 'CJPDAF') -exports = ''.join(f' "{name}": ".cheap_joint_probabilistic_data_association_filter",\n' for name in names) -anchor = ' "JPDAF": ".joint_probabilistic_data_association_filter",\n' -replace_once('src/pyrecest/filters/__init__.py', anchor, anchor + exports) - -capabilities = ''.join( - f''' "{name}": {{ - "numpy": "supported", - "pytorch": "unsupported", - "jax": "unsupported", - "notes": "Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration.", - }}, -''' for name in names) -replace_once('src/pyrecest/_backend/capabilities.py', 'API_BACKEND_CAPABILITIES: Final = {\n', - 'API_BACKEND_CAPABILITIES: Final = {\n' + capabilities) -registry = ''.join( - f''' "{name}": {{ - "module": "pyrecest.filters", - "category": "experimental", - "backend_contract": "{name}", - "notes": "Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only.", - }}, -''' for name in names) -replace_once('src/pyrecest/api_registry.py', 'PUBLIC_API_REGISTRY: Final = {\n', - 'PUBLIC_API_REGISTRY: Final = {\n' + registry) -for doc, script in renderers.items(): - new_table = subprocess.check_output([sys.executable, script], text=True) - replace_once(doc, old_tables[doc], new_table) - -with Path('docs/api-overview.md').open('a') as stream: - stream.write('\n## Cheap joint probabilistic data association\n\n' - '`CheapJPDAF` / `CJPDAF` provides normalized soft association for\n' - 'linear-Gaussian tracking without joint-event enumeration. It reuses\n' - 'the exact `JPDAF` Gaussian update and is explicitly NumPy-only.\n' - 'See [Cheap JPDAF](cheap-jpdaf.md) for usage, normalization,\n' - 'complexity, and the non-MAP greedy diagnostic.\n') -with Path('docs/backend-compatibility.md').open('a') as stream: - stream.write('\n### Cheap JPDAF\n\n' - '`CheapJPDAF`, `CJPDAF`, and\n' - '`CheapJointProbabilisticDataAssociationFilter` are NumPy-only.\n' - 'Association and measurement updates reject other backends explicitly.\n' - 'See [Cheap JPDAF](cheap-jpdaf.md) for the approximation contract.\n') -path = Path('tests/filters/test_cheap_joint_probabilistic_data_association_filter.py') -text = path.read_text() -path.write_text('# pylint: disable=protected-access,no-name-in-module,no-member\n' + text) -replace_once( - path, - ' tracker.filter_state = []\n', - ' # Isolate association caches from the inherited empty-bank history logger.\n' - ' tracker.log_prior_estimates = False\n' - ' tracker.filter_state = []\n', -) diff --git a/src/pyrecest/_backend/capabilities.py b/src/pyrecest/_backend/capabilities.py index 319f4db869..7dc3ba7e49 100644 --- a/src/pyrecest/_backend/capabilities.py +++ b/src/pyrecest/_backend/capabilities.py @@ -67,6 +67,24 @@ } API_BACKEND_CAPABILITIES: Final = { + "CheapJointProbabilisticDataAssociationFilter": { + "numpy": "supported", + "pytorch": "unsupported", + "jax": "unsupported", + "notes": "Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration.", + }, + "CheapJPDAF": { + "numpy": "supported", + "pytorch": "unsupported", + "jax": "unsupported", + "notes": "Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration.", + }, + "CJPDAF": { + "numpy": "supported", + "pytorch": "unsupported", + "jax": "unsupported", + "notes": "Normalized cheap JPDA for linear-Gaussian models; no joint-event enumeration.", + }, "KalmanFilter": { "numpy": "supported", "pytorch": "supported", diff --git a/src/pyrecest/api_registry.py b/src/pyrecest/api_registry.py index 803b56296b..135f6c7f49 100644 --- a/src/pyrecest/api_registry.py +++ b/src/pyrecest/api_registry.py @@ -12,6 +12,24 @@ ) PUBLIC_API_REGISTRY: Final = { + "CheapJointProbabilisticDataAssociationFilter": { + "module": "pyrecest.filters", + "category": "experimental", + "backend_contract": "CheapJointProbabilisticDataAssociationFilter", + "notes": "Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only.", + }, + "CheapJPDAF": { + "module": "pyrecest.filters", + "category": "experimental", + "backend_contract": "CheapJPDAF", + "notes": "Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only.", + }, + "CJPDAF": { + "module": "pyrecest.filters", + "category": "experimental", + "backend_contract": "CJPDAF", + "notes": "Fitzgerald-style cheap JPDA with complementary missed-detection mass; NumPy only.", + }, "BackendFacade": { "module": "pyrecest.backend", "category": "backend-specific", diff --git a/src/pyrecest/filters/__init__.py b/src/pyrecest/filters/__init__.py index ccd2dd74e9..ca330f50d8 100644 --- a/src/pyrecest/filters/__init__.py +++ b/src/pyrecest/filters/__init__.py @@ -99,6 +99,9 @@ "GGIWTracker": ".ggiw_tracker", "GlobalNearestNeighbor": ".global_nearest_neighbor", "JPDAF": ".joint_probabilistic_data_association_filter", + "CheapJointProbabilisticDataAssociationFilter": ".cheap_joint_probabilistic_data_association_filter", + "CheapJPDAF": ".cheap_joint_probabilistic_data_association_filter", + "CJPDAF": ".cheap_joint_probabilistic_data_association_filter", "JointProbabilisticDataAssociationFilter": ".joint_probabilistic_data_association_filter", "IMM": ".interacting_multiple_model_filter", "InteractingMultipleModelFilter": ".interacting_multiple_model_filter", diff --git a/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py b/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py index 70adea6173..ea182fd1d7 100644 --- a/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py +++ b/src/pyrecest/filters/cheap_joint_probabilistic_data_association_filter.py @@ -143,7 +143,9 @@ def _compute_association_probabilities( """Replace the exact event solver, leaving Gaussian updates unchanged.""" del eligible_measurements # The gated log-likelihood matrix encodes these. if np.any(np.isnan(log_likelihoods)) or np.any(np.isposinf(log_likelihoods)): - raise ValueError("Gated log likelihoods must be finite or negative infinity.") + raise ValueError( + "Gated log likelihoods must be finite or negative infinity." + ) log_weights = ( log_likelihoods + log(detection_probability) diff --git a/src/pyrecest/filters/joint_probabilistic_data_association_filter.py b/src/pyrecest/filters/joint_probabilistic_data_association_filter.py index a4426127e9..8f8da59ba4 100644 --- a/src/pyrecest/filters/joint_probabilistic_data_association_filter.py +++ b/src/pyrecest/filters/joint_probabilistic_data_association_filter.py @@ -269,6 +269,30 @@ def find_association_probabilities( "JPDAF: No measurement was within the gating threshold for at least one target." ) + association_probabilities, map_association = ( + self._compute_association_probabilities( + log_likelihoods, + eligible_measurements, + detection_probability, + clutter_intensity, + ) + ) + + self.latest_association_probabilities = association_probabilities + self.latest_map_association = map_association + self._latest_posterior_hypotheses = posterior_hypotheses + + return association_probabilities, map_association + + def _compute_association_probabilities( + self, + log_likelihoods, + eligible_measurements, + detection_probability, + clutter_intensity, + ): + """Solve the gated association problem by exact joint-event enumeration.""" + n_targets, n_meas = log_likelihoods.shape track_order = sorted( range(n_targets), key=lambda idx: len(eligible_measurements[idx]) ) @@ -333,10 +357,6 @@ def recurse(order_index, curr_log_weight): map_association = event_assignments[int(argmax(normalized_event_weights))] - self.latest_association_probabilities = association_probabilities - self.latest_map_association = map_association - self._latest_posterior_hypotheses = posterior_hypotheses - return association_probabilities, map_association def find_association( diff --git a/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py b/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py index e358e306ea..4482178054 100644 --- a/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py +++ b/tests/filters/test_cheap_joint_probabilistic_data_association_filter.py @@ -1,3 +1,4 @@ +# pylint: disable=protected-access,no-name-in-module,no-member """Cheap JPDA reference cases, normalization, and Gaussian-update contracts.""" import numpy as np @@ -93,8 +94,10 @@ def test_randomized_formula_rows_and_measurement_exclusivity(): weights[rng.random(weights.shape) < 0.5] = 0.0 beta = _marginals(weights) denominator = ( - 1 + weights.sum(axis=1, keepdims=True) - + weights.sum(axis=0, keepdims=True) - weights + 1 + + weights.sum(axis=1, keepdims=True) + + weights.sum(axis=0, keepdims=True) + - weights ) npt.assert_allclose(beta[:, 1:], weights / denominator, atol=1e-14) npt.assert_allclose(beta.sum(axis=1), 1.0, atol=1e-14) @@ -144,7 +147,9 @@ def test_greedy_diagnostic_is_feasible_but_not_claimed_to_be_map(): npt.assert_array_equal(tracker.latest_greedy_association, diagnostic) npt.assert_allclose(tracker.latest_association_probabilities, beta) assert tracker.latest_map_association is None - npt.assert_array_equal(tracker.find_association(np.zeros((2, 1)), np.eye(2), np.eye(2)), diagnostic) + npt.assert_array_equal( + tracker.find_association(np.zeros((2, 1)), np.eye(2), np.eye(2)), diagnostic + ) def test_exact_event_limit_is_not_used(monkeypatch): @@ -178,6 +183,8 @@ def test_no_measurements_leave_state_and_covariance_unchanged(): def test_empty_bank_clears_cached_diagnostics(): tracker = _tracker() tracker.find_association_probabilities(np.zeros((2, 1)), np.eye(2), np.eye(2)) + # Isolate association caches from the inherited empty-bank history logger. + tracker.log_prior_estimates = False tracker.filter_state = [] with pytest.warns(UserWarning, match="zero targets"): beta, diagnostic = tracker.find_association_probabilities( @@ -204,12 +211,16 @@ def test_all_measurements_outside_gate_leave_priors_unchanged(): def test_heterogeneous_covariances_and_vector_clutter_match_exact_single_track(): parameters = {"clutter_intensity": np.array([0.01, 0.04])} - cheap, exact = (_tracker(cls, means=(0.0,), **parameters) for cls in (CheapJPDAF, JPDAF)) + cheap, exact = ( + _tracker(cls, means=(0.0,), **parameters) for cls in (CheapJPDAF, JPDAF) + ) measurements = np.array([[-0.5, 1.0], [0.1, -0.2]]) covariances = np.stack([0.2 * np.eye(2), 2 * np.eye(2)], axis=2) for tracker in (cheap, exact): tracker.update_linear(measurements, np.eye(2), covariances) - npt.assert_allclose(cheap.latest_association_probabilities, exact.latest_association_probabilities) + npt.assert_allclose( + cheap.latest_association_probabilities, exact.latest_association_probabilities + ) npt.assert_allclose(cheap.get_point_estimate(), exact.get_point_estimate()) npt.assert_allclose(cheap.filter_state[0].C, exact.filter_state[0].C) @@ -271,4 +282,6 @@ def test_invalid_covariance_shape_and_pairwise_costs_rejected(): with pytest.raises(ValueError, match="cov_mats_meas must have shape"): tracker.find_association_probabilities(measurements, np.eye(2), np.ones((1, 1))) with pytest.raises(NotImplementedError, match="pairwise_cost_matrix"): - tracker.update_linear(measurements, np.eye(2), np.eye(2), pairwise_cost_matrix=np.zeros((2, 1))) + tracker.update_linear( + measurements, np.eye(2), np.eye(2), pairwise_cost_matrix=np.zeros((2, 1)) + )