diff --git a/monai/metrics/meaniou.py b/monai/metrics/meaniou.py index 0ff4f131f8..f9be38b38a 100644 --- a/monai/metrics/meaniou.py +++ b/monai/metrics/meaniou.py @@ -152,16 +152,10 @@ def compute_iou( if y.shape != y_pred.shape: raise ValueError(f"y_pred and y should have same shapes, got {y_pred.shape} and {y.shape}.") - if ignore_index is not None and 0 <= ignore_index < (y_pred.shape[1] + (0 if include_background else 1)): - ignore_channel = ignore_index if include_background else ignore_index - 1 - if 0 <= ignore_channel < y_pred.shape[1]: - y_pred = y_pred.clone() - y = y.clone() - y_pred[:, ignore_channel] = 0 - y[:, ignore_channel] = 0 - mask = None - else: - mask = create_ignore_mask(original_y if ignore_index is not None else y, ignore_index) + # Use the same spatial masking as DiceHelper so both metrics exclude the + # same voxels: zeroing the ignored channel alone would leave voxels of the + # ignored class counting as false positives for the other classes + mask = create_ignore_mask(original_y, ignore_index) if mask is not None: if mask.shape != y_pred.shape: mask = mask.expand_as(y_pred) diff --git a/tests/metrics/test_ignore_index_metrics.py b/tests/metrics/test_ignore_index_metrics.py index af5ffccee7..5353ab6618 100644 --- a/tests/metrics/test_ignore_index_metrics.py +++ b/tests/metrics/test_ignore_index_metrics.py @@ -24,6 +24,8 @@ MeanIoU, SurfaceDiceMetric, SurfaceDistanceMetric, + compute_dice, + compute_iou, ) from monai.utils import optional_import @@ -142,6 +144,42 @@ def test_metric_ignore_class_index_without_background(self, metric_class, kwargs torch.testing.assert_close(res1, res2, msg=f"Failed for {metric_class.__name__}") + def test_ignored_voxels_excluded_from_other_classes(self): + """Ignored voxels must be dropped from every class score, not just their own.""" + # 4 voxels, 3 one-hot classes; voxel 1 belongs to the ignored class 1 + y = torch.tensor([[[1.0, 0.0, 0.0, 1.0], [0.0, 1.0, 0.0, 0.0], [0.0, 0.0, 1.0, 0.0]]]) + # a perfect prediction except the ignored voxel is called class 0 + y_pred = torch.tensor([[[1.0, 1.0, 0.0, 1.0], [0.0, 0.0, 0.0, 0.0], [0.0, 0.0, 1.0, 0.0]]]) + + iou = compute_iou(y_pred, y, include_background=True, ignore_index=1) + dice = compute_dice(y_pred, y, include_background=True, ignore_index=1) + + # the mislabelled voxel is ignored, so class 0 is scored as perfect + self.assertEqual(iou[0, 0].item(), 1.0) + torch.testing.assert_close(iou, dice, equal_nan=True) + + def test_ignored_voxels_excluded_with_include_background_false(self): + """The ignore_index mask must line up with the ignore_background channel strip.""" + # 4 one-hot classes: 0=background, 1, 2=ignored, 3 + y = torch.zeros(1, 4, 4) + y[0, 0, 0] = 1 # voxel 0 -> background + y[0, 2, 1] = 1 # voxel 1 -> ignored class + y[0, 1, 2] = 1 # voxel 2 -> class 1 + y[0, 3, 3] = 1 # voxel 3 -> class 3 + + y_pred = y.clone() + # mislabel the ignored voxel as class 1 instead of leaving it unpredicted + y_pred[0, 2, 1] = 0 + y_pred[0, 1, 1] = 1 + + iou = compute_iou(y_pred, y, include_background=False, ignore_index=2) + dice = compute_dice(y_pred, y, include_background=False, ignore_index=2) + + # class 1's false positive at the ignored voxel must be dropped, not just + # its own (now background-stripped) channel + self.assertEqual(iou[0, 0].item(), 1.0) + torch.testing.assert_close(iou, dice, equal_nan=True) + @unittest.skipUnless(has_scipy, "Scipy required for surface metrics") class TestIgnoreIndexSurfaceMetrics(unittest.TestCase):