fix: keep it DRY

This commit is contained in:
Balthasar 2025-07-18 09:36:59 +02:00
parent 001a4254f8
commit afdbb0a75c
2 changed files with 34 additions and 32 deletions

View File

@ -918,6 +918,27 @@ class COCOEvaluator:
np.array(score_at_recall)
)
self.results = {
"params": self.params,
"counts": [num_iou_thresholds, num_recall_thresholds, num_categories, num_area_ranges, num_max_detections],
"date": datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
"precision": precision,
"recall": recall,
"scores": scores,
}
# Helper function to compute average precision while handling -1 sentinel values
def compute_average_precision(precision_slice):
"""Helper function to compute average precision while handling -1 sentinel values."""
masked = np.ma.masked_equal(precision_slice, -1)
if masked.count() == 0:
# All values are -1 (no data)
return np.full(num_iou_thresholds, -1), np.full((num_categories, num_iou_thresholds), -1)
else:
mAP_scores = np.ma.filled(masked.mean(axis=(1, 2)), -1)
ap_per_class = np.ma.filled(masked.mean(axis=1), -1).transpose(1, 0)
return mAP_scores, ap_per_class
# Average precision over all sizes, 100 max detections
area_range_idx = list(ObjectSize).index(ObjectSize.ALL)
max_100_dets_idx = self.params.max_dets.index(100)
@ -927,54 +948,28 @@ class COCOEvaluator:
]
# mAP over thresholds (dimension=num_thresholds)
# Use masked array to exclude -1 values when computing mean
masked = np.ma.masked_equal(average_precision_all_sizes, -1)
# Check if all values are masked (empty array)
if masked.count() == 0:
mAP_scores_all_sizes = np.full(num_iou_thresholds, -1)
ap_per_class_all_sizes = np.full((num_categories, num_iou_thresholds), -1)
else:
mAP_scores_all_sizes = np.ma.filled(masked.mean(axis=(1, 2)), -1)
# AP per class
ap_per_class_all_sizes = np.ma.filled(masked.mean(axis=1), -1).transpose(1, 0)
mAP_scores_all_sizes, ap_per_class_all_sizes = compute_average_precision(average_precision_all_sizes)
# Average precision for SMALL objects and 100 max detections
small_area_range_idx = list(ObjectSize).index(ObjectSize.SMALL)
average_precision_small = precision[
:, :, :, small_area_range_idx, max_100_dets_idx
]
masked_small = np.ma.masked_equal(average_precision_small, -1)
if masked_small.count() == 0:
mAP_scores_small = np.full(num_iou_thresholds, -1)
ap_per_class_small = np.full((num_categories, num_iou_thresholds), -1)
else:
mAP_scores_small = np.ma.filled(masked_small.mean(axis=(1, 2)), -1)
ap_per_class_small = np.ma.filled(masked_small.mean(axis=1), -1).transpose(1, 0)
mAP_scores_small, ap_per_class_small = compute_average_precision(average_precision_small)
# Average precision for MEDIUM objects and 100 max detections
medium_area_range_idx = list(ObjectSize).index(ObjectSize.MEDIUM)
average_precision_medium = precision[
:, :, :, medium_area_range_idx, max_100_dets_idx
]
masked_medium = np.ma.masked_equal(average_precision_medium, -1)
if masked_medium.count() == 0:
mAP_scores_medium = np.full(num_iou_thresholds, -1)
ap_per_class_medium = np.full((num_categories, num_iou_thresholds), -1)
else:
mAP_scores_medium = np.ma.filled(masked_medium.mean(axis=(1, 2)), -1)
ap_per_class_medium = np.ma.filled(masked_medium.mean(axis=1), -1).transpose(1, 0)
mAP_scores_medium, ap_per_class_medium = compute_average_precision(average_precision_medium)
# Average precision for LARGE objects and 100 max detections
large_area_range_idx = list(ObjectSize).index(ObjectSize.LARGE)
average_precision_large = precision[
:, :, :, large_area_range_idx, max_100_dets_idx
]
masked_large = np.ma.masked_equal(average_precision_large, -1)
if masked_large.count() == 0:
mAP_scores_large = np.full(num_iou_thresholds, -1)
ap_per_class_large = np.full((num_categories, num_iou_thresholds), -1)
else:
mAP_scores_large = np.ma.filled(masked_large.mean(axis=(1, 2)), -1)
ap_per_class_large = np.ma.filled(masked_large.mean(axis=1), -1).transpose(1, 0)
mAP_scores_large, ap_per_class_large = compute_average_precision(average_precision_large)
self.results = {
"params": self.params,

View File

@ -315,5 +315,12 @@ def test_empty_predictions_and_targets():
metric.update([Detections.empty()], [Detections.empty()])
result = metric.compute()
# Should handle empty case gracefully
assert result.map50_95 >= -1.0 # Can be -1 to indicate no data
# Should return -1 for no data (matching pycocotools behavior)
assert result.map50_95 == -1
assert result.map50 == -1
assert result.map75 == -1
# All object size categories should also be -1
assert result.small_objects.map50_95 == -1
assert result.medium_objects.map50_95 == -1
assert result.large_objects.map50_95 == -1