diff --git a/smauglab/transforms/cpu/contrast.py b/smauglab/transforms/cpu/contrast.py index 2c46128..f6b0fe9 100644 --- a/smauglab/transforms/cpu/contrast.py +++ b/smauglab/transforms/cpu/contrast.py @@ -2,15 +2,45 @@ import torch import torch.nn.functional as F +from batchgeneratorsv2.helpers.scalar_type import RandomScalar from batchgeneratorsv2.transforms.base.basic_transform import ImageOnlyTransform +from batchgeneratorsv2.transforms.intensity.gamma import GammaTransform from smauglab.transforms.kernels import laplace_kernel, scharr_kernels -class ConvTransform(ImageOnlyTransform): +class InvertedGammaTransform(GammaTransform): + """Gamma adjustment applied to the inverted image. + + batchgeneratorsv2 expresses this as `GammaTransform(p_invert_image=1)`, which the + old config spelled as a `GammaTransform_invert` key -- a key with no class behind + it. A real class keeps the config key 1:1 with a class on this backend too, and + means the inversion cannot be requested two different ways. + """ + + def __init__( + self, + gamma: RandomScalar = (0.7, 1.5), + synchronize_channels: bool = False, + p_per_channel: float = 1, + p_retain_stats: float = 1, + ): + super().__init__( + gamma=gamma, + p_invert_image=1, + synchronize_channels=synchronize_channels, + p_per_channel=p_per_channel, + p_retain_stats=p_retain_stats, + ) + + +class _ConvBaseTransform(ImageOnlyTransform): """ Applies a Laplace/Scharr filter to the image to highlight edges. + Shared implementation. Configs address the per-kernel leaves below, which mirror + the GPU split so both backends of one augmentation share an `aug_id`. + Based on https://github.com/spinalcordtoolbox/disc-labeling-playground/blob/main/src/ply/models/transform.py """ @@ -28,13 +58,13 @@ def get_parameters(self, **data_dict) -> dict: # _apply_to_image dispatches on kernel_type to tell the two apart. kernel: Union[torch.Tensor, list[torch.Tensor]] spatial_dims = len(data_dict["image"].shape) - 1 - if spatial_dims in (2, 3): - if self.kernel_type == "Laplace": - kernel = laplace_kernel(spatial_dims) - elif self.kernel_type == "Scharr": - kernel = scharr_kernels(spatial_dims) - else: + if spatial_dims not in (2, 3): raise ValueError(f"{self.__class__} can only handle 2D or 3D images.") + # Shared with the GPU backend. These tables used to be written out here and + # again in gpu/contrast.py, which is how the 2-D Scharr x-kernel came to have + # [-10, 0, -10] as its middle row on this side only -- summing to -20, so not + # a gradient operator at all. See smauglab/transforms/kernels.py. + kernel = laplace_kernel(spatial_dims) if self.kernel_type == "Laplace" else scharr_kernels(spatial_dims) return {"kernel_type": self.kernel_type, "kernel": kernel, "absolute": self.absolute, "retain_stats": self.retain_stats} @@ -65,6 +95,24 @@ def _apply_to_image(self, img: torch.Tensor, **params) -> torch.Tensor: return img +# One class per kernel, mirroring the GPU split so a config key names a class on +# either backend and `kernel_type` disappears from the config surface. + + +class LaplaceConvTransform(_ConvBaseTransform): + """Laplacian edge enhancement.""" + + def __init__(self, absolute: bool = False, retain_stats: bool = False): + super().__init__(kernel_type="Laplace", absolute=absolute, retain_stats=retain_stats) + + +class ScharrConvTransform(_ConvBaseTransform): + """Scharr gradient-magnitude edge filter.""" + + def __init__(self, absolute: bool = True, retain_stats: bool = False): + super().__init__(kernel_type="Scharr", absolute=absolute, retain_stats=retain_stats) + + class HistogramEqualTransform(ImageOnlyTransform): """ Update image intensity using histogram manipulations @@ -112,7 +160,7 @@ def _apply_to_image(self, img: torch.Tensor, **params) -> torch.Tensor: return img -class FunctionTransform(ImageOnlyTransform): +class _FunctionBaseTransform(ImageOnlyTransform): """ Apply different functions to image pixels @@ -148,6 +196,60 @@ def _apply_to_image(self, img: torch.Tensor, **params) -> torch.Tensor: return img +# One class per elementwise function; `function` is not expressible in JSON, so the +# old config had a single key that the builder fanned out over a hardcoded lambda +# list. Spelled out longhand rather than as torch.log1p / torch.sigmoid, which +# differ in the last ulp and would move the seeded determinism hashes. + + +def _log1p(x: torch.Tensor) -> torch.Tensor: + return torch.log(1 + x) + + +def _sigmoid(x: torch.Tensor) -> torch.Tensor: + return 1 / (1 + torch.exp(-x)) + + +class _NamedFunctionTransform(_FunctionBaseTransform): + """Shared constructor for the fixed-function leaves. Not registered itself.""" + + #: Set by each leaf. + function_impl: staticmethod + + def __init__(self, retain_stats: bool = False): + super().__init__(function=type(self).function_impl, retain_stats=retain_stats) + + +class Log1pTransform(_NamedFunctionTransform): + """Apply log(1 + x).""" + + function_impl = staticmethod(_log1p) + + +class SqrtTransform(_NamedFunctionTransform): + """Apply sqrt(x).""" + + function_impl = staticmethod(torch.sqrt) + + +class SinTransform(_NamedFunctionTransform): + """Apply sin(x).""" + + function_impl = staticmethod(torch.sin) + + +class ExpTransform(_NamedFunctionTransform): + """Apply exp(x).""" + + function_impl = staticmethod(torch.exp) + + +class SigmoidTransform(_NamedFunctionTransform): + """Apply the logistic sigmoid 1 / (1 + exp(-x)).""" + + function_impl = staticmethod(_sigmoid) + + def apply_filter(x: torch.Tensor, kernel: torch.Tensor, **kwargs) -> torch.Tensor: """ Copied from https://github.com/Project-MONAI/MONAI/blob/dev/monai/networks/layers/simplelayers.py @@ -213,3 +315,9 @@ def _apply_to_image(self, img: torch.Tensor, **params) -> torch.Tensor: std = torch.std(img[c]) img[c] = (img[c] - mean) / torch.clamp(std, min=1e-8) return img + + +# Temporary bridge for the CPU `if` ladder, which passes kernel_type from the config. +# Removed with that ladder; see the note in gpu/contrast.py. +ConvTransform = _ConvBaseTransform +FunctionTransform = _FunctionBaseTransform diff --git a/smauglab/transforms/cpu/transforms.py b/smauglab/transforms/cpu/transforms.py index 97cdaa3..c99860b 100644 --- a/smauglab/transforms/cpu/transforms.py +++ b/smauglab/transforms/cpu/transforms.py @@ -18,7 +18,7 @@ from batchgeneratorsv2.transforms.utils.random import RandomTransform from smauglab.transforms.cpu.artifact import ArtifactTransform -from smauglab.transforms.cpu.contrast import ConvTransform, FunctionTransform, HistogramEqualTransform +from smauglab.transforms.cpu.contrast import FunctionTransform, HistogramEqualTransform, _ConvBaseTransform from smauglab.transforms.cpu.fromSeg import RedistributeTransform from smauglab.transforms.cpu.spatial import ShapeTransform, SpatialCustomTransform @@ -63,11 +63,11 @@ def _build_transforms( transforms = [] # Scharr filter - conv_params = transform_params.get("ConvTransform") + conv_params = transform_params.get("_ConvBaseTransform") if conv_params is not None: transforms.append( RandomTransform( - ConvTransform( + _ConvBaseTransform( kernel_type=conv_params.get("kernel_type", "Scharr"), absolute=conv_params.get("absolute", True), retain_stats=transform_params.get("retain_stats", False), @@ -317,7 +317,7 @@ def _build_transforms(self): # Scharr filter transforms.append( RandomTransform( - ConvTransform( + _ConvBaseTransform( kernel_type="Scharr", absolute=True, ), diff --git a/smauglab/transforms/gpu/contrast.py b/smauglab/transforms/gpu/contrast.py index 49e0b7f..572ca09 100644 --- a/smauglab/transforms/gpu/contrast.py +++ b/smauglab/transforms/gpu/contrast.py @@ -1,5 +1,5 @@ import math -from collections.abc import Callable +from collections.abc import Callable, Sequence from typing import Any, Protocol, Union import torch @@ -102,7 +102,7 @@ def _foreground(mask: torch.Tensor, dim: int) -> torch.Tensor: This used to be `torch.argmax(mask, dim) > 0`, which is only "is anything labelled here" if class 0 is background -- and it is not: - * For a single-channel [B, 1, D, H, W] mask (an ordinary nnU-Net target, and what + * For a single-channel `[B, 1, D, H, W]` mask (an ordinary nnU-Net target, and what the tests build) `argmax` over a length-1 axis is always 0, so the result was **all False**. `in_seg` then applied the transform nowhere and `out_seg` applied it everywhere: the two knobs did nothing and the opposite of nothing. @@ -183,16 +183,20 @@ def _apply_region_mode( ## Convolution transform -class RandomConvTransformGPU(ImageOnlyTransform): +class _RandomConvBaseGPU(ImageOnlyTransform): """Apply convolution to image. If the image is torch Tensor, it is expected to have [N, C, X, Y] or [N, C, X, Y, Z] shape. Based on https://docs.pytorch.org/vision/0.9/transforms.html#torchvision.transforms.GaussianBlur Args: - kernel_type (str): Type of convolution kernel, either 'Laplace' or 'Scharr'. Default is 'Laplace'. - spatial_dims (int): Number of spatial dimensions of the input image, either 2 or 3. Default is 2. - absolute (bool): If True, take the absolute value of the convolution result. Default is False. - retain_stats (bool): If True, retain the original mean and standard deviation of the image after convolution. Default is False. + kernel_type (str): One of 'Laplace', 'Scharr', 'GaussianBlur', 'UnsharpMask', 'RandConv'. + apply_to_channel (list of int): Channel indices to convolve. Default is [0]. + absolute (bool): If True, take the absolute value of the result. Scharr only. + sigma (float): Gaussian width. GaussianBlur and UnsharpMask only. + unsharp_amount (float): Strength of the unsharp mask. UnsharpMask only. + kernel_sizes (list of int): Multi-scale kernel sizes to draw from. RandConv only. + mix_prob (float): Probability of blending the result back with the original. + retain_stats (bool): If True, restore the original mean and std afterwards. Returns: Tensor: Convolved version of the input image. @@ -202,35 +206,40 @@ class RandomConvTransformGPU(ImageOnlyTransform): def __init__( self, kernel_type: str = "Laplace", - apply_to_channel: list[int] | None = None, # Apply to first channel by default + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default same_on_batch: bool = False, retain_stats: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, + # Kernel-specific. These used to be read out of **kwargs, which meant they + # were invisible to `inspect.signature` and a typo in a config silently + # selected the default instead. Defaults here are the historical + # kwargs.get() ones, so behaviour is unchanged. + absolute: bool = False, + sigma: float = 1.0, + unsharp_amount: float = 1.0, + kernel_sizes: Sequence[int] = (1, 3, 5, 7), + mix_prob: float = 0.0, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) if kernel_type not in ["Laplace", "Scharr", "GaussianBlur", "UnsharpMask", "RandConv"]: raise NotImplementedError('Currently only "Laplace", "Scharr", "GaussianBlur", "UnsharpMask" and "RandConv" are supported.') else: self.kernel_type = kernel_type self.apply_to_channel = apply_to_channel - self.absolute = kwargs.get("absolute", False) - self.sigma = kwargs.get("sigma", 1.0) + self.absolute = absolute + self.sigma = sigma self.retain_stats = retain_stats self.in_seg = in_seg self.out_seg = out_seg self.mix_in_out = mix_in_out - # Unsharp mask parameters: amount controls strength of the mask - self.unsharp_amount = kwargs.get("unsharp_amount", 1.0) - # RandConv parameters - self.kernel_sizes = kwargs.get("kernel_sizes", [1, 3, 5, 7]) # multi-scale default - self.mix_prob = kwargs.get("mix_prob", 0.0) # probability to mix with original + self.unsharp_amount = unsharp_amount + self.kernel_sizes = kernel_sizes + self.mix_prob = mix_prob def get_kernel(self, device: torch.device) -> Union[Tensor, list[Tensor]]: # Scharr is the odd one out: it returns the three directional kernels as a @@ -273,7 +282,8 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ channel_data = input[:, c] # [N, ...spatial...] orig = channel_data.clone() - stats = _channel_stats(channel_data) if self.retain_stats else None + if self.retain_stats: + stats = _channel_stats(channel_data) # The asserts below restate what get_kernel guarantees per kernel_type: # only Scharr yields a list, and only its branch iterates. @@ -313,7 +323,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ alpha = torch.rand(1, device=input.device) x = alpha * orig + (1 - alpha) * x - if stats is not None: + if self.retain_stats: x = _restore_stats(x, stats) # Apply region selection @@ -325,6 +335,186 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ return input +# One class per convolution kernel. +# +# These used to be a single `kernel_type=` argument on the base, which meant four +# different augmentations shared one config key and every config had to repeat the +# kernel name redundantly. A class each keeps the config key 1:1 with the class, +# lets each expose only the parameters its kernel actually reads, and makes the +# CPU/GPU coverage matrix able to tell them apart. +# +# Defaults below are the values the old `_build_transforms` ladder passed for that +# kernel, NOT the base class defaults -- that is what keeps behaviour identical once +# the ladder is gone. + + +class RandomLaplaceGPU(_RandomConvBaseGPU): + """Laplacian edge enhancement.""" + + def __init__( + self, + absolute: bool = False, + mix_prob: float = 0.0, + apply_to_channel: Sequence[int] = (0,), + same_on_batch: bool = False, + retain_stats: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = True, + ) -> None: + super().__init__( + kernel_type="Laplace", + absolute=absolute, + mix_prob=mix_prob, + apply_to_channel=apply_to_channel, + same_on_batch=same_on_batch, + retain_stats=retain_stats, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + +class RandomScharrGPU(_RandomConvBaseGPU): + """Scharr gradient-magnitude edge filter.""" + + def __init__( + self, + absolute: bool = True, + retain_stats: bool = True, + mix_prob: float = 0.0, + apply_to_channel: Sequence[int] = (0,), + same_on_batch: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = True, + ) -> None: + super().__init__( + kernel_type="Scharr", + absolute=absolute, + retain_stats=retain_stats, + mix_prob=mix_prob, + apply_to_channel=apply_to_channel, + same_on_batch=same_on_batch, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + +class RandomGaussianBlurGPU(_RandomConvBaseGPU): + """Gaussian blur via separable convolution.""" + + def __init__( + self, + sigma: float = 1.0, + apply_to_channel: Sequence[int] = (0,), + same_on_batch: bool = False, + retain_stats: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + mix_prob: float = 0.0, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = True, + ) -> None: + super().__init__( + kernel_type="GaussianBlur", + sigma=sigma, + apply_to_channel=apply_to_channel, + same_on_batch=same_on_batch, + retain_stats=retain_stats, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + mix_prob=mix_prob, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + +class RandomUnsharpMaskGPU(_RandomConvBaseGPU): + """Unsharp masking: sharpen by subtracting a blurred copy.""" + + def __init__( + self, + sigma: float = 1.0, + unsharp_amount: float = 1.5, + mix_prob: float = 0.0, + apply_to_channel: Sequence[int] = (0,), + same_on_batch: bool = False, + retain_stats: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = True, + ) -> None: + super().__init__( + kernel_type="UnsharpMask", + sigma=sigma, + unsharp_amount=unsharp_amount, + mix_prob=mix_prob, + apply_to_channel=apply_to_channel, + same_on_batch=same_on_batch, + retain_stats=retain_stats, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + +class RandomRandConvGPU(_RandomConvBaseGPU): + """RandConv: convolution with a randomly drawn multi-scale kernel.""" + + def __init__( + self, + kernel_sizes: Sequence[int] = (1, 3, 5, 7), + mix_prob: float = 0.0, + apply_to_channel: Sequence[int] = (0,), + same_on_batch: bool = False, + retain_stats: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = True, + ) -> None: + super().__init__( + kernel_type="RandConv", + kernel_sizes=kernel_sizes, + mix_prob=mix_prob, + apply_to_channel=apply_to_channel, + same_on_batch=same_on_batch, + retain_stats=retain_stats, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + def apply_convolution(img: torch.Tensor, kernel: torch.Tensor, dim: int) -> torch.Tensor: """ Based on https://github.com/pytorch/vision/blob/e3b5d3a8bf5e8636462fd8bce9897bccc690b2a0/torchvision/transforms/_functional_tensor.py#L746 @@ -380,19 +570,17 @@ class RandomGaussianNoiseGPU(ImageOnlyTransform): def __init__( self, mean: float = 0.0, - std: float = 0.1, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + std: float = 1.0, + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default same_on_batch: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.apply_to_channel = apply_to_channel self.mean = mean self.std = std @@ -443,19 +631,17 @@ class RandomBrightnessGPU(ImageOnlyTransform): def __init__( self, - brightness_range: tuple[float, float] = (0.9, 1.1), - apply_to_channel: list[int] | None = None, # Apply to first channel by default + brightness_range: tuple[float, float] = (0.5, 1.5), + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default same_on_batch: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.brightness_range = brightness_range self.apply_to_channel = apply_to_channel self.in_seg = in_seg @@ -494,12 +680,12 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ ## Gamma transform -class RandomGammaGPU(ImageOnlyTransform): +class _RandomGammaBaseGPU(ImageOnlyTransform): """Apply random gamma adjustment to image. If the image is torch Tensor, it is expected to have [N, C, X, Y] or [N, C, X, Y, Z] shape. Args: - gamma_range (tuple of float): Range of gamma multipliers. Default is (0.9, 1.1). + gamma_range (tuple of float): Range of gamma multipliers. Default is (0.7, 1.5). invert_image (bool): If True, invert the image before and after gamma adjustment. Default is False. apply_to_channel (list of int): List of channel indices to apply the gamma adjustment to. Default is [0]. retain_stats (bool): If True, retain the original mean and standard deviation of the image after gamma adjustment. Default is False. @@ -513,21 +699,19 @@ class RandomGammaGPU(ImageOnlyTransform): def __init__( self, - gamma_range: tuple[float, float] = (0.9, 1.1), + gamma_range: tuple[float, float] = (0.7, 1.5), invert_image: bool = False, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default retain_stats: bool = False, same_on_batch: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = False, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.gamma_range = gamma_range self.invert_image = invert_image self.retain_stats = retain_stats @@ -546,7 +730,8 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ channel_data = -input[:, c] if self.invert_image else input[:, c] orig_full = input[:, c].clone() - stats = _channel_stats(channel_data) if self.retain_stats else None + if self.retain_stats: + stats = _channel_stats(channel_data) if self.same_on_batch: gamma = ( @@ -579,7 +764,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ # Apply gamma transform per batch element channel_data = torch.pow(((channel_data - minm) / (rnge + 1e-8)), gamma) * rnge + minm - if stats is not None: + if self.retain_stats: channel_data = _restore_stats(channel_data, stats) if self.invert_image: @@ -592,6 +777,73 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ return input +# Gamma, split so that "gamma" and "inverted gamma" are two config keys rather than +# one key plus an `invert_image` flag. Neither leaf exposes the flag, so a config +# cannot express the same augmentation two ways. + + +class RandomGammaGPU(_RandomGammaBaseGPU): + """Random gamma adjustment.""" + + def __init__( + self, + gamma_range: tuple[float, float] = (0.7, 1.5), + apply_to_channel: Sequence[int] = (0,), + retain_stats: bool = False, + same_on_batch: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = False, + ) -> None: + super().__init__( + gamma_range=gamma_range, + invert_image=False, + apply_to_channel=apply_to_channel, + retain_stats=retain_stats, + same_on_batch=same_on_batch, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + +class RandomInvGammaGPU(_RandomGammaBaseGPU): + """Random gamma adjustment applied to the inverted image.""" + + def __init__( + self, + gamma_range: tuple[float, float] = (0.7, 1.5), + apply_to_channel: Sequence[int] = (0,), + retain_stats: bool = False, + same_on_batch: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = False, + ) -> None: + super().__init__( + gamma_range=gamma_range, + invert_image=True, + apply_to_channel=apply_to_channel, + retain_stats=retain_stats, + same_on_batch=same_on_batch, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + ## nnunetv2 contrast transform class RandomContrastGPU(ImageOnlyTransform): """Apply random gamma adjustment to image. @@ -611,20 +863,18 @@ class RandomContrastGPU(ImageOnlyTransform): def __init__( self, - contrast_range: tuple[float, float] = (0.9, 1.1), - apply_to_channel: list[int] | None = None, # Apply to first channel by default + contrast_range: tuple[float, float] = (0.75, 1.25), + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default retain_stats: bool = False, same_on_batch: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.contrast_range = contrast_range self.apply_to_channel = apply_to_channel self.retain_stats = retain_stats @@ -640,7 +890,8 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ for c in self.apply_to_channel: channel_data = input[:, c] # [N, ...spatial...] orig = channel_data.clone() - stats = _channel_stats(channel_data) if self.retain_stats else None + if self.retain_stats: + stats = _channel_stats(channel_data) if self.same_on_batch: factor = ( @@ -661,7 +912,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ mean = x[i].mean() x[i] = (x[i] - mean) * factor[i] + mean - if stats is not None: + if self.retain_stats: x = _restore_stats(x, stats) checked = _select_and_check(self, orig, x, seg_mask) if checked is None: @@ -672,7 +923,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ ## Function transform -class RandomFunctionGPU(ImageOnlyTransform): +class _RandomFunctionBaseGPU(ImageOnlyTransform): """Apply function to the image based on probability. If the image is torch Tensor, it is expected to have [N, C, X, Y] or [N, C, X, Y, Z] shape. @@ -691,19 +942,17 @@ class RandomFunctionGPU(ImageOnlyTransform): def __init__( self, func: Callable[[Tensor], Tensor] = lambda x: x**2, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default retain_stats: bool = False, same_on_batch: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.func = func self.retain_stats = retain_stats self.apply_to_channel = apply_to_channel @@ -719,7 +968,8 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ for c in self.apply_to_channel: x = input[:, c] # shape [N, ...spatial...] orig = x.clone() - stats = _channel_stats(x) if self.retain_stats else None + if self.retain_stats: + stats = _channel_stats(x) # Normalize to make values >=0, per sample. # @@ -736,7 +986,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ # Apply function x = self.func(x) - if stats is not None: + if self.retain_stats: x = _restore_stats(x, stats) checked = _select_and_check(self, orig, x, seg_mask) if checked is None: @@ -746,6 +996,88 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ return input +# One class per elementwise function. +# +# `func` was a callable parameter, which no JSON config could ever express -- the old +# ladder worked around that by expanding a single "FunctionTransform" block into five +# transforms from a hardcoded lambda list. A class each makes every one addressable +# from a config, and removes the un-serialisable parameter entirely. +# +# Written out longhand rather than as torch.log1p / torch.sigmoid on purpose: those +# differ from the originals in the last ulp, which is enough to move the seeded +# determinism hashes and invalidate every published experiment. + + +def _log1p(x: Tensor) -> Tensor: + return torch.log(1 + x) + + +def _sigmoid(x: Tensor) -> Tensor: + return 1 / (1 + torch.exp(-x)) + + +class _RandomNamedFunctionGPU(_RandomFunctionBaseGPU): + """Shared constructor for the fixed-function leaves. Not registered itself.""" + + #: Set by each leaf; `func` is therefore absent from the config surface. + function: staticmethod + + def __init__( + self, + apply_to_channel: Sequence[int] = (0,), + retain_stats: bool = False, + same_on_batch: bool = False, + in_seg: float = 0.0, + out_seg: float = 0.0, + mix_in_out: bool = False, + p: float = 1.0, + p_batch: float = 1.0, + keepdim: bool = True, + ) -> None: + super().__init__( + func=type(self).function, + apply_to_channel=apply_to_channel, + retain_stats=retain_stats, + same_on_batch=same_on_batch, + in_seg=in_seg, + out_seg=out_seg, + mix_in_out=mix_in_out, + p=p, + p_batch=p_batch, + keepdim=keepdim, + ) + + +class RandomLog1pGPU(_RandomNamedFunctionGPU): + """Apply log(1 + x).""" + + function = staticmethod(_log1p) + + +class RandomSqrtGPU(_RandomNamedFunctionGPU): + """Apply sqrt(x).""" + + function = staticmethod(torch.sqrt) + + +class RandomSinGPU(_RandomNamedFunctionGPU): + """Apply sin(x).""" + + function = staticmethod(torch.sin) + + +class RandomExpGPU(_RandomNamedFunctionGPU): + """Apply exp(x).""" + + function = staticmethod(torch.exp) + + +class RandomSigmoidGPU(_RandomNamedFunctionGPU): + """Apply the logistic sigmoid 1 / (1 + exp(-x)).""" + + function = staticmethod(_sigmoid) + + ## Inverse transform class RandomInverseGPU(ImageOnlyTransform): """Inverse image based on probability. @@ -763,7 +1095,7 @@ class RandomInverseGPU(ImageOnlyTransform): def __init__( self, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default retain_stats: bool = False, same_on_batch: bool = False, in_seg: float = 0.0, @@ -771,12 +1103,10 @@ def __init__( mix_in_out: bool = False, mix_prob: float = 0.0, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.apply_to_channel = apply_to_channel self.retain_stats = retain_stats self.in_seg = in_seg @@ -811,14 +1141,10 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ alpha = torch.rand(1, device=input.device) x = alpha * orig + (1 - alpha) * x - if seg_mask is not None: - region_mode = _choose_region_mode(self.in_seg, self.out_seg, seg_mask[i]) - x = _apply_region_mode(orig, x, seg_mask[i], region_mode, mix_in_out=self.mix_in_out) - # Final safety: check if nan/inf appeared - if torch.isnan(x).any() or torch.isinf(x).any(): - print(f"Warning nan: {self.__class__.__name__}", flush=True) + checked = _select_and_check(self, orig, x, None if seg_mask is None else seg_mask[i]) + if checked is None: continue - input[i, c] = x + input[i, c] = checked return input @@ -841,7 +1167,7 @@ class RandomHistogramEqualizationGPU(ImageOnlyTransform): def __init__( self, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default retain_stats: bool = False, same_on_batch: bool = False, in_seg: float = 0.0, @@ -849,12 +1175,10 @@ def __init__( mix_in_out: bool = False, mix_prob: float = 0.0, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.retain_stats = retain_stats self.apply_to_channel = apply_to_channel self.in_seg = in_seg @@ -870,12 +1194,13 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ for c in self.apply_to_channel: # `.clone()`, not the bare `input[:, c]` view this used to take: the loop # below assigns into `channel_data[b]`, which through a view writes straight - # into `input`. The non-finite guard at the bottom would then `continue` - # over values that were already in the batch -- the guard skipped nothing. + # into `input`. The NaN guard at the bottom would then `continue` over + # values that were already in the batch -- the guard skipped nothing. channel_data = input[:, c].clone() # shape [N, ...spatial...] orig = channel_data.clone() - stats = _channel_stats(channel_data) if self.retain_stats else None + if self.retain_stats: + stats = _channel_stats(channel_data) # Process each batch element independently batch_size = channel_data.shape[0] @@ -908,7 +1233,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ alpha = torch.rand(1, device=input.device) channel_data[b] = alpha * orig[b] + (1 - alpha) * channel_data[b] - if stats is not None: + if self.retain_stats: channel_data = _restore_stats(channel_data, stats) checked = _select_and_check(self, orig, channel_data, seg_mask) @@ -944,7 +1269,7 @@ def __init__( self, coefficients: Union[float, tuple[float, float]] = 0.5, order: int = 3, - apply_to_channel: list[int] | None = None, + apply_to_channel: Sequence[int] = (0,), invert: bool = False, retain_stats: bool = False, in_seg: float = 0.0, @@ -952,12 +1277,10 @@ def __init__( mix_in_out: bool = False, same_on_batch: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) if isinstance(coefficients, (int, float)): self.coeff_range = (-float(coefficients), float(coefficients)) elif isinstance(coefficients, (tuple, list)) and len(coefficients) == 2: @@ -1120,20 +1443,18 @@ class RandomClampGPU(ImageOnlyTransform): def __init__( self, - max_clamp_amount: float = 0.2, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + max_clamp_amount: float = 0.0, + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default retain_stats: bool = False, same_on_batch: bool = False, in_seg: float = 0.0, out_seg: float = 0.0, mix_in_out: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.max_clamp_amount = max_clamp_amount self.apply_to_channel = apply_to_channel self.retain_stats = retain_stats @@ -1149,7 +1470,8 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ for c in self.apply_to_channel: channel_data = input[:, c] # [N, ...spatial...] orig = channel_data.clone() - stats = _channel_stats(channel_data) if self.retain_stats else None + if self.retain_stats: + stats = _channel_stats(channel_data) if self.same_on_batch: min_percentile = torch.rand(1, device=input.device, dtype=input.dtype) * self.max_clamp_amount @@ -1168,7 +1490,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ max_val = torch.quantile(x[i].flatten(), max_percentile) x[i] = torch.clamp(x[i], min_val, max_val) - if stats is not None: + if self.retain_stats: x = _restore_stats(x, stats) checked = _select_and_check(self, orig, x, seg_mask) if checked is None: @@ -1189,16 +1511,14 @@ class ZscoreNormalizationGPU(ImageOnlyTransform): def __init__( self, - apply_to_channel: list[int] | None = None, + apply_to_channel: Sequence[int] = (0,), keepdim: bool = True, in_seg: float = 0.0, out_seg: float = 0.0, p: float = 1.0, - **kwargs, + p_batch: float = 1.0, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=False, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=False, keepdim=keepdim) self.apply_to_channel = apply_to_channel self.in_seg = in_seg self.out_seg = out_seg @@ -1228,13 +1548,23 @@ def apply_transform( # use unbiased=False for stability, and clamp std to avoid division by ~0 std = channel.std(dim=reduce_dims, keepdim=True, unbiased=False).clamp_min(1e-8) channel = (channel - mean) / std - if seg_mask is not None: - region_mode = _choose_region_mode(self.in_seg, self.out_seg, seg_mask) - channel = _apply_region_mode(orig, channel, seg_mask, region_mode) - # Final safety: check if nan/inf appeared - if torch.isnan(channel).any() or torch.isinf(channel).any(): - print(f"Warning nan: {self.__class__.__name__}", flush=True) + # No mix_in_out here: z-scoring is applied whole, never to a random subset + # of the mask channels. + checked = _select_and_check(self, orig, channel, seg_mask) + if checked is None: continue - input[:, c] = channel + input[:, c] = checked return input + + +# --- temporary bridges for the hand-written pipelines ------------------------------ +# +# The three `if` ladders in gpu/transforms.py, gpu/transforms_list.py and +# cpu/transforms.py still read `kernel_type` and `func` out of the config and pass them +# in, which is exactly the dispatch the leaf classes above exist to remove. They are +# replaced by the registry-driven builder later in this series; until then these +# aliases keep them working without a second rewrite. Do not use them in new code. +RandomConvTransformGPU = _RandomConvBaseGPU +RandomFunctionGPU = _RandomFunctionBaseGPU +_RandomGammaWithInvertGPU = _RandomGammaBaseGPU diff --git a/smauglab/transforms/gpu/fromSeg.py b/smauglab/transforms/gpu/fromSeg.py index e4f5917..71ed118 100644 --- a/smauglab/transforms/gpu/fromSeg.py +++ b/smauglab/transforms/gpu/fromSeg.py @@ -1,3 +1,4 @@ +from collections.abc import Sequence from typing import Any import torch @@ -27,16 +28,14 @@ def _kmeans_1d(values: torch.Tensor, C: int, n_iter: int = 10) -> torch.Tensor: def _gaussian_blur_3d(x: torch.Tensor, sigma: float) -> torch.Tensor: - """Separable 3D Gaussian blur of a (B, 1, D, H, W) volume, clamped to [0, 1]. + """Separable 3D Gaussian blur of a [0, 1] volume. x: (B, 1, D, H, W). - Delegates to the shared implementation, which pads with `reflect`. This copy used - conv3d's implicit zero padding, which pulled the volume border towards 0 -- the - clamp below hid the top end of that but not the darkening. It also took its radius - from `round(3*sigma)` rather than `ceil`, so kernels can be one tap wider now. + The blur itself is the shared one; this wrapper only keeps the clamp, which is a + no-op guard for the already-normalised `synth_01` inputs. The previous local copy + zero-padded (via conv3d's `padding=`) rather than reflecting, which darkened the + volume border. """ - if sigma <= 0: - return x - return gaussian_blur3d(x, float(sigma)).clamp(0, 1) + return gaussian_blur3d(x, sigma).clamp(0, 1) def _voronoi_region_ids( @@ -45,7 +44,7 @@ def _voronoi_region_ids( fg: torch.Tensor, C: int, device: torch.device, - s_choices: list[int], + s_choices: Sequence[int], skip_sub_parc_prob: float, ) -> tuple[torch.Tensor, int]: """Spatially subdivide each K-means cluster into Voronoi sub-regions. @@ -100,22 +99,16 @@ class RandomRedistributeSegGPU(ImageOnlyTransform): def __init__( self, in_seg: float = 0.2, - apply_to_channel: list[int] | None = None, + apply_to_channel: Sequence[int] = (0,), retain_stats: bool = False, same_on_batch: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - std_noise_range: list[float] | None = None, - dilation_iterations_range: list[int] | None = None, - **kwargs, + std_noise_range: Sequence[float] = (0.1, 0.3), + dilation_iterations_range: Sequence[int] = (1, 3), ) -> None: - if dilation_iterations_range is None: - dilation_iterations_range = [1, 3] - if std_noise_range is None: - std_noise_range = [0.1, 0.3] - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.in_seg = in_seg self.apply_to_channel = apply_to_channel self.retain_stats = retain_stats @@ -265,7 +258,7 @@ def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[ return input -class RandomPALETTEGPU(ImageOnlyTransform): +class RandomPaletteGPU(ImageOnlyTransform): """ SmaugLab GPU augmentation implementing PALETTE synthesis. @@ -301,29 +294,27 @@ class RandomPALETTEGPU(ImageOnlyTransform): def __init__( self, - c_choices: list[int] | None = None, - s_choices: list[int] | None = None, - blur_sigmas: list[float] | None = None, + c_choices: Sequence[int] = (2, 3, 4, 5, 6), + s_choices: Sequence[int] = (2, 3, 4, 5, 6, 7, 8, 9, 10), + blur_sigmas: Sequence[float] = (0.0, 0.0, 0.0, 0.3, 0.5, 0.8), dark_threshold: float = 0.01, n_kmeans_subsample: int = 10_000, skip_parcellation_prob: float = 0.10, skip_sub_parc_prob: float = 0.40, - alpha_magnitude_range: list[float] | None = None, + alpha_magnitude_range: Sequence[float] = (0.5, 2.0), label_remap_prob: float = 0.5, min_label_voxels: int = 4, label_classes: list[int] | None = None, p: float = 1.0, - **kwargs: Any, + p_batch: float = 1.0, + same_on_batch: bool = False, + # Note the default is False, not the True its siblings use. This class + # previously forwarded **kwargs straight to super(), so keepdim fell through + # to kornia's own default -- and nothing ever passed it. Spelling that out + # rather than "fixing" it keeps the transform behaving exactly as before. + keepdim: bool = False, ) -> None: - if alpha_magnitude_range is None: - alpha_magnitude_range = [0.5, 2.0] - if blur_sigmas is None: - blur_sigmas = [0.0, 0.0, 0.0, 0.3, 0.5, 0.8] - if s_choices is None: - s_choices = [2, 3, 4, 5, 6, 7, 8, 9, 10] - if c_choices is None: - c_choices = [2, 3, 4, 5, 6] - super().__init__(p=p, **kwargs) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.c_choices = c_choices self.s_choices = s_choices self.blur_sigmas = blur_sigmas diff --git a/smauglab/transforms/gpu/spatial.py b/smauglab/transforms/gpu/spatial.py index 932900d..9d3cbd6 100644 --- a/smauglab/transforms/gpu/spatial.py +++ b/smauglab/transforms/gpu/spatial.py @@ -22,7 +22,7 @@ # Affine transform -class RandomAffine3DCustom(RigidAffineAugmentationBase3D): +class RandomAffineGPU(RigidAffineAugmentationBase3D): r"""Apply affine transformation 3D volumes (5D tensor). Based on :class:`kornia.augmentation.RandomAffine3D`. @@ -107,9 +107,9 @@ def __init__( tuple[float, float], tuple[float, float, float], tuple[tuple[float, float], tuple[float, float], tuple[float, float]], - ], - translate: Union[Tensor, tuple[float, float, float]] | None = None, - scale: Union[Tensor, tuple[float, float], tuple[tuple[float, float], tuple[float, float], tuple[float, float]]] | None = None, + ] = 10, + translate: Union[Tensor, tuple[float, float, float]] | None = (0.1, 0.1, 0.1), + scale: Union[Tensor, tuple[float, float], tuple[tuple[float, float], tuple[float, float], tuple[float, float]]] | None = (0.9, 1.1), shears: Union[ Tensor, float, @@ -124,14 +124,15 @@ def __init__( tuple[float, float], ], None, - ] = None, + ] = (-10, 10, -10, 10, -10, 10), resample: Union[str, int, Resample] = Resample.BILINEAR.name, same_on_batch: bool = False, align_corners: bool = True, p: float = 0.5, + p_batch: float = 1.0, keepdim: bool = True, ) -> None: - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.degrees = degrees self.shears = shears self.translate = translate @@ -209,10 +210,10 @@ def __init__( scale: tuple[float, float] = (0.3, 1.0), same_on_batch: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self._param_generator = ScaleGenerator3D(scale=scale) def compute_transformation(self, input: Tensor, params: dict[str, Tensor], flags: dict[str, Any]) -> Tensor: @@ -355,19 +356,19 @@ class RandomAcqTransformGPU(ImageOnlyTransform): def __init__( self, scale: tuple[float, float] = (0.3, 1.0), - one_dim: bool = False, same_on_batch: bool = False, - apply_to_channel: list[int] | None = None, # Apply to first channel by default + apply_to_channel: Sequence[int] = (0,), # Apply to first channel by default p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - if apply_to_channel is None: - apply_to_channel = [0] - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self.flags = {"resample": "trilinear"} self.apply_to_channel = apply_to_channel - self._param_generator = ScaleGenerator3D(scale=scale, one_dim=one_dim) + # one_dim is fixed rather than exposed: this class *is* the single-axis case, + # and RandomLowResTransformGPU is the isotropic one. Leaving it configurable + # meant two config keys could each produce either behaviour. + self._param_generator = ScaleGenerator3D(scale=scale, one_dim=True) @torch.no_grad() def apply_transform(self, input: Tensor, params: dict[str, Tensor], flags: dict[str, Any], transform: Tensor | None = None) -> Tensor: @@ -442,13 +443,13 @@ def __init__( self, # Both forms are accepted and normalised below; the annotation said `int` # while the default was a list and every caller passes a list. - flip_axis: Union[int, Sequence[int]] = [0, 1, 2], + flip_axis: Union[int, Sequence[int]] = (0,), same_on_batch: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) # normalize flip_axis into a list of ints if isinstance(flip_axis, int): self.flip_axis = [flip_axis] @@ -574,13 +575,13 @@ def __init__( # A (low, high) range like `crop`, not a per-axis triple: CropGenerator3D # feeds it to _tuple_range_reader(..., 3, ...), which broadcasts the range # across all three axes. The annotation said triple, the default was a pair. - pos: tuple[float, float] = (0.5, 1.0), # Fraction of the pos + pos: tuple[float, float] = (0.0, 1.0), # Fraction of the pos same_on_batch: bool = False, p: float = 1.0, + p_batch: float = 1.0, keepdim: bool = True, - **kwargs, ) -> None: - super().__init__(p=p, same_on_batch=same_on_batch, keepdim=keepdim) + super().__init__(p=p, p_batch=p_batch, same_on_batch=same_on_batch, keepdim=keepdim) self._param_generator = CropGenerator3D(crop=crop, pos=pos) def compute_transformation(self, input: Tensor, params: dict[str, Tensor], flags: dict[str, Any]) -> Tensor: diff --git a/smauglab/transforms/gpu/transforms.py b/smauglab/transforms/gpu/transforms.py index 05d1ca6..4424ce6 100644 --- a/smauglab/transforms/gpu/transforms.py +++ b/smauglab/transforms/gpu/transforms.py @@ -14,17 +14,17 @@ RandomContrastGPU, RandomConvTransformGPU, RandomFunctionGPU, - RandomGammaGPU, RandomGaussianNoiseGPU, RandomHistogramEqualizationGPU, RandomInverseGPU, ZscoreNormalizationGPU, + _RandomGammaWithInvertGPU, ) from smauglab.transforms.gpu.domain_transfer import RandomDomainTransferGPU -from smauglab.transforms.gpu.fromSeg import RandomPALETTEGPU, RandomRedistributeSegGPU +from smauglab.transforms.gpu.fromSeg import RandomPaletteGPU, RandomRedistributeSegGPU from smauglab.transforms.gpu.spatial import ( RandomAcqTransformGPU, - RandomAffine3DCustom, + RandomAffineGPU, RandomCropTransformGPU, RandomFlipTransformGPU, RandomLowResTransformGPU, @@ -75,7 +75,7 @@ def _build_transforms(self) -> list[TransformType]: affine_params = self.transform_params.get("AffineTransform") if affine_params is not None: transforms.append( - RandomAffine3DCustom( + RandomAffineGPU( degrees=affine_params.get("degrees", 10), translate=affine_params.get("translate", [0.1, 0.1, 0.1]), scale=affine_params.get("scale", [0.9, 1.1]), @@ -105,7 +105,7 @@ def _build_transforms(self) -> list[TransformType]: palette_params = self.transform_params.get("RandomPALETTETransform") if palette_params is not None: transforms.append( - RandomPALETTEGPU( + RandomPaletteGPU( p=palette_params.get("probability", 1.0), c_choices=palette_params.get("c_choices", [2, 3, 4, 5, 6]), s_choices=palette_params.get("s_choices", [2, 3, 4, 5, 6, 7, 8, 9, 10]), @@ -298,7 +298,7 @@ def _build_transforms(self) -> list[TransformType]: gamma_params = self.transform_params.get("GammaTransform") if gamma_params is not None: transforms.append( - RandomGammaGPU( + _RandomGammaWithInvertGPU( gamma_range=gamma_params.get("gamma_range", [0.7, 1.5]), p=gamma_params.get("probability", 0), invert_image=False, @@ -312,7 +312,7 @@ def _build_transforms(self) -> list[TransformType]: inv_gamma_params = self.transform_params.get("InvGammaTransform") if inv_gamma_params is not None: transforms.append( - RandomGammaGPU( + _RandomGammaWithInvertGPU( gamma_range=inv_gamma_params.get("gamma_range", [0.7, 1.5]), p=inv_gamma_params.get("probability", 0), in_seg=inv_gamma_params.get("in_seg", 0.0), @@ -376,7 +376,6 @@ def _build_transforms(self) -> list[TransformType]: RandomAcqTransformGPU( p=acq_params.get("probability", 0), scale=acq_params.get("scale", [0.3, 1.0]), - one_dim=True, same_on_batch=acq_params.get("same_on_batch", False), ) ) diff --git a/smauglab/transforms/gpu/transforms_list.py b/smauglab/transforms/gpu/transforms_list.py index 35f08f8..1f64f7c 100644 --- a/smauglab/transforms/gpu/transforms_list.py +++ b/smauglab/transforms/gpu/transforms_list.py @@ -14,14 +14,14 @@ RandomContrastGPU, RandomConvTransformGPU, RandomFunctionGPU, - RandomGammaGPU, RandomGaussianNoiseGPU, RandomHistogramEqualizationGPU, RandomInverseGPU, ZscoreNormalizationGPU, + _RandomGammaWithInvertGPU, ) from smauglab.transforms.gpu.fromSeg import RandomRedistributeSegGPU -from smauglab.transforms.gpu.spatial import RandomAcqTransformGPU, RandomAffine3DCustom, RandomFlipTransformGPU, RandomLowResTransformGPU +from smauglab.transforms.gpu.spatial import RandomAcqTransformGPU, RandomAffineGPU, RandomFlipTransformGPU, RandomLowResTransformGPU class AugTransformsGPURandomOrder(AugmentationSequentialCustom): @@ -67,7 +67,7 @@ def _build_transforms(self) -> list[TransformType]: affine_params = self.transform_params.get("AffineTransform") if affine_params is not None: transforms.append( - RandomAffine3DCustom( + RandomAffineGPU( degrees=affine_params.get("degrees", 10), translate=affine_params.get("translate", [0.1, 0.1, 0.1]), scale=affine_params.get("scale", [0.9, 1.1]), @@ -263,7 +263,7 @@ def _build_transforms(self) -> list[TransformType]: gamma_params = self.transform_params.get("GammaTransform") if gamma_params is not None: ge_transforms.append( - RandomGammaGPU( + _RandomGammaWithInvertGPU( gamma_range=gamma_params.get("gamma_range", [0.7, 1.5]), p=gamma_params.get("probability", 0), invert_image=False, @@ -277,7 +277,7 @@ def _build_transforms(self) -> list[TransformType]: inv_gamma_params = self.transform_params.get("InvGammaTransform") if inv_gamma_params is not None: ge_transforms.append( - RandomGammaGPU( + _RandomGammaWithInvertGPU( gamma_range=inv_gamma_params.get("gamma_range", [0.7, 1.5]), p=inv_gamma_params.get("probability", 0), in_seg=inv_gamma_params.get("in_seg", 0.0), @@ -309,7 +309,6 @@ def _build_transforms(self) -> list[TransformType]: RandomLowResTransformGPU( p=lowres_params.get("probability", 0), scale=lowres_params.get("scale", [0.3, 1.0]), - crop=lowres_params.get("crop", [1.0, 1.0]), same_on_batch=lowres_params.get("same_on_batch", False), ) ) @@ -320,8 +319,6 @@ def _build_transforms(self) -> list[TransformType]: RandomAcqTransformGPU( p=acq_params.get("probability", 0), scale=acq_params.get("scale", [0.3, 1.0]), - crop=acq_params.get("crop", [1.0, 1.0]), - one_dim=True, same_on_batch=acq_params.get("same_on_batch", False), ) ) @@ -395,7 +392,7 @@ def _build_transforms(self) -> list[TransformType]: affine_params = self.transform_params.get("AffineTransform") if affine_params is not None: transforms.append( - RandomAffine3DCustom( + RandomAffineGPU( degrees=affine_params.get("degrees", 10), translate=affine_params.get("translate", [0.1, 0.1, 0.1]), scale=affine_params.get("scale", [0.9, 1.1]), @@ -600,7 +597,7 @@ def _build_transforms(self) -> list[TransformType]: gamma_params = self.transform_params.get("GammaTransform") if gamma_params is not None: transforms.append( - RandomGammaGPU( + _RandomGammaWithInvertGPU( gamma_range=gamma_params.get("gamma_range", [0.7, 1.5]), p=gamma_params.get("probability", 0), invert_image=False, @@ -614,7 +611,7 @@ def _build_transforms(self) -> list[TransformType]: inv_gamma_params = self.transform_params.get("InvGammaTransform") if inv_gamma_params is not None: transforms.append( - RandomGammaGPU( + _RandomGammaWithInvertGPU( gamma_range=inv_gamma_params.get("gamma_range", [0.7, 1.5]), p=inv_gamma_params.get("probability", 0), in_seg=inv_gamma_params.get("in_seg", 0.0), @@ -646,7 +643,6 @@ def _build_transforms(self) -> list[TransformType]: RandomLowResTransformGPU( p=lowres_params.get("probability", 0), scale=lowres_params.get("scale", [0.3, 1.0]), - crop=lowres_params.get("crop", [1.0, 1.0]), same_on_batch=lowres_params.get("same_on_batch", False), ) ) @@ -657,8 +653,6 @@ def _build_transforms(self) -> list[TransformType]: RandomAcqTransformGPU( p=acq_params.get("probability", 0), scale=acq_params.get("scale", [0.3, 1.0]), - crop=acq_params.get("crop", [1.0, 1.0]), - one_dim=True, same_on_batch=acq_params.get("same_on_batch", False), ) ) diff --git a/smauglab/transforms/synthseg/README.md b/smauglab/transforms/synthseg/README.md index 2acabc7..67d225f 100644 --- a/smauglab/transforms/synthseg/README.md +++ b/smauglab/transforms/synthseg/README.md @@ -65,7 +65,7 @@ Defaults match `BrainGenerator.__init__` (which overrides several - **3D only** (5D `(B, C, D, H, W)` tensors), matching SmaugLab's GPU transforms. - Affine transforms are applied **about the volume centre** (like SmaugLab's - `RandomAffine3DCustom`), rather than the corner-origin used by neuron's + `RandomAffineGPU`), rather than the corner-origin used by neuron's `affine_to_shift`. This keeps the anatomy in frame and is the standard choice; the visual augmentation is equivalent. - The SVF is integrated at full resolution after upsampling the coarse velocity @@ -156,11 +156,11 @@ the deformed labels: ```python from smauglab.transforms.gpu.base import AugmentationSequentialCustom -from smauglab.transforms.gpu.spatial import RandomAffine3DCustom +from smauglab.transforms.gpu.spatial import RandomAffineGPU from smauglab.transforms.synthseg import RandomSynthSegGPU aug = AugmentationSequentialCustom( - RandomAffine3DCustom(degrees=15, scale=[0.8, 1.2], p=1.0), + RandomAffineGPU(degrees=15, scale=[0.8, 1.2], p=1.0), RandomSynthSegGPU(generation_labels=None, n_channels=1, bias_field_std=0.7, gamma_std=0.5, randomise_res=True, p=1.0), data_keys=["input", "mask"], same_on_batch=True, diff --git a/smauglab/transforms/synthseg/functional.py b/smauglab/transforms/synthseg/functional.py index 6a8ff49..8092f9e 100644 --- a/smauglab/transforms/synthseg/functional.py +++ b/smauglab/transforms/synthseg/functional.py @@ -20,7 +20,7 @@ the ``(x, y, z) = (W, H, D)`` order expected by ``F.grid_sample`` at the very end, with ``align_corners=True`` so that integer voxel indices map exactly. * Affine transforms are applied about the volume centre (standard practice and - matching SmaugLab's existing ``RandomAffine3DCustom``), so small rotations / + matching SmaugLab's existing ``RandomAffineGPU``), so small rotations / scalings keep the anatomy in frame. """ diff --git a/smauglab/transforms/synthseg/transforms.py b/smauglab/transforms/synthseg/transforms.py index ac31632..af8f381 100644 --- a/smauglab/transforms/synthseg/transforms.py +++ b/smauglab/transforms/synthseg/transforms.py @@ -5,7 +5,7 @@ * :class:`RandomSynthSegGPU` -- an :class:`ImageOnlyTransform` that *replaces* the image with a GMM-synthesised one derived from ``params['seg']``. It is intensity-only (no internal spatial deformation), so it composes with SmaugLab's - existing geometric transforms (``RandomAffine3DCustom``, ``RandomFlipTransformGPU``, + existing geometric transforms (``RandomAffineGPU``, ``RandomFlipTransformGPU``, ...) inside an :class:`AugmentationSequentialCustom`: place those *before* it so the mask is deformed first and SynthSeg generates from the deformed labels, keeping image and label aligned. Drop it into a ``transform_params_gpu.json`` diff --git a/unit_tests/test_bucket_and_rng.py b/unit_tests/test_bucket_and_rng.py index f4ce81c..1180b0c 100644 --- a/unit_tests/test_bucket_and_rng.py +++ b/unit_tests/test_bucket_and_rng.py @@ -12,7 +12,7 @@ import torch -from smauglab.transforms.gpu.contrast import RandomConvTransformGPU +from smauglab.transforms.gpu.contrast import RandomLaplaceGPU, RandomRandConvGPU from smauglab.transforms.gpu.spatial import RandomLowResTransformGPU from smauglab.transforms.gpu.transforms_list import RandomChooseXTransformsGPU from smauglab.transforms.rng import shared_choice, shared_rand @@ -76,7 +76,7 @@ def test_it_still_returns_something_transformed(self): """Cloning must not turn the bucket into a no-op.""" torch.manual_seed(0) bucket = RandomChooseXTransformsGPU( - transforms_list=[RandomConvTransformGPU(kernel_type="Laplace", p=1.0)], + transforms_list=[RandomLaplaceGPU(p=1.0)], num_transforms=1, p=1.0, same_on_batch=False, @@ -119,10 +119,10 @@ def test_randconv_is_reproducible_under_torch_seed_alone(self): outputs = [] for _ in range(2): torch.manual_seed(1234) - transform = RandomConvTransformGPU(kernel_type="RandConv", p=1.0, kernel_sizes=[1, 3, 5, 7]) + transform = RandomRandConvGPU(p=1.0, kernel_sizes=[1, 3, 5, 7]) outputs.append(transform.apply_transform(self.tiny_volume(), {}, {}, transform=None).clone()) - self.assertTrue(torch.equal(outputs[0], outputs[1]), "RandomConvTransformGPU drew its kernel size from an unseeded generator") + self.assertTrue(torch.equal(outputs[0], outputs[1]), "RandomRandConvGPU drew its kernel size from an unseeded generator") def test_shared_choice_covers_the_whole_sequence(self): torch.manual_seed(0) diff --git a/unit_tests/test_region_and_stats.py b/unit_tests/test_region_and_stats.py index c2c1d6e..e08e7f0 100644 --- a/unit_tests/test_region_and_stats.py +++ b/unit_tests/test_region_and_stats.py @@ -18,9 +18,9 @@ from smauglab.transforms.gpu.base import AugmentationSequentialCustom from smauglab.transforms.gpu.contrast import ( - RandomConvTransformGPU, - RandomFunctionGPU, RandomHistogramEqualizationGPU, + RandomScharrGPU, + RandomSqrtGPU, _apply_region_mode, _foreground, ) @@ -123,7 +123,7 @@ def test_in_seg_confines_a_scharr_transform_to_the_mask(self): # Driven through the container, as AugTransformsGPU does: that is what routes # the mask into params["seg"], which is where _apply_region_mode reads it. pipeline = AugmentationSequentialCustom( - RandomConvTransformGPU(kernel_type="Scharr", p=1.0, in_seg=1.0, out_seg=0.0, mix_prob=0.0), + RandomScharrGPU(p=1.0, in_seg=1.0, out_seg=0.0, mix_prob=0.0), data_keys=["input", "mask"], same_on_batch=True, ) @@ -145,7 +145,7 @@ class TestFunctionTransformIsPerSample(SmaugLabTestCase): def _run(self, volume: torch.Tensor) -> torch.Tensor: torch.manual_seed(0) - transform = RandomFunctionGPU(func=torch.sqrt, p=1.0) + transform = RandomSqrtGPU(p=1.0) return transform.apply_transform(volume.clone(), {}, {}, transform=None) def test_a_volume_is_augmented_the_same_alone_and_in_a_batch(self): diff --git a/unit_tests/test_spatial_sampling.py b/unit_tests/test_spatial_sampling.py index 592acf2..268f112 100644 --- a/unit_tests/test_spatial_sampling.py +++ b/unit_tests/test_spatial_sampling.py @@ -16,7 +16,7 @@ from kornia.constants import Resample from smauglab.transforms.gpu.base import AugmentationSequentialCustom -from smauglab.transforms.gpu.spatial import CropGenerator3D, RandomAffine3DCustom, RandomFlipTransformGPU, ScaleGenerator3D +from smauglab.transforms.gpu.spatial import CropGenerator3D, RandomAffineGPU, RandomFlipTransformGPU, ScaleGenerator3D from unit_tests.helpers import SmaugLabTestCase, first_output @@ -152,7 +152,7 @@ def test_the_crop_generator_neutralises_position_at_the_centre(self): class TestMaskResampleRestore(SmaugLabTestCase): - """Guard tests for `RandomAffine3DCustom.apply_transform_mask`. + """Guard tests for `RandomAffineGPU.apply_transform_mask`. Unlike the rest of this file these pass before the change too: `resample_method` was annotated but assigned only inside the `if`, so the restore below it could @@ -163,7 +163,7 @@ class TestMaskResampleRestore(SmaugLabTestCase): """ def _transform(self): - return RandomAffine3DCustom(p=1.0, degrees=5, align_corners=True) + return RandomAffineGPU(p=1.0, degrees=5, align_corners=True) def _flags(self, resample: str = "bilinear"): transform = self._transform() diff --git a/unit_tests/test_transforms_gpu.py b/unit_tests/test_transforms_gpu.py index f391ba7..9a5a0b1 100644 --- a/unit_tests/test_transforms_gpu.py +++ b/unit_tests/test_transforms_gpu.py @@ -39,6 +39,10 @@ def discover_transforms(): for name, obj in vars(module).items(): if not inspect.isclass(obj) or obj.__module__ != module_name: continue + if name.startswith("_"): + # Shared base classes (_RandomConvBaseGPU, _RandomNamedFunctionGPU, ...). + # They are not augmentations; their concrete leaves are discovered instead. + continue if not issubclass(obj, torch.nn.Module) or name in NOT_A_TRANSFORM: continue signature = inspect.signature(obj.__init__)