diff --git a/monai/metrics/panoptic_quality.py b/monai/metrics/panoptic_quality.py index 03cef9d566..386bf6ed87 100644 --- a/monai/metrics/panoptic_quality.py +++ b/monai/metrics/panoptic_quality.py @@ -244,38 +244,44 @@ def _get_id_list(gt: torch.Tensor) -> list[torch.Tensor]: return id_list +def _id_tensor(id_list: list[torch.Tensor], device: torch.device) -> torch.Tensor: + return torch.stack([torch.as_tensor(i, device=device) for i in id_list]).long() + + def _get_pairwise_iou( pred: torch.Tensor, gt: torch.Tensor, device: str | torch.device = "cpu" ) -> tuple[torch.Tensor, list[torch.Tensor], list[torch.Tensor]]: pred_id_list = _get_id_list(pred) true_id_list = _get_id_list(gt) - pairwise_iou = torch.zeros([len(true_id_list) - 1, len(pred_id_list) - 1], dtype=torch.float, device=device) - true_masks: list[torch.Tensor] = [] - pred_masks: list[torch.Tensor] = [] - - for t in true_id_list[1:]: - t_mask = torch.as_tensor(gt == t, device=device).int() - true_masks.append(t_mask) - - for p in pred_id_list[1:]: - p_mask = torch.as_tensor(pred == p, device=device).int() - pred_masks.append(p_mask) - - for true_id in range(1, len(true_id_list)): - t_mask = true_masks[true_id - 1] - pred_true_overlap = pred[t_mask > 0] - pred_true_overlap_id = list(pred_true_overlap.unique()) - for pred_id in pred_true_overlap_id: - if pred_id == 0: - continue - p_mask = pred_masks[pred_id - 1] - total = (t_mask + p_mask).sum() - inter = (t_mask * p_mask).sum() - iou = inter / (total - inter) - pairwise_iou[true_id - 1, pred_id - 1] = iou - - return pairwise_iou, true_id_list, pred_id_list + num_true = len(true_id_list) - 1 + num_pred = len(pred_id_list) - 1 + pairwise_iou = torch.zeros([num_true, num_pred], dtype=torch.float, device=device) + if num_true == 0 or num_pred == 0: + return pairwise_iou, true_id_list, pred_id_list + + gt_flat = gt.reshape(-1).long() + pred_flat = pred.reshape(-1).long().to(gt_flat.device) + # ids are contiguous after `remap_instance_id`; otherwise (`remap=False`) map each id to its + # position in the id list so it indexes the matching row/column + if int(true_id_list[-1]) != num_true: + gt_flat = torch.searchsorted(_id_tensor(true_id_list, gt_flat.device), gt_flat) + if int(pred_id_list[-1]) != num_pred: + pred_flat = torch.searchsorted(_id_tensor(pred_id_list, pred_flat.device), pred_flat) + + # count all pairwise intersections in one bincount + stride = num_pred + 1 + joint = gt_flat * stride + pred_flat + intersection = torch.bincount(joint, minlength=(num_true + 1) * stride).reshape(num_true + 1, stride) + true_area = torch.bincount(gt_flat, minlength=num_true + 1) + pred_area = torch.bincount(pred_flat, minlength=num_pred + 1) + + # keep counts integral through the union, as the per-pair loop did + inter = intersection[1:, 1:] # drop background row/column + union = true_area[1:, None] + pred_area[None, 1:] - inter + pairwise_iou = torch.where(inter > 0, inter / union.clamp_min(1), pairwise_iou.to(inter.device)) + + return pairwise_iou.to(device), true_id_list, pred_id_list def _get_paired_iou( diff --git a/tests/metrics/test_compute_panoptic_quality.py b/tests/metrics/test_compute_panoptic_quality.py index 3f8c21debb..c54a6ed0a7 100644 --- a/tests/metrics/test_compute_panoptic_quality.py +++ b/tests/metrics/test_compute_panoptic_quality.py @@ -18,7 +18,7 @@ from parameterized import parameterized from monai.metrics import PanopticQualityMetric, compute_panoptic_quality -from monai.metrics.panoptic_quality import compute_mean_iou +from monai.metrics.panoptic_quality import _get_pairwise_iou, compute_mean_iou from tests.test_utils import SkipIfNoModule _device = "cuda:0" if torch.cuda.is_available() else "cpu" @@ -215,6 +215,30 @@ def test_invalid_3d_shape(self): with self.assertRaises(ValueError): metric(invalid_pred, invalid_gt) + def test_pairwise_iou(self): + """`_get_pairwise_iou` returns the expected IoU matrix, including no-overlap and empty cases.""" + gt = torch.as_tensor([[1, 1, 0], [0, 2, 2], [0, 0, 0]], device=_device) + pred = torch.as_tensor([[1, 0, 0], [0, 2, 2], [0, 0, 0]], device=_device) + pairwise, _, _ = _get_pairwise_iou(pred, gt, device=_device) + np.testing.assert_allclose(pairwise.cpu().numpy(), [[0.5, 0.0], [0.0, 1.0]], atol=1e-6) + + # true instance with no overlapping prediction -> all-zero row, shape preserved + disjoint_pred = torch.as_tensor([[0, 0, 1], [0, 0, 0], [0, 0, 0]], device=_device) + pairwise, _, _ = _get_pairwise_iou(disjoint_pred, gt, device=_device) + self.assertEqual(pairwise.shape, torch.Size([2, 1])) + self.assertTrue(torch.all(pairwise == 0)) + + # empty gt / empty pred -> degenerate matrices, no error + empty = torch.zeros((3, 3), dtype=torch.int, device=_device) + self.assertEqual(_get_pairwise_iou(pred, empty, device=_device)[0].shape, torch.Size([0, 2])) + self.assertEqual(_get_pairwise_iou(empty, gt, device=_device)[0].shape, torch.Size([2, 0])) + + # non-contiguous ids (`remap=False`) index rows/columns by position in the id list + sparse_gt = torch.as_tensor([[3, 3, 0], [0, 7, 7], [0, 0, 0]], device=_device) + sparse_pred = torch.as_tensor([[5, 0, 0], [0, 9, 9], [0, 0, 0]], device=_device) + pairwise, _, _ = _get_pairwise_iou(sparse_pred, sparse_gt, device=_device) + np.testing.assert_allclose(pairwise.cpu().numpy(), [[0.5, 0.0], [0.0, 1.0]], atol=1e-6) + def test_compute_mean_iou_invalid_shape(self): """Test that compute_mean_iou raises ValueError for invalid shapes.""" from monai.metrics.panoptic_quality import compute_mean_iou