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
106 changes: 88 additions & 18 deletions monai/metrics/surface_distance.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
import torch

from monai.metrics.utils import (
compute_voronoi_regions_fast,
create_ignore_mask,
do_metric_reduction,
get_edge_surface_distance,
Expand All @@ -40,6 +41,19 @@ class SurfaceDistanceMetric(CumulativeIterationMetric):

Example of the typical execution steps of this metric class follows :py:class:`monai.metrics.metric.Cumulative`.


The ``per_component=True`` approach computes the Surface Distance on a per-connected component basis in the ground
truth segmentation. This ensures that each component contributes equally to the final metric, regardless of its size.
Traditional Surface Distance can be dominated by large structures, but the per-component method gives a more
balanced evaluation, particularly for small or fragmented objects. This provides a granular assessment of segmentation
quality, which is especially important in cases with multiple disconnected foreground components.
Note:
- The input prediction (`y_pred`) and ground truth (`y`) must both have 2 channels (foreground/background),
with binary segmentation (0 for background, 1 for foreground). That is, this assumes the shape of both prediction
and ground truth is B2HW[D].
- This method cannot be used with multiclass segmentation.
For more information, refer to the original paper: https://arxiv.org/abs/2410.18684

Args:
include_background: whether to include distance computation on the first channel of
the predicted output. Defaults to ``False``.
Expand All @@ -56,7 +70,7 @@ class SurfaceDistanceMetric(CumulativeIterationMetric):
Voxels with this label are excluded from the score, which is useful for padding, unlabeled regions,
or boundary artifacts. For federated or aggregated settings, ensure all clients use the same
ignore_index to keep score values comparable.

per_component: whether to compute the Surface Distance on a per-connected component basis. Defaults to ``False``.
"""

def __init__(
Expand All @@ -67,6 +81,7 @@ def __init__(
reduction: MetricReduction | str = MetricReduction.MEAN,
get_not_nans: bool = False,
ignore_index: int | None = None,
per_component: bool = False,
) -> None:
super().__init__()
self.include_background = include_background
Expand All @@ -75,6 +90,7 @@ def __init__(
self.reduction = reduction
self.get_not_nans = get_not_nans
self.ignore_index = ignore_index
self.per_component = per_component

def _compute_tensor(self, y_pred: torch.Tensor, y: torch.Tensor, **kwargs: Any) -> torch.Tensor: # type: ignore[override]
"""
Expand All @@ -100,6 +116,16 @@ def _compute_tensor(self, y_pred: torch.Tensor, y: torch.Tensor, **kwargs: Any)
"""
if y_pred.dim() < 3:
raise ValueError("y_pred should have at least three dimensions.")
if self.per_component:
same_rank = y_pred.ndim == y.ndim and y_pred.ndim in (4, 5)
binary_channels = y_pred.shape[1] == 2 and y.shape[1] == 2
same_shape = y_pred.shape == y.shape
if not (same_rank and binary_channels and same_shape):
raise ValueError(
"per_component requires matching 4D/5D binary tensors "
"(B, 2, H, W) or (B, 2, D, H, W). "
f"Got y_pred={tuple(y_pred.shape)}, y={tuple(y.shape)}."
)

mask = create_ignore_mask(y, self.ignore_index)
if mask is not None:
Expand All @@ -114,6 +140,7 @@ def _compute_tensor(self, y_pred: torch.Tensor, y: torch.Tensor, **kwargs: Any)
symmetric=self.symmetric,
distance_metric=self.distance_metric,
spacing=kwargs.get("spacing"),
per_component=self.per_component,
ignore_index=self.ignore_index,
)

Expand Down Expand Up @@ -146,6 +173,7 @@ def compute_average_surface_distance(
distance_metric: str = "euclidean",
spacing: int | float | np.ndarray | Sequence[int | float | np.ndarray | Sequence[int | float]] | None = None,
ignore_index: int | None = None,
per_component: bool = False,
) -> torch.Tensor:
"""
This function is used to compute the Average Surface Distance from `y_pred` to `y`
Expand Down Expand Up @@ -176,6 +204,7 @@ def compute_average_surface_distance(
ignore_index: optional class index used to suppress empty-mask warnings when that class is ignored.
For federated or aggregated settings, ensure all clients use the same ignore_index to keep score
values comparable.
per_component: whether to compute the Surface Distance on a per-connected component basis. Defaults to ``False``.
"""

if not include_background:
Expand All @@ -198,22 +227,63 @@ def compute_average_surface_distance(
for b, c in np.ndindex(batch_size, n_class):
yp = y_pred[b, c]
yt = y[b, c]

absolute_c = c + class_offset
warn_empty = ignore_index is None or absolute_c != ignore_index
_, distances, _ = get_edge_surface_distance(
yp,
yt,
distance_metric=distance_metric,
spacing=spacing_list[b],
symmetric=symmetric,
class_index=c,
warn_empty=warn_empty,
)

surface_distance = torch.cat(distances)
asd[b, c] = (
torch.tensor(float("nan"), device=asd.device) if surface_distance.numel() == 0 else surface_distance.mean()
)
if per_component:
pred_empty = yp.sum() == 0
label_empty = yt.sum() == 0
if pred_empty or label_empty:
asd[b, c] = 0.0 if (pred_empty and label_empty) else float("nan")
continue
Comment thread
VijayVignesh1 marked this conversation as resolved.
cc_assignment = compute_voronoi_regions_fast(yt.cpu().numpy())
if cc_assignment.device != yp.device:
cc_assignment = cc_assignment.to(yp.device)
component_scores = []
for cc_id in torch.unique(cc_assignment.view(-1)):
cc_mask = cc_assignment == cc_id
coords = torch.nonzero(cc_mask, as_tuple=False)
min_corner_idx = coords.min(dim=0).values
max_corner_idx = coords.max(dim=0).values

slices = tuple(
slice(min_corner_idx[i], max_corner_idx[i] + 1) for i in range(3 if y_pred.ndim == 5 else 2)
)
crop_pred = yp[slices]
crop_label = yt[slices]
cc_crop_mask = cc_mask[slices]

pred_masked = crop_pred * cc_crop_mask
label_masked = crop_label * cc_crop_mask

_, distances, _ = get_edge_surface_distance(
pred_masked,
label_masked,
distance_metric=distance_metric,
spacing=spacing_list[b],
symmetric=symmetric,
class_index=c,
)
surface_distance = torch.cat(distances)
component_scores.append(
torch.tensor(np.nan) if surface_distance.shape == (0,) else surface_distance.mean()
)
asd[b, c] = torch.nanmean(torch.stack(component_scores)) if component_scores else 0.0
else:
absolute_c = c + class_offset
warn_empty = ignore_index is None or absolute_c != ignore_index
_, distances, _ = get_edge_surface_distance(
yp,
yt,
distance_metric=distance_metric,
spacing=spacing_list[b],
symmetric=symmetric,
class_index=c,
warn_empty=warn_empty,
)

surface_distance = torch.cat(distances)
asd[b, c] = (
torch.tensor(float("nan"), device=asd.device)
if surface_distance.numel() == 0
else surface_distance.mean()
)

return convert_data_type(asd, output_type=torch.Tensor, device=y_pred.device, dtype=torch.float)[0]
54 changes: 54 additions & 0 deletions tests/metrics/test_surface_distance.py
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,46 @@ def create_spherical_seg_3d(
]


TEST_CASES_CC_METRICS = []
y = torch.zeros((2, 2, 32, 32, 32), device=_device)
y_pred = torch.zeros((2, 2, 32, 32, 32), device=_device)
TEST_CASES_CC_METRICS.append([[y_pred, y], [[0.0], [0.0]]])

y = torch.zeros((2, 2, 32, 32, 32), device=_device)
y_pred = torch.zeros((2, 2, 32, 32, 32), device=_device)
y_pred[0, 1, 5:10, 5:10, 5:10] = 1
y_pred[0, 0] = 1 - y_pred[0, 1]
TEST_CASES_CC_METRICS.append([[y_pred, y], [[float("nan")], [0.0]]])

y = torch.zeros((2, 2, 32, 32, 32), device=_device)
y_pred = torch.zeros((2, 2, 32, 32, 32), device=_device)
y[0, 1, 10:15, 10:15, 10:15] = 1
y[0, 0] = 1 - y[0, 1]
y_pred[0, 1, 10:15, 10:15, 10:15] = 1
y_pred[0, 0] = 1 - y_pred[0, 1]
TEST_CASES_CC_METRICS.append([[y_pred, y], [[0.0], [0.0]]])

y = torch.zeros((2, 2, 32, 32, 32), device=_device)
y_pred = torch.zeros((2, 2, 32, 32, 32), device=_device)
y[0, 1, 10:15, 10:15, 10:15] = 1
y[0, 1, 20:25, 20:25, 20:25] = 1
y[0, 0] = 1 - y[0, 1]
y_pred[0, 1, 11:16, 10:15, 10:15] = 1
y_pred[0, 1, 11:16, 19:24, 20:25] = 1
y_pred[0, 0] = 1 - y_pred[0, 1]
TEST_CASES_CC_METRICS.append([[y_pred, y], [[3.7987], [0.0]]])

y = torch.zeros((2, 2, 32, 32), device=_device)
y_pred = torch.zeros((2, 2, 32, 32), device=_device)
y[0, 1, 10:15, 10:15] = 1
y[0, 1, 20:25, 20:25] = 1
y[0, 0] = 1 - y[0, 1]
y_pred[0, 1, 10:15, 10:15] = 1
y_pred[0, 1, 21:26, 19:24] = 1
y_pred[0, 0] = 1 - y_pred[0, 1]
TEST_CASES_CC_METRICS.append([[y_pred, y], [[0.4504], [0.0]]])


class TestAllSurfaceMetrics(unittest.TestCase):

@parameterized.expand(TEST_CASES)
Expand Down Expand Up @@ -185,6 +225,20 @@ def test_nans(self, input_data):
np.testing.assert_allclose(0, result, rtol=1e-5)
np.testing.assert_allclose(0, not_nans, rtol=1e-5)

@parameterized.expand(TEST_CASES_CC_METRICS)
def test_cc_metrics(self, input_data, expected_value):
[seg_1, seg_2] = input_data
seg_1 = torch.tensor(seg_1)
seg_2 = torch.tensor(seg_2)
sd_metric = SurfaceDistanceMetric(per_component=True)
sd_metric(seg_1, seg_2)
result = sd_metric.aggregate(reduction="none")
np.testing.assert_allclose(result.cpu().numpy(), expected_value, atol=1e-4)

def test_channel_dimensions(self):
with self.assertRaises(ValueError):
SurfaceDistanceMetric(per_component=True)(torch.ones([3, 3, 144, 144]), torch.ones([3, 3, 144, 144]))


KDTREE_SPACINGS = [["isotropic_default", None], ["isotropic", (1.0, 1.0, 1.0)], ["anisotropic", (1.0, 2.5, 0.5)]]

Expand Down
Loading