diff --git a/src/pyrecest/evaluation/get_extract_mean.py b/src/pyrecest/evaluation/get_extract_mean.py index 381dda818..63e7497b4 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) diff --git a/tests/evaluation/test_get_extract_mean_mapping_tracks.py b/tests/evaluation/test_get_extract_mean_mapping_tracks.py index c684916ce..950143aff 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])