Skip to content

Commit 4674358

Browse files
committed
Add To transform
1 parent c95fa1b commit 4674358

3 files changed

Lines changed: 50 additions & 0 deletions

File tree

src/torchio/transforms/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,7 @@
5353
from .preprocessing import RescaleIntensity
5454
from .preprocessing import Resize
5555
from .preprocessing import SequentialLabels
56+
from .preprocessing import To
5657
from .preprocessing import ToCanonical
5758
from .preprocessing import ToOrientation
5859
from .preprocessing import ZNormalization
@@ -101,6 +102,7 @@
101102
'Crop',
102103
'Resize',
103104
'Resample',
105+
'To',
104106
'ToCanonical',
105107
'ToOrientation',
106108
'ZNormalization',

src/torchio/transforms/preprocessing/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from .intensity.histogram_standardization import HistogramStandardization
33
from .intensity.mask import Mask
44
from .intensity.rescale import RescaleIntensity
5+
from .intensity.to import To
56
from .intensity.z_normalization import ZNormalization
67
from .label.contour import Contour
78
from .label.keep_largest_component import KeepLargestComponent
@@ -31,6 +32,7 @@
3132
'EnsureShapeMultiple',
3233
'Mask',
3334
'RescaleIntensity',
35+
'To',
3436
'Clamp',
3537
'ZNormalization',
3638
'HistogramStandardization',
Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,46 @@
1+
from __future__ import annotations
2+
3+
from typing import Any
4+
5+
import torch
6+
7+
from ....data.image import ScalarImage
8+
from ....data.subject import Subject
9+
from ...intensity_transform import IntensityTransform
10+
11+
12+
class To(IntensityTransform):
13+
"""Convert the image tensor data type and/or device.
14+
15+
Args:
16+
target: First argument to :func:`torch.Tensor.to`.
17+
to_kwargs: Additional keyword arguments to pass to :func:`torch.Tensor.to`.
18+
19+
Example:
20+
>>> import torchio as tio
21+
>>> ct = tio.datasets.Slicer('CTChest').CT_chest
22+
>>> clamp = tio.Clamp(out_min=-1000, out_max=1000)
23+
>>> ct_clamped = clamp(ct)
24+
>>> rescale = tio.RescaleIntensity(in_min_max=(-1000, 1000), out_min_max=(0, 255))
25+
>>> ct_rescaled = rescale(ct_clamped)
26+
>>> to_uint8 = tio.To(torch.uint8)
27+
>>> ct_uint8 = to_uint8(ct_rescaled)
28+
29+
"""
30+
31+
def __init__(
32+
self,
33+
target: str | torch.dtype | torch.device,
34+
to_kwargs: dict[str, Any],
35+
**kwargs,
36+
):
37+
super().__init__(**kwargs)
38+
self.target = target
39+
self.to_kwargs = to_kwargs
40+
self.args_names = ['target', 'to_kwargs']
41+
42+
def apply_transform(self, subject: Subject) -> Subject:
43+
for image in self.get_images(subject):
44+
assert isinstance(image, ScalarImage)
45+
image.set_data(image.data.to(self.target, **self.to_kwargs))
46+
return subject

0 commit comments

Comments
 (0)