diff --git a/smauglab/trainers/nnUNetTrainerDAExt.py b/smauglab/trainers/nnUNetTrainerDAExt.py index 5926255..47098ab 100644 --- a/smauglab/trainers/nnUNetTrainerDAExt.py +++ b/smauglab/trainers/nnUNetTrainerDAExt.py @@ -8,15 +8,11 @@ import torch from batchgeneratorsv2.helpers.scalar_type import RandomScalar from batchgeneratorsv2.transforms.base.basic_transform import BasicTransform -from batchgeneratorsv2.transforms.nnunet.random_binary_operator import ApplyRandomBinaryOperatorTransform -from batchgeneratorsv2.transforms.nnunet.remove_connected_components import RemoveRandomConnectedComponentFromOneHotEncodingTransform from batchgeneratorsv2.transforms.nnunet.seg_to_onehot import MoveSegAsOneHotToDataTransform from batchgeneratorsv2.transforms.spatial.spatial import SpatialTransform from batchgeneratorsv2.transforms.utils.compose import ComposeTransforms from batchgeneratorsv2.transforms.utils.deep_supervision_downsampling import DownsampleSegForDSTransform -from batchgeneratorsv2.transforms.utils.nnunet_masking import MaskImageTransform from batchgeneratorsv2.transforms.utils.pseudo2d import Convert2DTo3DTransform, Convert3DTo2DTransform -from batchgeneratorsv2.transforms.utils.random import RandomTransform from batchgeneratorsv2.transforms.utils.remove_label import RemoveLabelTansform from batchgeneratorsv2.transforms.utils.seg_to_regions import ConvertSegmentationToRegionsTransform from nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer @@ -24,7 +20,7 @@ from torch import autocast from smauglab import configs -from smauglab.trainers.utils import DownsampleSegForDSTransformCustom +from smauglab.trainers.utils import DownsampleSegForDSTransformCustom, nnunet_tail_transforms from smauglab.transforms.cpu.transforms import AugTransforms from smauglab.transforms.gpu.transforms import AugTransformsGPU @@ -96,54 +92,16 @@ def get_training_transforms( if do_dummy_2d_data_aug: transforms.append(Convert2DTo3DTransform()) - if use_mask_for_norm is not None and any(use_mask_for_norm): - transforms.append( - MaskImageTransform( - apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]], - channel_idx_in_seg=0, - set_outside_to=0, - ) + transforms.extend( + nnunet_tail_transforms( + use_mask_for_norm=use_mask_for_norm, + deep_supervision_scales=deep_supervision_scales, + is_cascaded=is_cascaded, + foreground_labels=foreground_labels, + regions=regions, + ignore_label=ignore_label, ) - - transforms.append(RemoveLabelTansform(-1, 0)) - - # The following augmentations are related to special nnunet executions - if is_cascaded: - assert foreground_labels is not None, "We need foreground_labels for cascade augmentations" - transforms.append( - MoveSegAsOneHotToDataTransform(source_channel_idx=1, all_labels=foreground_labels, remove_channel_from_source=True) - ) - transforms.append( - RandomTransform( - ApplyRandomBinaryOperatorTransform( - channel_idx=list(range(-len(foreground_labels), 0)), strel_size=(1, 8), p_per_label=1 - ), - apply_probability=0.4, - ) - ) - transforms.append( - RandomTransform( - RemoveRandomConnectedComponentFromOneHotEncodingTransform( - channel_idx=list(range(-len(foreground_labels), 0)), - fill_with_other_class_p=0, - dont_do_if_covers_more_than_x_percent=0.15, - p_per_label=1, - ), - apply_probability=0.2, - ) - ) - - if regions is not None: - # the ignore label must also be converted - transforms.append( - ConvertSegmentationToRegionsTransform( - regions=[*list(regions), ignore_label] if ignore_label is not None else regions, channel_in_seg=0 - ) - ) - - if deep_supervision_scales is not None: - transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales)) - + ) return ComposeTransforms(transforms) @@ -220,57 +178,16 @@ def get_training_transforms( if do_dummy_2d_data_aug: transforms.append(Convert2DTo3DTransform()) - if use_mask_for_norm is not None and any(use_mask_for_norm): - transforms.append( - MaskImageTransform( - apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]], - channel_idx_in_seg=0, - set_outside_to=0, - ) - ) - - transforms.append(RemoveLabelTansform(-1, 0)) - - # The following augmentations are related to special nnunet executions - if is_cascaded: - assert foreground_labels is not None, "We need foreground_labels for cascade augmentations" - transforms.append( - MoveSegAsOneHotToDataTransform(source_channel_idx=1, all_labels=foreground_labels, remove_channel_from_source=True) - ) - transforms.append( - RandomTransform( - ApplyRandomBinaryOperatorTransform( - channel_idx=list(range(-len(foreground_labels), 0)), strel_size=(1, 8), p_per_label=1 - ), - apply_probability=0.4, - ) - ) - transforms.append( - RandomTransform( - RemoveRandomConnectedComponentFromOneHotEncodingTransform( - channel_idx=list(range(-len(foreground_labels), 0)), - fill_with_other_class_p=0, - dont_do_if_covers_more_than_x_percent=0.15, - p_per_label=1, - ), - apply_probability=0.2, - ) - ) - - if regions is not None: - # the ignore label must also be converted - transforms.append( - ConvertSegmentationToRegionsTransform( - regions=[*list(regions), ignore_label] if ignore_label is not None else regions, channel_in_seg=0 - ) + transforms.extend( + nnunet_tail_transforms( + use_mask_for_norm=use_mask_for_norm, + deep_supervision_scales=None, + is_cascaded=is_cascaded, + foreground_labels=foreground_labels, + regions=regions, + ignore_label=ignore_label, ) - - # transforms.append(ZscoreNormalization()) - - # NOTE: DownsampleSegForDSTransform is now handled in train_step for GPU augmentations - # if deep_supervision_scales is not None: - # transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales)) - + ) return ComposeTransforms(transforms) @staticmethod @@ -399,55 +316,16 @@ def get_training_transforms( ) ) - if use_mask_for_norm is not None and any(use_mask_for_norm): - transforms.append( - MaskImageTransform( - apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]], - channel_idx_in_seg=0, - set_outside_to=0, - ) + transforms.extend( + nnunet_tail_transforms( + use_mask_for_norm=use_mask_for_norm, + deep_supervision_scales=None, + is_cascaded=is_cascaded, + foreground_labels=foreground_labels, + regions=regions, + ignore_label=ignore_label, ) - - transforms.append(RemoveLabelTansform(-1, 0)) - - # The following augmentations are related to special nnunet executions - if is_cascaded: - assert foreground_labels is not None, "We need foreground_labels for cascade augmentations" - transforms.append( - MoveSegAsOneHotToDataTransform(source_channel_idx=1, all_labels=foreground_labels, remove_channel_from_source=True) - ) - transforms.append( - RandomTransform( - ApplyRandomBinaryOperatorTransform( - channel_idx=list(range(-len(foreground_labels), 0)), strel_size=(1, 8), p_per_label=1 - ), - apply_probability=0.4, - ) - ) - transforms.append( - RandomTransform( - RemoveRandomConnectedComponentFromOneHotEncodingTransform( - channel_idx=list(range(-len(foreground_labels), 0)), - fill_with_other_class_p=0, - dont_do_if_covers_more_than_x_percent=0.15, - p_per_label=1, - ), - apply_probability=0.2, - ) - ) - - if regions is not None: - # the ignore label must also be converted - transforms.append( - ConvertSegmentationToRegionsTransform( - regions=[*list(regions), ignore_label] if ignore_label is not None else regions, channel_in_seg=0 - ) - ) - - # NOTE: DownsampleSegForDSTransform is now handled in train_step for GPU augmentations - # if deep_supervision_scales is not None: - # transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales)) - + ) return ComposeTransforms(transforms) def train_step(self, batch: dict) -> dict: diff --git a/smauglab/trainers/nnUNetTrainerTest.py b/smauglab/trainers/nnUNetTrainerTest.py index 4b99d1a..a51f32f 100644 --- a/smauglab/trainers/nnUNetTrainerTest.py +++ b/smauglab/trainers/nnUNetTrainerTest.py @@ -5,23 +5,15 @@ import torch from batchgeneratorsv2.helpers.scalar_type import RandomScalar from batchgeneratorsv2.transforms.base.basic_transform import BasicTransform -from batchgeneratorsv2.transforms.nnunet.random_binary_operator import ApplyRandomBinaryOperatorTransform -from batchgeneratorsv2.transforms.nnunet.remove_connected_components import RemoveRandomConnectedComponentFromOneHotEncodingTransform -from batchgeneratorsv2.transforms.nnunet.seg_to_onehot import MoveSegAsOneHotToDataTransform from batchgeneratorsv2.transforms.spatial.spatial import SpatialTransform from batchgeneratorsv2.transforms.utils.compose import ComposeTransforms -from batchgeneratorsv2.transforms.utils.deep_supervision_downsampling import DownsampleSegForDSTransform -from batchgeneratorsv2.transforms.utils.nnunet_masking import MaskImageTransform from batchgeneratorsv2.transforms.utils.pseudo2d import Convert2DTo3DTransform, Convert3DTo2DTransform -from batchgeneratorsv2.transforms.utils.random import RandomTransform -from batchgeneratorsv2.transforms.utils.remove_label import RemoveLabelTansform -from batchgeneratorsv2.transforms.utils.seg_to_regions import ConvertSegmentationToRegionsTransform from nnunetv2.training.nnUNetTrainer.nnUNetTrainer import nnUNetTrainer from nnunetv2.utilities.helpers import dummy_context from torch import autocast from smauglab import configs -from smauglab.trainers.utils import DownsampleSegForDSTransformCustom +from smauglab.trainers.utils import DownsampleSegForDSTransformCustom, nnunet_tail_transforms from smauglab.transforms.cpu.transforms import AugTransformsTest from smauglab.transforms.gpu.transforms import AugTransformsGPU @@ -74,54 +66,16 @@ def get_training_transforms( if do_dummy_2d_data_aug: transforms.append(Convert2DTo3DTransform()) - if use_mask_for_norm is not None and any(use_mask_for_norm): - transforms.append( - MaskImageTransform( - apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]], - channel_idx_in_seg=0, - set_outside_to=0, - ) + transforms.extend( + nnunet_tail_transforms( + use_mask_for_norm=use_mask_for_norm, + deep_supervision_scales=deep_supervision_scales, + is_cascaded=is_cascaded, + foreground_labels=foreground_labels, + regions=regions, + ignore_label=ignore_label, ) - - transforms.append(RemoveLabelTansform(-1, 0)) - - # The following augmentations are related to special nnunet executions - if is_cascaded: - assert foreground_labels is not None, "We need foreground_labels for cascade augmentations" - transforms.append( - MoveSegAsOneHotToDataTransform(source_channel_idx=1, all_labels=foreground_labels, remove_channel_from_source=True) - ) - transforms.append( - RandomTransform( - ApplyRandomBinaryOperatorTransform( - channel_idx=list(range(-len(foreground_labels), 0)), strel_size=(1, 8), p_per_label=1 - ), - apply_probability=0.4, - ) - ) - transforms.append( - RandomTransform( - RemoveRandomConnectedComponentFromOneHotEncodingTransform( - channel_idx=list(range(-len(foreground_labels), 0)), - fill_with_other_class_p=0, - dont_do_if_covers_more_than_x_percent=0.15, - p_per_label=1, - ), - apply_probability=0.2, - ) - ) - - if regions is not None: - # the ignore label must also be converted - transforms.append( - ConvertSegmentationToRegionsTransform( - regions=[*list(regions), ignore_label] if ignore_label is not None else regions, channel_in_seg=0 - ) - ) - - if deep_supervision_scales is not None: - transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales)) - + ) return ComposeTransforms(transforms) @@ -175,55 +129,16 @@ def get_training_transforms( if do_dummy_2d_data_aug: transforms.append(Convert2DTo3DTransform()) - if use_mask_for_norm is not None and any(use_mask_for_norm): - transforms.append( - MaskImageTransform( - apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]], - channel_idx_in_seg=0, - set_outside_to=0, - ) + transforms.extend( + nnunet_tail_transforms( + use_mask_for_norm=use_mask_for_norm, + deep_supervision_scales=None, + is_cascaded=is_cascaded, + foreground_labels=foreground_labels, + regions=regions, + ignore_label=ignore_label, ) - - transforms.append(RemoveLabelTansform(-1, 0)) - - # The following augmentations are related to special nnunet executions - if is_cascaded: - assert foreground_labels is not None, "We need foreground_labels for cascade augmentations" - transforms.append( - MoveSegAsOneHotToDataTransform(source_channel_idx=1, all_labels=foreground_labels, remove_channel_from_source=True) - ) - transforms.append( - RandomTransform( - ApplyRandomBinaryOperatorTransform( - channel_idx=list(range(-len(foreground_labels), 0)), strel_size=(1, 8), p_per_label=1 - ), - apply_probability=0.4, - ) - ) - transforms.append( - RandomTransform( - RemoveRandomConnectedComponentFromOneHotEncodingTransform( - channel_idx=list(range(-len(foreground_labels), 0)), - fill_with_other_class_p=0, - dont_do_if_covers_more_than_x_percent=0.15, - p_per_label=1, - ), - apply_probability=0.2, - ) - ) - - if regions is not None: - # the ignore label must also be converted - transforms.append( - ConvertSegmentationToRegionsTransform( - regions=[*list(regions), ignore_label] if ignore_label is not None else regions, channel_in_seg=0 - ) - ) - - # NOTE: DownsampleSegForDSTransform is now handled in train_step for GPU augmentations - # if deep_supervision_scales is not None: - # transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales)) - + ) return ComposeTransforms(transforms) def train_step(self, batch: dict) -> dict: diff --git a/smauglab/trainers/utils.py b/smauglab/trainers/utils.py index 0e4e6bb..e7a511e 100644 --- a/smauglab/trainers/utils.py +++ b/smauglab/trainers/utils.py @@ -1,10 +1,96 @@ from collections.abc import Sequence -from typing import Union +from typing import Any, Union import torch from torch.nn.functional import interpolate +def nnunet_tail_transforms( + *, + use_mask_for_norm: list[bool] | None = None, + deep_supervision_scales: Union[list, tuple, None] = None, + is_cascaded: bool = False, + foreground_labels: Union[tuple[int, ...], list[int], None] = None, + regions: list[Union[list[int], tuple[int, ...], int]] | None = None, + ignore_label: int | None = None, +) -> list[Any]: + """The nnU-Net transforms that follow SmaugLab's augmentations, in order. + + Five `get_training_transforms` methods across two modules ended with a + character-identical copy of this -- intensity masking, the -1 label removal, the + two cascade transforms, region conversion and deep-supervision downsampling. + + `deep_supervision_scales=None` skips the downsampling, which is how the GPU + trainers had it: they carry the block commented out, because with GPU + augmentations the mask is still being deformed after this point and the + multi-scale targets have to be built from the augmented mask in `train_step`. + + `get_validation_transforms` deliberately does not use this. Its cascade branch + adds only MoveSegAsOneHotToDataTransform, without the two RandomTransform + wrappers, so it is a different sequence rather than another copy of this one. + """ + from batchgeneratorsv2.transforms.nnunet.random_binary_operator import ApplyRandomBinaryOperatorTransform + from batchgeneratorsv2.transforms.nnunet.remove_connected_components import ( + RemoveRandomConnectedComponentFromOneHotEncodingTransform, + ) + from batchgeneratorsv2.transforms.nnunet.seg_to_onehot import MoveSegAsOneHotToDataTransform + from batchgeneratorsv2.transforms.utils.deep_supervision_downsampling import DownsampleSegForDSTransform + from batchgeneratorsv2.transforms.utils.nnunet_masking import MaskImageTransform + from batchgeneratorsv2.transforms.utils.random import RandomTransform + from batchgeneratorsv2.transforms.utils.remove_label import RemoveLabelTansform + from batchgeneratorsv2.transforms.utils.seg_to_regions import ConvertSegmentationToRegionsTransform + + transforms: list[Any] = [] + + if use_mask_for_norm is not None and any(use_mask_for_norm): + transforms.append( + MaskImageTransform( + apply_to_channels=[i for i in range(len(use_mask_for_norm)) if use_mask_for_norm[i]], + channel_idx_in_seg=0, + set_outside_to=0, + ) + ) + + transforms.append(RemoveLabelTansform(-1, 0)) + + # The following augmentations are related to special nnunet executions + if is_cascaded: + assert foreground_labels is not None, "We need foreground_labels for cascade augmentations" + transforms.append( + MoveSegAsOneHotToDataTransform(source_channel_idx=1, all_labels=foreground_labels, remove_channel_from_source=True) + ) + transforms.append( + RandomTransform( + ApplyRandomBinaryOperatorTransform(channel_idx=list(range(-len(foreground_labels), 0)), strel_size=(1, 8), p_per_label=1), + apply_probability=0.4, + ) + ) + transforms.append( + RandomTransform( + RemoveRandomConnectedComponentFromOneHotEncodingTransform( + channel_idx=list(range(-len(foreground_labels), 0)), + fill_with_other_class_p=0, + dont_do_if_covers_more_than_x_percent=0.15, + p_per_label=1, + ), + apply_probability=0.2, + ) + ) + + if regions is not None: + # the ignore label must also be converted + transforms.append( + ConvertSegmentationToRegionsTransform( + regions=[*list(regions), ignore_label] if ignore_label is not None else regions, channel_in_seg=0 + ) + ) + + if deep_supervision_scales is not None: + transforms.append(DownsampleSegForDSTransform(ds_scales=deep_supervision_scales)) + + return transforms + + class DownsampleSegForDSTransformCustom: """ Custom deep supervision downsampling transform that handles batched tensors properly.