From c79e102dce82ed3e180860116fec8fbdcaebf605 Mon Sep 17 00:00:00 2001 From: iback Date: Thu, 20 Aug 2026 12:42:02 +0000 Subject: [PATCH] refactor: extract the nnU-Net transform tail shared by five trainers Five get_training_transforms methods across nnUNetTrainerDAExt.py and nnUNetTrainerTest.py ended with a character-identical copy of the same nnU-Net sequence: intensity masking, the -1 label removal, the two cascade transforms, region conversion and deep-supervision downsampling. That is trainers/utils.py::nnunet_tail_transforms now. -255 lines, +134. The four full copies differed in exactly one thing: whether the DownsampleSegForDSTransform block was live or commented out. The GPU trainers carry it commented, 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. That is expressed by passing deep_supervision_scales=None rather than by a comment. get_validation_transforms is deliberately left alone. Its cascade branch adds only MoveSegAsOneHotToDataTransform, without the two RandomTransform wrappers, so it is a different sequence rather than a sixth copy of this one. Forcing it through the helper would need a flag that changes that branch, which costs more than it saves. Verified by enumerating the transform list every trainer builds across 160 argument combinations -- dummy 2D on/off, mask, cascade, regions, deep supervision -- and comparing the type and repr of each entry before and after. All 160 are identical once function memory addresses are normalised out. That check is local-only and not repeatable in CI: nnunetv2 is an optional extra and `pip install -e ".[dev]"` does not pull it, so nothing under trainers/ is imported by the test suite at all. It is the reason this is a separate commit from the gpu/contrast.py extraction, which CI does cover. Co-Authored-By: Claude Opus 5 --- smauglab/trainers/nnUNetTrainerDAExt.py | 178 ++++-------------------- smauglab/trainers/nnUNetTrainerTest.py | 123 +++------------- smauglab/trainers/utils.py | 88 +++++++++++- 3 files changed, 134 insertions(+), 255 deletions(-) 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.