diff --git a/.gitignore b/.gitignore index 6f81433b..48d7e7c7 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ # Byte-compiled / optimized / DLL files __pycache__/ +notebooks/medical_seg/data/ *.py[cod] *$py.class *.vscode @@ -7,13 +8,18 @@ __pycache__ *.pyc *.idea *.DS_Store - +tmp* +tests/integration/ +zz* # data which is automatically downloaded inside notebooks notebooks/ExBox3 +notebooks/MedNIST.tar.gz +notebooks/MedNIST/ +notebooks/demo_classification.py # C extensions *.so - +notebooks/MedNIST # Distribution / packaging .Python build/ diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 0e405be8..6e698572 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,5 +1,5 @@ default_language_version: - python: python3.8 + python: python3.9 ci: autofix_prs: true @@ -19,6 +19,7 @@ repos: - id: pretty-format-json - id: check-added-large-files exclude: .*\.ipynb + args: ['--maxkb=1000'] - id: check-docstring-first - id: detect-private-key @@ -35,7 +36,7 @@ repos: - id: yesqa - repo: https://github.com/psf/black - rev: 21.7b0 + rev: 22.3.0 hooks: - id: black name: Format code @@ -47,7 +48,7 @@ repos: name: imports - repo: https://github.com/executablebooks/mdformat - rev: 0.7.7 + rev: 0.7.10 hooks: - id: mdformat additional_dependencies: diff --git a/CITATION.cff b/CITATION.cff new file mode 100644 index 00000000..d245c58f --- /dev/null +++ b/CITATION.cff @@ -0,0 +1,18 @@ +cff-version: 1.2.0 +title: >- + PhoenixDL/rising: High-Performance Differentiable + Medical Data Augmentation +message: >- + If you use this software, please cite it using the + metadata from this file. +type: software +authors: + - given-names: Justus + family-names: Schock + affiliation: Contributed Equally + orcid: 'https://orcid.org/0000-0003-0512-3053' + - given-names: Michael + family-names: Baumgartner + affiliation: Contributed Equally + - given-names: Leon + family-names: Weninger diff --git a/README.md b/README.md index a2b073d5..d0841ad5 100644 --- a/README.md +++ b/README.md @@ -7,7 +7,7 @@ ![PyPI - License](https://img.shields.io/pypi/l/rising) [![Chat](https://img.shields.io/badge/Slack-PhoenixDL-orange)](https://join.slack.com/t/phoenixdl/shared_invite/enQtODgwODI0MTE1MjgzLTJkZDE4N2NhM2VmNzVhYTEyMzI3NzFmMDY0NjM3MzJlZWRmMTk5ZWM1YzY2YjY5ZGQ1NWI1YmJmOTdiYTdhYTE) [![Documentation Status](https://readthedocs.org/projects/rising/badge/?version=latest)](https://rising.readthedocs.io/en/latest/?badge=latest) -[![pre-commit.ci status](https://results.pre-commit.ci/badge/github/PhoenixDL/rising/master.svg)](https://results.pre-commit.ci/latest/github/PhoenixDL/rising/master) +[![pre-commit.ci status](https://results.pre-commit.ci/badge/github/PhoenixDL/rising/master.svg)](https://results.pre-commit.ci/latest/github/PhoenixDL/rising/master)[![DOI](https://zenodo.org/badge/222287924.svg)](https://zenodo.org/badge/latestdoi/222287924) @@ -17,6 +17,19 @@ | ![Python](https://img.shields.io/badge/python-3.6/3.7/3.8-orange) | ![System](https://img.shields.io/badge/Windows-blue) | ![Unittests Windows](https://github.com/PhoenixDL/rising/workflows/Unittests%20Windows/badge.svg) | | ![Python](https://img.shields.io/badge/python-3.6/3.7/3.8-orange) | ![System](https://img.shields.io/badge/MacOS-blue) | ![Unittests macOS](https://github.com/PhoenixDL/rising/workflows/Unittests%20MacOS/badge.svg) | +## What I have modified from the original repo? + +I noticed that rising under its current form has small bugs and inconsistency in terms of design and performance, especially for medical image segmentation task with 3D volumes. Based on project requirement, several improvements have been made so that it can now enjoy a better compatibility with 3D datas and intenstive augmentation operations using GPUs. Issues and PRs are welcome. + +Thanks a lot @Yuxiang1990 for the kind PR. + +## Working example on ACDC dataset + +[ACDC example](notebooks/medical_seg/acdc_seg.py) should provide a resonable example of how to use rising besides the below +explanations. + +______________________________________________________________________ + ## What is `rising`? Rising is a high-performance data loading and augmentation library for 2D *and* 3D data completely written in PyTorch. diff --git a/notebooks/__init__.py b/notebooks/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/notebooks/medical_seg/__init__.py b/notebooks/medical_seg/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/notebooks/medical_seg/acdc_seg.py b/notebooks/medical_seg/acdc_seg.py new file mode 100644 index 00000000..0f05d5e5 --- /dev/null +++ b/notebooks/medical_seg/acdc_seg.py @@ -0,0 +1,66 @@ +from tqdm import tqdm + +from notebooks.medical_seg.dataset import ACDCDataset, InfiniteRandomSampler +from rising import transforms +from rising.constants import FInterpolation +from rising.loading import DataLoader, default_transform_call +from rising.random import UniformParameter + +tra_dataset = ACDCDataset(root="/home/jizong/Workspace/rising/notebooks/medical_seg/data", train=True) + +seq_tra_augment = transforms.Compose( + transforms.NormPercentile(keys=("image",), min=0.02, max=0.98), + transforms.ResizeNative( + size=(10, 224, 224), + mode=(FInterpolation.trilinear, FInterpolation.nearest), + preserve_range=True, + keys=("image", "label"), + ), + # transforms.PadRandomCrop(size=(10, 224, 224), pad_size=2, pad_value=(0, 0), keys=("image", "label")), +) + +batch_augment = transforms.Compose( + # transforms.ToDtype(keys=("image", "label"), dtype=torch.half), + # transforms.ToDevice(keys=("image", "label"), device=torch.device("cuda")), + transforms.GammaCorrection(gamma=UniformParameter(0.8, 2), keys=("image",)), + transforms.RicianNoiseTransform(keys=("image",), std=0.05, keep_range=False), + transforms.BaseAffine( + scale=(1, UniformParameter(0.5, 4), UniformParameter(0.9, 1.2)), + rotation=(0, UniformParameter(-10, 10), UniformParameter(-10, 10)), + degree=True, + p=1, + per_sample=True, + interpolation_mode=("bilinear", "nearest"), + keys=("image", "label"), + ), + transforms.RandomCrop(size=(8, 192, 168), keys=("image", "label")), + transforms.ElasticDistortion( + std=20, + alpha=0.2, + dim=3, + keys=("image", "label"), + interpolation_mode=("bilinear", "nearest"), + ), + transform_call=default_transform_call, +) + +tra_loader = DataLoader( + tra_dataset, + sampler=InfiniteRandomSampler(tra_dataset, shuffle=False), + batch_size=6, + sample_transforms=seq_tra_augment, + gpu_transforms=batch_augment, + pseudo_batch_dim=True, + num_workers=16, +) + +for data in tqdm(tra_loader): + image, label = data["image"], data["label"] + from tests.realtime_viewer import multi_slice_viewer_debug + + multi_slice_viewer_debug([*image.squeeze()], *label.squeeze(), block=True, no_contour=True) + # from tests.realtime_viewer import multi_slice_viewer_debug + # + # for img, lab in zip(image, label): + # multi_slice_viewer_debug(img.squeeze(), lab.squeeze(), block=False, no_contour=True) + # plt.show() diff --git a/notebooks/medical_seg/data/train/patient001/Info.cfg b/notebooks/medical_seg/data/train/patient001/Info.cfg new file mode 100644 index 00000000..70809255 --- /dev/null +++ b/notebooks/medical_seg/data/train/patient001/Info.cfg @@ -0,0 +1,6 @@ +ED: 1 +ES: 12 +Group: DCM +Height: 184.0 +NbFrame: 30 +Weight: 95.0 diff --git a/notebooks/medical_seg/data/train/patient001/patient001_frame01.nii.gz b/notebooks/medical_seg/data/train/patient001/patient001_frame01.nii.gz new file mode 100644 index 00000000..9ba4d2b2 Binary files /dev/null and b/notebooks/medical_seg/data/train/patient001/patient001_frame01.nii.gz differ diff --git a/notebooks/medical_seg/data/train/patient001/patient001_frame01_gt.nii.gz b/notebooks/medical_seg/data/train/patient001/patient001_frame01_gt.nii.gz new file mode 100644 index 00000000..dea134d8 Binary files /dev/null and b/notebooks/medical_seg/data/train/patient001/patient001_frame01_gt.nii.gz differ diff --git a/notebooks/medical_seg/data/train/patient001/patient001_frame12.nii.gz b/notebooks/medical_seg/data/train/patient001/patient001_frame12.nii.gz new file mode 100644 index 00000000..d1f5d11b Binary files /dev/null and b/notebooks/medical_seg/data/train/patient001/patient001_frame12.nii.gz differ diff --git a/notebooks/medical_seg/data/train/patient001/patient001_frame12_gt.nii.gz b/notebooks/medical_seg/data/train/patient001/patient001_frame12_gt.nii.gz new file mode 100644 index 00000000..55fa3773 Binary files /dev/null and b/notebooks/medical_seg/data/train/patient001/patient001_frame12_gt.nii.gz differ diff --git a/notebooks/medical_seg/dataset.py b/notebooks/medical_seg/dataset.py new file mode 100644 index 00000000..b37db2a0 --- /dev/null +++ b/notebooks/medical_seg/dataset.py @@ -0,0 +1,72 @@ +import re +from collections import Iterator +from pathlib import Path + +import numpy as np +import SimpleITK as sitk +import torch +from torch.utils.data import Sampler +from torch.utils.data.dataset import Dataset, T_co + + +class ACDCDataset(Dataset): + def __init__(self, *, root: str, train: bool = True) -> None: + self.root = root + self.train = train + self._root = Path(root, "train" if train else "val") + + self.images = filter( + self.image_filter, [str(x.relative_to(self._root)) for x in Path(self._root).rglob("*.nii.gz")] + ) + self.images = sorted(self.images) + + def __getitem__(self, index) -> T_co: + image_path = str(self._root / self.images[index]) + gt_path = image_path.replace(".nii.gz", "_gt.nii.gz") + image = sitk.GetArrayFromImage(sitk.ReadImage(image_path)).astype(float, copy=False)[None, ...] + gt = sitk.GetArrayFromImage(sitk.ReadImage(gt_path)).astype(float, copy=False)[None, ...] + return {"image": torch.from_numpy(image), "label": torch.from_numpy(gt)} + + def __len__(self): + return len(self.images) + + @staticmethod + def image_filter(path: str): + _match = re.compile(r"patient\d+_frame\d+.nii.gz").search(str(path)) + if _match is None: + return False + return True + + +class _InfiniteRandomIterator(Iterator): + def __init__(self, data_source, shuffle=True): + self.data_source = data_source + self.shuffle = shuffle + if self.shuffle: + self.iterator = iter(torch.randperm(len(self.data_source)).tolist()) + else: + self.iterator = iter(torch.arange(start=0, end=len(self.data_source)).tolist()) + + def __next__(self): + try: + idx = next(self.iterator) + except StopIteration: + if self.shuffle: + self.iterator = iter(torch.randperm(len(self.data_source)).tolist()) + else: + self.iterator = iter(torch.arange(start=0, end=len(self.data_source)).tolist()) + idx = next(self.iterator) + return idx + + +class InfiniteRandomSampler(Sampler): + def __init__(self, data_source, shuffle=True): + super().__init__(data_source) + self.data_source = data_source + self.shuffle = shuffle + + def __iter__(self): + return _InfiniteRandomIterator(self.data_source, shuffle=self.shuffle) + + def __len__(self): + return len(self.data_source) diff --git a/notebooks/medical_seg/unet_3d.py b/notebooks/medical_seg/unet_3d.py new file mode 100644 index 00000000..906127f6 --- /dev/null +++ b/notebooks/medical_seg/unet_3d.py @@ -0,0 +1,130 @@ +# 3D-UNet model. +# x: 128x128 resolution for 32 frames. +import torch +import torch.nn as nn + + +def conv_block_3d(in_dim, out_dim, activation): + return nn.Sequential( + nn.Conv3d(in_dim, out_dim, kernel_size=3, stride=1, padding=1), + nn.BatchNorm3d(out_dim), + activation, + ) + + +def conv_trans_block_3d(in_dim, out_dim, activation): + return nn.Sequential( + nn.ConvTranspose3d(in_dim, out_dim, kernel_size=3, stride=2, padding=1, output_padding=1), + nn.BatchNorm3d(out_dim), + activation, + ) + + +def max_pooling_3d(): + return nn.MaxPool3d(kernel_size=2, stride=2, padding=0) + + +def conv_block_2_3d(in_dim, out_dim, activation, kernel_size=(3, 3, 3)): + return nn.Sequential( + conv_block_3d(in_dim, out_dim, activation), + nn.Conv3d(out_dim, out_dim, kernel_size=kernel_size, stride=1, padding=1), + nn.BatchNorm3d(out_dim), + ) + + +class UNet(nn.Module): + def __init__(self, in_dim, out_dim, num_filters): + super(UNet, self).__init__() + + self.in_dim = in_dim + self.out_dim = out_dim + self.num_filters = num_filters + activation = nn.LeakyReLU(0.2, inplace=True) + + # Down sampling + self.down_1 = conv_block_2_3d(self.in_dim, self.num_filters, activation, kernel_size=(1, 3, 3)) + self.pool_1 = max_pooling_3d() + self.down_2 = conv_block_2_3d(self.num_filters, self.num_filters * 2, activation, kernel_size=(1, 3, 3)) + self.pool_2 = max_pooling_3d() + self.down_3 = conv_block_2_3d(self.num_filters * 2, self.num_filters * 4, activation, kernel_size=(1, 3, 3)) + self.pool_3 = max_pooling_3d() + self.down_4 = conv_block_2_3d(self.num_filters * 4, self.num_filters * 8, activation) + self.pool_4 = max_pooling_3d() + self.down_5 = conv_block_2_3d(self.num_filters * 8, self.num_filters * 16, activation) + self.pool_5 = max_pooling_3d() + + # Bridge + self.bridge = conv_block_2_3d(self.num_filters * 16, self.num_filters * 32, activation) + + # Up sampling + self.trans_1 = conv_trans_block_3d(self.num_filters * 32, self.num_filters * 32, activation) + self.up_1 = conv_block_2_3d(self.num_filters * 48, self.num_filters * 16, activation) + self.trans_2 = conv_trans_block_3d(self.num_filters * 16, self.num_filters * 16, activation) + self.up_2 = conv_block_2_3d(self.num_filters * 24, self.num_filters * 8, activation) + self.trans_3 = conv_trans_block_3d(self.num_filters * 8, self.num_filters * 8, activation) + self.up_3 = conv_block_2_3d(self.num_filters * 12, self.num_filters * 4, activation) + self.trans_4 = conv_trans_block_3d(self.num_filters * 4, self.num_filters * 4, activation) + self.up_4 = conv_block_2_3d(self.num_filters * 6, self.num_filters * 2, activation) + self.trans_5 = conv_trans_block_3d(self.num_filters * 2, self.num_filters * 2, activation) + self.up_5 = conv_block_2_3d(self.num_filters * 3, self.num_filters * 1, activation) + + # Output + self.out = conv_block_3d(self.num_filters, out_dim, activation) + + def forward(self, x): + # Down sampling + down_1 = self.down_1(x) # -> [1, 4, 128, 128, 128] + pool_1 = self.pool_1(down_1) # -> [1, 4, 64, 64, 64] + + down_2 = self.down_2(pool_1) # -> [1, 8, 64, 64, 64] + pool_2 = self.pool_2(down_2) # -> [1, 8, 32, 32, 32] + + down_3 = self.down_3(pool_2) # -> [1, 16, 32, 32, 32] + pool_3 = self.pool_3(down_3) # -> [1, 16, 16, 16, 16] + + down_4 = self.down_4(pool_3) # -> [1, 32, 16, 16, 16] + pool_4 = self.pool_4(down_4) # -> [1, 32, 8, 8, 8] + + down_5 = self.down_5(pool_4) # -> [1, 64, 8, 8, 8] + pool_5 = self.pool_5(down_5) # -> [1, 64, 4, 4, 4] + + # Bridge + bridge = self.bridge(pool_5) # -> [1, 128, 4, 4, 4] + + # Up sampling + trans_1 = self.trans_1(bridge) # -> [1, 128, 8, 8, 8] + concat_1 = torch.cat([trans_1, down_5], dim=1) # -> [1, 192, 8, 8, 8] + up_1 = self.up_1(concat_1) # -> [1, 64, 8, 8, 8] + + trans_2 = self.trans_2(up_1) # -> [1, 64, 16, 16, 16] + concat_2 = torch.cat([trans_2, down_4], dim=1) # -> [1, 96, 16, 16, 16] + up_2 = self.up_2(concat_2) # -> [1, 32, 16, 16, 16] + + trans_3 = self.trans_3(up_2) # -> [1, 32, 32, 32, 32] + concat_3 = torch.cat([trans_3, down_3], dim=1) # -> [1, 48, 32, 32, 32] + up_3 = self.up_3(concat_3) # -> [1, 16, 32, 32, 32] + + trans_4 = self.trans_4(up_3) # -> [1, 16, 64, 64, 64] + concat_4 = torch.cat([trans_4, down_2], dim=1) # -> [1, 24, 64, 64, 64] + up_4 = self.up_4(concat_4) # -> [1, 8, 64, 64, 64] + + trans_5 = self.trans_5(up_4) # -> [1, 8, 128, 128, 128] + concat_5 = torch.cat([trans_5, down_1], dim=1) # -> [1, 12, 128, 128, 128] + up_5 = self.up_5(concat_5) # -> [1, 4, 128, 128, 128] + + # Output + out = self.out(up_5) # -> [1, 3, 128, 128, 128] + return out + + +if __name__ == "__main__": + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + image_size = 32 + x = torch.Tensor(2, 3, image_size, image_size, image_size) + x = x.to(device) + print("x size: {}".format(x.size())) + + model = UNet(in_dim=3, out_dim=3, num_filters=16).to(device) + + out = model(x) + print("out size: {}".format(out.size())) diff --git a/requirements/install.txt b/requirements/install.txt index 900ba94d..3a706f6a 100644 --- a/requirements/install.txt +++ b/requirements/install.txt @@ -2,3 +2,4 @@ numpy torch>=1.6 # before 1.6 torch.searchsorted is not present threadpoolctl tqdm +SimpleITK diff --git a/rising/constants.py b/rising/constants.py new file mode 100644 index 00000000..007ed725 --- /dev/null +++ b/rising/constants.py @@ -0,0 +1,16 @@ +from enum import Enum + + +class FInterpolation(Enum): + bilinear = "bilinear" + nearest = "nearest" + trilinear = "trilinear" + + +class Interpolation(Enum): + pass + + +class AffineInterpolation(Enum): + linear = "bilinear" + nearest = "nearest" diff --git a/rising/loading/loader.py b/rising/loading/loader.py index a3de6dae..b018d460 100644 --- a/rising/loading/loader.py +++ b/rising/loading/loader.py @@ -12,6 +12,8 @@ from torch.utils.data.dataloader import _MultiProcessingDataLoaderIter as __MultiProcessingDataLoaderIter from torch.utils.data.dataloader import _SingleProcessDataLoaderIter as __SingleProcessDataLoaderIter +from rising.utils.transforms import get_dtype_from_transforms, get_keys_from_transforms + try: import numpy as np @@ -194,11 +196,17 @@ def __init__( if device is None: device = torch.cuda.current_device() - to_gpu_trafo = ToDevice(device=device, non_blocking=pin_memory) + keys = get_keys_from_transforms(gpu_transforms) + to_gpu_trafo = ToDevice(device=device, non_blocking=pin_memory, keys=keys) - gpu_transforms = Compose(to_gpu_trafo, gpu_transforms) + gpu_transforms = Compose(to_gpu_trafo, gpu_transforms, transform_call=default_transform_call) gpu_transforms = gpu_transforms.to(device) + # check the dtype from the gpu compose + dtype = get_dtype_from_transforms(gpu_transforms) + if dtype: + gpu_transforms = gpu_transforms.to(dtype) + self.device = device self.sample_transforms = sample_transforms self.pseudo_batch_dim = pseudo_batch_dim and sample_transforms is not None @@ -385,7 +393,7 @@ def __call__(self, *args, **kwargs) -> Any: batch = self._transform_call(batch, self._transforms) if self._auto_convert: - batch = default_convert(batch) + batch = default_convert(batch) # convert to tensor return batch diff --git a/rising/random/abstract.py b/rising/random/abstract.py index 4cd2ef67..f9f55c61 100644 --- a/rising/random/abstract.py +++ b/rising/random/abstract.py @@ -1,11 +1,24 @@ -from abc import abstractmethod +import typing as t +from abc import ABC, abstractmethod from typing import Optional, Sequence, Union import torch +from torch.distributions import Distribution from rising.utils.shape import reshape -__all__ = ["AbstractParameter"] +__all__ = ["AbstractParameter", "ConstantParameter"] + + +def get_name(func): + from inspect import isclass, isfunction + + if isclass(func): + return func.__class__.__name__ + elif isfunction(func): + return func.__name__ + else: + return str(func) class AbstractParameter(torch.nn.Module): @@ -13,6 +26,11 @@ class AbstractParameter(torch.nn.Module): Abstract Parameter class to inject randomness to transforms """ + def __repr__(self): + name = self.__class__.__name__ + params = {k: get_name(v) for k, v in self.__dict__.items() if not k.startswith("_")} + return f"{name}({params})" + @staticmethod def _get_n_samples(size: Union[Sequence, torch.Size] = (1,)): """ @@ -29,7 +47,7 @@ def _get_n_samples(size: Union[Sequence, torch.Size] = (1,)): return size.numel() @abstractmethod - def sample(self, n_samples: int) -> Union[torch.Tensor, list]: + def sample(self, n_samples: int) -> Union[torch.Tensor, t.List[torch.Tensor]]: """ Abstract sampling function @@ -69,7 +87,9 @@ def forward( if the parameter ``tensor_like`` is given, it overwrites the parameters ``dtype`` and ``device`` """ - n_samples = self._get_n_samples(size if size is not None else (1,)) + if size is None: + size = (1,) + n_samples = self._get_n_samples(size) samples = self.sample(n_samples) if any([s is None for s in samples]): @@ -78,8 +98,7 @@ def forward( if not isinstance(samples, torch.Tensor): samples = torch.tensor(samples).flatten() - if size is not None: - samples = reshape(samples, size) + samples = reshape(samples, size) if isinstance(samples, torch.Tensor): if tensor_like is not None: @@ -87,3 +106,12 @@ def forward( else: samples = samples.to(device=device, dtype=dtype) return samples + + +class ConstantParameter(Distribution, ABC): + def __init__(self, constant: t.Union[int, float]): + super().__init__(validate_args=False) + self.constant = constant + + def sample(self, n_samples: Union[t.Tuple[int], torch.Size] = torch.Size()): + return torch.as_tensor([self.constant for _ in range(n_samples[0])]) diff --git a/rising/random/continuous.py b/rising/random/continuous.py index c5712cbd..b30e6acc 100644 --- a/rising/random/continuous.py +++ b/rising/random/continuous.py @@ -1,9 +1,10 @@ -from typing import Union +from typing import Union, cast import torch from torch.distributions import Distribution as TorchDistribution -from rising.random.abstract import AbstractParameter +from rising.random.abstract import AbstractParameter, ConstantParameter +from rising.utils import check_scalar __all__ = ["ContinuousParameter", "NormalParameter", "UniformParameter"] @@ -29,6 +30,8 @@ def sample(self, n_samples: int) -> torch.Tensor: Returns torch.Tensor: samples """ + + # input should be a tuple or torch.Size tuple. return self.dist.sample((n_samples,)) @@ -36,6 +39,7 @@ class NormalParameter(ContinuousParameter): """ Samples Parameters from a normal distribution. For details have a look at :class:`torch.distributions.Normal` + if sigma is 0, return a ConstantParameter """ def __init__(self, mu: Union[float, torch.Tensor], sigma: Union[float, torch.Tensor]): @@ -44,19 +48,33 @@ def __init__(self, mu: Union[float, torch.Tensor], sigma: Union[float, torch.Ten mu : the distributions mean sigma : the distributions standard deviation """ - super().__init__(torch.distributions.Normal(loc=mu, scale=sigma)) + assert check_scalar(mu) and check_scalar(sigma) + if sigma == 0: + dist = cast(torch.distributions.Distribution, ConstantParameter(constant=mu)) + else: + dist = torch.distributions.Normal(mu, sigma) + super().__init__(dist) class UniformParameter(ContinuousParameter): """ Samples Parameters from a uniform distribution. For details have a look at :class:`torch.distributions.Uniform` + if `low`==`high` , return a ConstantParameter """ - def __init__(self, low: Union[float, torch.Tensor], high: Union[float, torch.Tensor]): + def __init__(self, low: Union[float, int, torch.Tensor], high: Union[float, int, torch.Tensor]): """ Args: low : the lower range (inclusive) high : the higher range (exclusive) """ - super().__init__(torch.distributions.Uniform(low=low, high=high)) + assert check_scalar(low) and check_scalar(high) + + if low == high: + dist = cast(torch.distributions.Distribution, ConstantParameter(low)) + elif low < high: + dist = torch.distributions.Uniform(low=low, high=high) + else: + raise ValueError("low must be smaller than high, given: low={} and high={}".format(low, high)) + super().__init__(dist) diff --git a/rising/random/discrete.py b/rising/random/discrete.py index 64bf6d15..9dd748d1 100644 --- a/rising/random/discrete.py +++ b/rising/random/discrete.py @@ -2,12 +2,14 @@ from itertools import combinations from random import choices as sample_with_replacement from random import sample as sample_without_replacement -from typing import List, Sequence +from typing import List, Sequence, TypeVar from rising.random.abstract import AbstractParameter __all__ = ["DiscreteParameter", "DiscreteCombinationsParameter"] +T = TypeVar("T") + def combinations_all(data: Sequence) -> List: """ @@ -32,28 +34,31 @@ class DiscreteParameter(AbstractParameter): """ def __init__( - self, population: Sequence, replacement: bool = False, weights: Sequence = None, cum_weights: Sequence = None + self, + population: Sequence[T], + replacement: bool = False, + weights: Sequence[T] = None, + cum_weights: Sequence[T] = None, ): """ Args: population : the parameter population to sample from - replacement : whether or not to sample with replacement + replacement : whether to sample with replacement weights : relative sampling weights cum_weights : cumulative sampling weights """ super().__init__() if replacement: sample_fn = partial(sample_with_replacement, weights=weights, cum_weights=cum_weights) + elif weights is not None or cum_weights is not None: + raise ValueError("weights and cum_weights should only be specified if " "replacement is set to True!") else: - if weights is not None or cum_weights is not None: - raise ValueError("weights and cum_weights should only be specified if " "replacement is set to True!") - sample_fn = sample_without_replacement self.sample_fn = sample_fn self.population = population - def sample(self, n_samples: int) -> list: + def sample(self, n_samples: int) -> List[T]: """ Samples from the discrete internal population @@ -73,11 +78,15 @@ class DiscreteCombinationsParameter(DiscreteParameter): possible combinations of the given population """ - def __init__(self, population: Sequence, replacement: bool = False): + def __init__(self, population: Sequence[T], replacement: bool = False): """ Args: population : population to build combination of - replacement : whether or not to sample with replacement + replacement : whether to sample with replacement """ population = combinations_all(population) super().__init__(population=population, replacement=replacement) + + +if __name__ == "__main__": + print(DiscreteParameter(population=[1, 2, 3, 4, 5], replacement=True)) diff --git a/rising/random/utils.py b/rising/random/utils.py new file mode 100644 index 00000000..3a1247c3 --- /dev/null +++ b/rising/random/utils.py @@ -0,0 +1,24 @@ +import torch +import functools +from torch import Tensor +from typing import Union + + +class fix_random_seed_ctx: + def __init__(self, seed: Union[Tensor, int]) -> None: + super().__init__() + self._seed = int(seed) + + def __enter__(self): + self._prev_seed = torch.random.get_rng_state() + + def __exit__(self, exc_type, exc_val, exc_tb): + torch.random.set_rng_state(self._prev_seed) # noqa + + def __call__(self, func): + @functools.wraps(func) + def wrapped_func(*args, **kwargs): + with self: + return func(*args, **kwargs) + + return wrapped_func diff --git a/rising/transforms/__init__.py b/rising/transforms/__init__.py index d37b4ed1..bb6e36e2 100644 --- a/rising/transforms/__init__.py +++ b/rising/transforms/__init__.py @@ -20,35 +20,46 @@ """ from rising.transforms.abstract import ( - AbstractTransform, BaseTransform, - BaseTransformSeeded, - PerChannelTransform, - PerSampleTransform, + BaseTransformMixin, + PerChannelTransformMixin, + PerSampleTransformMixin, + _AbstractTransform, ) -from rising.transforms.affine import Affine, BaseAffine, Resize, Rotate, Scale, StackedAffine, Translate +from rising.transforms.affine import BaseAffine, Resize, Rotate, Scale, Translate, _Affine, _StackedAffine from rising.transforms.channel import ArgMax, OneHot from rising.transforms.compose import Compose, DropoutCompose, OneOf -from rising.transforms.crop import CenterCrop, RandomCrop +from rising.transforms.crop import CenterCrop, PadRandomCrop, RandomCrop, PadCenterCrop from rising.transforms.format import FilterKeys, MapToSeq, PopKeys, RenameKeys, SeqToMap +from rising.transforms.grid import ( + CenterCropGrid, + ElasticDistortion, + GridTransform, + RadialDistortion, + RandomCropGrid, + StackedGridTransform, +) from rising.transforms.intensity import ( Clamp, ExponentialNoise, GammaCorrection, GaussianNoise, InvertAmplitude, - Noise, NormMeanStd, NormMinMax, + NormPercentile, NormRange, NormZeroMeanUnitStd, RandomAddValue, RandomBezierTransform, RandomScaleValue, - RandomValuePerChannel, + RicianNoiseTransform, ) from rising.transforms.kernel import GaussianSmoothing, KernelTransform +from rising.transforms.pad import Pad from rising.transforms.painting import LocalPixelShuffle, RandomInOrOutpainting, RandomInpainting, RandomOutpainting -from rising.transforms.spatial import Mirror, ProgressiveResize, ResizeNative, Rot90, SizeStepScheduler, Zoom -from rising.transforms.tensor import Permute, TensorOp, ToDevice, ToDeviceDtype, ToDtype, ToTensor +from rising.transforms.sitk import SITK2Tensor, SITKResample, SITKWindows +from rising.transforms.spatial import Mirror, ProgressiveResize, ResizeNative, Rot90, SizeStepScheduler, Zoom, \ + ResizeNativeCentreCrop +from rising.transforms.tensor import Permute, TensorOp, ToDevice, ToDtype, ToTensor, _ToDeviceDtype from rising.transforms.utility import BoxToSeg, DoNothing, InstanceToSemantic, SegToBox diff --git a/rising/transforms/abstract.py b/rising/transforms/abstract.py index acc2c17b..629045ca 100644 --- a/rising/transforms/abstract.py +++ b/rising/transforms/abstract.py @@ -1,26 +1,44 @@ -from typing import Any, Callable, Sequence, Tuple, Union +from abc import ABC, abstractmethod + +try: + from typing import Any, Callable, Dict, List, Optional, Sequence, TypeVar, Union, final +except ImportError: + from typing import Any, Callable, Dict, List, Optional, Sequence, TypeVar, Union + from typing_extensions import final import torch +from torch import nn from rising.random import AbstractParameter, DiscreteParameter +from rising.utils.mise import fix_seed_cxm, ntuple, nullcxm + +__all__ = [ + "_AbstractTransform", + "ItemSeq", + "BaseTransform", + "PerChannelTransformMixin", + "PerSampleTransformMixin", + "BaseTransformMixin", +] -__all__ = ["AbstractTransform", "BaseTransform", "PerSampleTransform", "PerChannelTransform", "BaseTransformSeeded"] +T = TypeVar("T") +ItemSeq = Union[T, Sequence[T]] -augment_callable = Callable[[torch.Tensor], Any] +augment_callable = Callable[..., Any] augment_axis_callable = Callable[[torch.Tensor, Union[float, Sequence]], Any] -class AbstractTransform(torch.nn.Module): +class _AbstractTransform(nn.Module): """Base class for all transforms""" - def __init__(self, grad: bool = False, **kwargs): + def __init__(self, *, grad: bool = False, **kwargs): """ Args: grad: enable gradient computation inside transformation """ super().__init__() self.grad = grad - self._registered_samplers = [] + self._registered_samplers: List[str] = [] for key, item in kwargs.items(): setattr(self, key, item) @@ -39,9 +57,10 @@ def register_sampler(self, name: str, sampler: Union[Sequence, AbstractParameter **kwargs : additional keyword arguments (will be forwarded to sampler call) """ + if name in self._registered_samplers: + raise ValueError(f"{name} has been registered as sampler.") self._registered_samplers.append(name) - if hasattr(self, name): - raise NameError("Name %s already exists" % name) + if not isinstance(sampler, (tuple, list)): sampler = [sampler] @@ -60,11 +79,23 @@ def sample(self): if len(sample_result) == 1: return sample_result[0] - else: - return sample_result + return sample_result + if hasattr(self, name): + delattr(self, name) setattr(self, name, property(sample)) + def need_sampler(self, value) -> bool: + if isinstance(value, AbstractParameter): + return True + if isinstance(value, (list, tuple)): + return any([self.need_sampler(x) for x in value]) + if isinstance(value, dict): + raise NotImplementedError(value) + else: + return False + + @final def __getattribute__(self, item) -> Any: """ Automatically dereference registered samplers @@ -83,6 +114,7 @@ def __getattribute__(self, item) -> Any: else: return res + @final def __call__(self, *args, **kwargs) -> Any: """ Call super class with correct torch context @@ -103,6 +135,7 @@ def __call__(self, *args, **kwargs) -> Any: with context: return super().__call__(*args, **kwargs) + @abstractmethod def forward(self, **data) -> dict: """ Implement transform functionality here @@ -116,27 +149,44 @@ def forward(self, **data) -> dict: raise NotImplementedError -class BaseTransform(AbstractTransform): +class BaseTransform(_AbstractTransform, ABC): """ Transform to apply a functional interface to given keys .. warning:: This transform should not be used with functions which have randomness build in because it will result in different augmentations per key. + Modifications: + Three kinds of attributes can be found here. + 1. The attribute which can be sampled (include AbstractParameters, Sequence[Abstract], int, float, Tensor etc.) + These attributes can be sampled when calling __get__ magic function. stored in _registered_samplers. + 2. The attribute which is specific to each key, such as interpolation, mode, etc. These keys must be get + per key. stored in _registered_key_pairs. + 3. The attribute which is normal. + + + Since we have modified the seed manager with context manager, + we can release the requirement of seeding. + + We need to pass a list of attributes to be passed to augment_fn ??? + """ def __init__( self, + *, augment_fn: augment_callable, - *args, - keys: Sequence = ("data",), + keys: Sequence[str] = ("data",), grad: bool = False, - property_names: Sequence[str] = (), - **kwargs + paired_kw_names: Sequence[str] = (), + augment_fn_names: Sequence[str] = (), + per_sample: bool = True, + **kwargs, ): """ Args: augment_fn: function for augmentation + Modification made here: all augment_fu accept data under form of BCHW(D) form. *args: positional arguments passed to augment_fn keys: keys which should be augmented grad: enable gradient computation inside transformation @@ -144,46 +194,100 @@ def __init__( during forward pass **kwargs: keyword arguments passed to augment_fn """ - sampler_vals = [kwargs.pop(name) for name in property_names] super().__init__(grad=grad, **kwargs) + sampler_vals = {k: v for k, v in kwargs.items() if self.need_sampler(v)} + + self._paired_kw_names: List[str] = [] # hidden list + self._augment_fn_names = augment_fn_names # kwargs passed to the augment_fn + self.paired_kw_names = paired_kw_names + self.augment_fn = augment_fn + + assert isinstance(keys, Sequence), keys self.keys = keys - self.property_names = property_names - self.args = args - self.kwargs = kwargs - for name, val in zip(property_names, sampler_vals): - self.register_sampler(name, val) + self.tuple_generator = ntuple(len(self.keys)) - def forward(self, **data) -> dict: + self.per_sample = per_sample + + for kwarg_name in self.paired_kw_names: + self.register_paired_attribute(kwarg_name, getattr(self, kwarg_name)) + + for name, val in sampler_vals.items(): + self.register_sampler(name, val) # lazy sampling values + + def sample_for_batch(self, name: str, batch_size: int) -> Optional[Union[Any, Sequence[Any]]]: """ - Apply transformation + Sample elements for batch Args: - data: dict with tensors + name: name of parameter + batch_size: batch size Returns: - dict: dict with augmented data + Optional[Union[Any, Sequence[Any]]]: sampled elements """ - kwargs = {} - for k in self.property_names: - kwargs[k] = getattr(self, k) + elem = getattr(self, name) + if elem is not None and self.per_sample: + return [elem] + [getattr(self, name) for _ in range(batch_size - 1)] + else: + return elem # either a single scalar value or None + + def register_paired_attribute(self, name: str, value: ItemSeq[T]): + if name in self._paired_kw_names: + raise ValueError(f"{name} has been registered in self._pair_kwarg_names") + if name not in self._augment_fn_names: + raise ValueError(f"{name} must be provided in `augment_fn_names`") + self._paired_kw_names.append(name) + setattr(self, name, self.tuple_generator(value)) + + def get_pair_kwargs(self, key: str) -> Dict[str, Any]: + assert key in self.keys, key + index = self.keys.index(key) + return {k: getattr(self, k)[index] for k in self._paired_kw_names} + + @abstractmethod + def forward(self, **data) -> dict: + """ + implementation override by mixin + """ + raise NotImplementedError - kwargs.update(self.kwargs) - for _key in self.keys: - data[_key] = self.augment_fn(data[_key], *self.args, **kwargs) - return data +class _BaseMixin(ABC): + """ + base mixin class to perform the forward function. + """ + + per_sample: bool # we have per_sample by default. + get_pair_kwargs: Callable[[str], Dict[str, Any]] + kwargs: Dict[str, Any] + keys: Sequence[str] + augment_fn: Callable # all augment_fn take BCHWD as data input. + _augment_fn_names: Sequence[str] + _paired_kw_names: List[str] -class BaseTransformSeeded(BaseTransform): + +class BaseTransformMixin(_BaseMixin): """ - Transform to apply a functional interface to given keys and use the same - pytorch(!) seed for every key. + this mixin call augment_fun and put all batch data into the data. + you don't care about per_sample option. """ - def forward(self, **data) -> dict: + def __init__(self, *, seeded: bool = True, p: float = 1, **kwargs) -> None: """ - Apply transformation and use same seed for every key + Args: + seeded: bool, default False. if the transformation need to fix the random seed for each key + p: float: probability of applying augment_fn per batch + """ + super().__init__(**kwargs) + self.seeded = seeded + assert 0 <= p <= 1, p + self.p = p + + def forward(self, **data) -> Dict[str, Any]: + """ + Apply transformation Args: data: dict with tensors @@ -191,30 +295,41 @@ def forward(self, **data) -> dict: Returns: dict: dict with augmented data """ - kwargs = {} - for k in self.property_names: - kwargs[k] = getattr(self, k) + seed = int(torch.randint(0, int(1e16), (1,))) - kwargs.update(self.kwargs) - - seed = torch.random.get_rng_state() for _key in self.keys: - torch.random.set_rng_state(seed) - data[_key] = self.augment_fn(data[_key], *self.args, **kwargs) + with self.random_cxm(seed): + kwargs = {k: getattr(self, k) for k in self._augment_fn_names if k not in self._paired_kw_names} + kwargs.update(self.get_pair_kwargs(_key)) + + if torch.rand(1).item() < self.p: + data[_key] = self.augment_fn(data[_key], **kwargs) return data + @property + def random_cxm(self): + """random seed control context manager, if self.seeded.""" + return fix_seed_cxm if self.seeded else nullcxm + -class PerSampleTransform(BaseTransform): +class PerSampleTransformMixin(BaseTransformMixin): """ - Apply transformation to each sample in batch individually - :attr:`augment_fn` must be callable with option :attr:`out` - where results are saved in. + Transfer to a mixed in to override the forward of `BasesTransform`. + + 1. random draw a seed as int + 3. fixed the seed + i as the beginning of the sample 1CBHWD - .. warning:: This transform should not be used - with functions which have randomness build in because it will - result in different augmentations per sample and key. """ + def __init__(self, *, p: float = 1, **kwargs): + """ + Args: + p: probability of applying the transform per sample. + """ + super(PerSampleTransformMixin, self).__init__(**kwargs) + assert 0 <= p <= 1, p + self.p = p + def forward(self, **data) -> dict: """ Args: @@ -223,21 +338,34 @@ def forward(self, **data) -> dict: Returns: dict: dict with augmented data """ - kwargs = {} - for k in self.property_names: - kwargs[k] = getattr(self, k) + if not self.per_sample: + return super(PerSampleTransformMixin, self).forward(**data) - kwargs.update(self.kwargs) - for _key in self.keys: - out = torch.empty_like(data[_key]) - for _i in range(data[_key].shape[0]): - out[_i] = self.augment_fn(data[_key][_i], out=out[_i], **kwargs) - data[_key] = out + seed = int(torch.randint(0, int(1e16), (1,))) + + for key in self.keys: + batch_size = data[key].shape[0] + out = [] + for b in range(batch_size): + with self.random_cxm(seed + b): + kwargs = {k: getattr(self, k) for k in self._augment_fn_names if k not in self._paired_kw_names} + kwargs.update(self.get_pair_kwargs(key)) + + if torch.rand(1).item() < self.p: + out.append(self.augment_fn(data[key][b][None, ...], **kwargs)) + else: + out.append(data[key][b][None, ...]) + + data[key] = torch.cat(out, dim=0) return data -class PerChannelTransform(BaseTransform): +class PerChannelTransformMixin(BaseTransformMixin): """ + Transfer to a mixed in to override the forward of `BasesTransform`. + + This mixin gives augment_fn without per_channel attribute a chance to perform channel-wise operation. + Apply transformation per channel (but still to whole batch) .. warning:: This transform should not be used @@ -245,25 +373,15 @@ class PerChannelTransform(BaseTransform): result in different augmentations per channel and key. """ - def __init__( - self, - augment_fn: augment_callable, - per_channel: bool = False, - keys: Sequence = ("data",), - grad: bool = False, - property_names: Tuple[str] = (), - **kwargs - ): + def __init__(self, *, per_channel: bool, p: float = 1, **kwargs): """ Args: - augment_fn: function for augmentation - per_channel: enable transformation per channel - keys: keys which should be augmented - grad: enable gradient computation inside transformation - kwargs: keyword arguments passed to augment_fn + per_channel:bool parameter to perform per_channel operation + kwargs: base parameters """ - super().__init__(augment_fn=augment_fn, keys=keys, grad=grad, property_names=property_names, **kwargs) + super().__init__(**kwargs) self.per_channel = per_channel + self.p = p def forward(self, **data) -> dict: """ @@ -275,17 +393,73 @@ def forward(self, **data) -> dict: Returns: dict: dict with augmented data """ - if self.per_channel: - kwargs = {} - for k in self.property_names: - kwargs[k] = getattr(self, k) - - kwargs.update(self.kwargs) - for _key in self.keys: - out = torch.empty_like(data[_key]) - for _i in range(data[_key].shape[1]): - out[:, _i] = self.augment_fn(data[_key][:, _i], out=out[:, _i], **kwargs) - data[_key] = out - return data - else: + if not self.per_channel: return super().forward(**data) + + seed = int(torch.randint(0, int(1e16), (1,))) + + for key in self.keys: + out = [] + channel_dim = data[key].shape[1] + for c in range(channel_dim): + with self.random_cxm(seed + c): + kwargs = {k: getattr(self, k) for k in self._augment_fn_names if k not in self._paired_kw_names} + kwargs.update(self.get_pair_kwargs(key)) + if torch.rand(1).item() < self.p: + out.append(self.augment_fn(data[key][:, c].unsqueeze(1), **kwargs)) + else: + out.append(data[key][:, c].unsqueeze(1)) + data[key] = torch.cat(out, dim=1) + + return data + + +class PerSamplePerChannelTransformMixin(BaseTransformMixin): + def __init__(self, *, per_channel: bool, p_channel: float = 1, per_sample: bool, p_sample: float = 1, **kwargs): + """ + Args: + per_channel:bool parameter to perform per_channel operation + kwargs: base parameters + """ + super().__init__(**kwargs) + self.per_channel = per_channel + self.p_channel = p_channel + + self.per_sample = per_sample + self.p_sample = p_sample + + def forward(self, **data) -> dict: + """ + Apply transformation + + Args: + data: dict with tensors + + Returns: + dict: dict with augmented data + """ + if not self.per_channel: + self.p = self.p_sample + return PerSampleTransformMixin.forward(self, **data) + if not self.per_sample: + self.p = self.p_channel + return PerChannelTransformMixin.forward(self, **data) + + seed = int(torch.randint(0, int(1e16), (1,))) + + for key in self.keys: + batch_size, channel_dim = data[key].shape[0:2] + for b in range(batch_size): + cur_data = data[key] + processed_batch = [] + for c in range(channel_dim): + with self.random_cxm(seed + b + c): + kwargs = {k: getattr(self, k) for k in self._augment_fn_names if k not in self._paired_kw_names} + kwargs.update(self.get_pair_kwargs(key)) + if torch.rand(1).item() < self.p: + processed_batch.append(self.augment_fn(cur_data[b, c][None, None, ...], **kwargs)) + else: + processed_batch.append(data[key][:, c][None, None, ...]) + data[key][b] = torch.cat(processed_batch, dim=1)[0] + + return data diff --git a/rising/transforms/affine.py b/rising/transforms/affine.py index 1fc463fd..4c95e762 100644 --- a/rising/transforms/affine.py +++ b/rising/transforms/affine.py @@ -2,15 +2,13 @@ import torch -from rising.transforms.abstract import BaseTransform +from rising.transforms.abstract import BaseTransform, BaseTransformMixin, ItemSeq from rising.transforms.functional.affine import AffineParamType, affine_image_transform, parametrize_matrix from rising.utils.affine import matrix_to_cartesian, matrix_to_homogeneous from rising.utils.checktype import check_scalar __all__ = [ - "Affine", "BaseAffine", - "StackedAffine", "Rotate", "Scale", "Translate", @@ -18,7 +16,7 @@ ] -class Affine(BaseTransform): +class _Affine(BaseTransform): """ Class Performing an Affine Transformation on a given sample dict. The transformation will be applied to all the dict-entries specified @@ -32,12 +30,11 @@ def __init__( grad: bool = False, output_size: Optional[tuple] = None, adjust_size: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, reverse_order: bool = False, per_sample: bool = True, - **kwargs, ): """ Args: @@ -69,18 +66,20 @@ def __init__( batch order [(D,)H,W] per_sample: sample different values for each element in the batch. The transform is still applied in a batched wise fashion. - **kwargs: additional keyword arguments passed to the - affine transform """ - super().__init__(augment_fn=affine_image_transform, keys=keys, grad=grad, **kwargs) + super().__init__( + augment_fn=affine_image_transform, + keys=keys, + per_sample=per_sample, + grad=grad, + ) self.matrix = matrix self.register_sampler("output_size", output_size) self.adjust_size = adjust_size - self.interpolation_mode = interpolation_mode - self.padding_mode = padding_mode - self.align_corners = align_corners + self.interpolation_mode = self.tuple_generator(interpolation_mode) + self.padding_mode = self.tuple_generator(padding_mode) + self.align_corners = self.tuple_generator(align_corners) self.reverse_order = reverse_order - self.per_sample = per_sample def assemble_matrix(self, **data) -> torch.Tensor: """ @@ -100,22 +99,22 @@ def assemble_matrix(self, **data) -> torch.Tensor: self.matrix = torch.tensor(self.matrix) self.matrix = self.matrix.to(data[self.keys[0]]) - batchsize = data[self.keys[0]].shape[0] + batch_size = data[self.keys[0]].shape[0] ndim = len(data[self.keys[0]].shape) - 2 # channel and batch dim # batch dimension missing -> Replicate for each sample in batch if len(self.matrix.shape) == 2: - self.matrix = self.matrix[None].expand(batchsize, -1, -1).clone() - if self.matrix.shape == (batchsize, ndim, ndim + 1): + self.matrix = self.matrix[None].expand(batch_size, -1, -1).clone() + if self.matrix.shape == (batch_size, ndim, ndim + 1): return self.matrix - elif self.matrix.shape == (batchsize, ndim, ndim): + elif self.matrix.shape == (batch_size, ndim, ndim): return matrix_to_homogeneous(self.matrix)[:, :-1] - elif self.matrix.shape == (batchsize, ndim + 1, ndim + 1): + elif self.matrix.shape == (batch_size, ndim + 1, ndim + 1): return matrix_to_cartesian(self.matrix) raise ValueError( "Invalid Shape for affine transformation matrix. " - "Got %s but expected %s" % (str(tuple(self.matrix.shape)), str((batchsize, ndim, ndim + 1))) + "Got %s but expected %s" % (str(tuple(self.matrix.shape)), str((batch_size, ndim, ndim + 1))) ) def forward(self, **data) -> dict: @@ -130,17 +129,18 @@ def forward(self, **data) -> dict: """ matrix = self.assemble_matrix(**data) - for key in self.keys: + for key, interpolation, padding, align_corners in zip( + self.keys, self.interpolation_mode, self.padding_mode, self.align_corners + ): data[key] = self.augment_fn( data[key], matrix_batch=matrix, output_size=self.output_size, adjust_size=self.adjust_size, - interpolation_mode=self.interpolation_mode, - padding_mode=self.padding_mode, - align_corners=self.align_corners, + interpolation_mode=interpolation, # this can be different + padding_mode=padding, # this can be different + align_corners=align_corners, # this can be different reverse_order=self.reverse_order, - **self.kwargs, ) return data @@ -154,10 +154,10 @@ def __add__(self, other: Any) -> BaseTransform: other: the other transformation Returns: - StackedAffine: a stacked affine transformation + _StackedAffine: a stacked affine transformation """ - if not isinstance(other, Affine): - other = Affine( + if not isinstance(other, _Affine): + other = _Affine( matrix=other, keys=self.keys, grad=self.grad, @@ -166,10 +166,11 @@ def __add__(self, other: Any) -> BaseTransform: interpolation_mode=self.interpolation_mode, padding_mode=self.padding_mode, align_corners=self.align_corners, + per_sample=self.per_sample, **self.kwargs, ) - return StackedAffine( + return _StackedAffine( self, other, keys=self.keys, @@ -179,6 +180,7 @@ def __add__(self, other: Any) -> BaseTransform: interpolation_mode=self.interpolation_mode, padding_mode=self.padding_mode, align_corners=self.align_corners, + per_sample=self.per_sample, **self.kwargs, ) @@ -191,10 +193,10 @@ def __radd__(self, other) -> BaseTransform: other: the other transformation Returns: - StackedAffine: a stacked affine transformation + _StackedAffine: a stacked affine transformation """ - if not isinstance(other, Affine): - other = Affine( + if not isinstance(other, _Affine): + other = _Affine( matrix=other, keys=self.keys, grad=self.grad, @@ -203,10 +205,9 @@ def __radd__(self, other) -> BaseTransform: interpolation_mode=self.interpolation_mode, padding_mode=self.padding_mode, align_corners=self.align_corners, - **self.kwargs, ) - return StackedAffine( + return _StackedAffine( other, self, grad=other.grad, @@ -215,11 +216,10 @@ def __radd__(self, other) -> BaseTransform: interpolation_mode=other.interpolation_mode, padding_mode=other.padding_mode, align_corners=other.align_corners, - **other.kwargs, ) -class StackedAffine(Affine): +class _StackedAffine(_Affine): """ Class to stack multiple affines with dynamic ensembling by matrix multiplication to avoid multiple interpolations. @@ -227,16 +227,16 @@ class StackedAffine(Affine): def __init__( self, - *transforms: Union[Affine, Sequence[Union[Sequence[Affine], Affine]]], + *transforms: Union[_Affine, Sequence[Union[Sequence[_Affine], _Affine]]], keys: Sequence = ("data",), grad: bool = False, output_size: Optional[tuple] = None, adjust_size: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, reverse_order: bool = False, - **kwargs, + per_sample=True, ): """ Args: @@ -276,7 +276,7 @@ def __init__( transforms = transforms[0] # ensure trafos are Affines and not raw matrices - transforms = tuple([trafo if isinstance(trafo, Affine) else Affine(matrix=trafo) for trafo in transforms]) + transforms = tuple([trafo if isinstance(trafo, _Affine) else _Affine(matrix=trafo) for trafo in transforms]) super().__init__( keys=keys, @@ -287,7 +287,7 @@ def __init__( padding_mode=padding_mode, align_corners=align_corners, reverse_order=reverse_order, - **kwargs, + per_sample=per_sample, ) self.transforms = transforms @@ -315,7 +315,7 @@ def assemble_matrix(self, **data) -> torch.Tensor: return matrix_to_cartesian(whole_trafo) -class BaseAffine(Affine): +class BaseAffine(_Affine, BaseTransformMixin): """ Class performing a basic Affine Transformation on a given sample dict. The transformation will be applied to all the dict-entries specified @@ -332,12 +332,12 @@ def __init__( grad: bool = False, output_size: Optional[tuple] = None, adjust_size: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[Optional[bool]] = False, reverse_order: bool = False, per_sample: bool = True, - **kwargs, + p: float = 1, ): """ Args: @@ -373,6 +373,14 @@ def __init__( calculated dynamically to ensure that the whole image fits. interpolation_mode: interpolation mode to calculate output values ``'bilinear'`` | ``'nearest'``. Default: ``'bilinear'`` + documents from PyTorch: + mode (str) – interpolation mode to calculate output values 'bilinear' + | 'nearest' | 'bicubic'. Default: 'bilinear' + Note: mode='bicubic' supports only 4-D input. + When mode='bilinear' and the input is 5-D, the interpolation mode used + internally will actually be trilinear. However, when the input is 4-D, + the interpolation mode will legitimately be bilinear. + padding_mode: padding mode for outside grid values ``'zeros'`` | ``'border'`` | ``'reflection'``. Default: ``'zeros'`` @@ -389,8 +397,8 @@ def __init__( batch order [(D,)H,W] per_sample: sample different values for each element in the batch. The transform is still applied in a batched wise fashion. - **kwargs: additional keyword arguments passed to the - affine transform + p: float, the probability of applying the transformation on batches. + """ super().__init__( keys=keys, @@ -402,8 +410,9 @@ def __init__( align_corners=align_corners, reverse_order=reverse_order, per_sample=per_sample, - **kwargs, ) + BaseTransformMixin.__init__(self, seeded=True, p=p) + self.p = p self.register_sampler("scale", scale) self.register_sampler("rotation", rotation) self.register_sampler("translation", translation) @@ -423,16 +432,22 @@ def assemble_matrix(self, **data) -> torch.Tensor: Returns: torch.Tensor: the (batched) transformation matrix """ - batchsize = data[self.keys[0]].shape[0] + batch_size = data[self.keys[0]].shape[0] ndim = len(data[self.keys[0]].shape) - 2 # channel and batch dim device = data[self.keys[0]].device dtype = data[self.keys[0]].dtype + seed = int(torch.randint(0, int(1e6), (1,))) + + scale = self.sample_for_batch_with_prob("scale", batch_size, default_value=1.0, seed=seed) + rotation = self.sample_for_batch_with_prob("rotation", batch_size, default_value=0.0, seed=seed) + translation = self.sample_for_batch_with_prob("translation", batch_size, default_value=0.0, seed=seed) + self.matrix = parametrize_matrix( - scale=self.sample_for_batch("scale", batchsize), - rotation=self.sample_for_batch("rotation", batchsize), - translation=self.sample_for_batch("translation", batchsize), - batchsize=batchsize, + scale=scale, + rotation=rotation, + translation=translation, + batchsize=batch_size, ndim=ndim, degree=self.degree, device=device, @@ -441,22 +456,25 @@ def assemble_matrix(self, **data) -> torch.Tensor: ) return self.matrix - def sample_for_batch(self, name: str, batchsize: int) -> Optional[Union[Any, Sequence[Any]]]: + def sample_for_batch_with_prob(self, name: str, batch_size: int, *, default_value: float, seed: int): """ - Sample elements for batch - - Args: - name: name of parameter - batchsize: batch size - - Returns: - Optional[Union[Any, Sequence[Any]]]: sampled elements + sampling batch with self.p and self.per_sample """ - elem = getattr(self, name) - if elem is not None and self.per_sample: - return [elem] + [getattr(self, name) for _ in range(batchsize - 1)] - else: - return elem # either a single scalar value or None + batch_element = self.sample_for_batch(name, batch_size) + if batch_element is None: + return batch_element + with self.random_cxm(seed=int(seed)): + if not self.per_sample: + batch_element = batch_element if torch.rand(1) < self.p else default_value + return batch_element + elem_length = len(batch_element[0]) + if elem_length == 1: + return [v if torch.rand(1) < self.p else torch.as_tensor(default_value) for v in batch_element] + else: + return [ + v if torch.rand(1) < self.p else tuple(torch.as_tensor(default_value) for _ in range(elem_length)) + for v in batch_element + ] class Rotate(BaseAffine): @@ -476,11 +494,12 @@ def __init__( degree: bool = False, output_size: Optional[tuple] = None, adjust_size: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, reverse_order: bool = False, - **kwargs, + per_sample: bool = True, + p: float = 1, ): """ Args: @@ -516,14 +535,11 @@ def __init__( transformation to conform to the pytorch convention: transformation params order [W,H(,D)] and batch order [(D,)H,W] - **kwargs: additional keyword arguments passed to the - affine transform """ super().__init__( scale=None, rotation=rotation, translation=None, - matrix=None, keys=keys, grad=grad, degree=degree, @@ -533,7 +549,8 @@ def __init__( padding_mode=padding_mode, align_corners=align_corners, reverse_order=reverse_order, - **kwargs, + per_sample=per_sample, + p=p, ) @@ -552,11 +569,13 @@ def __init__( grad: bool = False, output_size: Optional[tuple] = None, adjust_size: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, unit: str = "pixel", reverse_order: bool = False, + per_sample: bool = True, + p: float = 1, **kwargs, ): """ @@ -600,7 +619,6 @@ def __init__( scale=None, rotation=None, translation=translation, - matrix=None, keys=keys, grad=grad, degree=False, @@ -610,6 +628,8 @@ def __init__( padding_mode=padding_mode, align_corners=align_corners, reverse_order=reverse_order, + per_sample=per_sample, + p=p, **kwargs, ) self.unit = unit @@ -647,10 +667,12 @@ def __init__( grad: bool = False, output_size: Optional[tuple] = None, adjust_size: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, reverse_order: bool = False, + per_sample: bool = True, + p: float = 1, **kwargs, ): """ @@ -700,7 +722,6 @@ def __init__( scale=scale, rotation=None, translation=None, - matrix=None, keys=keys, grad=grad, degree=False, @@ -710,6 +731,8 @@ def __init__( padding_mode=padding_mode, align_corners=align_corners, reverse_order=reverse_order, + per_sample=per_sample, + p=p, **kwargs, ) @@ -718,11 +741,11 @@ class Resize(Scale): def __init__( self, size: Union[int, Tuple[int]], - keys: Sequence = ("data",), + keys: Sequence[str] = ("data",), grad: bool = False, - interpolation_mode: str = "bilinear", - padding_mode: str = "zeros", - align_corners: bool = False, + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, reverse_order: bool = False, **kwargs, ): diff --git a/rising/transforms/channel.py b/rising/transforms/channel.py index 603d091d..5a3d67a7 100644 --- a/rising/transforms/channel.py +++ b/rising/transforms/channel.py @@ -2,13 +2,13 @@ import torch -from rising.transforms import BaseTransform +from rising.transforms import BaseTransform, BaseTransformMixin from rising.transforms.functional import one_hot_batch __all__ = ["OneHot", "ArgMax"] -class OneHot(BaseTransform): +class OneHot(BaseTransformMixin, BaseTransform): """ Convert to one hot encoding. One hot encoding is applied in first dimension which results in shape N x NumClasses x [same as input] while input is expected to @@ -38,10 +38,21 @@ def __init__( Input tensor needs to be of type torch.long. This could be achieved by applying `TenorOp("long", keys=("seg",))`. """ - super().__init__(augment_fn=one_hot_batch, keys=keys, grad=grad, num_classes=num_classes, dtype=dtype, **kwargs) + super().__init__( + augment_fn=one_hot_batch, + keys=keys, + grad=grad, + num_classes=num_classes, + dtype=dtype, + augment_fn_names=( + "num_classes", + "dtype", + ), + **kwargs + ) -class ArgMax(BaseTransform): +class ArgMax(BaseTransformMixin, BaseTransform): """ Compute argmax along given dimension. Can be used to revert OneHot encoding. @@ -60,4 +71,12 @@ def __init__(self, dim: int, keepdim: bool = True, keys: Sequence = ("seg",), gr Warnings The output of the argmax function is always a tensor of dtype long. """ - super().__init__(augment_fn=torch.argmax, keys=keys, grad=grad, dim=dim, keepdim=keepdim, **kwargs) + super().__init__( + augment_fn=torch.argmax, + keys=keys, + grad=grad, + dim=dim, + keepdim=keepdim, + augment_fn_names=("dim", "keepdim"), + **kwargs + ) diff --git a/rising/transforms/compose.py b/rising/transforms/compose.py index e7336358..77beada9 100644 --- a/rising/transforms/compose.py +++ b/rising/transforms/compose.py @@ -1,13 +1,18 @@ -from random import shuffle from typing import Any, Callable, Mapping, Optional, Sequence, Union +import numpy as np import torch from rising.random import ContinuousParameter, UniformParameter -from rising.transforms import AbstractTransform +from rising.transforms import _AbstractTransform +from rising.transforms.sitk import _ITKTransform from rising.utils import check_scalar -__all__ = ["Compose", "DropoutCompose", "OneOf"] +__all__ = [ + "Compose", + "DropoutCompose", + "OneOf", +] def dict_call(batch: dict, transform: Callable) -> Any: @@ -21,6 +26,10 @@ def dict_call(batch: dict, transform: Callable) -> Any: Returns: Any: transformed batch """ + if not isinstance(batch, Mapping): + raise RuntimeError( + "You may want to pass `default_transform_call` from rising.loading as `transform_call` to `Compose`" + ) return transform(**batch) @@ -56,14 +65,14 @@ def forward(self, *args, **kwargs) -> Any: return self.trafo(*args, **kwargs) -class Compose(AbstractTransform): +class Compose(_AbstractTransform): """ Compose multiple transforms """ def __init__( self, - *transforms: Union[AbstractTransform, Sequence[AbstractTransform]], + *transforms: Union[_AbstractTransform, Sequence[_AbstractTransform]], shuffle: bool = False, transform_call: Callable[[Any, Callable], Any] = dict_call, ): @@ -103,10 +112,13 @@ def forward(self, *seq_like, **map_like) -> Union[Sequence, Mapping]: assert len(self.transforms) == len(self.transform_order) data = seq_like if seq_like else map_like + seed = int(torch.randint(0, int(1e6), (1,))) + torch.manual_seed(seed) if self.shuffle: - shuffle(self.transform_order) + self.shuffle_transform() for idx in self.transform_order: + torch.manual_seed(seed + idx) data = self.transform_call(data, self.transforms[idx]) return data @@ -121,7 +133,7 @@ def transforms(self) -> torch.nn.ModuleList: return self._transforms @transforms.setter - def transforms(self, transforms: Union[AbstractTransform, Sequence[AbstractTransform]]): + def transforms(self, transforms: Union[_AbstractTransform, Sequence[_AbstractTransform]]): """ Transforms setter @@ -163,6 +175,20 @@ def shuffle(self, shuffle: bool): self._shuffle = shuffle self.transform_order = list(range(len(self.transforms))) + def shuffle_transform(self): + """ + this function shuffles the Tensor transformation, and exclude all others such as ITK based ones. + """ + tensor_transform_indicator = [] + for i, trans in enumerate(self.transforms): + assert isinstance(trans, _AbstractTransform) + if not isinstance(trans, _ITKTransform): + tensor_transform_indicator.append(i) + transform_mapping = { + k: v for k, v in zip(tensor_transform_indicator, np.random.permutation(tensor_transform_indicator)) + } + self.transform_order = [transform_mapping.get(i, i) for i in self.transform_order] + class DropoutCompose(Compose): """ @@ -171,7 +197,7 @@ class DropoutCompose(Compose): def __init__( self, - *transforms: Union[AbstractTransform, Sequence[AbstractTransform]], + *transforms: Union[_AbstractTransform, Sequence[_AbstractTransform]], dropout: Union[float, Sequence[float]] = 0.5, shuffle: bool = False, random_sampler: ContinuousParameter = None, @@ -239,14 +265,14 @@ def forward(self, *seq_like, **map_like) -> Union[Sequence, Mapping]: return data -class OneOf(AbstractTransform): +class OneOf(_AbstractTransform): """ Apply one of the given transforms. """ def __init__( self, - *transforms: Union[AbstractTransform, Sequence[AbstractTransform]], + *transforms: Union[_AbstractTransform, Sequence[_AbstractTransform]], weights: Optional[Sequence[float]] = None, p: float = 1.0, transform_call: Callable[[Any, Callable], Any] = dict_call, @@ -265,7 +291,7 @@ def __init__( transforms = transforms[0] if not transforms: raise ValueError("At least one transformation needs to be selected.") - self.transforms = transforms + self.transforms = torch.nn.ModuleList(transforms) if weights is not None and len(weights) != len(transforms): raise ValueError( @@ -285,3 +311,6 @@ def forward(self, **data) -> dict: index = torch.multinomial(self.weights, 1) data = self.transform_call(data, self.transforms[int(index)]) return data + + def extra_repr(self) -> str: + return ", ".join([str(x) for x in self.transforms]) diff --git a/rising/transforms/crop.py b/rising/transforms/crop.py index cddb4768..ae026c06 100644 --- a/rising/transforms/crop.py +++ b/rising/transforms/crop.py @@ -1,18 +1,18 @@ from typing import Sequence, Union -import torch - -from rising.random import AbstractParameter -from rising.transforms.abstract import BaseTransform, BaseTransformSeeded +from rising.transforms.abstract import BaseTransform, BaseTransformMixin, PerSampleTransformMixin from rising.transforms.functional.crop import center_crop, random_crop +from rising.transforms.functional.crop_pad import pad_random_crop, pad_center_crop -__all__ = ["CenterCrop", "RandomCrop"] +__all__ = ["CenterCrop", "RandomCrop", "PadRandomCrop", "PadCenterCrop"] -class CenterCrop(BaseTransform): - def __init__( - self, size: Union[int, Sequence, AbstractParameter], keys: Sequence = ("data",), grad: bool = False, **kwargs - ): +class CenterCrop(BaseTransformMixin, BaseTransform): + """ + CenterCrop input image given size. + """ + + def __init__(self, *, size: Union[int, Sequence], keys: Sequence = ("data",), grad: bool = False, **kwargs): """ Args: size: size of crop @@ -20,17 +20,21 @@ def __init__( grad: enable gradient computation inside transformation **kwargs: keyword arguments passed to augment_fn """ - super().__init__(augment_fn=center_crop, keys=keys, grad=grad, property_names=("size",), size=size, **kwargs) + super().__init__(augment_fn=center_crop, keys=keys, grad=grad, augment_fn_names=("size",), size=size, **kwargs) -class RandomCrop(BaseTransformSeeded): +class RandomCrop(BaseTransformMixin, BaseTransform): + """ + RandomCrop images given size and dist (distance to the border) + """ + def __init__( - self, - size: Union[int, Sequence, AbstractParameter], - dist: Union[int, Sequence, AbstractParameter] = 0, - keys: Sequence = ("data",), - grad: bool = False, - **kwargs + self, + *, + size: Union[int, Sequence], + dist: Union[int, Sequence] = 0, + keys: Sequence = ("data",), + grad: bool = False, ): """ Args: @@ -38,14 +42,72 @@ def __init__( dist: minimum distance to border. By default zero keys: keys which should be augmented grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn """ super().__init__( augment_fn=random_crop, keys=keys, + grad=grad, size=size, dist=dist, + seeded=True, + augment_fn_names=("size", "dist"), + ) + + +class PadRandomCrop(PerSampleTransformMixin, BaseTransform): + """ + This operation pad the image and crop to the desired size. + """ + + def __init__( + self, + size: Union[int, Sequence], + pad_size: Union[int, Sequence[int]] = 0, + pad_value: Union[int, float, Sequence[int], Sequence[float]] = 0, + keys: Sequence = ("data",), + grad: bool = False, + ): + """ + Args: + size: random crop size + pad_size: int, sequence[int], padding image to size+pad + pad_value: int, float or a list of them. the value to pad + """ + super(PadRandomCrop, self).__init__( + augment_fn=pad_random_crop, + keys=keys, grad=grad, - property_names=("size", "dist"), - **kwargs + size=size, + pad_size=pad_size, + pad_value=pad_value, + augment_fn_names=("size", "pad_size", "pad_value"), + paired_kw_names=("pad_value",), + ) + + +class PadCenterCrop(PerSampleTransformMixin, BaseTransform): + + def __init__( + self, + size: Union[int, Sequence], + pad_size: Union[int, Sequence[int]] = 0, + pad_value: Union[int, float, Sequence[int], Sequence[float]] = 0, + keys: Sequence = ("data",), + grad: bool = False, + ): + """ + Args: + size: random crop size + pad_size: int, sequence[int], padding image to size+pad + pad_value: int, float or a list of them. the value to pad + """ + super(PadCenterCrop, self).__init__( + augment_fn=pad_center_crop, + keys=keys, + grad=grad, + size=size, + pad_size=pad_size, + pad_value=pad_value, + augment_fn_names=("size", "pad_size", "pad_value"), + paired_kw_names=("pad_value",), ) diff --git a/rising/transforms/format.py b/rising/transforms/format.py index 0c943232..e5697cba 100644 --- a/rising/transforms/format.py +++ b/rising/transforms/format.py @@ -1,13 +1,12 @@ from typing import Callable, Dict, Hashable, Mapping, Sequence, Tuple, Union +from rising.transforms.abstract import _AbstractTransform from rising.transforms.functional.utility import filter_keys, pop_keys -from .abstract import AbstractTransform - __all__ = ["MapToSeq", "SeqToMap", "PopKeys", "FilterKeys", "RenameKeys"] -class MapToSeq(AbstractTransform): +class MapToSeq(_AbstractTransform): """ Convert dict to sequence """ @@ -37,7 +36,7 @@ def forward(self, **data) -> tuple: return tuple(data[_k] for _k in self.keys) -class SeqToMap(AbstractTransform): +class SeqToMap(_AbstractTransform): """Convert sequence to dict""" def __init__(self, *keys, grad: bool = False, **kwargs): @@ -65,7 +64,7 @@ def forward(self, *data, **kwargs) -> dict: return {_key: data[_idx] for _idx, _key in enumerate(self.keys)} -class PopKeys(AbstractTransform): +class PopKeys(_AbstractTransform): """ Pops keys from a given data dict """ @@ -88,7 +87,7 @@ def forward(self, **data) -> Union[dict, Tuple[dict, dict]]: return pop_keys(data=data, keys=self.keys, return_popped=self.return_popped) -class FilterKeys(AbstractTransform): +class FilterKeys(_AbstractTransform): """ Filters keys from a given data dict """ @@ -111,7 +110,7 @@ def forward(self, **data) -> Union[dict, Tuple[dict, dict]]: return filter_keys(data=data, keys=self.keys, return_popped=self.return_popped) -class RenameKeys(AbstractTransform): +class RenameKeys(_AbstractTransform): """Rename keys inside batch""" def __init__(self, keys: Mapping[Hashable, Hashable]): diff --git a/rising/transforms/functional/__init__.py b/rising/transforms/functional/__init__.py index f894ab6b..b830de50 100644 --- a/rising/transforms/functional/__init__.py +++ b/rising/transforms/functional/__init__.py @@ -15,7 +15,7 @@ """ from rising.transforms.functional.channel import one_hot_batch -from rising.transforms.functional.crop import center_crop, crop, random, random_crop +from rising.transforms.functional.crop import center_crop, crop, random_crop from rising.transforms.functional.intensity import ( add_noise, add_value, @@ -30,6 +30,7 @@ scale_by_value, ) from rising.transforms.functional.painting import local_pixel_shuffle, random_inpainting, random_outpainting +from rising.transforms.functional.sitk import itk_resample, itk_clip, itk2tensor from rising.transforms.functional.spatial import mirror, resize_native, rot90 from rising.transforms.functional.tensor import tensor_op, to_device_dtype from rising.transforms.functional.utility import box_to_seg, filter_keys, instance_to_semantic, pop_keys, seg_to_box diff --git a/rising/transforms/functional/affine.py b/rising/transforms/functional/affine.py index 1c03be6e..a0dced5a 100644 --- a/rising/transforms/functional/affine.py +++ b/rising/transforms/functional/affine.py @@ -23,6 +23,7 @@ "create_scale", "create_translation", "parametrize_matrix", + "AffineParamType", ] from rising.utils.inverse import orthogonal_inverse @@ -205,7 +206,7 @@ def create_rotation( """ if rotation is None: rotation = 0 - num_rot_params = 1 if ndim == 2 else ndim + num_rot_params = 1 if ndim == 2 else ndim # this prevents to put 2 dimensional input for 2d images. rotation = expand_scalar_param(rotation, batchsize, num_rot_params).to(device=device, dtype=dtype) if degree: @@ -446,6 +447,7 @@ def affine_image_transform( pixels. If set to False, they are instead considered as referring to the corner points of the input’s corner pixels, making the sampling more resolution agnostic. + reverse_order: todo to add the specific details. Returns: torch.Tensor: transformed image diff --git a/rising/transforms/functional/crop.py b/rising/transforms/functional/crop.py index ccc96512..93e0d9b7 100644 --- a/rising/transforms/functional/crop.py +++ b/rising/transforms/functional/crop.py @@ -1,5 +1,4 @@ -import random -from typing import List, Sequence, Tuple, Union +from typing import Sequence, Union import torch @@ -8,76 +7,121 @@ __all__ = ["crop", "center_crop", "random_crop"] -def crop(data: torch.Tensor, corner: Sequence[int], size: Sequence[int]) -> torch.Tensor: +def crop(data: torch.Tensor, corner: Sequence[int], size: Sequence[int], grid_crop: bool = False): """ Extract crop from last dimensions of data - Args: - data: input tensor - corner: top left corner point - size: size of patch - - Returns: - torch.Tensor: cropped data + Parameters + ---------- + data: torch.Tensor + input tensor [... , spatial dims] spatial dims can be arbitrary + spatial dimensions. Leading dimensions will be preserved. + corner: Sequence[int] + top left corner point + size: Sequence[int] + size of patch + grid_crop: bool + crop from grid of shape [N, spatial dims, NDIM], where N is the batch + size, spatial dims can be arbitrary spatial dimensions and NDIM + is the number of spatial dimensions + + Returns + ------- + torch.Tensor + cropped data """ _slices = [] - if len(corner) < data.ndim: - for i in range(data.ndim - len(corner)): + ndim = data.ndimension() - int(bool(grid_crop)) # alias for ndim() + if len(corner) < ndim: + for i in range(ndim - len(corner)): _slices.append(slice(0, data.shape[i])) _slices = _slices + [slice(c, c + s) for c, s in zip(corner, size)] + if grid_crop: + _slices.append(slice(0, data.shape[-1])) return data[_slices] -def center_crop(data: torch.Tensor, size: Union[int, Sequence[int]]) -> torch.Tensor: +def center_crop(data: torch.Tensor, size: Union[int, Sequence[int]], grid_crop: bool = False) -> torch.Tensor: """ Crop patch from center - Args: - data: input tensor - size: size of patch - - Returns: - torch.Tensor: output tensor cropped from input tensor + Parameters + ---------- + data: torch.Tensor + input tensor [... , spatial dims] spatial dims can be arbitrary + spatial dimensions. Leading dimensions will be preserved. + size: Union[int, Sequence[int]] + size of patch + grid_crop: bool + crop from grid of shape [N, spatial dims, NDIM], where N is the batch + size, spatial dims can be arbitrary spatial dimensions and NDIM + is the number of spatial dimensions + + Returns + ------- + torch.Tensor + output tensor cropped from input tensor """ if check_scalar(size): size = [size] * (data.ndim - 2) if not isinstance(size[0], int): size = [int(s) for s in size] - corner = [int(round((img_dim - crop_dim) / 2.0)) for img_dim, crop_dim in zip(data.shape[2:], size)] - return crop(data, corner, size) + if grid_crop: + data_shape = data.shape[1:-1] # N, H, W, (D), NDIM + else: + data_shape = data.shape[2:] + + corner = [int(round((img_dim - crop_dim) / 2.0)) for img_dim, crop_dim in zip(data_shape, size)] + return crop(data, corner, size, grid_crop=grid_crop) def random_crop( - data: torch.Tensor, size: Union[int, Sequence[int]], dist: Union[int, Sequence[int]] = 0 + data: torch.Tensor, size: Union[int, Sequence[int]], dist: Union[int, Sequence[int]] = 0, grid_crop: bool = False ) -> torch.Tensor: """ Crop random patch/volume from input tensor - - Args: - data: input tensor - size: size of patch/volume - dist: minimum distance to border. By default zero - - Returns: - torch.Tensor: cropped output - List[int]: top left corner used for crop + This function crop images by batch with the same random state. + + Parameters + ---------- + data: torch.Tensor + input tensor [... , spatial dims] spatial dims can be arbitrary + spatial dimensions. Leading dimensions will be preserved. + size: Union[int, Sequence[int]] + size of patch/volume + dist: Union[int, Sequence[int]] + minimum distance to border. By default zero + grid_crop: bool + crop from grid of shape [N, spatial dims, NDIM], where N is the batch + size, spatial dims can be arbitrary spatial dimensions and NDIM + is the number of spatial dimensions + + Returns + ------- + torch.Tensor + cropped output """ if check_scalar(dist): dist = [dist] * (data.ndim - 2) - if isinstance(dist[0], torch.Tensor): - dist = [int(i) for i in dist] if check_scalar(size): size = [size] * (data.ndim - 2) if not isinstance(size[0], int): size = [int(s) for s in size] - if any([crop_dim + dist_dim >= img_dim for img_dim, crop_dim, dist_dim in zip(data.shape[2:], size, dist)]): - raise TypeError(f"Crop can not be realized with given size {size} and dist {dist}.") + if grid_crop: + data_shape = data.shape[1:-1] + else: + data_shape = data.shape[2:] + + if any([crop_dim + dist_dim > img_dim for img_dim, crop_dim, dist_dim in zip(data_shape, size, dist)]): + raise TypeError( + f"Crop can not be realized with given size {size} and dist {dist}, " f"given input shape {data_shape}" + ) corner = [ - torch.randint(0, img_dim - crop_dim - dist_dim, (1,)).item() - for img_dim, crop_dim, dist_dim in zip(data.shape[2:], size, dist) + int(torch.randint(0, max(int(img_dim - crop_dim - dist_dim), 1), (1,))) + for img_dim, crop_dim, dist_dim in zip(data_shape, size, dist) ] - return crop(data, corner, size) + return crop(data, corner, size, grid_crop=grid_crop) diff --git a/rising/transforms/functional/crop_pad.py b/rising/transforms/functional/crop_pad.py new file mode 100644 index 00000000..9e17f251 --- /dev/null +++ b/rising/transforms/functional/crop_pad.py @@ -0,0 +1,44 @@ +from typing import Sequence, Union + +import torch + +from rising.transforms.functional import random_crop, center_crop + + +def pad_random_crop( + data: torch.Tensor, size: Union[int, Sequence[int]], pad_size=Union[int, Sequence[int]], pad_value=0 +): + ndim = data.dim() - 2 + if isinstance(size, (float, int)): + size = [size] * ndim + if isinstance(pad_size, (float, int)): + pad_size = [pad_size] * ndim + assert len(size) == len(pad_size) == ndim + + from rising.transforms import Pad + + data = Pad(pad_size=[x + y for x, y in zip(size, pad_size)], pad_value=pad_value, keys=("data",))(data=data)["data"] + return random_crop( + data, + size=size, + dist=0, + ) + + +def pad_center_crop( + data: torch.Tensor, size: Union[int, Sequence[int]], pad_size=Union[int, Sequence[int]], pad_value=0 +): + ndim = data.dim() - 2 + if isinstance(size, (float, int)): + size = [size] * ndim + if isinstance(pad_size, (float, int)): + pad_size = [pad_size] * ndim + assert len(size) == len(pad_size) == ndim + + from rising.transforms import Pad + + data = Pad(pad_size=[x + y for x, y in zip(size, pad_size)], pad_value=pad_value, keys=("data",))(data=data)["data"] + return center_crop( + data, + size=size, + ) diff --git a/rising/transforms/functional/intensity.py b/rising/transforms/functional/intensity.py index 29ffe3b2..f027470c 100644 --- a/rising/transforms/functional/intensity.py +++ b/rising/transforms/functional/intensity.py @@ -1,3 +1,4 @@ +import warnings from typing import Optional, Sequence, Union import torch @@ -17,6 +18,7 @@ "clamp", "bezier_3rd_order", "random_inversion", + "norm_min_max_percentile", ] @@ -43,7 +45,7 @@ def norm_range( Scale range of tensor Args: - data: input data. Per channel option supports [C,H,W] and [C,H,W,D]. + data: input data. Per channel option supports [B,C,H,W] and [B,C,H,W,D]. min: minimal value max: maximal value per_channel: range is normalized per channel @@ -52,6 +54,7 @@ def norm_range( Returns: torch.Tensor: normalized data """ + assert data.shape[0] == 1, f"per sample example (batch_size = 1) as the input data, given {data.shape}" if out is None: out = torch.zeros_like(data) @@ -68,7 +71,7 @@ def norm_min_max( Scale range to [0,1] Args: - data: input data. Per channel option supports [C,H,W] and [C,H,W,D]. + data: input data without batch dimension. Per channel option supports [B,C,H,W] and [B,C,H,W,D]. per_channel: range is normalized per channel out: if provided, result is saved in here eps: small constant for numerical stability. @@ -77,6 +80,7 @@ def norm_min_max( Returns: torch.Tensor: scaled data """ + assert data.shape[0] == 1, f"per sample example (batch_size = 1) as the input data, given {data.shape}" def _norm(_data: torch.Tensor, _out: torch.Tensor): _min = _data.min() @@ -90,13 +94,58 @@ def _norm(_data: torch.Tensor, _out: torch.Tensor): out = torch.zeros_like(data) if per_channel: - for _c in range(data.shape[0]): - out[_c] = _norm(data[_c], out[_c]) + for _c in range(data.shape[1]): + out[:, _c] = _norm(data[:, _c], out[:, _c]) else: out = _norm(data, out) return out +def norm_min_max_percentile( + data: torch.Tensor, + min: float, + max: float, + per_channel: bool = True, + out: Optional[torch.Tensor] = None, + eps=1e-8, +): + """ + Normalize data based on min and max percentile between 0 - 100. + Args: + data: input data, under form of [B,C,H,W] and [B,C,H,W,D]. + min: low percentile to clamp, between 0 and 100 + max: high percentile to clam, between 0 and 100 + per_channel: if compute the percentile on channel + out: + eps: small eps to prevent division by 0.l + Returns: + torch.Tensor: normalized data + """ + + if out is None: + out = torch.empty_like(data) + + if per_channel: + for i in range(data.shape[1]): + min_ = torch.quantile(data[:, i].float(), float(min)) + max_ = torch.quantile(data[:, i].float().float(), float(max)) + out[:, i] = clamp( + data[:, i], + min=float(min_), + max=float(max_), + ) + else: + min_ = torch.quantile(data, float(min)) + max_ = torch.quantile(data, float(max)) + out = clamp( + data, + min=float(min_), + max=float(max_), + ) + + return norm_min_max(out, per_channel=per_channel, out=out, eps=eps) + + def norm_zero_mean_unit_std( data: torch.Tensor, per_channel: bool = True, out: Optional[torch.Tensor] = None, eps: Optional[float] = 1e-8 ) -> torch.Tensor: @@ -113,6 +162,7 @@ def norm_zero_mean_unit_std( Returns: torch.Tensor: normalized data """ + assert data.shape[0] == 1, f"per sample example (batch_size = 1) as the input data, given {data.shape}" def _norm(_data: torch.Tensor, _out: torch.Tensor): denom = _data.std() @@ -152,6 +202,8 @@ def norm_mean_std( Returns: torch.Tensor: normalized data """ + assert data.shape[0] == 1, f"per sample example (batch_size = 1) as the input data, given {data.shape}" + if out is None: out = torch.zeros_like(data) @@ -204,6 +256,10 @@ def gamma_correction(data: torch.Tensor, gamma: float) -> torch.Tensor: Returns: torch.Tensor: gamma corrected data """ + min_, max_ = data.min().detach(), data.max().detach() + if min_ < 0 or max_ > 1: + warnings.warn("`data` range not in [0, 1]", RuntimeWarning) + if torch.is_tensor(gamma): gamma = gamma.to(data) return data.pow(gamma) @@ -246,12 +302,13 @@ def scale_by_value(data: torch.Tensor, value: float, out: Optional[torch.Tensor] def bezier_3rd_order( data: torch.Tensor, maxv: float = 1.0, minv: float = 0.0, out: Optional[torch.Tensor] = None ) -> torch.Tensor: - p0 = torch.zeros((1, 2)) - p1 = torch.rand((1, 2)) - p2 = torch.rand((1, 2)) - p3 = torch.ones((1, 2)) + device, dtype = data.device, data.dtype + p0 = torch.zeros((1, 2), device=device, dtype=dtype) + p1 = torch.rand((1, 2), device=device, dtype=dtype) + p2 = torch.rand((1, 2), device=device, dtype=dtype) + p3 = torch.ones((1, 2), device=device, dtype=dtype) - t = torch.linspace(0.0, 1.0, 1000).unsqueeze(1) + t = torch.linspace(0.0, 1.0, 1000, device=device, dtype=dtype).unsqueeze(1) points = (1 - t * t * t) * p0 + 3 * (1 - t) * (1 - t) * t * p1 + 3 * (1 - t) * t * t * p2 + t * t * t * p3 @@ -273,7 +330,6 @@ def random_inversion( minv: float = 0.0, out: Optional[torch.Tensor] = None, ) -> torch.Tensor: - if torch.rand((1)) < prob_inversion: # Inversion of curve out = maxv + minv - data @@ -282,3 +338,19 @@ def random_inversion( out = data return out + + +def augment_rician_noise(data: torch.Tensor, std: float, keep_range=False): + """augment with rician noise + Args: + data:Tensor input data having dimension of BCHW(D) + std: the std of the noise, float + keep_range: if keep the image range, default False + """ + min_, max_ = data.min(), data.max() + data = torch.sqrt( + (data + torch.randn_like(data) * float(std)).pow(2) + (torch.randn_like(data) * float(std)) ** 2 + ) * torch.sign(data) + if keep_range: + data = clamp(data, min=min_, max=max_) + return data diff --git a/rising/transforms/functional/pad.py b/rising/transforms/functional/pad.py new file mode 100644 index 00000000..3855f865 --- /dev/null +++ b/rising/transforms/functional/pad.py @@ -0,0 +1,41 @@ +from typing import Sequence, Union + +from torch import Tensor +from torch.nn import functional as F + +from rising.utils import check_scalar +from rising.utils.mise import ntuple + + +def pad(data: Tensor, pad_size: Union[int, Sequence[int]], grid_pad=False, mode="constant", value: float = 0.0): + """ + Args: + data: input data with size [B,C,H,W,(D)] + pad_size: int or seq of int. the dimension to pad, following functional.pad function convention, where + order is inversed. + grid_pad: bool must be False, True is not implemented. + mode: str, padding mode, following functional.pad + ``'constant'``, ``'reflect'``, ``'replicate'`` or ``'circular'`` + value: float, padding value. + """ + n_dim = data.dim() - 2 + # padding parameters + if check_scalar(pad_size): + pad_size = ntuple(n_dim * 2)(pad_size) + elif isinstance(pad_size, Sequence): + pad_size = tuple(pad_size) + if not (len(pad_size) == 0 or len(pad_size) != n_dim or len(pad_size) != n_dim * 2): + raise TypeError(pad_size) + if len(pad_size) == n_dim: + pad_size = tuple((z for double in zip(pad_size, pad_size) for z in double)) + elif (len(pad_size)) == n_dim * 2: + pad_size = pad_size + else: + raise RuntimeError(pad_size) + assert isinstance(pad_size, tuple) and len(pad_size) == 2 * n_dim + + # todo: understand the grid_pad for affine distribution + if grid_pad is False: + return F.pad(data, pad=pad_size, mode=mode, value=value) + else: + raise NotImplementedError(grid_pad) diff --git a/rising/transforms/functional/sitk.py b/rising/transforms/functional/sitk.py new file mode 100644 index 00000000..0d045862 --- /dev/null +++ b/rising/transforms/functional/sitk.py @@ -0,0 +1,54 @@ +from typing import Union, Tuple + +import SimpleITK as sitk +import numpy as np +import torch + +from rising.utils import check_scalar + + +def itk_resample(image: sitk.Image, spacing: Union[float, Tuple[float, float, float]], *, + interpolation: str = "nearest", pad_value: int) -> sitk.Image: + """ + resample sitk image given spacing, pad value and interpolation. + + Args: + image: sitk image + spacing: new spacing, either a scalar or a tuple of three scalars. + interpolation: interpolation method, "linear" or "nearest". + pad_value: pad value for out of space pixels. + + Returns: + torch.Tensor: affine params in correct shape + """ + if check_scalar(spacing): + spacing: Tuple[float, float, float] = (spacing, spacing, spacing) # noqa + + ori_spacing = image.GetSpacing() + ori_size = image.GetSize() + new_size = (round(ori_size[0] * (ori_spacing[0] / spacing[0])), + round(ori_size[1] * (ori_spacing[1] / spacing[1])), + round(ori_size[2] * (ori_spacing[2] / spacing[2]))) + interp = {"linear": sitk.sitkLinear, "nearest": sitk.sitkNearestNeighbor, "cosine": sitk.sitkCosineWindowedSinc}[ + interpolation] + return sitk.Resample(image, new_size, sitk.Transform(), interp, image.GetOrigin(), spacing, image.GetDirection(), + pad_value, image.GetPixelID()) + + +def itk_clip(image: sitk.Image, low: int, high: int) -> sitk.Image: + """ + clamp sitk image given low and high values, used to windows CT images. + Args: + image: sitk image + low: low threshold to clip, type: int + high: high threshold to clip, type: int + Returns: + sitk.Image + """ + assert low < high, (low, high) + return sitk.Clamp(image, sitk.sitkInt16, int(low), int(high)) + + +def itk2tensor(image: sitk.Image, *, dtype=torch.float): + np_array = sitk.GetArrayFromImage(image).astype(float, copy=False) + return torch.from_numpy(np_array).to(dtype)[None, ...] diff --git a/rising/transforms/functional/spatial.py b/rising/transforms/functional/spatial.py index 2b47c051..c33b1051 100644 --- a/rising/transforms/functional/spatial.py +++ b/rising/transforms/functional/spatial.py @@ -6,13 +6,18 @@ __all__ = ["mirror", "rot90", "resize_native"] +from rising.constants import FInterpolation -def mirror(data: torch.Tensor, dims: Union[int, Sequence[int]]) -> torch.Tensor: + +def mirror( + data: torch.Tensor, + dims: Union[int, Sequence[int]], +) -> torch.Tensor: """ Mirror data at dims Args: - data: input data + data: input data [B,C,H,W,D] dims: dimensions to mirror Returns: @@ -30,24 +35,25 @@ def rot90(data: torch.Tensor, k: int, dims: Union[int, Sequence[int]]): Rotate 90 degrees around dims Args: - data: input data + data: input data [B,C,H,W,D] k: number of times to rotate dims: dimensions to mirror Returns: torch.Tensor: tensor with mirrored dimensions """ + dims = [dims, ] if check_scalar(dims) else dims # type: ignore dims = [int(d + 2) for d in dims] return torch.rot90(data, int(k), dims) def resize_native( - data: torch.Tensor, - size: Optional[Union[int, Sequence[int]]] = None, - scale_factor: Optional[Union[float, Sequence[float]]] = None, - mode: str = "nearest", - align_corners: Optional[bool] = None, - preserve_range: bool = False, + data: torch.Tensor, + size: Optional[Union[int, Sequence[int]]] = None, + scale_factor: Optional[Union[float, Sequence[float]]] = None, + mode: str = "nearest", + align_corners: Optional[bool] = None, + preserve_range: bool = False, ): """ Down/up-sample sample to either the given :attr:`size` or the given @@ -76,8 +82,11 @@ def resize_native( if check_scalar(scale_factor): # pytorch internally checks for an iterable. Single value tensors are still iterable scale_factor = float(scale_factor) + elif scale_factor is not None: + assert len(scale_factor) == len(data.shape) - 2, "scale_factor must be scalar or have same length as data" + scale_factor = [float(f) for f in scale_factor] out = torch.nn.functional.interpolate( - data, size=size, scale_factor=scale_factor, mode=mode, align_corners=align_corners + data, size=size, scale_factor=scale_factor, mode=FInterpolation(mode).value, align_corners=align_corners ) if preserve_range: diff --git a/rising/transforms/functional/tensor.py b/rising/transforms/functional/tensor.py index 05502bce..7d1e74bf 100644 --- a/rising/transforms/functional/tensor.py +++ b/rising/transforms/functional/tensor.py @@ -41,6 +41,7 @@ def to_device_dtype( Args: data: data which should be pushed to device. Sequence and mapping items are mapping individually to gpu + dtype: target dtype device: target device kwargs: keyword arguments passed to assigning function diff --git a/rising/transforms/grid.py b/rising/transforms/grid.py new file mode 100644 index 00000000..7b9d5db3 --- /dev/null +++ b/rising/transforms/grid.py @@ -0,0 +1,254 @@ +from abc import abstractmethod +from typing import Dict, Optional, Sequence, Tuple, Union + +import torch +from torch import Tensor +from torch.nn import functional as F + +from rising.random.utils import fix_random_seed_ctx +from rising.transforms.abstract import ItemSeq, _AbstractTransform +from rising.transforms.functional import center_crop, random_crop +from rising.transforms.kernel import GaussianSmoothing +from rising.utils.affine import get_batched_eye, matrix_to_homogeneous +from rising.utils.mise import ntuple + +__all__ = [ + "GridTransform", + "StackedGridTransform", + "CenterCropGrid", + "RandomCropGrid", + "ElasticDistortion", + "RadialDistortion", +] + + +class GridTransform(_AbstractTransform): + """ + Abstract class for grid transformation. + """ + + def __init__( + self, + keys: Sequence[str] = ("data",), + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, + grad: bool = False, + **kwargs, + ): + super().__init__(grad=grad) + self.keys = keys + self._tuple_generator = ntuple(len(self.keys)) + self.interpolation_mode: Sequence[str] = self._tuple_generator(interpolation_mode) + self.padding_mode: Sequence[str] = self._tuple_generator(padding_mode) + self.align_corners: Sequence[bool] = self._tuple_generator(align_corners) + self.kwargs = kwargs + + self.grid: Optional[Dict[str, Tensor]] = None + + def forward(self, **data) -> dict: + device, dtype = data[self.keys[0]].device, data[self.keys[0]].dtype + + if self.grid is None: + self.grid = self.create_grid(data, device=device, dtype=dtype) + + self.grid = self.augment_grid(self.grid, device=device, dtype=dtype) + + for key, interpol, padding_mode, align_corners in zip( + self.keys, self.interpolation_mode, self.padding_mode, self.align_corners + ): + data[key] = F.grid_sample( + data[key], self.grid[key], mode=interpol, padding_mode=padding_mode, align_corners=align_corners + ) + self.grid = None + return data + + @abstractmethod + def augment_grid(self, grid: Dict[str, Tensor], *, device, dtype) -> Dict[str, Tensor]: + """ + this functions modifies the grid + """ + raise NotImplementedError + + def create_grid( + self, data: Dict[str, Tensor], matrix: Tensor = None, *, device: torch.device, dtype: torch.dtype + ) -> Dict[str, Tensor]: + grid = {} + for key, align_corners in zip(self.keys, self.align_corners): + cur_data = data[key] + batch_size = cur_data.shape[0] + ndim = cur_data.dim() - 2 + if matrix is None: + matrix = get_batched_eye(batchsize=batch_size, ndim=ndim, device=device, dtype=dtype) + matrix = matrix_to_homogeneous(matrix)[:, :-1] + + grid[key] = F.affine_grid(matrix, size=list(cur_data.shape), align_corners=align_corners) + return grid + + def __add__(self, other): + if not isinstance(other, GridTransform): + raise ValueError("Concatenation is only supported for grid transforms.") + return StackedGridTransform(self, other) + + def __radd__(self, other): + if not isinstance(other, GridTransform): + raise ValueError("Concatenation is only supported for grid transforms.") + return StackedGridTransform(other, self) + + +class StackedGridTransform(GridTransform): + def __init__(self, *transforms: Union[GridTransform, Sequence[GridTransform]]): + super().__init__(keys=None, interpolation_mode=None, padding_mode=None, align_corners=None) + if isinstance(transforms, (tuple, list)): + if isinstance(transforms[0], (tuple, list)): + transforms = transforms[0] + self.transforms = transforms + + def create_grid( + self, data: Dict[str, Tensor], matrix: Tensor = None, *, device: torch.device, dtype: torch.dtype + ) -> Dict[str, Tensor]: + return self.transforms[0].create_grid(data=data, matrix=matrix, device=device, dtype=dtype) + + def augment_grid(self, grid: Dict[str, Tensor], *, device, dtype) -> Dict[str, Tensor]: + for transform in self.transforms: + grid = transform.augment_grid(grid, device=device, dtype=dtype) + return grid + + +class CenterCropGrid(GridTransform): + def __init__( + self, + *, + size: Union[int, Sequence[int]], + keys: Sequence[str] = ("data",), + interpolation_mode: str = "bilinear", + padding_mode: str = "zeros", + align_corners: bool = False, + grad: bool = False, + **kwargs, + ): + super().__init__( + keys=keys, + interpolation_mode=interpolation_mode, + padding_mode=padding_mode, + align_corners=align_corners, + grad=grad, + **kwargs, + ) + self.size = size + + def augment_grid(self, grid: Dict[str, Tensor], **kwargs) -> Dict[str, Tensor]: + return {key: center_crop(cur_grid, size=self.size, grid_crop=True) for key, cur_grid in grid.items()} + + +class RandomCropGrid(GridTransform): + def __init__( + self, + size: Union[int, Sequence[int]], + dist: Union[int, Sequence[int]] = 0, + keys: Sequence[str] = ("data",), + interpolation_mode: str = "bilinear", + padding_mode: str = "zeros", + align_corners: bool = False, + grad: bool = False, + **kwargs, + ): + super().__init__( + keys=keys, + interpolation_mode=interpolation_mode, + padding_mode=padding_mode, + align_corners=align_corners, + grad=grad, + **kwargs, + ) + self.size = size + self.dist = dist + + def augment_grid(self, grid: Dict[str, Tensor], **kwargs) -> Dict[str, Tensor]: + return {key: random_crop(item, size=self.size, dist=self.dist, grid_crop=True) for key, item in grid.items()} + + +class ElasticDistortion(GridTransform): + """ + ElasticDistortion transformation + """ + + def __init__( + self, + std: Union[float, Sequence[float]], + alpha: float, + dim: int = 2, + keys: Sequence[str] = ("data",), + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, + grad: bool = False, + per_sample: bool = True, + **kwargs, + ): + """ + std: std of the gaussian smooth + """ + super().__init__( + keys=keys, + interpolation_mode=interpolation_mode, + padding_mode=padding_mode, + align_corners=align_corners, + grad=grad, + **kwargs, + ) + self.std = std + self.alpha = alpha + self.per_sample = per_sample + self.gaussian = GaussianSmoothing(in_channels=1, kernel_size=7, std=self.std, dim=dim, stride=1, padding=3) + + def augment_grid(self, grid: Dict[Tuple, Tensor], *, device, dtype) -> Dict[Tuple, Tensor]: + seed = torch.randint(0, int(1e6), size=(1,)) + + def get_perturb_grid(batch_size: int = 1) -> Tensor: + random_offsets = torch.rand(batch_size, 1, *grid[key].shape[1:-1], device=device, dtype=dtype) * 2 - 1 + return self.gaussian(data=random_offsets)["data"] * self.alpha + + for key in grid.keys(): + cur_data = grid[key] + batch_size = cur_data.shape[0] + with fix_random_seed_ctx(seed): + if self.per_sample: + random_offsets = get_perturb_grid(batch_size)[:, 0, ..., None] + else: + random_offsets = get_perturb_grid(1)[:, 0, ..., None] + grid[key] += random_offsets + return grid + + +class RadialDistortion(GridTransform): + def __init__( + self, + scale: Tuple[float, float, float], + keys: Sequence[str] = ("data",), + interpolation_mode: ItemSeq[str] = "bilinear", + padding_mode: ItemSeq[str] = "zeros", + align_corners: ItemSeq[bool] = False, + grad: bool = False, + **kwargs, + ): + super().__init__( + keys=keys, + interpolation_mode=interpolation_mode, + padding_mode=padding_mode, + align_corners=align_corners, + grad=grad, + **kwargs, + ) + self.scale = scale + + def augment_grid(self, grid: Dict[Tuple, Tensor], **kwargs) -> Dict[Tuple, Tensor]: + new_grid = {key: radial_distortion_grid(cur_grid, scale=self.scale) for key, cur_grid in grid.items()} + return new_grid + + +def radial_distortion_grid(grid: Tensor, scale: Tuple[float, float, float]) -> Tensor: + dist = torch.norm(grid, p=2, dim=-1, keepdim=True) + dist = dist / dist.max() + distortion = (scale[0] * dist.pow(3) + scale[1] * dist.pow(2) + scale[2] * dist) / 3 + return grid * (1 - distortion) diff --git a/rising/transforms/intensity.py b/rising/transforms/intensity.py index f3ee73e0..80a9c07c 100644 --- a/rising/transforms/intensity.py +++ b/rising/transforms/intensity.py @@ -1,17 +1,24 @@ from typing import Optional, Sequence, Union -import torch - from rising.random import AbstractParameter -from rising.transforms.abstract import BaseTransform, PerChannelTransform, PerSampleTransform +from rising.transforms.abstract import ( + BaseTransform, + BaseTransformMixin, + PerChannelTransformMixin, + PerSampleTransformMixin, + ItemSeq, + augment_callable, +) from rising.transforms.functional.intensity import ( add_noise, add_value, + augment_rician_noise, bezier_3rd_order, clamp, gamma_correction, norm_mean_std, norm_min_max, + norm_min_max_percentile, norm_range, norm_zero_mean_unit_std, random_inversion, @@ -21,56 +28,60 @@ __all__ = [ "Clamp", "NormRange", + "NormPercentile", "NormMinMax", "NormZeroMeanUnitStd", "NormMeanStd", - "Noise", + "_Noise", "GaussianNoise", "ExponentialNoise", "GammaCorrection", - "RandomValuePerChannel", + "_RandomValuePerChannel", "RandomAddValue", "RandomScaleValue", "RandomBezierTransform", "InvertAmplitude", + "RicianNoiseTransform", ] -class Clamp(BaseTransform): +class Clamp(BaseTransformMixin, BaseTransform): """Apply augment_fn to keys""" def __init__( self, - min: Union[float, AbstractParameter], - max: Union[float, AbstractParameter], + min: ItemSeq[Union[float, AbstractParameter]], + max: ItemSeq[Union[float, AbstractParameter]], keys: Sequence = ("data",), grad: bool = False, - **kwargs ): """ - - Args: min: minimal value max: maximal value keys: the keys corresponding to the values to clamp grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn """ super().__init__( - augment_fn=clamp, keys=keys, grad=grad, min=min, max=max, property_names=("min", "max"), **kwargs + augment_fn=clamp, + keys=keys, + grad=grad, + paired_kw_names=("min", "max"), + augment_fn_names=("min", "max"), + min=min, + max=max, ) -class NormRange(PerSampleTransform): +class NormRange(PerSampleTransformMixin, BaseTransform): def __init__( self, - min: Union[float, AbstractParameter], - max: Union[float, AbstractParameter], + min: ItemSeq[Union[float, AbstractParameter]], + max: ItemSeq[Union[float, AbstractParameter]], keys: Sequence = ("data",), per_channel: bool = True, + per_sample=True, grad: bool = False, - **kwargs ): """ Args: @@ -79,79 +90,129 @@ def __init__( keys: keys to normalize per_channel: normalize per channel grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to normalization function """ super().__init__( augment_fn=norm_range, keys=keys, grad=grad, + paired_kw_names=("min", "max"), + augment_fn_names=("min", "max", "per_channel"), + min=min, + max=max, + per_sample=per_sample, + per_channel=per_channel, + ) + + +class NormPercentile(PerSampleTransformMixin, BaseTransform): + """clamp the distribution based on percentile and normalize it between 0 and 1""" + + def __init__( + self, + min: ItemSeq[Union[float, AbstractParameter]], + max: ItemSeq[Union[float, AbstractParameter]], + keys: Sequence[str] = ("data",), + grad: bool = False, + per_channel: bool = True, + per_sample=True, + ): + """ + Args: + min: min percentile, between 0 and 1 + max: max percentile, between 0 and 1 + per_channel: if normalize per channel + per_sample: if normalize per sample or per batch + """ + super().__init__( + augment_fn=norm_min_max_percentile, + keys=keys, + grad=grad, min=min, max=max, per_channel=per_channel, - property_names=("min", "max"), - **kwargs + paired_kw_names=("min", "max"), + augment_fn_names=("min", "max", "per_channel"), + per_sample=per_sample, ) -class NormMinMax(PerSampleTransform): +class NormMinMax(PerSampleTransformMixin, BaseTransform): """Norm to [0, 1]""" def __init__( self, keys: Sequence = ("data",), per_channel: bool = True, + per_sample=True, grad: bool = False, eps: Optional[float] = 1e-8, - **kwargs ): """ Args: keys: keys to normalize per_channel: normalize per channel + per_sample: normalize per sample or per batch grad: enable gradient computation inside transformation eps: small constant for numerical stability. If None, no factor constant will be added - **kwargs: keyword arguments passed to normalization function """ - super().__init__(augment_fn=norm_min_max, keys=keys, grad=grad, per_channel=per_channel, eps=eps, **kwargs) + super().__init__( + augment_fn=norm_min_max, + keys=keys, + grad=grad, + per_channel=per_channel, + per_sample=per_sample, + eps=eps, + augment_fn_names=( + "per_channel", + "eps", + ), + ) -class NormZeroMeanUnitStd(PerSampleTransform): +class NormZeroMeanUnitStd(PerSampleTransformMixin, BaseTransform): """Normalize mean to zero and std to one""" def __init__( self, keys: Sequence = ("data",), per_channel: bool = True, + per_sample=True, grad: bool = False, eps: Optional[float] = 1e-8, - **kwargs ): """ Args: keys: keys to normalize per_channel: normalize per channel + per_sample: normalize per sample or per batch. grad: enable gradient computation inside transformation eps: small constant for numerical stability. If None, no factor constant will be added - **kwargs: keyword arguments passed to normalization function """ super().__init__( - augment_fn=norm_zero_mean_unit_std, keys=keys, grad=grad, per_channel=per_channel, eps=eps, **kwargs + augment_fn=norm_zero_mean_unit_std, + keys=keys, + grad=grad, + per_channel=per_channel, + augment_fn_names=("eps", "per_channel"), + eps=eps, + per_sample=per_sample, ) -class NormMeanStd(PerSampleTransform): +class NormMeanStd(PerSampleTransformMixin, BaseTransform): """Normalize mean and std with provided values""" def __init__( self, - mean: Union[float, Sequence[float]], - std: Union[float, Sequence[float]], + mean: ItemSeq[Union[float, Sequence[float]]], + std: ItemSeq[Union[float, Sequence[float]]], keys: Sequence[str] = ("data",), per_channel: bool = True, + per_sample=True, grad: bool = False, - **kwargs + **kwargs, ): """ Args: @@ -163,11 +224,42 @@ def __init__( **kwargs: keyword arguments passed to normalization function """ super().__init__( - augment_fn=norm_mean_std, keys=keys, grad=grad, mean=mean, std=std, per_channel=per_channel, **kwargs + augment_fn=norm_mean_std, + keys=keys, + grad=grad, + mean=mean, + std=std, + per_channel=per_channel, + paired_kw_names=( + "mean", + "std", + ), + augment_fn_names=("mean", "std", "per_channel"), + per_sample=per_sample, + **kwargs, + ) + + +class GammaCorrection(BaseTransformMixin, BaseTransform): + """Apply Gamma correction""" + + def __init__(self, gamma: Union[float, AbstractParameter], keys: Sequence = ("data",), grad: bool = False): + """ + Args: + gamma: define gamma + keys: keys to normalize + grad: enable gradient computation inside transformation + """ + super().__init__( + augment_fn=gamma_correction, + keys=keys, + grad=grad, + augment_fn_names=("gamma",), + gamma=gamma, ) -class Noise(PerChannelTransform): +class _Noise(PerChannelTransformMixin, BaseTransform): """ Add noise to data @@ -194,62 +286,55 @@ def __init__( ) -class ExponentialNoise(Noise): +class ExponentialNoise(_Noise): """ Add exponential noise to data .. warning:: This transform will apply different noise patterns to different keys. """ - def __init__(self, lambd: float, keys: Sequence = ("data",), grad: bool = False, **kwargs): + def __init__(self, lambd: float, keys: Sequence = ("data",), grad: bool = False): """ Args: lambd: lambda of exponential distribution keys: keys to normalize grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to noise function """ - super().__init__(noise_type="exponential_", lambd=lambd, keys=keys, grad=grad, **kwargs) + super().__init__( + noise_type="exponential_", + lambd=lambd, + keys=keys, + grad=grad, + augment_fn_names=("noise_type", "lambd"), + ) -class GaussianNoise(Noise): +class GaussianNoise(_Noise): """ Add gaussian noise to data .. warning:: This transform will apply different noise patterns to different keys. """ - def __init__(self, mean: float, std: float, keys: Sequence = ("data",), grad: bool = False, **kwargs): + def __init__(self, mean: float, std: float, keys: Sequence = ("data",), grad: bool = False): """ Args: mean: mean of normal distribution std: std of normal distribution keys: keys to normalize grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to noise function - """ - super().__init__(noise_type="normal_", mean=mean, std=std, keys=keys, grad=grad, **kwargs) - - -class GammaCorrection(BaseTransform): - """Apply Gamma correction""" - - def __init__( - self, gamma: Union[float, AbstractParameter], keys: Sequence = ("data",), grad: bool = False, **kwargs - ): - """ - Args: - gamma: define gamma - keys: keys to normalize - grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to superclass """ super().__init__( - augment_fn=gamma_correction, gamma=gamma, property_names=("gamma",), keys=keys, grad=grad, **kwargs + noise_type="normal_", + mean=mean, + std=std, + keys=keys, + grad=grad, + augment_fn_names=("noise_type", "mean", "std"), ) -class RandomValuePerChannel(PerChannelTransform): +class _RandomValuePerChannel(PerChannelTransformMixin, BaseTransform): """ Apply augmentations which take random values as input by keyword :attr:`value` @@ -258,62 +343,31 @@ class RandomValuePerChannel(PerChannelTransform): def __init__( self, - augment_fn: callable, + augment_fn: augment_callable, + augment_fn_names: Sequence[str], random_sampler: AbstractParameter, per_channel: bool = False, - keys: Sequence = ("data",), + keys: Sequence[str] = ("data",), grad: bool = False, - **kwargs ): """ Args: augment_fn: augmentation function - random_mode: specifies distribution which should be used to - sample additive value. All function from python's random - module are supported - random_args: positional arguments passed for random function per_channel: enable transformation per channel keys: keys which should be augmented grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn """ super().__init__( augment_fn=augment_fn, + augment_fn_names=augment_fn_names, per_channel=per_channel, keys=keys, grad=grad, - random_sampler=random_sampler, - property_names=("random_sampler",), - **kwargs ) + self.register_sampler("value", random_sampler) - def forward(self, **data) -> dict: - """ - Perform Augmentation. - - Args: - data: dict with data - Returns: - dict: augmented data - """ - if self.per_channel: - seed = torch.random.get_rng_state() - for _key in self.keys: - torch.random.set_rng_state(seed) - out = torch.empty_like(data[_key]) - for _i in range(data[_key].shape[1]): - rand_value = self.random_sampler - out[:, _i] = self.augment_fn(data[_key][:, _i], value=rand_value, out=out[:, _i], **self.kwargs) - data[_key] = out - else: - rand_value = self.random_sampler - for _key in self.keys: - data[_key] = self.augment_fn(data[_key], value=rand_value, **self.kwargs) - return data - - -class RandomAddValue(RandomValuePerChannel): +class RandomAddValue(_RandomValuePerChannel): """ Increase values additively @@ -326,7 +380,6 @@ def __init__( per_channel: bool = False, keys: Sequence = ("data",), grad: bool = False, - **kwargs ): """ Args: @@ -334,14 +387,18 @@ def __init__( per_channel: enable transformation per channel keys: keys which should be augmented grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn """ super().__init__( - augment_fn=add_value, random_sampler=random_sampler, per_channel=per_channel, keys=keys, grad=grad, **kwargs + augment_fn=add_value, + augment_fn_names=("value",), + random_sampler=random_sampler, + per_channel=per_channel, + keys=keys, + grad=grad, ) -class RandomScaleValue(RandomValuePerChannel): +class RandomScaleValue(_RandomValuePerChannel): """ Scale Values @@ -354,7 +411,7 @@ def __init__( per_channel: bool = False, keys: Sequence = ("data",), grad: bool = False, - **kwargs + **kwargs, ): """ Args: @@ -366,28 +423,88 @@ def __init__( """ super().__init__( augment_fn=scale_by_value, + augment_fn_names=("value",), random_sampler=random_sampler, per_channel=per_channel, keys=keys, grad=grad, - **kwargs ) -class RandomBezierTransform(BaseTransform): +class RandomBezierTransform(BaseTransformMixin, BaseTransform): """Apply a random 3rd order bezier spline to the intensity values, as proposed in Models Genesis.""" - def __init__(self, maxv: float = 1.0, minv: float = 0.0, keys: Sequence = ("data",), **kwargs): - super().__init__(augment_fn=bezier_3rd_order, maxv=maxv, minv=minv, keys=keys, grad=False, **kwargs) + def __init__( + self, + maxv: float = 1.0, + minv: float = 0.0, + keys: Sequence = ("data",), + ): + super().__init__( + augment_fn=bezier_3rd_order, + keys=keys, + grad=False, + maxv=maxv, + augment_fn_names=("maxv", "minv"), + minv=minv, + ) -class InvertAmplitude(BaseTransform): +class InvertAmplitude(BaseTransformMixin, BaseTransform): """ Inverts the amplitude with probability p according to the following formula: out = maxv + minv - data """ - def __init__(self, prob: float = 0.5, maxv: float = 1.0, minv: float = 0.0, keys: Sequence = ("data",), **kwargs): + def __init__( + self, + prob: float = 0.5, + maxv: float = 1.0, + minv: float = 0.0, + keys: Sequence = ("data",), + ): + super().__init__( + augment_fn=random_inversion, + keys=keys, + grad=False, + prob_inversion=prob, + maxv=maxv, + minv=minv, + augment_fn_names=("prob_inversion", "maxv", "minv"), + ) + + +class RicianNoiseTransform(PerSampleTransformMixin, BaseTransform): + def __init__( + self, + *, + keys: Sequence[str], + grad: bool = False, + std: Union[float, AbstractParameter], + per_sample=True, + p: float = 1, + keep_range: bool = True, + ): + """Adds rician noise with the given std. + The Noise of MRI data tends to have a rician distribution: https://www.ncbi.nlm.nih.gov/pmc/articles/PMC2254141/ + Args: + std : Union[float, AbstractParameter], samples std of Gaussian distribution used to calculate + per_sample: if apply the noise per sample + p: the probability of applying the transform. + keep_range: if keep range of the image. + CAREFUL: This transform will modify the value range of your data! + + adapted from batchgenerators: + https://github.com/MIC-DKFZ/batchgenerators/blob/master/batchgenerators/transforms/noise_transforms.py + """ + super().__init__( - augment_fn=random_inversion, prob_inversion=prob, maxv=maxv, minv=minv, keys=keys, grad=False, **kwargs + augment_fn=augment_rician_noise, + p=p, + keys=keys, + grad=grad, + augment_fn_names=("std", "keep_range"), + per_sample=per_sample, ) + self.register_sampler("std", std) + self.keep_range = keep_range diff --git a/rising/transforms/kernel.py b/rising/transforms/kernel.py index 53a6af92..72562018 100644 --- a/rising/transforms/kernel.py +++ b/rising/transforms/kernel.py @@ -1,16 +1,18 @@ import math -from typing import Callable, Sequence, Union +from typing import Callable, Sequence import torch +from torch.nn import functional as F from rising.utils import check_scalar +from rising.utils.mise import ntuple -from .abstract import AbstractTransform +from .abstract import ItemSeq, _AbstractTransform __all__ = ["KernelTransform", "GaussianSmoothing"] -class KernelTransform(AbstractTransform): +class KernelTransform(_AbstractTransform): """ Baseclass for kernel based transformations (kernel is applied to each channel individually) @@ -19,12 +21,12 @@ class KernelTransform(AbstractTransform): def __init__( self, in_channels: int, - kernel_size: Union[int, Sequence], + kernel_size: ItemSeq[int], dim: int = 2, - stride: Union[int, Sequence] = 1, - padding: Union[int, Sequence] = 0, + stride: ItemSeq[int] = 1, + padding: ItemSeq[int] = 0, padding_mode: str = "zero", - keys: Sequence = ("data",), + keys: Sequence[str] = ("data",), grad: bool = False, **kwargs ): @@ -45,6 +47,9 @@ def __init__( :func:`torch.functional.pad` """ super().__init__(grad=grad, **kwargs) + self.keys = keys + self._tuple_generator = ntuple(len(self.keys)) + self.in_channels = in_channels if check_scalar(kernel_size): @@ -59,8 +64,7 @@ def __init__( padding = [padding] * dim * 2 self.padding = padding - self.padding_mode = padding_mode - self.keys = keys + self.padding_mode = self._tuple_generator(padding_mode) kernel = self.create_kernel() self.register_buffer("weight", kernel) @@ -79,11 +83,11 @@ def get_conv(dim) -> Callable: Callable: the suitable convolutional function """ if dim == 1: - return torch.nn.functional.conv1d + return F.conv1d elif dim == 2: - return torch.nn.functional.conv2d + return F.conv2d elif dim == 3: - return torch.nn.functional.conv3d + return F.conv3d else: raise TypeError("Only 1, 2 and 3 dimensions are supported. Received {}.".format(dim)) @@ -103,8 +107,11 @@ def forward(self, **data) -> dict: Returns: dict: dict with transformed data """ - for key in self.keys: - inp_pad = torch.nn.functional.pad(data[key], self.padding, mode=self.padding_mode) + # dtype, device = data[self.keys[0]].dtype, data[self.keys[0]].device + # self.to(dtype) + for key, padding_mode in zip(self.keys, self.padding_mode): + inp_pad = F.pad(data[key], self.padding, mode=padding_mode) + data[key] = self.conv(inp_pad, weight=self.weight, groups=self.groups, stride=self.stride) return data @@ -112,7 +119,7 @@ def forward(self, **data) -> dict: class GaussianSmoothing(KernelTransform): """ Perform Gaussian Smoothing. - Filtering is performed seperately for each channel in the input using a + Filtering is performed separately for each channel in the input using a depthwise convolution. This code is adapted from: 'https://discuss.pytorch.org/t/is-there-anyway-to-do-' @@ -122,13 +129,13 @@ class GaussianSmoothing(KernelTransform): def __init__( self, in_channels: int, - kernel_size: Union[int, Sequence], - std: Union[int, Sequence], + kernel_size: ItemSeq[int], + std: ItemSeq[float], dim: int = 2, - stride: Union[int, Sequence] = 1, - padding: Union[int, Sequence] = 0, - padding_mode: str = "reflect", - keys: Sequence = ("data",), + stride: ItemSeq[int] = 1, + padding: ItemSeq[int] = 0, + padding_mode: ItemSeq[str] = "constant", + keys: Sequence[str] = ("data",), grad: bool = False, **kwargs ): @@ -177,7 +184,7 @@ def create_kernel(self) -> torch.Tensor: kernel *= 1 / (std * math.sqrt(2 * math.pi)) * torch.exp(-(((mgrid - mean) / std) ** 2) / 2) # Make sure sum of values in gaussian kernel equals 1. - kernel = kernel / kernel.sum() + kernel /= kernel.sum() # Reshape to depthwise convolutional weight kernel = kernel.view(1, 1, *kernel.size()) diff --git a/rising/transforms/pad.py b/rising/transforms/pad.py new file mode 100644 index 00000000..62ce6c7f --- /dev/null +++ b/rising/transforms/pad.py @@ -0,0 +1,66 @@ +from typing import Sequence, Union + +import torch + +from rising.transforms.abstract import BaseTransform, BaseTransformMixin, ItemSeq +from rising.transforms.functional.pad import pad as _pad +from rising.utils.mise import ntuple + + +class Pad(BaseTransformMixin, BaseTransform): + def __init__( + self, + *, + pad_size: ItemSeq[int], + mode: str = "constant", + pad_value: ItemSeq[float], + keys: Sequence[str] = ("data",), + grad: bool = False, + **kwargs + ): + """ + padding_size should be the same for all the keys. + pad_value can be different for different keys + """ + super().__init__( + augment_fn=_pad, + keys=keys, + grad=grad, + paired_kw_names=("value", "mode"), + augment_fn_names=("pad_size", "mode", "value"), + pad_size=pad_size, + mode=mode, + value=pad_value, + **kwargs + ) + + def forward(self, **data) -> dict: + """ + Apply transformation + + Args: + data: dict with tensors + + Returns: + dict: dict with augmented data + """ + seed = int(torch.randint(0, int(1e16), (1,))) + + for _key in self.keys: + with self.random_cxm(seed): + kwargs = {k: getattr(self, k) for k in self._augment_fn_names if k not in self._paired_kw_names} + kwargs.update(self.get_pair_kwargs(_key)) + input_shape = data[_key].shape[2:] + pad_size = self.pad_parameters(input_shape, self.pad_size, ndim=(data[_key].dim() - 2)) + kwargs.update({"pad_size": pad_size}) + + if torch.rand(1).item() < self.p: + data[_key] = self.augment_fn(data[_key], **kwargs) + return data + + @staticmethod + def pad_parameters(input_shape: Union[int, Sequence[int]], resize_shape: Union[int, Sequence[int]], *, ndim: int): + input_shape = ntuple(ndim)(input_shape) + resize_shape = ntuple(ndim)(resize_shape) + shape_difference = (max(0, r - i) for r, i in zip(resize_shape, input_shape)) + return [sub for y in [(x // 2, x - (x // 2)) for x in shape_difference] for sub in y][::-1] diff --git a/rising/transforms/painting.py b/rising/transforms/painting.py index eafc486a..f681e5cc 100644 --- a/rising/transforms/painting.py +++ b/rising/transforms/painting.py @@ -2,7 +2,7 @@ import torch -from rising.transforms.abstract import AbstractTransform, BaseTransform +from rising.transforms.abstract import BaseTransform, _AbstractTransform from rising.transforms.functional.painting import local_pixel_shuffle, random_inpainting, random_outpainting __all__ = ["RandomInpainting", "RandomOutpainting", "RandomInOrOutpainting", "LocalPixelShuffle"] @@ -20,7 +20,7 @@ def __init__(self, n: int = -1, keys: Sequence = ("data",), grad: bool = False, grad: enable gradient computation inside transformation **kwargs: keyword arguments passed to augment_fn """ - super().__init__(augment_fn=local_pixel_shuffle, n=n, keys=keys, grad=grad, **kwargs) + super().__init__(augment_fn=local_pixel_shuffle, keys=keys, grad=grad, n=n, **kwargs) class RandomInpainting(BaseTransform): @@ -38,10 +38,10 @@ def __init__( grad: enable gradient computation inside transformation **kwargs: keyword arguments passed to augment_fn """ - super().__init__(augment_fn=random_inpainting, n=n, maxv=maxv, minv=minv, keys=keys, grad=grad, **kwargs) + super().__init__(augment_fn=random_inpainting, keys=keys, grad=grad, n=n, maxv=maxv, minv=minv, **kwargs) -class RandomOutpainting(AbstractTransform): +class RandomOutpainting(_AbstractTransform): """The border of the images will be replaced by uniform noise, as proposed in Models Genesis""" @@ -75,7 +75,7 @@ def forward(self, **data) -> dict: return data -class RandomInOrOutpainting(AbstractTransform): +class RandomInOrOutpainting(_AbstractTransform): """Applies either random inpainting or random outpainting to the image, as proposed in Models Genesis""" diff --git a/rising/transforms/sitk.py b/rising/transforms/sitk.py new file mode 100644 index 00000000..2e7224da --- /dev/null +++ b/rising/transforms/sitk.py @@ -0,0 +1,116 @@ +from typing import Sequence, Tuple, Union + +import torch + +from rising.random.abstract import AbstractParameter +from rising.transforms.abstract import BaseTransform, ItemSeq +from rising.transforms.functional.sitk import itk2tensor, itk_clip, itk_resample + +SpacingParamType = Union[ + int, Sequence[int], float, Sequence[float], torch.Tensor, AbstractParameter, Sequence[AbstractParameter] +] +SpacingTypeOrTuple = Union[SpacingParamType, Tuple[SpacingParamType, SpacingParamType, SpacingParamType]] + +IntNumType = Union[int, AbstractParameter] + + +class _ITKTransform: + """ + this mixin indicates if the transform is Tensor-based, use to not shuffle in Compose. + """ + + pass + + +class SITKResample(_ITKTransform, BaseTransform): + """ + simpleitk resampling class + """ + + def __init__( + self, + spacing: SpacingParamType, + *, + pad_value: Union[int, float], + keys: Sequence = ("data",), + interpolation: ItemSeq[str] = "nearest", + ): + """ + resample simpleitk image given new spacing and padding value + Args: + spacing: float or tuple of three floats, indicating spacing for each dimension. + pad_value: padding values + interpolation: str or sequence of str to indicate the interpolation for different keys. + """ + super().__init__( + augment_fn=itk_resample, + keys=keys, + grad=False, + spacing=spacing, + pad_value=pad_value, + ) + self.interpolation_mode = self.tuple_generator(interpolation) + + def forward(self, **data) -> dict: + for key, interpolation in zip(self.keys, self.interpolation_mode): + data[key] = self.augment_fn( + data[key], + spacing=self.spacing, + interpolation=interpolation, + pad_value=self.pad_value, + **self.kwargs, + ) + + return data + + +class SITKWindows(_ITKTransform, BaseTransform): + """ + simpleitk windows class + """ + + def __init__(self, low: IntNumType, high: IntNumType, *, keys: Sequence[str] = ("data",), **kwargs): + super().__init__(augment_fn=itk_clip, keys=keys, grad=False, property_names=("low", "high"), low=low, high=high, **kwargs) + + def forward(self, **data) -> dict: + kwargs = {} + for k in self.property_names: + kwargs[k] = getattr(self, k) + + kwargs.update(self.kwargs) + + # to make sure that in the sampling, there is case where `high` is lower than `low`. + if kwargs["low"] > kwargs["high"]: + kwargs["low"], kwargs["high"] = kwargs["high"], kwargs["low"] + + for _key in self.keys: + data[_key] = self.augment_fn(data[_key], *self.args, **kwargs) + return data + + +class SITK2Tensor(_ITKTransform, BaseTransform): + def __init__( + self, + *, + keys: Sequence = ("data",), + dtype: ItemSeq[torch.dtype] = torch.float, + insert_dim: int = None, + grad: bool = False, + **kwargs, + ): + """ + Convert sitk image to Tensor + Args: + dtype: tensor's dtype + insert_dim: type: int, if you need to expand the tensor given specific dimension, default None, + """ + super().__init__(augment_fn=itk2tensor, keys=keys, grad=grad, **kwargs) + self.dtype = self.tuple_generator(dtype) + self.insert_dim = insert_dim + + def forward(self, **data) -> dict: + for key, dtype in zip(self.keys, self.dtype): + data[key] = self.augment_fn(data[key], dtype=dtype) + if self.insert_dim is not None: + data[key] = data[key].unsqueeze(self.insert_dim) + return data diff --git a/rising/transforms/spatial.py b/rising/transforms/spatial.py index f5b26c3c..a57d1689 100644 --- a/rising/transforms/spatial.py +++ b/rising/transforms/spatial.py @@ -2,83 +2,62 @@ from itertools import combinations from typing import Callable, Optional, Sequence, Union -import torch from torch.multiprocessing import Value +from rising.constants import FInterpolation from rising.random import AbstractParameter, DiscreteParameter -from rising.transforms.abstract import AbstractTransform, BaseTransform - -__all__ = ["Mirror", "Rot90", "ResizeNative", "Zoom", "ProgressiveResize", "SizeStepScheduler"] - -from rising.transforms.functional import mirror, resize_native, rot90 +from rising.transforms.abstract import BaseTransform, BaseTransformMixin, PerSampleTransformMixin, ItemSeq +from rising.transforms.functional import mirror, resize_native, rot90, center_crop +from rising.utils import check_scalar scheduler_type = Callable[[int], Union[int, Sequence[int]]] +__all__ = ["Mirror", "Rot90", "ResizeNative", "Zoom", "ProgressiveResize", "SizeStepScheduler", + "ResizeNativeCentreCrop"] -class Mirror(AbstractTransform): + +class Mirror(PerSampleTransformMixin, BaseTransform): """Random mirror transform""" def __init__( - self, - dims: Union[int, DiscreteParameter, Sequence[Union[int, DiscreteParameter]]], - keys: Sequence[str] = ("data",), - prob: float = 0.5, - grad: bool = False, - **kwargs + self, + *, + dims: ItemSeq[Union[int, DiscreteParameter]], + p_sample: float = 0.5, + keys: Sequence[str] = ("data",), + grad: bool = False, + per_sample: bool = True, ): """ Args: - dims: axes which should be mirrored - keys: keys which should be mirrored - prob: probability for mirror. If float value is provided, - it is used for all dims - grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to superclass - - Examples: - >>> # Use mirror transform for augmentations - >>> from rising.random import DiscreteCombinationsParameter - >>> # We sample from all possible mirror combination for - >>> # volumetric data - >>> trafo = Mirror(DiscreteCombinationsParameter((0, 1, 2))) - """ - super().__init__(grad=grad, **kwargs) - self.keys = keys - self.prob = prob - if not isinstance(dims, DiscreteParameter): - if len(dims) > 2: - dims = list(combinations(dims, 2)) - else: - dims = (dims,) - dims = DiscreteParameter(dims) - self.register_sampler("dims", dims) - - def forward(self, **data) -> dict: - """ - Apply transformation - - Args: - data: dict with tensors - Returns: - dict: dict with augmented data + dims: dimensions to apply random mirror + p_sample: the probability of applying the mirror (0<=p_sample<=1), default=0.5 + per_sample: if applied per sample. + keys: attributes to be applied. """ - if torch.rand(1) < self.prob: - for key in self.keys: - data[key] = mirror(data[key], self.dims) - return data + super().__init__( + augment_fn=mirror, + keys=keys, + grad=grad, + augment_fn_names=("dims",), + per_sample=per_sample, + dims=dims, + p=p_sample, + seeded=True, + ) -class Rot90(AbstractTransform): +class Rot90(PerSampleTransformMixin, BaseTransform): """Rotate 90 degree around dims""" def __init__( - self, - dims: Union[Sequence[int], DiscreteParameter], - keys: Sequence[str] = ("data",), - num_rots: Sequence[int] = (0, 1, 2, 3), - prob: float = 0.5, - grad: bool = False, - **kwargs + self, + dims: ItemSeq[Union[Sequence[int], DiscreteParameter]], + keys: Sequence[str] = ("data",), + num_rots: DiscreteParameter = DiscreteParameter((0, 1, 2, 3)), + p_sample: float = 0.5, + per_sample: bool = True, + grad: bool = False, ): """ Args: @@ -86,56 +65,41 @@ def __init__( provided, 2 dimensions are randomly chosen at each call keys: keys which should be rotated num_rots: possible values for number of rotations - prob: probability for rotation grad: enable gradient computation inside transformation - kwargs: keyword arguments passed to superclass See Also: :func:`torch.Tensor.rot90` """ - super().__init__(grad=grad, **kwargs) - self.keys = keys - self.prob = prob if not isinstance(dims, DiscreteParameter): - if len(dims) > 2: + if len(dims) >= 2: dims = list(combinations(dims, 2)) else: - dims = (dims,) + raise RuntimeError(f"dims must be at least 2 dims, given {dims}.") dims = DiscreteParameter(dims) - self.register_sampler("dims", dims) - self.register_sampler("num_rots", DiscreteParameter(num_rots)) - - def forward(self, **data) -> dict: - """ - Apply transformation - - Args: - data: dict with tensors - - Returns: - dict: dict with augmented data - """ - if torch.rand(1) < self.prob: - num_rots = self.num_rots - rand_dims = self.dims - - for key in self.keys: - data[key] = rot90(data[key], k=num_rots, dims=rand_dims) - return data + super().__init__( + augment_fn=rot90, + keys=keys, + grad=grad, + augment_fn_names=("dims", "k"), + per_sample=per_sample, + dims=dims, + p=p_sample, + seeded=True, + k=num_rots, + ) -class ResizeNative(BaseTransform): +class ResizeNative(BaseTransformMixin, BaseTransform): """Resize data to given size""" def __init__( - self, - size: Union[int, Sequence[int]], - mode: str = "nearest", - align_corners: Optional[bool] = None, - preserve_range: bool = False, - keys: Sequence = ("data",), - grad: bool = False, - **kwargs + self, + size: Union[int, Sequence[int]], + mode: ItemSeq[FInterpolation] = FInterpolation.nearest, + align_corners: ItemSeq[bool] = None, + preserve_range: bool = False, + keys: Sequence[str] = ("data",), + grad: bool = False, ): """ Args: @@ -143,42 +107,91 @@ def __init__( number of channels) mode: one of ``nearest``, ``linear``, ``bilinear``, ``bicubic``, ``trilinear``, ``area`` (for more inforamtion see - :func:`torch.nn.functional.interpolate`) + :func:`torch.nn.functional.interpolate`) or their sequence, for different keys. align_corners: input and output tensors are aligned by the center \ points of their corners pixels, preserving the values at the - corner pixels. + corner pixels. Input can be sequence, for different keys. preserve_range: output tensor has same range as input tensor keys: keys which should be augmented grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn """ super().__init__( augment_fn=resize_native, + keys=keys, + grad=grad, size=size, mode=mode, align_corners=align_corners, preserve_range=preserve_range, - keys=keys, - grad=grad, - **kwargs + p=1, + augment_fn_names=("size", "mode", "align_corners", "preserve_range"), + per_sample=False, + paired_kw_names=("mode", "align_corners"), ) + assert self.p == 1.0 + + +class ResizeNativeCentreCrop(ResizeNative): + + def __init__(self, size: Union[int, Sequence[int]], mode: ItemSeq[FInterpolation] = FInterpolation.nearest, + align_corners: ItemSeq[bool] = None, preserve_range: bool = False, margin: ItemSeq[int] = 0, + keys: Sequence[str] = ("data",), + grad: bool = False): + """ + This class extends the ResizeNative class to crop the image in the center, useful to exclude the empty borders + + image -> resize with size + margin -> center crop with size + + Args: + size: spatial output size (excluding batch size and + number of channels) + mode: one of ``nearest``, ``linear``, ``bilinear``, ``bicubic``, + ``trilinear``, ``area`` (for more inforamtion see + :func:`torch.nn.functional.interpolate`) or their sequence, for different keys. + align_corners: input and output tensors are aligned by the center \ + points of their corners pixels, preserving the values at the + corner pixels. Input can be sequence, for different keys. + preserve_range: output tensor has same range as input tensor + keys: keys which should be augmented + grad: enable gradient computation inside transformation + """ + if check_scalar(margin) ^ check_scalar(size): + raise ValueError("`size` and `margin` must be scalar or sequence of scalars in the same time.") + if not check_scalar(margin) and len(margin) != len(size): + raise ValueError("`size` and `margin` must have the same length.") + + if check_scalar(margin): + resized_size = size + margin + else: + resized_size = [x + y for x, y in zip(size, margin)] + + super().__init__(size, mode, align_corners, preserve_range, keys, grad) + self._resized_size = resized_size + self._crop_size = size + self._margin = margin + + def forward(self, **data): + self.size = self._resized_size + data = super().forward(**data) + for _key in self.keys: + data[_key] = center_crop(data[_key], self._crop_size) + return data -class Zoom(BaseTransform): +class Zoom(BaseTransformMixin, BaseTransform): """Apply augment_fn to keys. By default the scaling factor is sampled from a uniform distribution with the range specified by :attr:`random_args` """ def __init__( - self, - scale_factor: Union[Sequence, AbstractParameter] = (0.75, 1.25), - mode: str = "nearest", - align_corners: bool = None, - preserve_range: bool = False, - keys: Sequence = ("data",), - grad: bool = False, - **kwargs + self, + scale_factor: Union[Sequence, AbstractParameter] = (0.75, 1.25), + mode: ItemSeq[FInterpolation] = FInterpolation.nearest, + align_corners: ItemSeq[bool] = None, + preserve_range: bool = False, + keys: Sequence[str] = ("data",), + grad: bool = False, ): """ Args: @@ -195,36 +208,37 @@ def __init__( preserve_range: output tensor has same range as input tensor keys: keys which should be augmented grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn See Also: :func:`random.uniform`, :func:`torch.nn.functional.interpolate` """ super().__init__( augment_fn=resize_native, + keys=keys, + grad=grad, scale_factor=scale_factor, mode=mode, align_corners=align_corners, preserve_range=preserve_range, - keys=keys, - grad=grad, - property_names=("scale_factor",), - **kwargs + p=1, + augment_fn_names=("scale_factor", "mode", "align_corners", "preserve_range"), + per_sample=False, + paired_kw_names=("mode", "align_corners"), ) + assert self.p == 1.0 class ProgressiveResize(ResizeNative): """Resize data to sizes specified by scheduler""" def __init__( - self, - scheduler: scheduler_type, - mode: str = "nearest", - align_corners: bool = None, - preserve_range: bool = False, - keys: Sequence = ("data",), - grad: bool = False, - **kwargs + self, + scheduler: scheduler_type, + mode: ItemSeq[FInterpolation] = FInterpolation.nearest, + align_corners: ItemSeq[Optional[bool]] = None, + preserve_range: bool = False, + keys: Sequence = ("data",), + grad: bool = False, ): """ Args: @@ -240,7 +254,6 @@ def __init__( preserve_range: output tensor has same range as input tensor keys: keys which should be augmented grad: enable gradient computation inside transformation - **kwargs: keyword arguments passed to augment_fn Warnings: When this transformations is used in combination with @@ -256,7 +269,6 @@ def __init__( preserve_range=preserve_range, keys=keys, grad=grad, - **kwargs ) self.scheduler = scheduler self._step = Value("i", 0) @@ -303,7 +315,7 @@ def forward(self, **data) -> dict: Returns: dict: augmented batch """ - self.kwargs["size"] = self.scheduler(self.step) + self.size = self.scheduler(self.step) self.increment() return super().forward(**data) diff --git a/rising/transforms/tensor.py b/rising/transforms/tensor.py index 56dc6c1a..a1332b5d 100644 --- a/rising/transforms/tensor.py +++ b/rising/transforms/tensor.py @@ -3,13 +3,20 @@ import torch from torch.utils.data._utils.collate import default_convert -from rising.transforms import AbstractTransform, BaseTransform +from rising.transforms import BaseTransform, BaseTransformMixin from rising.transforms.functional import tensor_op, to_device_dtype -__all__ = ["ToTensor", "ToDeviceDtype", "ToDevice", "ToDtype", "TensorOp", "Permute"] +__all__ = [ + "ToTensor", + "_ToDeviceDtype", + "ToDevice", + "ToDtype", + "TensorOp", + "Permute", +] -class ToTensor(BaseTransform): +class ToTensor(BaseTransformMixin, BaseTransform): """Transform Input Collection to Collection of :class:`torch.Tensor`""" def __init__(self, keys: Sequence = ("data",), grad: bool = False, **kwargs): @@ -22,7 +29,7 @@ def __init__(self, keys: Sequence = ("data",), grad: bool = False, **kwargs): super().__init__(augment_fn=default_convert, keys=keys, grad=grad, **kwargs) -class ToDeviceDtype(BaseTransform): +class _ToDeviceDtype(BaseTransformMixin, BaseTransform): """Push data to device and convert to tdype""" def __init__( @@ -33,7 +40,7 @@ def __init__( copy: bool = False, keys: Sequence = ("data",), grad: bool = False, - **kwargs + **kwargs, ): """ Args: @@ -55,11 +62,11 @@ def __init__( dtype=dtype, non_blocking=non_blocking, copy=copy, - **kwargs + **kwargs, ) -class ToDevice(ToDeviceDtype): +class ToDevice(_ToDeviceDtype): """Push data to device""" def __init__( @@ -69,7 +76,7 @@ def __init__( copy: bool = False, keys: Sequence = ("data",), grad: bool = False, - **kwargs + **kwargs, ): """ Args: @@ -82,10 +89,18 @@ def __init__( grad: enable gradient computation inside transformation **kwargs: keyword arguments passed to function """ - super().__init__(device=device, non_blocking=non_blocking, copy=copy, keys=keys, grad=grad, **kwargs) + super().__init__( + device=device, + non_blocking=non_blocking, + copy=copy, + keys=keys, + grad=grad, + augment_fn_names=("device",), + **kwargs, + ) -class ToDtype(ToDeviceDtype): +class ToDtype(_ToDeviceDtype): """Convert data to dtype""" def __init__(self, dtype: torch.dtype, keys: Sequence = ("data",), grad: bool = False, **kwargs): @@ -96,10 +111,10 @@ def __init__(self, dtype: torch.dtype, keys: Sequence = ("data",), grad: bool = grad: enable gradient computation inside transformation kwargs: keyword arguments passed to function """ - super().__init__(dtype=dtype, keys=keys, grad=grad, **kwargs) + super().__init__(dtype=dtype, keys=keys, grad=grad, augment_fn_names=("dtype",), **kwargs) -class TensorOp(BaseTransform): +class TensorOp(BaseTransformMixin, BaseTransform): """Apply function which are supported by the `torch.Tensor` class""" def __init__(self, op_name: str, *args, keys: Sequence = ("data",), grad: bool = False, **kwargs): @@ -111,10 +126,12 @@ def __init__(self, op_name: str, *args, keys: Sequence = ("data",), grad: bool = grad: enable gradient computation inside transformation **kwargs: keyword arguments passed to function """ - super().__init__(tensor_op, op_name, *args, keys=keys, grad=grad, **kwargs) + super().__init__( + augment_fn=tensor_op, fn=op_name, *args, keys=keys, grad=grad, augment_fn_names=("fn",), **kwargs + ) -class Permute(BaseTransform): +class Permute(BaseTransformMixin, BaseTransform): """Permute dimensions of tensor""" def __init__(self, dims: Dict[str, Sequence[int]], grad: bool = False, **kwargs): diff --git a/rising/transforms/utility.py b/rising/transforms/utility.py index 0e3f31e1..50a5818d 100644 --- a/rising/transforms/utility.py +++ b/rising/transforms/utility.py @@ -2,13 +2,13 @@ import torch -from rising.transforms.abstract import AbstractTransform +from rising.transforms.abstract import _AbstractTransform from rising.transforms.functional.utility import box_to_seg, instance_to_semantic, seg_to_box __all__ = ["DoNothing", "SegToBox", "BoxToSeg", "InstanceToSemantic"] -class DoNothing(AbstractTransform): +class DoNothing(_AbstractTransform): """Transform that returns the input as is""" def __init__(self, grad: bool = False, **kwargs): @@ -32,7 +32,7 @@ def forward(self, **data) -> dict: return data -class SegToBox(AbstractTransform): +class SegToBox(_AbstractTransform): """Convert instance segmentation to bounding boxes""" def __init__(self, keys: Mapping[Hashable, Hashable], grad: bool = False, **kwargs): @@ -61,7 +61,7 @@ def forward(self, **data) -> dict: return data -class BoxToSeg(AbstractTransform): +class BoxToSeg(_AbstractTransform): """Convert bounding boxes to instance segmentation""" def __init__( @@ -108,7 +108,7 @@ def forward(self, **data) -> dict: return data -class InstanceToSemantic(AbstractTransform): +class InstanceToSemantic(_AbstractTransform): """Convert an instance segmentation to a semantic segmentation""" def __init__(self, keys: Mapping[str, str], cls_key: Hashable, grad: bool = False, **kwargs): diff --git a/rising/utils/affine.py b/rising/utils/affine.py index 02def8cb..a858e514 100644 --- a/rising/utils/affine.py +++ b/rising/utils/affine.py @@ -106,7 +106,7 @@ def get_batched_eye( batchsize: int, ndim: int, device: Optional[Union[torch.device, str]] = None, - dtype: Optional[Union[torch.dtype, str]] = None, + dtype: torch.dtype = torch.float, ) -> torch.Tensor: """ Produces a batched matrix containing 1s on the diagonal @@ -119,14 +119,14 @@ def get_batched_eye( device : torch.device, str, optional the device to put the resulting tensor to. Defaults to the default device - dtype : torch.dtype, str, optional - the dtype of the resulting trensor. Defaults to the default dtype + dtype : torch.dtype, optional + the dtype of the resulting tensor. Defaults to the default torch.float Returns: torch.Tensor: batched eye matrix """ - return torch.eye(ndim, device=device, dtype=dtype).view(1, ndim, ndim).expand(batchsize, -1, -1).clone() + return torch.eye(ndim, device=device, dtype=dtype)[None, ...].expand(batchsize, -1, -1).clone() def deg_to_rad(angles: Union[torch.Tensor, float, int]) -> Union[torch.Tensor, float, int]: diff --git a/rising/utils/mise.py b/rising/utils/mise.py new file mode 100644 index 00000000..8b6c2446 --- /dev/null +++ b/rising/utils/mise.py @@ -0,0 +1,120 @@ +import collections +import functools +import random +import typing as t +from contextlib import contextmanager, nullcontext +from itertools import repeat + +import numpy as np +import torch +from torch import Tensor +from torch import multiprocessing as mp +from torch import nn + +from rising.random.abstract import AbstractParameter + +T = t.TypeVar("T") + +__all__ = ["ntuple", "single", "pair", "triple", "quadruple", "fix_seed_cxm", "nullcxm"] + +nullcxm = nullcontext + + +def ntuple(n: int) -> t.Callable[[t.Union[T, t.Sequence[T]]], t.Sequence[T]]: + def parse(x: t.Union[T, t.Sequence[T]]) -> t.Sequence[T]: + if isinstance(x, AbstractParameter): + return nn.ModuleList([x]) + if isinstance(x, (Tensor, np.ndarray, str)): + return tuple(repeat(x, n)) + if isinstance(x, collections.abc.Iterable): + item_list = tuple(x) + if len(item_list) == n: + return item_list + if len(item_list) == 1: + return tuple(repeat(item_list[0], n)) + if len(list(x)) != n: + raise RuntimeError(f"Iterable shape inconsistent, n = {n}, given {len(list(x))}") + + return tuple(repeat(x, n)) + + return parse + + +single = ntuple(1) +pair = ntuple(2) +triple = ntuple(3) +quadruple = ntuple(4) + + +class fixed_torch_seed: + """ + fixed random seed for torch module + """ + + def __init__(self, seed: int = 10, cuda: bool = True) -> None: + super().__init__() + self.seed = seed + self.cuda_flag = cuda and torch.cuda.is_available() + + def __enter__(self): + seed = self.seed + self.__pre_state = torch.get_rng_state() + self.__pre_cuda_state_all = None + if self.cuda_flag: + self.__pre_cuda_state_all = torch.cuda.get_rng_state_all() + + torch.manual_seed(seed) + if self.cuda_flag: + torch.cuda.manual_seed_all(seed) + + def __exit__(self, exc_type, exc_val, exc_tb): + torch.set_rng_state(self.__pre_state) + if self.cuda_flag: + torch.cuda.set_rng_state_all(self.__pre_cuda_state_all) + + def __call__(self, func): + @functools.wraps(func) + def generator_context(*args, **kwargs): + with self: + gen = func(*args, **kwargs) + return gen + + return generator_context + + +class fixed_random_seed: + """ + fixed random seed for random module + """ + + def __init__(self, seed: int = 10, **kwargs) -> None: + super().__init__() + self.seed = seed + + def __enter__(self): + seed = self.seed + self.__pre_state = random.getstate() + random.seed(seed) + + def __exit__(self, exc_type, exc_val, exc_tb): + random.setstate(self.__pre_state) + + def __call__(self, func): + @functools.wraps(func) + def generator_context(*args, **kwargs): + with self: + gen = func(*args, **kwargs) + return gen + + return generator_context + + +def on_main_process(): + return mp.current_process().name == "MainProcess" + + +@contextmanager +def fix_seed_cxm(seed: int = 10): + cuda = on_main_process() and torch.cuda.is_available() + with fixed_torch_seed(seed=seed, cuda=cuda), fixed_random_seed(seed=seed): + yield diff --git a/rising/utils/transforms.py b/rising/utils/transforms.py new file mode 100644 index 00000000..c6b28e0b --- /dev/null +++ b/rising/utils/transforms.py @@ -0,0 +1,47 @@ +import typing +from typing import Iterable, Optional, Sequence, Union + +import torch +from torch import nn + +if typing.TYPE_CHECKING: + from rising.transforms.abstract import _AbstractTransform + from rising.transforms.compose import Compose + + +def iter_transform( + transforms: Union["Compose", "_AbstractTransform", Sequence["_AbstractTransform"], nn.ModuleList] +) -> Iterable["_AbstractTransform"]: + from rising.transforms.abstract import _AbstractTransform + from rising.transforms.compose import Compose + + if isinstance(transforms, _AbstractTransform) and not isinstance(transforms, Compose): + yield transforms + elif isinstance(transforms, Compose): + yield from iter_transform(transforms.transforms) + elif isinstance(transforms, Sequence): + for x in transforms: + yield from iter_transform(x) + elif isinstance(transforms, nn.ModuleList): + for x in transforms: + yield from iter_transform(x) + else: + raise TypeError(transforms) + + +def get_keys_from_transforms(transforms) -> Sequence[str]: + _keys = [transform.keys for transform in iter_transform(transforms) if hasattr(transform, "keys")] + keys = tuple(set([item for sublist in _keys for item in sublist])) + return keys + + +def get_dtype_from_transforms(transforms) -> Optional[torch.dtype]: + """ + this gives the dtype that the gpu transform should be converted, if ToDtypeTransform is given. + """ + from rising.transforms import ToDtype + + for transform in iter_transform(transforms): + if isinstance(transform, ToDtype): + dtype = transform.dtype + return dtype diff --git a/rising/viewer.py b/rising/viewer.py new file mode 100644 index 00000000..ca098c81 --- /dev/null +++ b/rising/viewer.py @@ -0,0 +1,228 @@ +# Copyright 2017 Fabian Isensee +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import sys + +import numpy as np +import pyqtgraph as pg +from pyqtgraph.Qt import QtCore, QtGui + + +class SliceViewer(QtGui.QWidget): + def __init__(self): + super(SliceViewer, self).__init__() + self.initUI() + + def wheelEvent(self, event): + for v in [self.imageViewer, self.imageViewer2]: + v.setSlice(v.getSlice() + event.delta() / 120) + + def initUI(self): + image1 = np.random.uniform(-0.5, 255.0, (100, 100, 100)).astype(float) + image2 = np.random.uniform(-100.0, 255.0, (100, 100, 100)).astype(float) + + self.imageViewer = ImageSlicingWidget() + self.imageViewer2 = ImageSlicingWidget() + + self.imageViewer.setImage(image1) + self.imageViewer2.setImage(image2) + + hLayout = QtGui.QHBoxLayout() + hLayout.addWidget(self.imageViewer) + hLayout.addWidget(self.imageViewer2) + hLayout.addStretch() + + self.setLayout(hLayout) + + # self.setGeometry(0, 0, 1200, 1200) + self.setWindowTitle("QtGui.QCheckBox") + + self.show() + + +class ImageViewer2DWidget(QtGui.QWidget): + def __init__(self, width=300, height=300): + QtGui.QWidget.__init__(self) + self.image = None + self.lut = None + self.init(width, height) + + def init(self, width, height): + self.rect = QtCore.QRect(0, 0, width, height) + self.imageItem = pg.ImageItem() + self.imageItem.setImage(None) + + self.graphicsScene = pg.GraphicsScene() + self.graphicsScene.addItem(self.imageItem) + + self.graphicsView = pg.GraphicsView() + self.graphicsView.setRenderHint(QtGui.QPainter.Antialiasing) + self.graphicsView.setScene(self.graphicsScene) + + layout = QtGui.QHBoxLayout() + layout.addWidget(self.graphicsView) + self.setLayout(layout) + self.setMaximumSize(width, height) + self.setMinimumSize(width - 10, height - 10) + + def setImage(self, image): + assert len(image.shape) == 2 or len(image.shape) == 3 + if len(image.shape) == 4: + assert image.shape[-1] == 4 + self.image = image + self.imageItem.setImage(self.image) + self.imageItem.setRect(self.rect) + if self.lut is not None: + self.imageItem.setLookupTable(self.lut) + + def setLevels(self, levels, update=True): + self.imageItem.setLevels(levels, update) + + def setLUT(self, lut): + self.lut = lut + + +class ImageSlicingWidget(ImageViewer2DWidget): + def __init__(self, width=300, height=300): + self.slice = 0 + ImageViewer2DWidget.__init__(self, width, height) + + def setImage(self, image): + assert len(image.shape) == 3 or len(image.shape) == 4 + if len(image.shape) == 4: + assert image.shape[-1] == 4 + self.image3D = np.array(image) + self._updateImageSlice() + + def _updateImageSlice(self): + self.imageItem.setImage(self.image3D[self.slice]) + self.imageItem.setRect(self.rect) + if self.lut is not None: + self.imageItem.setLookupTable(self.lut) + + def setSlice(self, slice): + slice = np.max((0, slice)) + slice = np.min((slice, self.image3D.shape[0] - 1)) + self.slice = slice + self._updateImageSlice() + + def getSlice(self): + return self.slice + + +class BatchViewer(QtGui.QWidget): + def __init__(self, parent=None, width=300, height=300): + QtGui.QWidget.__init__(self, parent) + self.batch = None + self.width = width + self.height = height + self.slicingWidgets = {} + self._init_gui() + + def setBatch(self, batch, lut={}): + assert len(batch.shape) == 4 + batch = np.copy(batch) + for v in self.slicingWidgets.values(): + self._my_layout.removeWidget(v) + self.slicingWidgets = {} + + if not isinstance(lut, dict): + lut = {i: lut for i in range(self.batch.shape[0])} + + for b in range(batch.shape[0]): + mn = batch[b].min() + mx = batch[b].max() + batch[b, :, 0, 0] = mn + batch[b, :, 0, 1] = mx + + self.batch = batch + num_col = int(np.ceil(np.sqrt(batch.shape[0]))) + col = 0 + row = 0 + for i in range(self.batch.shape[0]): + w = ImageSlicingWidget(self.width, self.height) + if lut is not None and i in lut.keys(): + w.setLUT(lut[i]) + w.setImage(self.batch[i]) + w.setLevels([self.batch[i].min(), self.batch[i].max()]) + w.setSlice(0) + self._my_layout.addWidget(w, row, col) + col += 1 + if col >= num_col: + col = 0 + row += 1 + self.slicingWidgets[i] = w + + def _init_gui(self): + self._my_layout = QtGui.QGridLayout() + self.slicingWidgets = {} + self.setLayout(self._my_layout) + self.setWindowTitle("Batch Viewer") + + def wheelEvent(self, QWheelEvent): + for v in self.slicingWidgets.values(): + offset = np.sign(QWheelEvent.angleDelta().y()) + v.setSlice(v.getSlice() + np.sign(offset)) + + +def view_batch(*args, width=300, height=300, lut={}): + use_these = args + if not isinstance(use_these, (np.ndarray, np.memmap)): + use_these = list(use_these) + for i in range(len(use_these)): + item = use_these[i] + try: + import torch + + if isinstance(item, torch.Tensor): + item = item.detach().cpu().numpy() + except ImportError: + pass + while len(item.shape) < 4: + item = item[None] + use_these[i] = item + use_these = np.concatenate(use_these, 0) + else: + while len(use_these.shape) < 4: + use_these = use_these[None] + + global app + app = QtGui.QApplication.instance() + if app is None: + app = QtGui.QApplication(sys.argv) + sv = BatchViewer(width=width, height=height) + sv.setBatch(use_these, lut) + sv.show() + app.exit(app.exec_()) + + +if __name__ == "__main__": + import matplotlib.pyplot as plt + + global app + app = QtGui.QApplication.instance() + if app is None: + app = QtGui.QApplication(sys.argv) + sv = BatchViewer() + batch = np.random.uniform(0, 3, (6, 100, 100, 100)).astype(int) + lut = { + 2: np.array([[0, 0.5, 0, 1], [0, 0, 0.5, 1], [0.5, 0, 0, 1], [0.5, 0.5, 0, 1]]) * 255, + 1: np.array([[1, 0.5, 0, 1], [0, 0, 0.5, 1], [0.5, 0, 0, 1], [0.5, 0.5, 0, 1]]) * 255, + } + sv.setBatch(batch, lut) + sv.show() + app.exec_() + app.deleteLater() + sys.exit() + # IPython.embed() diff --git a/tests/data/patient004_frame01.nii.gz b/tests/data/patient004_frame01.nii.gz new file mode 100644 index 00000000..efa8fc25 Binary files /dev/null and b/tests/data/patient004_frame01.nii.gz differ diff --git a/tests/data/patient004_frame01_gt.nii.gz b/tests/data/patient004_frame01_gt.nii.gz new file mode 100644 index 00000000..bf2fa11b Binary files /dev/null and b/tests/data/patient004_frame01_gt.nii.gz differ diff --git a/tests/rand/test_abstract.py b/tests/rand/test_abstract.py index 79f2f1e1..203dc431 100644 --- a/tests/rand/test_abstract.py +++ b/tests/rand/test_abstract.py @@ -21,7 +21,7 @@ def test_tensor_like(self): def test_abstract_error(self): param = AbstractParameter() with self.assertRaises(NotImplementedError): - param.sample((1,)) + param.sample(n_samples=1) if __name__ == "__main__": diff --git a/tests/rand/test_continuous.py b/tests/rand/test_continuous.py index 5fb40b96..e66f140e 100644 --- a/tests/rand/test_continuous.py +++ b/tests/rand/test_continuous.py @@ -8,9 +8,19 @@ class TestContinuous(unittest.TestCase): def test_uniform(self): self.check_distribution(UniformParameter(0, 2), torch.distributions.Uniform(0, 2)) + self.check_constant(UniformParameter(0, 0), 0) + with self.assertRaises(AssertionError): + self.check_constant(UniformParameter(0.0, 0.0), 1.0) + with self.assertRaises(RuntimeError): + self.check_constant(UniformParameter(0.0, 0.0), 0) def test_normal(self): self.check_distribution(NormalParameter(0, 2), torch.distributions.Normal(0, 2)) + self.check_constant(NormalParameter(0, 0), 0) + with self.assertRaises(AssertionError): + self.check_constant(NormalParameter(0.0, 0.0), 1.0) + with self.assertRaises(RuntimeError): + self.check_constant(NormalParameter(0.0, 0.0), 0) def check_distribution(self, param, dist, size=(10,)): state = torch.random.get_rng_state() @@ -20,6 +30,11 @@ def check_distribution(self, param, dist, size=(10,)): res_dist = dist.sample(size) self.assertTrue(res_dist.allclose(res_param)) + def check_constant(self, param, constant, size=(10,)): + res_params = param(size) + assert res_params.shape == torch.Size(size) + self.assertTrue(res_params.allclose(torch.ones(*size, dtype=torch.long) * constant)) + if __name__ == "__main__": unittest.main() diff --git a/tests/rand/test_discrete.py b/tests/rand/test_discrete.py index 8154e8e1..89a494d2 100644 --- a/tests/rand/test_discrete.py +++ b/tests/rand/test_discrete.py @@ -1,18 +1,27 @@ import unittest +import torch + from rising.random import DiscreteCombinationsParameter, DiscreteParameter from rising.random.discrete import combinations_all class TestDiscrete(unittest.TestCase): def test_discrete_error(self): - with self.assertRaises(ValueError): - param = DiscreteParameter((1.0, 2.0), replacement=False, weights=(0.3, 0.7)) + with self.assertRaises(ValueError) as e: + param = DiscreteParameter((1.0, 2.0), replacement=False, weights=("yes", 0.7)) + expected_msg = "weights and cum_weights should only be specified if replacement is set to True!" + assert e.exception.args[0] == expected_msg def test_discrete_parameter(self): - param = DiscreteParameter((1,)) + param = DiscreteParameter((1,), replacement=True) + sample = param(size=(100, 100)) + assert sample.allclose(torch.ones_like(sample) * 1) + + def test_discrete_parameter2(self): + param = DiscreteParameter((1, 2), replacement=True) sample = param() - self.assertEqual(sample, 1) + assert sample in (1, 2) def test_discrete_combinations_parameter(self): param = DiscreteCombinationsParameter((1,)) diff --git a/tests/realtime_viewer.py b/tests/realtime_viewer.py new file mode 100644 index 00000000..e3ad1235 --- /dev/null +++ b/tests/realtime_viewer.py @@ -0,0 +1,154 @@ +# this is the viewer script for 3D volumns visualization +from typing import List, Tuple, Union + +import matplotlib.pyplot as plt +import numpy as np +import torch + +Tensor = Union[np.ndarray, torch.Tensor] + + +def _empty_iterator(tensor) -> bool: + """ + check if a list (tuple) is empty + """ + from collections.abc import Iterable + + if isinstance(tensor, Iterable): + if len(tensor) == 0: + return True + return False + + +def _is_tensor(tensor) -> bool: + """ + return bool indicating if an input is a tensor of numpy or torch. + """ + if torch.is_tensor(tensor): + return True + if isinstance(tensor, np.ndarray): + return True + return False + + +def _is_iterable_tensor(tensor) -> bool: + """ + return bool indicating if an punt is a list or a tuple of tensor + """ + from collections.abc import Iterable + + if isinstance(tensor, Iterable): + if len(tensor) > 0: + if _is_tensor(tensor[0]): + return True + return False + + +def tensor2plotable(tensor) -> np.ndarray: + if isinstance(tensor, np.ndarray): + return tensor + elif isinstance(tensor, torch.Tensor): + return tensor.detach().cpu().numpy() + else: + raise TypeError(f"tensor should be an instance of Tensor, given {type(tensor)}") + + +# below are the functions mostly utilized. +def multi_slice_viewer_debug( + img_volume: Union[Tensor, List[Tensor], Tuple[Tensor, ...]], + *gt_volumes: Tensor, + no_contour=False, + block=False, + alpha=0.2, +) -> None: + def process_mouse_wheel(event): + fig = event.canvas.figure + for i, ax in enumerate(fig.axes): + if event.button == "up": + previous_slice(ax) + elif event.button == "down": + next_slice(ax) + fig.canvas.draw() + + def process_key(event): + fig = event.canvas.figure + ax = fig.axes[0] + if event.key == "j": + previous_slice(ax) + elif event.key == "k": + next_slice(ax) + fig.canvas.draw() + + def previous_slice(ax): + img_volume = ax.img_volume + ax.index = (ax.index - 1) if (ax.index - 1) >= 0 else 0 # wrap around using % + ax.images[0].set_array(img_volume[ax.index]) + + if ax.gt_volume is not None: + if not no_contour: + for con in ax.con.collections: + con.remove() + ax.con = ax.contour(ax.gt_volume[ax.index]) + else: + ax.con.remove() + ax.con = ax.imshow(ax.gt_volume[ax.index], alpha=alpha, cmap="rainbow") + # ax.set_title(f"plane = {ax.index}") + + def next_slice(ax): + img_volume = ax.img_volume + ax.index = (ax.index + 1) if (ax.index + 1) < img_volume.shape[0] else img_volume.shape[0] - 1 + ax.images[0].set_array(img_volume[ax.index]) + + if ax.gt_volume is not None: + if not no_contour: + for con in ax.con.collections: + con.remove() + ax.con = ax.contour(ax.gt_volume[ax.index]) + else: + ax.con.remove() + ax.con = ax.imshow(ax.gt_volume[ax.index], alpha=alpha, cmap="rainbow") + # ax.set_title(f"plane = {ax.index}") + + try: + import matplotlib + + matplotlib.use("tkagg", force=True) + except Exception as e: + print(e) + + # assertion part: + assert _is_tensor(img_volume) or _is_iterable_tensor(img_volume), f"input wrong for img_volume, given {img_volume}." + assert _is_iterable_tensor(gt_volumes) or gt_volumes == (), f"input wrong for gt_volumes, given {gt_volumes}." + if _is_tensor(img_volume): + img_volume = [img_volume] + row_num, col_num = len(img_volume), max(len(gt_volumes), 1) + + fig, axs = plt.subplots(row_num, col_num) + if not isinstance(axs, np.ndarray): + # lack of numpy wrapper + axs = np.array([axs]) + axs = axs.reshape((row_num, col_num)) + + for _row_num, row_axs in enumerate(axs): + # each row + assert len(row_axs) == col_num + for _col_num, ax in enumerate(row_axs): + ax.img_volume = tensor2plotable(img_volume[_row_num]) + min_value, max_value = ax.img_volume.min(), ax.img_volume.max() + + ax.index = ax.img_volume.shape[0] // 2 + ax.imshow(ax.img_volume[ax.index], cmap="gray", vmin=min_value, vmax=max_value) + ax.gt_volume = None if _empty_iterator(gt_volumes) else tensor2plotable(gt_volumes[_col_num]) + try: + if not no_contour: + ax.con = ax.contour(ax.gt_volume[ax.index]) + else: + ax.con = ax.imshow(ax.gt_volume[ax.index], alpha=alpha, cmap="rainbow") + except Exception as e: + pass + ax.axis("off") + # ax.set_title(f"plane = {ax.index}") + + fig.canvas.mpl_connect("key_press_event", process_key) + fig.canvas.mpl_connect("scroll_event", process_mouse_wheel) + plt.show(block=block) diff --git a/tests/transforms/_helpers.py b/tests/transforms/_helpers.py index e4458f13..7577fb32 100644 --- a/tests/transforms/_helpers.py +++ b/tests/transforms/_helpers.py @@ -1,7 +1,7 @@ -from rising.transforms.abstract import AbstractTransform +from rising.transforms.abstract import _AbstractTransform -def chech_data_preservation(trafo: AbstractTransform, batch: dict, key: str = "data") -> bool: +def chech_data_preservation(trafo: _AbstractTransform, batch: dict, key: str = "data") -> bool: """ Checks for inplace modification of input data diff --git a/tests/transforms/functional/test_pad.py b/tests/transforms/functional/test_pad.py new file mode 100644 index 00000000..44b0c104 --- /dev/null +++ b/tests/transforms/functional/test_pad.py @@ -0,0 +1,31 @@ +import unittest + +import torch + +from rising.transforms.functional.pad import pad + + +class MyTestCase(unittest.TestCase): + def setUp(self) -> None: + self.data = torch.ones(1, 1, 10, 10) + self.complex_data = torch.ones(10, 3, 100, 100, 100) + + def test_constant_pad(self): + for p in range(1, 4): + padded = pad(self.data, pad_size=p, mode="constant", value=1.23) + expected = torch.ones(1, 1, 10 + p * 2, 10 + p * 2) + expected[:, :, 0:p, :] = 1.23 + expected[:, :, :, 0:p] = 1.23 + expected[:, :, -p:, :] = 1.23 + expected[:, :, :, -p:] = 1.23 + self.assertTrue((padded == expected).all()) + + def test_complex_pad(self): + for p in range(1, 4): + padded = pad(self.complex_data, pad_size=p, mode="circular", ) + expected = torch.ones(10, 3, 100 + p * 2, 100 + p * 2, 100 + p * 2) + self.assertTrue((padded == expected).all()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/transforms/functional/test_sitk.py b/tests/transforms/functional/test_sitk.py new file mode 100644 index 00000000..8719782c --- /dev/null +++ b/tests/transforms/functional/test_sitk.py @@ -0,0 +1,25 @@ +from rising.transforms.functional.sitk import itk_resample, itk_clip +import SimpleITK as sitk +import unittest +import numpy as np +from deepclustering2.viewer import multi_slice_viewer_debug + + +class SITKTestCase(unittest.TestCase): + + def setUp(self) -> None: + super().setUp() + self._image_path = "../../data/patient004_frame01.nii.gz" + self._mask_path = "../../data/patient004_frame01_gt.nii.gz" + self._sitk_image = sitk.ReadImage(self._image_path) + self._sitk_mask = sitk.ReadImage(self._mask_path) + + def test_resampling(self): + resampled_image = itk_resample(self._sitk_image, spacing=(0.1, 0.1, 10), interpolation="linear", pad_value=0) + resampled_mask = itk_resample(self._sitk_mask, spacing=(0.1, 0.1, 10), interpolation="nearest", pad_value=0) + assert np.allclose(np.unique(sitk.GetArrayFromImage(resampled_mask)), np.array([0, 1, 2, 3])) + assert resampled_image.GetSpacing() == (0.1, 0.1, 10) + + def test_clip(self): + clipped_image = itk_clip(self._sitk_image, 10, 600) + multi_slice_viewer_debug(sitk.GetArrayFromImage(clipped_image).squeeze(), block=True) diff --git a/tests/transforms/test_abstract_transform.py b/tests/transforms/test_abstract_transform.py index be96cb8c..8ee7f2c0 100644 --- a/tests/transforms/test_abstract_transform.py +++ b/tests/transforms/test_abstract_transform.py @@ -1,19 +1,35 @@ import unittest +from typing import Sequence from unittest.mock import Mock, call import torch -from rising.transforms import AbstractTransform, BaseTransform, PerChannelTransform, PerSampleTransform - - -class AddTransform(AbstractTransform): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.grad_tensor = torch.rand(1, 1, 32, 32, requires_grad=True) +from rising.transforms import BaseTransform, PerChannelTransformMixin, PerSampleTransformMixin, _AbstractTransform + + +class AddTransform(PerSampleTransformMixin, BaseTransform): + def __init__( + self, + *, + dims=[1], + keys: Sequence[str] = ("data",), + grad: bool = False, + property_names: Sequence[str] = (), + key_associate_kwargs_names: Sequence[str] = ("dims",) + ): + super(AddTransform, self).__init__( + augment_fn=sum_dim, + keys=keys, + grad=grad, + property_names=property_names, + key_associate_kwargs_names=key_associate_kwargs_names, + per_sample=True, + seeded=True, + dims=dims, + ) def forward(self, **data) -> dict: - data["data"] = data["data"] + self.grad_tensor - return data + return super().forward(**data) def sum_dim(data, dims, **kwargs): @@ -30,7 +46,7 @@ def setUp(self) -> None: self.batch_dict = {"data": torch.rand(1, 1, 32, 32), "seg": torch.rand(1, 1, 32, 32), "label": torch.arange(3)} def test_abstract_transform(self): - trafo = AbstractTransform(grad=False, internal0=True) + trafo = _AbstractTransform(grad=False, internal0=True) self.assertTrue(trafo.internal0) with self.assertRaises(NotImplementedError): trafo(**self.batch_dict) @@ -69,7 +85,7 @@ def test_per_sample_transform(self): def augment_fn(inp, *args, **kwargs): return mock(inp) - trafo = PerSampleTransform(augment_fn, keys=("label",)) + trafo = PerSampleTransformMixin(augment_fn, keys=("label",)) output = trafo(**self.batch_dict) calls = [ call(torch.tensor([0])), @@ -84,7 +100,7 @@ def test_per_channel_transform_per_channel_true(self): def augment_fn(inp, *args, **kwargs): return mock(inp) - trafo = PerChannelTransform(augment_fn, per_channel=True, keys=("label",)) + trafo = PerChannelTransformMixin(augment_fn, per_channel=True, keys=("label",)) self.batch_dict["label"] = self.batch_dict["label"][None] output = trafo(**self.batch_dict) calls = [ @@ -100,7 +116,7 @@ def test_per_channel_transform_per_channel_false(self): def augment_fn(inp, *args, **kwargs): return mock(inp) - trafo = PerChannelTransform(augment_fn, per_channel=False, keys=("label",)) + trafo = PerChannelTransformMixin(augment_fn, per_channel=False, keys=("label",)) self.batch_dict["label"] = self.batch_dict["label"][None] output = trafo(**self.batch_dict) mock.assert_called_once() diff --git a/tests/transforms/test_affine.py b/tests/transforms/test_affine.py index 009deade..b1680ad3 100644 --- a/tests/transforms/test_affine.py +++ b/tests/transforms/test_affine.py @@ -2,7 +2,8 @@ import torch -from rising.transforms.affine import Affine, BaseAffine, Resize, Rotate, Scale, StackedAffine, Translate +from rising.random import UniformParameter +from rising.transforms.affine import BaseAffine, Resize, Rotate, Scale, Translate, _Affine, _StackedAffine from rising.utils.affine import matrix_to_cartesian, matrix_to_homogeneous @@ -19,7 +20,7 @@ def test_affine(self): target_size = target_sizes.pop(0) with self.subTest(adjust_size=adjust_size, target_size=target_size, output_size=output_size): - trafo = Affine(matrix=matrix, adjust_size=adjust_size, output_size=output_size) + trafo = _Affine(matrix=matrix, adjust_size=adjust_size, output_size=output_size) sample = {"data": image_batch, "label": 4} if output_size is not None and adjust_size: with self.assertWarns(UserWarning): @@ -50,7 +51,7 @@ def test_affine_assemble_matrix(self): for matrix, expected, ve in zip(matrices, expected_matrices, value_error): with self.subTest(matrix=matrix, expected=expected): - trafo = Affine(matrix=matrix) + trafo = _Affine(matrix=matrix) if ve: with self.assertRaises(ValueError): assembled = trafo.assemble_matrix(**batch) @@ -60,15 +61,15 @@ def test_affine_assemble_matrix(self): def test_affine_stacking(self): affines = [ - Affine(scale=1), + _Affine(scale=1), [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], torch.tensor([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]]), - StackedAffine(Affine(scale=1), Affine(scale=1)), + _StackedAffine(_Affine(scale=1), _Affine(scale=1)), ] for first_affine in affines: for second_affine in affines: - if not isinstance(first_affine, Affine) and not isinstance(second_affine, Affine): + if not isinstance(first_affine, _Affine) and not isinstance(second_affine, _Affine): continue if torch.is_tensor(first_affine): @@ -78,12 +79,12 @@ def test_affine_stacking(self): with self.subTest(first_affine=first_affine, second_affine=second_affine): result = first_affine + second_affine - self.assertIsInstance(result, StackedAffine) + self.assertIsInstance(result, _StackedAffine) def test_stacked_transformation_assembly(self): first_matrix = torch.tensor([[[2.0, 0.0, 1.0], [0.0, 3.0, 2.0]]]) second_matrix = torch.tensor([[[4.0, 0.0, 3.0], [0.0, 5.0, 4.0]]]) - trafo = StackedAffine([first_matrix, second_matrix]) + trafo = _StackedAffine([first_matrix, second_matrix]) sample = {"data": torch.rand(1, 3, 25, 25)} @@ -159,6 +160,18 @@ def test_translation_assemble_matrix_with_pixel(self): trafo.assemble_matrix(**sample) self.assertTrue(expected.allclose(expected)) + def test_affine_prob(self): + image = torch.randn(13, 1, 224, 224, 224) + prob = torch.randn(13, 1, 224, 224, 224) + translate = BaseAffine( + rotation=UniformParameter(-10, 10), + degree=True, + scale=UniformParameter(-10, 10), + keys=("data", "prob"), + ) + image_, prob_ = translate(data=image, prob=prob).values() + pass + if __name__ == "__main__": unittest.main() diff --git a/tests/transforms/test_affine_transform.py b/tests/transforms/test_affine_transform.py new file mode 100644 index 00000000..e8b6d4b9 --- /dev/null +++ b/tests/transforms/test_affine_transform.py @@ -0,0 +1,209 @@ +import unittest + +import numpy as np +import SimpleITK as sitk +import torch + +from rising.random import UniformParameter +from rising.transforms.affine import BaseAffine, Resize, Rotate, Scale, Translate, _Affine, _StackedAffine +from rising.utils.affine import matrix_to_cartesian, matrix_to_homogeneous + + +class AffineTestCase(unittest.TestCase): + def setUp(self) -> None: + torch.manual_seed(0) + self.batch_dict = { + "data": self.load_nii_data("/home/jizong/Workspace/rising/tests/data/patient004_frame01.nii.gz"), + "label": self.load_nii_data("/home/jizong/Workspace/rising/tests/data/patient004_frame01_gt.nii.gz"), + } + + def load_nii_data(self, path): + return torch.from_numpy( + sitk.GetArrayFromImage(sitk.ReadImage(str(path))).astype(float, copy=False) + ).unsqueeze(1) + + def test_affine(self): + matrix = torch.tensor([[4.0, 0.0, 0.0], [0.0, 5.0, 0.0]]) + image_batch = torch.zeros(10, 3, 25, 25, dtype=torch.float, device="cpu") + matrix = matrix.expand(image_batch.size(0), -1, -1).clone() + + target_sizes = [(100, 125), image_batch.shape[2:], (50, 50), (50, 50), (45, 50), (45, 50)] + + for output_size in [None, 50, (45, 50)]: + for adjust_size in [True, False]: + target_size = target_sizes.pop(0) + + with self.subTest(adjust_size=adjust_size, target_size=target_size, output_size=output_size): + trafo = _Affine(matrix=matrix, adjust_size=adjust_size, output_size=output_size) + sample = {"data": image_batch, "label": 4} + if output_size is not None and adjust_size: + with self.assertWarns(UserWarning): + result = trafo(**sample) + else: + result = trafo(**sample) + self.assertTupleEqual(result["data"].shape[2:], target_size) + + self.assertEqual(sample["label"], result["label"]) + + def test_affine_assemble_matrix(self): + matrices = [ + [[1.0, 0.0], [0.0, 1.0]], + [[1.0, 0.0, 1.0], [0.0, 1.0, 1.0]], + [[1.0, 0.0, 1.0], [0.0, 1.0, 1.0], [0.0, 0.0, 1.0]], + None, + [0.0, 1.0, 1.0, 0.0], + ] + expected_matrices = [ + torch.tensor([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]])[None], + torch.tensor([[1.0, 0.0, 1.0], [0.0, 1.0, 1.0]])[None], + torch.tensor([[1.0, 0.0, 1.0], [0.0, 1.0, 1.0]])[None], + None, + None, + ] + value_error = [False, False, False, True, True] + batch = {"data": torch.zeros(1, 1, 10, 10)} + + for matrix, expected, ve in zip(matrices, expected_matrices, value_error): + with self.subTest(matrix=matrix, expected=expected): + trafo = _Affine(matrix=matrix) + if ve: + with self.assertRaises(ValueError): + assembled = trafo.assemble_matrix(**batch) + else: + assembled = trafo.assemble_matrix(**batch) + self.assertTrue(expected.allclose(assembled)) + + def test_affine_stacking(self): + affines = [ + _Affine(scale=1), + [[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], + torch.tensor([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]]), + _StackedAffine(_Affine(scale=1), _Affine(scale=1)), + ] + + for first_affine in affines: + for second_affine in affines: + if not isinstance(first_affine, _Affine) and not isinstance(second_affine, _Affine): + continue + + if torch.is_tensor(first_affine): + # TODO: Remove this, once this has been fixed in PyTorch: + # PR: https://github.com/pytorch/pytorch/pull/31769 + continue + + with self.subTest(first_affine=first_affine, second_affine=second_affine): + result = first_affine + second_affine + self.assertIsInstance(result, _StackedAffine) + + def test_stacked_transformation_assembly(self): + first_matrix = torch.tensor([[[2.0, 0.0, 1.0], [0.0, 3.0, 2.0]]]) + second_matrix = torch.tensor([[[4.0, 0.0, 3.0], [0.0, 5.0, 4.0]]]) + trafo = _StackedAffine([first_matrix, second_matrix]) + + sample = {"data": torch.rand(1, 3, 25, 25)} + + matrix = trafo.assemble_matrix(**sample) + + target_matrix = matrix_to_cartesian( + torch.bmm(matrix_to_homogeneous(first_matrix), matrix_to_homogeneous(second_matrix)) + ) + + self.assertTrue(torch.allclose(matrix, target_matrix)) + + def test_affine_subtypes(self): + sample = {"data": torch.rand(1, 3, 25, 30)} + + trafos = [ + BaseAffine(), + BaseAffine(adjust_size=True), + Scale(5, adjust_size=True), + Scale([5, 3], adjust_size=True), + Scale(5, adjust_size=False), + Scale([5, 3], adjust_size=False), + Resize(50), + Resize((50, 90)), + Rotate([90], adjust_size=True, degree=True), + Rotate([90], adjust_size=False, degree=True), + Translate(10, adjust_size=True, unit="pixel"), + Translate(10, adjust_size=False, unit="pixel"), + Translate([5, 10], adjust_size=False, unit="pixel"), + Scale(5, adjust_size=False, per_sample=False), + Rotate([90], adjust_size=False, degree=True, per_sample=False), + Translate(10, adjust_size=False, unit="pixel", per_sample=False), + ] + + expected_sizes = [ + (25, 30), + (25, 30), + (5, 6), + (5, 10), + (25, 30), + (25, 30), + (50, 50), + (50, 90), + (30, 25), + (25, 30), + (25, 30), + (25, 30), + (25, 30), + (25, 30), + (25, 30), + (25, 30), + ] + + for trafo, expected_size in zip(trafos, expected_sizes): + with self.subTest(trafo=trafo, exp_size=expected_size): + result = trafo(**sample)["data"] + self.assertIsInstance(result, torch.Tensor) + self.assertTupleEqual(expected_size, result.shape[-2:]) + + def test_translation_assemble_matrix_with_pixel(self): + trafo = Translate([10, 100], unit="pixel") + sample = {"data": torch.rand(3, 3, 100, 100)} + expected = torch.tensor( + [ + [1.0, 0.0, -0.1], + [0.0, 1.0, -0.01], + [1.0, 0.0, -0.1], + [0.0, 1.0, -0.01], + [1.0, 0.0, -0.1], + [0.0, 1.0, -0.01], + ] + ) + + trafo.assemble_matrix(**sample) + self.assertTrue(expected.allclose(expected)) + + def test_affine_prob(self): + image = torch.randn(13, 1, 224, 224) + prob = torch.randn(13, 1, 224, 224) + translate = BaseAffine( + rotation=(UniformParameter(-10, 10)), + degree=True, + scale=(UniformParameter(-10, 10), UniformParameter(-10, 10)), + keys=("data", "prob"), + ) + image_, prob_ = translate(data=image, prob=prob).values() + + def test_again(self): + random_affine = BaseAffine( + rotation=UniformParameter(-50, 50), + degree=True, + scale=(UniformParameter(0.7, 0.95), UniformParameter(2, 2.8)), + translation=UniformParameter(-0.2, 0.2), + p=1, + keys=("data", "label"), + interpolation_mode=("bilinear", "nearest"), + ) + image, label = self.batch_dict.values() + for i in range(100): + output_image, output_label = random_affine(**self.batch_dict).values() + from tests.realtime_viewer import multi_slice_viewer_debug + + multi_slice_viewer_debug( + [image.squeeze(), output_image.squeeze()], label.squeeze(), output_label.squeeze(), block=True + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/transforms/test_compose.py b/tests/transforms/test_compose.py index d34b63b6..418bef3b 100644 --- a/tests/transforms/test_compose.py +++ b/tests/transforms/test_compose.py @@ -4,9 +4,10 @@ import torch -from rising.transforms import AbstractTransform +from rising.transforms import _AbstractTransform from rising.transforms.compose import Compose, DropoutCompose, OneOf, _TransformWrapper from rising.transforms.spatial import Mirror +from rising.utils.transforms import iter_transform class TestCompose(unittest.TestCase): @@ -71,7 +72,7 @@ def test_dropout_compose_error(self): compose = DropoutCompose(self.transforms, dropout=[1.0]) def test_device_dtype_change(self): - class DummyTrafo(AbstractTransform): + class DummyTrafo(_AbstractTransform): def __init__(self, a): super().__init__(False) self.register_buffer("tmp", a) @@ -134,6 +135,11 @@ def test_no_trafo_error(self): with self.assertRaises(ValueError): comp = trafo_cls() + def test_iter_transform(self): + tras = iter_transform(self.transforms) + tras2 = iter_transform(Compose(self.transforms)) + assert list(tras) == list(tras2) + if __name__ == "__main__": unittest.main() diff --git a/tests/transforms/test_format_transforms.py b/tests/transforms/test_format_transforms.py index 245ccf3b..5683d9bf 100644 --- a/tests/transforms/test_format_transforms.py +++ b/tests/transforms/test_format_transforms.py @@ -2,6 +2,8 @@ import torch +from rising.loading import default_transform_call +from rising.transforms import Compose from rising.transforms.format import MapToSeq, RenameKeys, SeqToMap @@ -17,7 +19,7 @@ def test_map_to_seq(self): self.assertEqual(out[2], 2) def test_seq_to_map(self): - trafo = SeqToMap(("data", "seg", "label")) + trafo = Compose(SeqToMap(("data", "seg", "label")), transform_call=default_transform_call) out = trafo(0, 1, 2) self.assertEqual(out["data"], 0) self.assertEqual(out["seg"], 1) diff --git a/tests/transforms/test_grid.py b/tests/transforms/test_grid.py new file mode 100644 index 00000000..bedde078 --- /dev/null +++ b/tests/transforms/test_grid.py @@ -0,0 +1,45 @@ +import typing as t +import unittest + +import numpy as np +import torch +from PIL import Image +from deepclustering2.viewer import multi_slice_viewer_debug + +from rising.transforms.grid import ElasticDistortion, RadialDistortion +from tests.utils.mise import gpu_timeit + +T = t.TypeVar("T") +item_or_seq = t.Union[T, t.Sequence[T]] +float_or_seq = item_or_seq[int] + + +class TestGridCase(unittest.TestCase): + def setUp(self) -> None: + super().setUp() + + self.device = "cuda" + self.dtype = torch.float16 + self._image = Image.open("/home/jizong/Workspace/rising/notebooks/MedNIST/BreastMRI/000000.jpeg").convert("L") + self._image = torch.from_numpy(np.asarray(self._image), ).float()[None, None, ...]. \ + repeat(100, 1, 1, 1).to(self.device).to(self.dtype) + self._target = (self._image > 0.5).to(self.dtype) + + def test_elastic_transform(self): + transform = ElasticDistortion(std=2, alpha=0.1, keys=("data", "target"), + interpolation_mode=("bilinear", "nearest"), per_sample=False).to(self.device).to( + self.dtype) + with gpu_timeit(): + for _ in range(100): + output, target = transform(data=self._image, target=self._target).values() + + def test_radial_distortion(self): + transform = RadialDistortion(scale=(0, 0, 0.3), keys=("data", "target"), + interpolation_mode=("bilinear", "nearest")).to(self.device).to(self.dtype) + + with gpu_timeit(): + for _ in range(100): + output, target = transform(data=self._image, target=self._target).values() + multi_slice_viewer_debug([self._image.squeeze().float(), output.float().squeeze()], + self._target.squeeze().float(), target.squeeze().float(), block=True, + no_contour=True) diff --git a/tests/transforms/test_intensity_transforms.py b/tests/transforms/test_intensity_transforms.py index 03260056..96bffaf4 100644 --- a/tests/transforms/test_intensity_transforms.py +++ b/tests/transforms/test_intensity_transforms.py @@ -1,7 +1,6 @@ import random import unittest from math import isclose -from unittest.mock import Mock, call import torch @@ -11,14 +10,13 @@ ExponentialNoise, GammaCorrection, GaussianNoise, - Noise, NormMeanStd, NormMinMax, + NormPercentile, NormRange, NormZeroMeanUnitStd, RandomAddValue, RandomScaleValue, - RandomValuePerChannel, ) from tests.transforms import chech_data_preservation @@ -52,6 +50,20 @@ def test_norm_range_transform(self): self.assertTrue(isclose(outp["data"].min().item(), 0.1, abs_tol=1e-6)) self.assertTrue(isclose(outp["data"].max().item(), 0.2, abs_tol=1e-6)) + def test_norm_percentile_transform(self): + trafo = NormPercentile(0.001, 0.99, per_channel=False) + # out = trafo(**self.batch_dict) + self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) + + trafo = NormPercentile(0.001, 0.99, per_channel=True) + self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) + + outp = trafo(**self.batch_dict) + self.assertTrue( + isclose(outp["data"].min().item(), torch.quantile(self.batch_dict["data"], 0.001), abs_tol=1e-6) + ) + self.assertTrue(isclose(outp["data"].max().item(), torch.quantile(self.batch_dict["data"], 0.99), abs_tol=1e-6)) + def test_norm_min_max_transform(self): trafo = NormMinMax(per_channel=False) self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) @@ -87,11 +99,6 @@ def test_norm_std_transform(self): self.assertTrue(isclose(outp["data"].mean().item(), 0.0, abs_tol=1e-6)) self.assertTrue(isclose(outp["data"].std().item(), 1.0, abs_tol=1e-6)) - def test_noise_transform(self): - trafo = Noise("normal", mean=75, std=1) - self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) - self.check_noise_distance(trafo) - def test_expoential_noise_transform(self): trafo = ExponentialNoise(lambd=0.0001) self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) @@ -107,25 +114,6 @@ def check_noise_distance(self, trafo, min_diff=50): comp_diff = (outp["data"] - self.batch_dict["data"]).mean().item() self.assertTrue(comp_diff > min_diff) - def test_per_channel_transform_per_channel_true(self): - # TODO: check why sometimes an overflow occurs - mock = Mock(return_value=0) - - def augment_fn(inp, *args, **kwargs): - return mock(inp) - - trafo = RandomValuePerChannel( - augment_fn, random_sampler=DiscreteParameter((1,)), per_channel=True, keys=("label",) - ) - self.batch_dict["label"] = self.batch_dict["label"][None] - output = trafo(**self.batch_dict) - calls = [ - call(torch.tensor([0])), - call(torch.tensor([1])), - call(torch.tensor([2])), - ] - mock.assert_has_calls(calls) - def test_random_add_value(self): trafo = RandomAddValue(DiscreteParameter((2,))) self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) diff --git a/tests/transforms/test_pad.py b/tests/transforms/test_pad.py new file mode 100644 index 00000000..36f5f03a --- /dev/null +++ b/tests/transforms/test_pad.py @@ -0,0 +1,33 @@ +import unittest + +import torch + +from rising.transforms import Pad + + +class TestPad(unittest.TestCase): + def setUp(self) -> None: + super().setUp() + image = torch.zeros(1, 1, 10, 10) + self._data = {"data": image, "seg": image.clone()} + + def test_pad_no_pad(self): + for pad_size in range(5, 10): + transform = Pad(pad_size=pad_size, pad_value=(0, -1), keys=("data", "seg")) + cropped = transform(**self._data) + expected = torch.zeros(1, 1, 10, 10) + assert torch.allclose(cropped["data"], expected) + assert torch.allclose(cropped["seg"], expected) + + def test_pad(self): + for pad_size in range(11, 15): + transform = Pad(pad_size=pad_size, pad_value=(0, -1), keys=("data", "seg")) + cropped = transform(**self._data) + expected = torch.zeros(1, 1, pad_size, pad_size) + expected_mask = torch.zeros_like(expected) + expected_mask[:, :, 0 : (pad_size - 10) // 2, :] = -1 + expected_mask[:, :, :, 0 : (pad_size - 10) // 2] = -1 + expected_mask[:, :, :, -(pad_size - (pad_size - 10) // 2 - 10) :] = -1 + expected_mask[:, :, -(pad_size - (pad_size - 10) // 2 - 10) :, :] = -1 + assert torch.allclose(cropped["data"], expected) + assert torch.allclose(cropped["seg"], expected_mask) diff --git a/tests/transforms/test_spatial_transforms.py b/tests/transforms/test_spatial_transforms.py index 2b3410e4..241b86a2 100644 --- a/tests/transforms/test_spatial_transforms.py +++ b/tests/transforms/test_spatial_transforms.py @@ -1,13 +1,16 @@ import random import unittest +import SimpleITK as sitk import torch +from matplotlib import pyplot as plt +from rising.constants import FInterpolation from rising.loading import DataLoader -from rising.random import UniformParameter -from rising.transforms import Mirror, ProgressiveResize, ResizeNative, Rot90, SizeStepScheduler, Zoom -from rising.transforms.functional import resize_native -from tests.transforms import chech_data_preservation +from rising.random import DiscreteParameter, UniformParameter +from rising.transforms import Mirror, ProgressiveResize, ResizeNative, Rot90, SizeStepScheduler, Zoom, \ + ResizeNativeCentreCrop +from tests.realtime_viewer import multi_slice_viewer_debug class TestSpatialTransforms(unittest.TestCase): @@ -15,58 +18,79 @@ def setUp(self) -> None: torch.manual_seed(0) random.seed(0) self.batch_dict = { - "data": torch.arange(1, 10).reshape(1, 1, 3, 3).float(), - "seg": torch.randint(0, 3, (1, 1, 3, 3)).long(), - "label": torch.arange(3), + "data": self.load_nii_data("../../tests/data/patient004_frame01.nii.gz"), + "label": self.load_nii_data("../../tests/data/patient004_frame01_gt.nii.gz"), } + def load_nii_data(self, path): + return torch.from_numpy( + sitk.GetArrayFromImage(sitk.ReadImage(str(path))).astype(float, copy=False) + ).unsqueeze(1) + def test_mirror_transform(self): - trafo = Mirror((0, 1)) + trafo = Mirror(dims=DiscreteParameter((0, 1, (0, 1))), p_sample=0.5, keys=("data", "label")) outp = trafo(**self.batch_dict) - - self.assertTrue(outp["data"][0, 0].allclose(torch.tensor([[9, 8, 7], [6, 5, 4], [3, 2, 1]]).float())) - self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) + image1, target1 = self.batch_dict.values() + image2, target2 = outp.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze()) + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=True) def test_rot90_transform(self): - random.seed(0) - trafo = Rot90((0, 1), prob=1, num_rots=(1,)) + trafo = Rot90(dims=[0, 1], num_rots=DiscreteParameter((2,)), p_sample=0.5, keys=("data", "label")) outp = trafo(**self.batch_dict) - self.assertTrue((outp["data"][0, 0] == torch.tensor([[3, 6, 9], [2, 5, 8], [1, 4, 7]])).all()) - self.assertTrue(chech_data_preservation(trafo, self.batch_dict)) + image1, target1 = self.batch_dict.values() + image2, target2 = outp.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze()) + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=True) - trafo = Rot90((0, 1), prob=0) - data_orig = self.batch_dict["data"].clone() + trafo = Rot90(dims=[0, 1], num_rots=DiscreteParameter((2,)), p_sample=1, keys=("data", "label")) outp = trafo(**self.batch_dict) - self.assertTrue((outp["data"] == data_orig).all()) + image1, target1 = self.batch_dict.values() + image2, target2 = outp.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze()) + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=True) def test_resize_transform(self): - trafo = ResizeNative((2, 2)) + trafo = ResizeNative( + (128, 256), + keys=( + "data", + "label", + ), + mode=(FInterpolation.bilinear, FInterpolation.nearest), + align_corners=(False, None), + ) out = trafo(**self.batch_dict) - expected = torch.tensor([[1, 2], [4, 5]]) - self.assertTrue((out["data"] == expected).all()) + image1, target1 = self.batch_dict.values() + image2, target2 = out.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze()) + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=True, no_contour=True) def test_zoom_transform(self): - _range = (2.0, 3.0) - torch.manual_seed(0) - scale_factor = UniformParameter(*_range)() + _range = (1.5, 2.0) + # scale_factor = UniformParameter(*_range)() - trafo = Zoom(scale_factor=UniformParameter(*_range)) - torch.manual_seed(0) - out = trafo(**self.batch_dict) + trafo = Zoom(scale_factor=[UniformParameter(*_range), UniformParameter(*_range)], keys=("data", "label")) - expected = resize_native(self.batch_dict["data"], mode="nearest", scale_factor=scale_factor) - self.assertTrue((out["data"] == expected).all()) + out = trafo(**self.batch_dict) + image1, target1 = self.batch_dict.values() + image2, target2 = out.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze(), block=False, no_contour=True) + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=True, no_contour=True) def test_progressive_resize(self): + image1, target1 = self.batch_dict.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze(), no_contour=True) + sizes = [1, 3, 6] - scheduler = SizeStepScheduler([1, 2], [1, 3, 6]) - trafo = ProgressiveResize(scheduler) + scheduler = SizeStepScheduler([1, 2], [112, 224, 336]) + trafo = ProgressiveResize(scheduler, keys=("data", "label")) for i in range(3): outp = trafo(**self.batch_dict) - self.assertTrue(all([s == sizes[i] for s in outp["data"].shape[2:]])) + image2, target2 = outp.values() + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=False, no_contour=True) - trafo.reset_step() - self.assertEqual(trafo.step, 0) + plt.show() def test_size_step_scheduler(self): scheduler = SizeStepScheduler([10, 20], [16, 32, 64]) @@ -90,9 +114,18 @@ def test_progressive_resize_integration(self): data_shape = [tuple(i["data"].shape) for i in loader] - self.assertIn((1, 1, 1, 1, 1), data_shape) - # self.assertIn((1, 1, 3, 3, 3), data_shape) - self.assertIn((1, 1, 6, 6, 6), data_shape) + self.assertIn((1, 10, 1, 1, 1), data_shape) + self.assertIn((1, 10, 3, 3, 3), data_shape) + self.assertIn((1, 10, 6, 6, 6), data_shape) + + def test_resize_native_center_crop(self): + trafo = ResizeNativeCentreCrop(size=(1000, 2000), margin=(10, 15), keys=("data", "label"), + mode=(FInterpolation.bilinear, FInterpolation.nearest)) + outp = trafo(**self.batch_dict) + image1, target1 = self.batch_dict.values() + image2, target2 = outp.values() + multi_slice_viewer_debug(image1.squeeze(), target1.squeeze()) + multi_slice_viewer_debug(image2.squeeze(), target2.squeeze(), block=True) if __name__ == "__main__": diff --git a/tests/utils/mise.py b/tests/utils/mise.py new file mode 100644 index 00000000..75f2bb22 --- /dev/null +++ b/tests/utils/mise.py @@ -0,0 +1,27 @@ +import functools + +import torch +from loguru import logger + + +class gpu_timeit: + + def __enter__(self): + self.start = torch.cuda.Event(enable_timing=True) + self.end = torch.cuda.Event(enable_timing=True) + self.start.record() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + self.end.record() + torch.cuda.synchronize() + elapsed_time = self.start.elapsed_time(self.end) + logger.opt(depth=1).info(f"operation time: {elapsed_time / 1000:.4f}.") + + def __call__(self, func): + @functools.wraps(func) + def wrapped_func(*args, **kwargs): + with self.__class__(): + return func(*args, **kwargs) + + return wrapped_func