Skip to content
This repository was archived by the owner on Jan 12, 2024. It is now read-only.
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
482 changes: 482 additions & 0 deletions notebooks/modelsGenesis_in_dMRI.ipynb

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion requirements/install.txt
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
numpy
torch>=1.3 # before 1.3 torch.nn.functional.affine_grid has no argument align_corners
torch>=1.6 # before 1.6 torch.searchsorted is not present
threadpoolctl
tqdm

2 changes: 2 additions & 0 deletions rising/transforms/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
* Spatial Transforms
* Tensor Transforms
* Utility Transforms
* Painting Transforms
"""

from rising.transforms.abstract import *
Expand All @@ -29,3 +30,4 @@
from rising.transforms.utility import *
from rising.transforms.tensor import *
from rising.transforms.affine import *
from rising.transforms.painting import *
1 change: 1 addition & 0 deletions rising/transforms/functional/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,3 +20,4 @@
from rising.transforms.functional.tensor import *
from rising.transforms.functional.utility import *
from rising.transforms.functional.channel import *
from rising.transforms.functional.painting import *
40 changes: 39 additions & 1 deletion rising/transforms/functional/intensity.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,11 @@
from typing import Union, Sequence, Optional

from rising.utils import check_scalar
from rising.utils.torchinterp1d import Interp1d

__all__ = ["norm_range", "norm_min_max", "norm_zero_mean_unit_std", "norm_mean_std",
"add_noise", "add_value", "gamma_correction", "scale_by_value", "clamp"]
"add_noise", "add_value", "gamma_correction", "scale_by_value", "clamp",
"bezier_3rd_order", "random_inversion"]


def clamp(data: torch.Tensor, min: float, max: float,
Expand Down Expand Up @@ -227,3 +229,39 @@ def scale_by_value(data: torch.Tensor, value: float,
torch.Tensor: augmented data
"""
return torch.mul(data, value, out=out)


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))

t = torch.linspace(0.0, 1.0, 1000).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

# scaling according to maxv,minv
points = points*(maxv-minv) + minv

xvals = points[:,0]
yvals = points[:,1]

out_flat = Interp1d.apply(xvals, yvals, data.view(-1))

return out_flat.view(data.shape)


def random_inversion(data: torch.Tensor, prob_inversion: float=0.5,
maxv: float=1.0, minv: float=0.0,
out: Optional[torch.Tensor] = None) -> torch.Tensor:

if torch.rand((1)) < prob_inversion:
# Inversion of curve
out = maxv + minv - data
else:
# do nothing
out = data

return out
85 changes: 85 additions & 0 deletions rising/transforms/functional/painting.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
import torch

__all__ = ["local_pixel_shuffle", "random_inpainting", "random_outpainting"]


def local_pixel_shuffle(data: torch.Tensor, n: int = -1, block_size: tuple=(0,0,0), rel_block_size: float = 0.1) -> torch.Tensor:

batch_size, channels, img_rows, img_cols, img_deps = data.size()

if n < 0:
n = int(1000*channels) # changes ~ 12.5% of voxels
for b in range(batch_size):
for _ in range(n):
c = torch.randint(0,channels-1, (1,))

(block_size_x, block_size_y, block_size_z) = (torch.tensor([size]) for size in block_size)

if rel_block_size > 0:
block_size_x = torch.randint(2, int(img_rows * rel_block_size), (1,))
block_size_y = torch.randint(2, int(img_cols * rel_block_size), (1,))
block_size_z = torch.randint(2, int(img_deps * rel_block_size), (1,))

x = torch.randint(0, int(img_rows - block_size_x), (1,))
y = torch.randint(0, int(img_cols - block_size_y), (1,))
z = torch.randint(0, int(img_deps - block_size_z), (1,))

window = data[b, c, x:x + block_size_x,
y:y + block_size_y,
z:z + block_size_z,
]
idx = torch.randperm(window.numel())
window = window.view(-1)[idx].view(window.size())

data[b, c, x:x + block_size_x,
y:y + block_size_y,
z:z + block_size_z] = window

return data


def random_inpainting(data: torch.Tensor, n: int = 5, maxv: float = 1.0, minv: float = 0.0) -> torch.Tensor:

batch_size, channels, img_rows, img_cols, img_deps = data.size()

while n > 0 and torch.rand((1)) < 0.95:
for b in range(batch_size):
block_size_x = torch.randint(img_rows // 10, img_rows // 4, (1,))
block_size_y = torch.randint(img_rows // 10, img_rows // 4, (1,))
block_size_z = torch.randint(img_rows // 10, img_rows // 4, (1,))
x = torch.randint(3, int(img_rows - block_size_x - 3), (1,))
y = torch.randint(3, int(img_cols - block_size_y - 3), (1,))
z = torch.randint(3, int(img_deps - block_size_z - 3), (1,))

block = torch.rand((1, channels, block_size_x, block_size_y, block_size_z)) \
* (maxv-minv) + minv

data[b, :, x:x + block_size_x,
y:y + block_size_y,
z:z + block_size_z] = block

n = n - 1

return data


def random_outpainting(data: torch.Tensor, maxv: float = 1.0, minv: float = 0.0) -> torch.Tensor:

batch_size, channels, img_rows, img_cols, img_deps = data.size()

out = torch.rand(data.size()) * (maxv - minv) + minv

block_size_x = torch.randint(5*img_rows // 7, 6*img_rows // 7, (1,))
block_size_y = torch.randint(5*img_cols // 7, 6*img_cols // 7, (1,))
block_size_z = torch.randint(5*img_deps // 7, 6*img_deps // 7, (1,))
x = torch.randint(3, int(img_rows - block_size_x - 3), (1,))
y = torch.randint(3, int(img_cols - block_size_y - 3), (1,))
z = torch.randint(3, int(img_deps - block_size_z - 3), (1,))

out[:, :, x:x + block_size_x,
y:y + block_size_y,
z:z + block_size_z] = data[:, :, x:x + block_size_x,
y:y + block_size_y,
z:z + block_size_z]

return out
2 changes: 1 addition & 1 deletion rising/transforms/functional/spatial.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ def mirror(data: torch.Tensor, dims: Union[int, Sequence[int]]) -> torch.Tensor:
"""
if check_scalar(dims):
dims = (dims,)
# batch and channel dims
# batch and channel dims
dims = [d + 2 for d in dims]
return data.flip(dims)

Expand Down
31 changes: 29 additions & 2 deletions rising/transforms/intensity.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,13 +12,18 @@
gamma_correction,
add_value,
scale_by_value,
clamp)
clamp,
bezier_3rd_order,
random_inversion,
)

from rising.random import AbstractParameter

__all__ = ["Clamp", "NormRange", "NormMinMax",
"NormZeroMeanUnitStd", "NormMeanStd", "Noise",
"GaussianNoise", "ExponentialNoise", "GammaCorrection",
"RandomValuePerChannel", "RandomAddValue", "RandomScaleValue"]
"RandomValuePerChannel", "RandomAddValue", "RandomScaleValue",
"RandomBezierTransform", "InvertAmplitude"]


class Clamp(BaseTransform):
Expand Down Expand Up @@ -303,3 +308,25 @@ def __init__(self, random_sampler: AbstractParameter,
"""
super().__init__(augment_fn=scale_by_value, random_sampler=random_sampler,
per_channel=per_channel, keys=keys, grad=grad, **kwargs)


class RandomBezierTransform(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)


class InvertAmplitude(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):

super().__init__(augment_fn=random_inversion, prob_inversion=prob, maxv=maxv, minv=minv,
keys=keys, grad=False, **kwargs)

106 changes: 106 additions & 0 deletions rising/transforms/painting.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
import torch
from typing import Sequence

from rising.transforms.abstract import AbstractTransform, BaseTransform
from rising.transforms.functional.painting import (
local_pixel_shuffle, random_inpainting, random_outpainting
)


__all__ = ["RandomInpainting", "RandomOutpainting", "RandomInOrOutpainting", "LocalPixelShuffle"]


class LocalPixelShuffle(BaseTransform):
""" Shuffels Pixels locally in n patches,
as proposed in Models Genesis """

def __init__(self, n: int=-1,
keys: Sequence = ('data',), grad: bool = False, **kwargs):
"""
Args:
n: number of local patches to shuffle, default = 1000*channels
keys: the keys corresponding to the values to distort
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)


class RandomInpainting(BaseTransform):
""" In n local areas, the image is replaced by uniform noise in range (minv, maxv),
as proposed in Models Genesis """

def __init__(self, n: int = 5,
maxv: float=1.0, minv: float = 0.0,
keys: Sequence = ('data',), grad: bool = False, **kwargs):
"""
Args:
minv, maxv: range of uniform noise
n: number of local patches to randomize
keys: the keys corresponding to the values to distort
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)


class RandomOutpainting(AbstractTransform):
""" The border of the images will be replaced by uniform noise,
as proposed in Models Genesis """

def __init__(self, prob: float = 0.5, maxv: float=1.0, minv: float = 0.0,
keys: Sequence = ('data',), grad: bool = False, **kwargs):
"""
Args:
minv, maxv: range of uniform noise
prob: probability of outpainting. For prob<1.0, not all images will be augmented
keys: the keys corresponding to the values to distort
grad: enable gradient computation inside transformation
**kwargs: keyword arguments passed to augment_fn
"""
super().__init__(grad=grad, **kwargs)
self.prob = prob
self.maxv = maxv
self.minv = minv
self.keys = keys

def forward(self, **data) -> dict:
if torch.rand(1) < self.prob:
for key in self.keys:
data[key] = random_outpainting(data[key], maxv=self.maxv, minv=self.minv)
return data


class RandomInOrOutpainting(AbstractTransform):
"""Applies either random inpainting or random outpainting to the image,
as proposed in Models Genesis """

def __init__(self, prob: float = 0.5, n: int = 5,
maxv: float=1.0, minv: float = 0.0,
keys: Sequence = ('data',), grad: bool = False, **kwargs):
"""
Args:
minv, maxv: range of uniform noise
prob: probability of outpainting, probability of inpainting is 1-prob.
n: number of local patches to randomize in case of inpainting
keys: the keys corresponding to the values to distort
grad: enable gradient computation inside transformation
**kwargs: keyword arguments passed to augment_fn
"""
super().__init__(grad=grad, **kwargs)
self.prob = prob
self.maxv = maxv
self.minv = minv
self.keys = keys
self.n = n

def forward(self, **data) -> dict:
if torch.rand(1) < self.prob:
for key in self.keys:
data[key] = random_outpainting(data[key], maxv=self.maxv, minv=self.minv)
else:
for key in self.keys:
data[key] = random_inpainting(data[key], n=self.n, maxv=self.maxv, minv=self.minv)
return data
35 changes: 28 additions & 7 deletions rising/transforms/spatial.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,13 +16,12 @@
scheduler_type = Callable[[int], Union[int, Sequence[int]]]


class Mirror(BaseTransform):
class Mirror(AbstractTransform):
"""Random mirror transform"""

def __init__(self,
dims: Union[int, DiscreteParameter,
Sequence[Union[int, DiscreteParameter]]],
keys: Sequence[str] = ('data',), grad: bool = False, **kwargs):
def __init__(self, dims: Union[int, DiscreteParameter, Sequence[Union[int, DiscreteParameter]]],
keys: Sequence[str] = ('data',), prob: float = 0.5,
grad: bool = False, **kwargs):
"""
Args:
dims: axes which should be mirrored
Expand All @@ -39,8 +38,30 @@ def __init__(self,
>>> # volumetric data
>>> trafo = Mirror(DiscreteCombinationsParameter((0, 1, 2)))
"""
super().__init__(augment_fn=mirror, dims=dims, keys=keys, grad=grad,
property_names=('dims',), **kwargs)
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
"""
if torch.rand(1) < self.prob:
for key in self.keys:
data[key] = mirror(data[key], self.dims)
return data


class Rot90(AbstractTransform):
Expand Down
Loading