Skip to content
Merged
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
4 changes: 3 additions & 1 deletion smauglab/configs/all_augmentations.json
Original file line number Diff line number Diff line change
Expand Up @@ -586,6 +586,7 @@
"noise": false,
"p": 1.0,
"random_pick": false,
"second_channel_is_labels": true,
"spike": false,
"swap": false
},
Expand All @@ -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,
Expand Down
18 changes: 16 additions & 2 deletions smauglab/transforms/cpu/artifact.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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)
Expand All @@ -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)
7 changes: 5 additions & 2 deletions smauglab/transforms/cpu/spatial.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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)
Expand All @@ -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
Expand Down
30 changes: 23 additions & 7 deletions smauglab/transforms/cpu/torchio_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)),
Expand All @@ -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


Expand Down
54 changes: 54 additions & 0 deletions unit_tests/test_torchio_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Loading