diff --git a/smauglab/configs/all_augmentations.json b/smauglab/configs/all_augmentations.json index 9c3d28f..6262de5 100644 --- a/smauglab/configs/all_augmentations.json +++ b/smauglab/configs/all_augmentations.json @@ -586,6 +586,7 @@ "noise": false, "p": 1.0, "random_pick": false, + "second_channel_is_labels": true, "spike": false, "swap": false }, @@ -595,7 +596,8 @@ "elastic": false, "flip": false, "p": 1.0, - "random_pick": false + "random_pick": false, + "second_channel_is_labels": true }, "SpatialTransform": { "align_corners": false, diff --git a/smauglab/transforms/cpu/artifact.py b/smauglab/transforms/cpu/artifact.py index 6f6469d..cf9d0ee 100644 --- a/smauglab/transforms/cpu/artifact.py +++ b/smauglab/transforms/cpu/artifact.py @@ -24,7 +24,18 @@ group=AugType.TA, ) class ArtifactTransform(BasicTransform): - def __init__(self, motion=False, ghosting=False, spike=False, bias_field=False, blur=False, noise=False, swap=False, random_pick=False): + def __init__( + self, + motion=False, + ghosting=False, + spike=False, + bias_field=False, + blur=False, + noise=False, + swap=False, + random_pick=False, + second_channel_is_labels=True, + ): """ Apply all selected artifacts (motion, ghosting, spike, bias field, blur, noise, and swap) to the image if they are enabled (set to True). If `random_pick` is True, randomly select and apply ONE of the enabled artifacts. @@ -40,6 +51,9 @@ def __init__(self, motion=False, ghosting=False, spike=False, bias_field=False, self.noise = noise self.swap = swap self.random_pick = random_pick + # Channel count alone cannot tell an image+labels pair from a + # two-modality image; see apply_tio. + self.second_channel_is_labels = second_channel_is_labels def get_parameters(self, **data_dict) -> dict: return select({name: getattr(self, name) for name in ARTIFACTS}, self.random_pick) @@ -50,4 +64,4 @@ def apply(self, data_dict: dict, **params) -> dict: return data_dict def _apply_to_image(self, img: torch.Tensor, seg: torch.Tensor, **params) -> tuple[torch.Tensor, torch.Tensor]: - return apply_enabled(ARTIFACTS, img, seg, params) + return apply_enabled(ARTIFACTS, img, seg, params, second_channel_is_labels=self.second_channel_is_labels) diff --git a/smauglab/transforms/cpu/spatial.py b/smauglab/transforms/cpu/spatial.py index 1cb9f54..d53ee92 100644 --- a/smauglab/transforms/cpu/spatial.py +++ b/smauglab/transforms/cpu/spatial.py @@ -23,7 +23,7 @@ group=AugType.GEO, ) class SpatialCustomTransform(BasicTransform): - def __init__(self, flip=False, affine=False, elastic=False, anisotropy=False, random_pick=False): + def __init__(self, flip=False, affine=False, elastic=False, anisotropy=False, random_pick=False, second_channel_is_labels=True): """ Apply all selected spatial transformation (flip, affine, elastic and anisotropy) to the image if they are enabled (set to True). If `random_pick` is True, randomly select and apply ONE of the enabled transformation. @@ -36,6 +36,9 @@ def __init__(self, flip=False, affine=False, elastic=False, anisotropy=False, ra self.elastic = elastic self.anisotropy = anisotropy self.random_pick = random_pick + # Channel count alone cannot tell an image+labels pair from a + # two-modality image; see apply_tio. + self.second_channel_is_labels = second_channel_is_labels def get_parameters(self, **data_dict) -> dict: return select({name: getattr(self, name) for name in SPATIAL_TRANSFORMS}, self.random_pick) @@ -46,7 +49,7 @@ def apply(self, data_dict: dict, **params) -> dict: return data_dict def _apply_to_image(self, img: torch.Tensor, seg: torch.Tensor, **params) -> tuple[torch.Tensor, torch.Tensor]: - return apply_enabled(SPATIAL_TRANSFORMS, img, seg, params) + return apply_enabled(SPATIAL_TRANSFORMS, img, seg, params, second_channel_is_labels=self.second_channel_is_labels) ### Shape transform diff --git a/smauglab/transforms/cpu/torchio_ops.py b/smauglab/transforms/cpu/torchio_ops.py index c3b7c0f..fdc4d88 100644 --- a/smauglab/transforms/cpu/torchio_ops.py +++ b/smauglab/transforms/cpu/torchio_ops.py @@ -39,22 +39,36 @@ def _image_data(subject: tio.Subject, key: str) -> torch.Tensor: return cast(tio.Image, subject[key]).data -def apply_tio(transform: tio.Transform, img: torch.Tensor, seg: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: +def apply_tio( + transform: tio.Transform, + img: torch.Tensor, + seg: torch.Tensor, + *, + second_channel_is_labels: bool = True, +) -> tuple[torch.Tensor, torch.Tensor]: """Apply a torchio transform to an image/segmentation pair. Two input layouts, both from totalspineseg's augment.py: - * `img` with two channels is the "step 2" layout -- channel 0 is the image and - channel 1 an odd-disc segmentation. The second channel goes in as a `LabelMap` - so torchio resamples it nearest-neighbour rather than interpolating labels. - * Anything else is a plain image plus its segmentation. + * the "step 2" layout -- channel 0 is the image and channel 1 an odd-disc + segmentation. The second channel goes in as a `LabelMap` so torchio + resamples it nearest-neighbour rather than interpolating labels. + * anything else is a plain image plus its segmentation. + + Which one applies used to be decided by channel count alone, so a genuine + two-modality input (T1+T2, in/out-of-phase) had its second modality + registered as a LabelMap: nearest-neighbour resampling, and every intensity + artifact skipped on it, so the two modalities diverged in both. Channel count + cannot tell those apart, so `second_channel_is_labels` says which it is. It + defaults to True, the long-standing behaviour; a multi-modality caller sets it + to False. The explicit `del` and `gc.collect()` are inherited: torchio subjects hold the whole volume several times over and these run inside dataloader workers. """ # Images come back through `_image_data`; see its docstring for why neither key # nor attribute access type-checks on its own. - if img.shape[0] == 2: + if img.shape[0] == 2 and second_channel_is_labels: subject = transform( tio.Subject( image=tio.ScalarImage(tensor=torch.unsqueeze(img[0], dim=0)), @@ -77,11 +91,13 @@ def apply_enabled( img: torch.Tensor, seg: torch.Tensor, enabled: Mapping[str, bool], + *, + second_channel_is_labels: bool = True, ) -> tuple[torch.Tensor, torch.Tensor]: """Apply each enabled transform in `factories` order, chaining the result.""" for name, factory in factories.items(): if enabled.get(name): - img, seg = apply_tio(factory(), img, seg) + img, seg = apply_tio(factory(), img, seg, second_channel_is_labels=second_channel_is_labels) return img, seg diff --git a/unit_tests/test_torchio_ops.py b/unit_tests/test_torchio_ops.py index ae3f28c..f8001bb 100644 --- a/unit_tests/test_torchio_ops.py +++ b/unit_tests/test_torchio_ops.py @@ -169,3 +169,57 @@ def test_random_pick_applies_exactly_one_artifact(self): params = transform.get_parameters(image=None) self.assertEqual(sum(bool(v) for v in params.values()), 1) + + +class TestSecondChannelLayoutIsExplicit(SmaugLabTestCase): + """Channel count cannot tell an image+labels pair from two modalities. + + `apply_tio` dispatched on `img.shape[0] == 2` alone, so a genuine + two-modality input (T1+T2, in-phase/out-of-phase) had its second modality + registered as a `tio.LabelMap`: resampled nearest-neighbour, and skipped by + every intensity artifact, so the two modalities diverged in both respects. + + The layout is now stated rather than guessed. The default is unchanged, so + the totalspineseg "step 2" pipeline behaves exactly as before. + """ + + def test_the_default_still_treats_the_second_channel_as_labels(self): + img, seg = _pair(channels=2) + img[1] = (img[1] > 0.5).float() + + img_out, _ = apply_tio(tio.RandomAffine(degrees=15), img, seg) + + self.assertTrue(bool(torch.isin(img_out[1], torch.tensor([0.0, 1.0])).all())) + + def test_a_second_modality_can_be_kept_as_an_image(self): + """The case that used to be impossible to express.""" + img, seg = _pair(channels=2) + + img_out, _ = apply_tio( + tio.RandomBlur(std=(1.0, 1.0)), + img, + seg, + second_channel_is_labels=False, + ) + + self.assertFalse( + bool(torch.equal(img_out[1], img[1])), + "an intensity artifact skipped the second channel, so it is still a LabelMap", + ) + + def test_as_labels_the_second_channel_is_left_alone_by_an_intensity_artifact(self): + """The control for the test above: the same call with the default.""" + img, seg = _pair(channels=2) + + img_out, _ = apply_tio(tio.RandomBlur(std=(1.0, 1.0)), img, seg) + + self.assertTrue(bool(torch.equal(img_out[1], img[1]))) + + def test_a_single_channel_image_is_unaffected_by_the_flag(self): + img, seg = _pair(channels=1) + + for flag in (True, False): + with self.subTest(second_channel_is_labels=flag): + img_out, _ = apply_tio(tio.RandomAffine(degrees=5), img, seg, second_channel_is_labels=flag) + + self.assertEqual(img_out.shape, img.shape)