From a5044b106fa8fcdaac219e34ec1f825109a91319 Mon Sep 17 00:00:00 2001 From: Richard Brown <33289025+rijobro@users.noreply.github.com> Date: Tue, 13 Sep 2022 10:41:14 +0100 Subject: [PATCH 1/2] keep largest n components Signed-off-by: Richard Brown <33289025+rijobro@users.noreply.github.com> --- monai/transforms/post/array.py | 9 +++-- monai/transforms/post/dictionary.py | 8 ++++- monai/transforms/utils.py | 35 ++++++++++++++----- .../test_keep_largest_connected_component.py | 33 +++++++++++++++++ 4 files changed, 72 insertions(+), 13 deletions(-) diff --git a/monai/transforms/post/array.py b/monai/transforms/post/array.py index e12ccc26c7..fb3e019f63 100644 --- a/monai/transforms/post/array.py +++ b/monai/transforms/post/array.py @@ -273,6 +273,7 @@ def __init__( is_onehot: Optional[bool] = None, independent: bool = True, connectivity: Optional[int] = None, + num_components: int = 1, ) -> None: """ Args: @@ -290,6 +291,7 @@ def __init__( Accepted values are ranging from 1 to input.ndim. If ``None``, a full connectivity of ``input.ndim`` is used. for more details: https://scikit-image.org/docs/dev/api/skimage.measure.html#skimage.measure.label. + num_components: The number of largest components to preserve. """ super().__init__() @@ -297,6 +299,7 @@ def __init__( self.is_onehot = is_onehot self.independent = independent self.connectivity = connectivity + self.num_components = num_components def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor: """ @@ -316,7 +319,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor: if self.independent: for i in applied_labels: foreground = img_[i] > 0 if is_onehot else img_[0] == i - mask = get_largest_connected_component_mask(foreground, self.connectivity) + mask = get_largest_connected_component_mask(foreground, self.connectivity, self.num_components) if is_onehot: img_[i][foreground != mask] = 0 else: @@ -325,12 +328,12 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor: if not is_onehot: # not one-hot, union of labels labels, *_ = convert_to_dst_type(applied_labels, dst=img_, wrap_sequence=True) foreground = (img_[..., None] == labels).any(-1)[0] - mask = get_largest_connected_component_mask(foreground, self.connectivity) + mask = get_largest_connected_component_mask(foreground, self.connectivity, self.num_components) img_[0][foreground != mask] = 0 return convert_to_dst_type(img_, dst=img)[0] # one-hot, union of labels foreground = (img_[applied_labels, ...] == 1).any(0) - mask = get_largest_connected_component_mask(foreground, self.connectivity) + mask = get_largest_connected_component_mask(foreground, self.connectivity, self.num_components) for i in applied_labels: img_[i][foreground != mask] = 0 return convert_to_dst_type(img_, dst=img)[0] diff --git a/monai/transforms/post/dictionary.py b/monai/transforms/post/dictionary.py index decf64bda9..b5439d0c1b 100644 --- a/monai/transforms/post/dictionary.py +++ b/monai/transforms/post/dictionary.py @@ -204,6 +204,7 @@ def __init__( is_onehot: Optional[bool] = None, independent: bool = True, connectivity: Optional[int] = None, + num_components: int = 1, allow_missing_keys: bool = False, ) -> None: """ @@ -224,12 +225,17 @@ def __init__( Accepted values are ranging from 1 to input.ndim. If ``None``, a full connectivity of ``input.ndim`` is used. for more details: https://scikit-image.org/docs/dev/api/skimage.measure.html#skimage.measure.label. + num_components: The number of largest components to preserve. allow_missing_keys: don't raise exception if key is missing. """ super().__init__(keys, allow_missing_keys) self.converter = KeepLargestConnectedComponent( - applied_labels=applied_labels, is_onehot=is_onehot, independent=independent, connectivity=connectivity + applied_labels=applied_labels, + is_onehot=is_onehot, + independent=independent, + connectivity=connectivity, + num_components=num_components, ) def __call__(self, data: Mapping[Hashable, NdarrayOrTensor]) -> Dict[Hashable, NdarrayOrTensor]: diff --git a/monai/transforms/utils.py b/monai/transforms/utils.py index e32ada29ed..43387a07e9 100644 --- a/monai/transforms/utils.py +++ b/monai/transforms/utils.py @@ -59,7 +59,7 @@ from monai.utils.enums import TransformBackends from monai.utils.type_conversion import convert_data_type, convert_to_dst_type, convert_to_tensor -measure, _ = optional_import("skimage.measure", "0.14.2", min_version) +measure, has_measure = optional_import("skimage.measure", "0.14.2", min_version) morphology, has_morphology = optional_import("skimage.morphology") ndimage, _ = optional_import("scipy.ndimage") cp, has_cp = optional_import("cupy") @@ -951,7 +951,9 @@ def generate_spatial_bounding_box( return box_start, box_end -def get_largest_connected_component_mask(img: NdarrayTensor, connectivity: Optional[int] = None) -> NdarrayTensor: +def get_largest_connected_component_mask( + img: NdarrayTensor, connectivity: Optional[int] = None, num_components: int = 1 +) -> NdarrayTensor: """ Gets the largest connected component mask of an image. @@ -961,8 +963,9 @@ def get_largest_connected_component_mask(img: NdarrayTensor, connectivity: Optio Accepted values are ranging from 1 to input.ndim. If ``None``, a full connectivity of ``input.ndim`` is used. for more details: https://scikit-image.org/docs/dev/api/skimage.measure.html#skimage.measure.label. + num_components: The number of largest components to preserve. """ - if isinstance(img, torch.Tensor) and has_cp and has_cucim: + if isinstance(img, torch.Tensor) and has_cp and has_cucim and num_components == 1: x_cupy = monai.transforms.ToCupy()(img.short()) x_label = cucim.skimage.measure.label(x_cupy, connectivity=connectivity) vals, counts = cp.unique(x_label[cp.nonzero(x_label)], return_counts=True) @@ -972,13 +975,27 @@ def get_largest_connected_component_mask(img: NdarrayTensor, connectivity: Optio return out_tensor # type: ignore - img_arr = convert_data_type(img, np.ndarray)[0] - largest_cc: np.ndarray = np.zeros(shape=img_arr.shape, dtype=img_arr.dtype) - img_arr = measure.label(img_arr, connectivity=connectivity) - if img_arr.max() != 0: - largest_cc[...] = img_arr == (np.argmax(np.bincount(img_arr.flat)[1:]) + 1) + if not has_measure: + raise RuntimeError("Skimage.measure required.") - return convert_to_dst_type(largest_cc, dst=img, dtype=largest_cc.dtype)[0] + img_np, *_ = convert_data_type(img, np.ndarray) + # features will be an image, 0 for background and then each different + # feature will have its own index. + features, num_features = measure.label(img_np, connectivity=connectivity, return_num=True) + # # if num features less than max desired, nothing to do. + if num_features <= num_components: + out = img_np.astype(bool) + else: + # get number voxels per feature (bincount). np.argsort[::-1] to get indices + # of largest components. Convert to list for ease + features_to_keep = list(np.argsort(np.bincount(features.flat))[::-1]) + # remove 0 (background) + features_to_keep.remove(0) + # only keep the first n non-background indices + features_to_keep = features_to_keep[:num_components] + # generate labelfield. True if in list of features to keep + out = np.isin(features, features_to_keep) + return convert_to_dst_type(out, dst=img, dtype=out.dtype)[0] def remove_small_objects( diff --git a/tests/test_keep_largest_connected_component.py b/tests/test_keep_largest_connected_component.py index 80dbc1c51d..a0e309f2d7 100644 --- a/tests/test_keep_largest_connected_component.py +++ b/tests/test_keep_largest_connected_component.py @@ -78,6 +78,8 @@ def to_onehot(x): ] grid_5 = [[[0, 0, 1, 0, 0], [0, 1, 1, 1, 1], [1, 1, 1, 0, 0], [1, 1, 0, 1, 0], [1, 1, 0, 0, 1]]] +grid_6 = [[[0, 0, 1, 1, 0, 0, 1], [0, 0, 0, 1, 0, 0, 1], [1, 1, 0, 0, 1, 0, 1], [0, 0, 0, 1, 0, 0, 1]]] + TESTS = [] for p in TEST_NDARRAYS: TESTS.append( @@ -343,6 +345,37 @@ def to_onehot(x): torch.tensor([[[0, 0, 1, 0, 0], [0, 2, 1, 1, 1], [0, 2, 1, 0, 0], [0, 2, 0, 1, 0], [2, 2, 0, 0, 0]]]), ] ) + # no connected regions + TESTS.append(["0 regions", {"num_components": 0}, p(grid_6), p(torch.zeros(1, 4, 7))]) + # 1 connected region + TESTS.append( + [ + "1 region", + {"num_components": 1}, + p(grid_6), + p( + torch.tensor( + [[[0, 0, 1, 1, 0, 0, 0], [0, 0, 0, 1, 0, 0, 0], [0, 0, 0, 0, 1, 0, 0], [0, 0, 0, 1, 0, 0, 0]]] + ) + ), + ] + ) + # 2 connected regions + TESTS.append( + [ + "2 regions", + {"num_components": 2}, + p(grid_6), + p( + torch.tensor( + [[[0, 0, 1, 1, 0, 0, 1], [0, 0, 0, 1, 0, 0, 1], [0, 0, 0, 0, 1, 0, 1], [0, 0, 0, 1, 0, 0, 1]]] + ) + ), + ] + ) + # 3+ connected regions unchanged (as input has 3) + for num_connected in (3, 4): + TESTS.append([f"{num_connected} regions", {"num_components": num_connected}, p(grid_6), p(grid_6)]) class TestKeepLargestConnectedComponent(unittest.TestCase): From 3f6a3537339cfdb9e1afcbe53133cbd63756e5cb Mon Sep 17 00:00:00 2001 From: Richard Brown <33289025+rijobro@users.noreply.github.com> Date: Tue, 13 Sep 2022 12:13:50 +0100 Subject: [PATCH 2/2] cupy keep largest Signed-off-by: Richard Brown <33289025+rijobro@users.noreply.github.com> --- monai/transforms/utils.py | 49 ++++++++++++++++++++------------------- 1 file changed, 25 insertions(+), 24 deletions(-) diff --git a/monai/transforms/utils.py b/monai/transforms/utils.py index 43387a07e9..3096d76889 100644 --- a/monai/transforms/utils.py +++ b/monai/transforms/utils.py @@ -57,7 +57,7 @@ optional_import, ) from monai.utils.enums import TransformBackends -from monai.utils.type_conversion import convert_data_type, convert_to_dst_type, convert_to_tensor +from monai.utils.type_conversion import convert_data_type, convert_to_cupy, convert_to_dst_type, convert_to_tensor measure, has_measure = optional_import("skimage.measure", "0.14.2", min_version) morphology, has_morphology = optional_import("skimage.morphology") @@ -965,36 +965,37 @@ def get_largest_connected_component_mask( https://scikit-image.org/docs/dev/api/skimage.measure.html#skimage.measure.label. num_components: The number of largest components to preserve. """ - if isinstance(img, torch.Tensor) and has_cp and has_cucim and num_components == 1: - x_cupy = monai.transforms.ToCupy()(img.short()) - x_label = cucim.skimage.measure.label(x_cupy, connectivity=connectivity) - vals, counts = cp.unique(x_label[cp.nonzero(x_label)], return_counts=True) - comp = x_label == vals[cp.ndarray.argmax(counts)] - out_tensor = monai.transforms.ToTensor(device=img.device)(comp) - out_tensor = out_tensor.bool() - - return out_tensor # type: ignore - - if not has_measure: - raise RuntimeError("Skimage.measure required.") + # use skimage/cucim.skimage and np/cp depending on whether packages are + # available and input is non-cpu torch.tensor + use_cp = has_cp and has_cucim and isinstance(img, torch.Tensor) and img.device != torch.device("cpu") + if use_cp: + img_ = convert_to_cupy(img.short()) # type: ignore + label = cucim.skimage.measure.label + lib = cp + else: + if not has_measure: + raise RuntimeError("Skimage.measure required.") + img_, *_ = convert_data_type(img, np.ndarray) + label = measure.label + lib = np - img_np, *_ = convert_data_type(img, np.ndarray) - # features will be an image, 0 for background and then each different + # features will be an image -- 0 for background and then each different # feature will have its own index. - features, num_features = measure.label(img_np, connectivity=connectivity, return_num=True) - # # if num features less than max desired, nothing to do. + features, num_features = label(img_, connectivity=connectivity, return_num=True) + # if num features less than max desired, nothing to do. if num_features <= num_components: - out = img_np.astype(bool) + out = img_.astype(bool) else: - # get number voxels per feature (bincount). np.argsort[::-1] to get indices - # of largest components. Convert to list for ease - features_to_keep = list(np.argsort(np.bincount(features.flat))[::-1]) - # remove 0 (background) - features_to_keep.remove(0) + # ignore background + nonzeros = features[lib.nonzero(features)] + # get number voxels per feature (bincount). argsort[::-1] to get indices + # of largest components. + features_to_keep = lib.argsort(lib.bincount(nonzeros))[::-1] # only keep the first n non-background indices features_to_keep = features_to_keep[:num_components] # generate labelfield. True if in list of features to keep - out = np.isin(features, features_to_keep) + out = lib.isin(features, features_to_keep) + return convert_to_dst_type(out, dst=img, dtype=out.dtype)[0]