From 78f1b98a6e955eae2a0b0ed8701dfbc9f08b5137 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Thu, 6 Aug 2026 14:58:01 +0800 Subject: [PATCH 1/2] Fix direct MTT mapping mean extraction --- src/pyrecest/evaluation/get_extract_mean.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/pyrecest/evaluation/get_extract_mean.py b/src/pyrecest/evaluation/get_extract_mean.py index 381dda818d..63e7497b42 100644 --- a/src/pyrecest/evaluation/get_extract_mean.py +++ b/src/pyrecest/evaluation/get_extract_mean.py @@ -95,7 +95,7 @@ def _extract_mtt_mean(filter_state): return _extract_track_collection_mean(filter_state.tracks) if hasattr(filter_state, "single_target_filters"): return _extract_track_collection_mean(filter_state.single_target_filters) - if isinstance(filter_state, (list, tuple)): + if isinstance(filter_state, (Mapping, list, tuple)): return _extract_track_collection_mean(filter_state) return _point_estimate_or_mean(filter_state) From a9322eb2b378817e97979af2dba6a38c10f52cf8 Mon Sep 17 00:00:00 2001 From: Florian Pfaff <6773539+FlorianPfaff@users.noreply.github.com> Date: Thu, 6 Aug 2026 14:58:20 +0800 Subject: [PATCH 2/2] Test direct MTT mapping state extraction --- .../test_get_extract_mean_mapping_tracks.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/tests/evaluation/test_get_extract_mean_mapping_tracks.py b/tests/evaluation/test_get_extract_mean_mapping_tracks.py index c684916ce2..950143affb 100644 --- a/tests/evaluation/test_get_extract_mean_mapping_tracks.py +++ b/tests/evaluation/test_get_extract_mean_mapping_tracks.py @@ -41,3 +41,18 @@ def test_mtt_mean_extracts_mapping_values_from_get_tracks(): assert len(means) == 2 np.testing.assert_array_equal(means[0], [5.0, 6.0]) np.testing.assert_array_equal(means[1], [7.0, 8.0]) + + +def test_mtt_mean_extracts_mapping_values_from_plain_mapping(): + extract_mean = get_extract_mean("euclidean", mtt_scenario=True) + tracker_state = { + "track-a": Track([9.0, 10.0]), + "track-b": Track([11.0, 12.0]), + } + + means = extract_mean(tracker_state) + + assert isinstance(means, list) + assert len(means) == 2 + np.testing.assert_array_equal(means[0], [9.0, 10.0]) + np.testing.assert_array_equal(means[1], [11.0, 12.0])