Skip to content

Commit e8a5344

Browse files
engmohamedsalahMohamed Abdeltawab
andauthored
Fix ClipIntensityPercentiles metadata accumulation (Project-MONAI#9003)
### Description `ClipIntensityPercentiles` stored returned clipping values on the transform instance. Reusing the same transform therefore accumulated values from earlier calls and also mutated the metadata of earlier outputs. This change keeps the clipping-value accumulator local to each `__call__`. Each result now receives only its own clipping values, and later calls cannot change metadata already returned to a caller. Regression tests cover repeated channel-wise and non-channel-wise calls. ### Types of changes - [x] Non-breaking change (fix or new feature that would not break existing functionality). - [ ] Breaking change (fix or new feature that would cause existing functionality to change). - [x] New tests added to cover the changes. - [ ] Integration tests passed locally by running `./runtests.sh -f -u --net --coverage`. - [ ] Quick tests passed locally by running `./runtests.sh --quick --unittests --disttests`. - [ ] In-line docstrings updated. - [ ] Documentation updated, tested `make html` command in the `docs/` folder. ### Testing - `python -m unittest tests.transforms.test_clip_intensity_percentiles tests.transforms.test_clip_intensity_percentilesd` — 96 tests passed - `python -m ruff check monai/transforms/intensity/array.py tests/transforms/test_clip_intensity_percentiles.py` - `python -m black --check monai/transforms/intensity/array.py tests/transforms/test_clip_intensity_percentiles.py` Signed-off-by: Mohamed Abdeltawab <mohamed.abdeltawab@integrant.com> Co-authored-by: Mohamed Abdeltawab <mohamed.abdeltawab@integrant.com>
1 parent c1240a2 commit e8a5344

2 files changed

Lines changed: 34 additions & 9 deletions

File tree

monai/transforms/intensity/array.py

Lines changed: 10 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1128,12 +1128,12 @@ def __init__(
11281128
self.upper = upper
11291129
self.sharpness_factor = sharpness_factor
11301130
self.channel_wise = channel_wise
1131-
if return_clipping_values:
1132-
self.clipping_values: list[tuple[float | None, float | None]] = []
11331131
self.return_clipping_values = return_clipping_values
11341132
self.dtype = dtype
11351133

1136-
def _clip(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
1134+
def _clip(
1135+
self, img: NdarrayOrTensor, clipping_values: list[tuple[float | None, float | None]] | None = None
1136+
) -> NdarrayOrTensor:
11371137
if self.sharpness_factor is not None:
11381138
lower_percentile = percentile(img, self.lower) if self.lower is not None else None
11391139
upper_percentile = percentile(img, self.upper) if self.upper is not None else None
@@ -1143,8 +1143,8 @@ def _clip(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
11431143
upper_percentile = percentile(img, self.upper) if self.upper is not None else percentile(img, 100)
11441144
img = clip(img, lower_percentile, upper_percentile)
11451145

1146-
if self.return_clipping_values:
1147-
self.clipping_values.append(
1146+
if clipping_values is not None:
1147+
clipping_values.append(
11481148
(
11491149
(
11501150
lower_percentile
@@ -1165,16 +1165,17 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
11651165
"""
11661166
Apply the transform to `img`.
11671167
"""
1168+
clipping_values: list[tuple[float | None, float | None]] | None = [] if self.return_clipping_values else None
11681169
img = convert_to_tensor(img, track_meta=get_track_meta())
11691170
img_t = convert_to_tensor(img, track_meta=False)
11701171
if self.channel_wise:
1171-
img_t = torch.stack([self._clip(img=d) for d in img_t]) # type: ignore
1172+
img_t = torch.stack([self._clip(img=d, clipping_values=clipping_values) for d in img_t]) # type: ignore
11721173
else:
1173-
img_t = self._clip(img=img_t)
1174+
img_t = self._clip(img=img_t, clipping_values=clipping_values)
11741175

11751176
img = convert_to_dst_type(img_t, dst=img)[0]
1176-
if self.return_clipping_values:
1177-
img.meta["clipping_values"] = self.clipping_values # type: ignore
1177+
if clipping_values is not None:
1178+
img.meta["clipping_values"] = clipping_values # type: ignore
11781179

11791180
return img
11801181

tests/transforms/test_clip_intensity_percentiles.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -192,5 +192,29 @@ def test_channel_wise(self, p):
192192
assert_allclose(result[i], p(expected), type_test="tensor", rtol=1e-4, atol=0)
193193

194194

195+
class TestClipIntensityPercentilesClippingValues(unittest.TestCase):
196+
def test_clipping_values_repeated_channel_wise_calls(self):
197+
clipper = ClipIntensityPercentiles(lower=0, upper=100, channel_wise=True, return_clipping_values=True)
198+
first = clipper(torch.tensor([[[0.0, 1.0]], [[10.0, 20.0]]]))
199+
first_clipping_values = list(first.meta["clipping_values"])
200+
201+
second = clipper(torch.tensor([[[100.0, 200.0]], [[1000.0, 2000.0]]]))
202+
203+
self.assertEqual(first_clipping_values, [(0.0, 1.0), (10.0, 20.0)])
204+
self.assertEqual(first.meta["clipping_values"], first_clipping_values)
205+
self.assertEqual(second.meta["clipping_values"], [(100.0, 200.0), (1000.0, 2000.0)])
206+
207+
def test_clipping_values_repeated_non_channel_wise_calls(self):
208+
clipper = ClipIntensityPercentiles(lower=0, upper=100, return_clipping_values=True)
209+
first = clipper(torch.tensor([[[0.0, 1.0]]]))
210+
first_clipping_values = list(first.meta["clipping_values"])
211+
212+
second = clipper(torch.tensor([[[100.0, 200.0]]]))
213+
214+
self.assertEqual(first_clipping_values, [(0.0, 1.0)])
215+
self.assertEqual(first.meta["clipping_values"], first_clipping_values)
216+
self.assertEqual(second.meta["clipping_values"], [(100.0, 200.0)])
217+
218+
195219
if __name__ == "__main__":
196220
unittest.main()

0 commit comments

Comments
 (0)