diff --git a/monai/transforms/intensity/array.py b/monai/transforms/intensity/array.py index d941f43ad7..a23c867d60 100644 --- a/monai/transforms/intensity/array.py +++ b/monai/transforms/intensity/array.py @@ -1128,12 +1128,12 @@ def __init__( self.upper = upper self.sharpness_factor = sharpness_factor self.channel_wise = channel_wise - if return_clipping_values: - self.clipping_values: list[tuple[float | None, float | None]] = [] self.return_clipping_values = return_clipping_values self.dtype = dtype - def _clip(self, img: NdarrayOrTensor) -> NdarrayOrTensor: + def _clip( + self, img: NdarrayOrTensor, clipping_values: list[tuple[float | None, float | None]] | None = None + ) -> NdarrayOrTensor: if self.sharpness_factor is not None: lower_percentile = percentile(img, self.lower) if self.lower is not None else None upper_percentile = percentile(img, self.upper) if self.upper is not None else None @@ -1143,8 +1143,8 @@ def _clip(self, img: NdarrayOrTensor) -> NdarrayOrTensor: upper_percentile = percentile(img, self.upper) if self.upper is not None else percentile(img, 100) img = clip(img, lower_percentile, upper_percentile) - if self.return_clipping_values: - self.clipping_values.append( + if clipping_values is not None: + clipping_values.append( ( ( lower_percentile @@ -1165,16 +1165,17 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor: """ Apply the transform to `img`. """ + clipping_values: list[tuple[float | None, float | None]] | None = [] if self.return_clipping_values else None img = convert_to_tensor(img, track_meta=get_track_meta()) img_t = convert_to_tensor(img, track_meta=False) if self.channel_wise: - img_t = torch.stack([self._clip(img=d) for d in img_t]) # type: ignore + img_t = torch.stack([self._clip(img=d, clipping_values=clipping_values) for d in img_t]) # type: ignore else: - img_t = self._clip(img=img_t) + img_t = self._clip(img=img_t, clipping_values=clipping_values) img = convert_to_dst_type(img_t, dst=img)[0] - if self.return_clipping_values: - img.meta["clipping_values"] = self.clipping_values # type: ignore + if clipping_values is not None: + img.meta["clipping_values"] = clipping_values # type: ignore return img diff --git a/tests/transforms/test_clip_intensity_percentiles.py b/tests/transforms/test_clip_intensity_percentiles.py index 18ed47dbaa..12d93da47d 100644 --- a/tests/transforms/test_clip_intensity_percentiles.py +++ b/tests/transforms/test_clip_intensity_percentiles.py @@ -192,5 +192,29 @@ def test_channel_wise(self, p): assert_allclose(result[i], p(expected), type_test="tensor", rtol=1e-4, atol=0) +class TestClipIntensityPercentilesClippingValues(unittest.TestCase): + def test_clipping_values_repeated_channel_wise_calls(self): + clipper = ClipIntensityPercentiles(lower=0, upper=100, channel_wise=True, return_clipping_values=True) + first = clipper(torch.tensor([[[0.0, 1.0]], [[10.0, 20.0]]])) + first_clipping_values = list(first.meta["clipping_values"]) + + second = clipper(torch.tensor([[[100.0, 200.0]], [[1000.0, 2000.0]]])) + + self.assertEqual(first_clipping_values, [(0.0, 1.0), (10.0, 20.0)]) + self.assertEqual(first.meta["clipping_values"], first_clipping_values) + self.assertEqual(second.meta["clipping_values"], [(100.0, 200.0), (1000.0, 2000.0)]) + + def test_clipping_values_repeated_non_channel_wise_calls(self): + clipper = ClipIntensityPercentiles(lower=0, upper=100, return_clipping_values=True) + first = clipper(torch.tensor([[[0.0, 1.0]]])) + first_clipping_values = list(first.meta["clipping_values"]) + + second = clipper(torch.tensor([[[100.0, 200.0]]])) + + self.assertEqual(first_clipping_values, [(0.0, 1.0)]) + self.assertEqual(first.meta["clipping_values"], first_clipping_values) + self.assertEqual(second.meta["clipping_values"], [(100.0, 200.0)]) + + if __name__ == "__main__": unittest.main()