Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 4 additions & 10 deletions monai/metrics/meaniou.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
38 changes: 38 additions & 0 deletions tests/metrics/test_ignore_index_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,8 @@
MeanIoU,
SurfaceDiceMetric,
SurfaceDistanceMetric,
compute_dice,
compute_iou,
)
from monai.utils import optional_import

Expand Down Expand Up @@ -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):
Expand Down
Loading