mrid-python 0.1.3__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
mrid/__init__.py ADDED
@@ -0,0 +1,5 @@
1
+ from . import utils
2
+ from .atlas import get_mni152, get_sri24
3
+ from .preprocessing import *
4
+ from .loading import *
5
+ from .study import Study
@@ -0,0 +1,80 @@
1
+ # https://zenodo.org/api/records/15470657/files-archive
2
+
3
+ import os
4
+ import shutil
5
+ import tempfile
6
+ from pathlib import Path
7
+ from typing import Literal
8
+
9
+ __all__ = [
10
+ "get_mni152",
11
+ ]
12
+
13
+ _ROOT = Path(os.path.dirname(__file__))
14
+
15
+ _URLS = {
16
+ "2006 T1w symmetric": (
17
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t1_06_sym.nii.gz?download=1",
18
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t1_06_sym_bet.nii.gz?download=1",
19
+ ),
20
+ "2009a T1w symmetric": (
21
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t1_09a_sym.nii.gz?download=1",
22
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t1_09a_sym_bet.nii.gz?download=1",
23
+ ),
24
+ "2009a T2w symmetric": (
25
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t2_09a_sym.nii.gz?download=1",
26
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t2_09a_sym_bet.nii.gz?download=1",
27
+ ),
28
+ "2009a T1w asymmetric": (
29
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t1_09a_asym.nii.gz?download=1",
30
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t1_09a_asym_bet.nii.gz?download=1",
31
+ ),
32
+ "2009a T2w asymmetric": (
33
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t2_09a_asym.nii.gz?download=1",
34
+ "https://zenodo.org/records/15470657/files/icbm_mni152_t2_09a_asym_bet.nii.gz?download=1",
35
+ ),
36
+ }
37
+
38
+ _mni152_url = "https://zenodo.org/records/15470657/files/icbm_mni152_t1_06_sym_bet.nii.gz?download=1"
39
+
40
+ def _download_template(type: str, bet:bool):
41
+ filename = f"{type} {bool(bet)}.nii.gz"
42
+ if filename in os.listdir(_ROOT):
43
+ raise RuntimeError(f"Template {type} is already downloaded")
44
+
45
+ import requests
46
+
47
+ response = requests.get(_URLS[type][bet], stream=True, timeout=30)
48
+ response.raise_for_status()
49
+
50
+ with open(_ROOT / f"{filename}", 'wb') as file:
51
+ shutil.copyfileobj(response.raw, file) # type:ignore
52
+
53
+ def get_mni152(
54
+ type: Literal[
55
+ "2006 T1w symmetric",
56
+ "2009a T1w symmetric",
57
+ "2009a T2w symmetric",
58
+ "2009a T1w asymmetric",
59
+ "2009a T2w asymmetric",
60
+ ],
61
+ skullstripped: bool = False,
62
+ ):
63
+ """Returns path to .nii.gz file of specified MNI-152 template.
64
+
65
+ Descriptions of templates are available here https://zenodo.org/records/15470657
66
+ """
67
+ filename = f"{type} {bool(skullstripped)}.nii.gz"
68
+
69
+ if filename in os.listdir(_ROOT):
70
+ return str(_ROOT / filename)
71
+
72
+ print(f"{filename} will be downloaded from https://zenodo.org/records/15470657, this may take a few minutes.")
73
+ _download_template(type, skullstripped)
74
+
75
+ if filename not in os.listdir(_ROOT):
76
+ raise RuntimeError(
77
+ f"Failed to download {filename}; try downloading it manually from https://zenodo.org/records/15470657"
78
+ )
79
+
80
+ return str(_ROOT / filename)
@@ -0,0 +1,77 @@
1
+ """This subpackage allows one to download SRI-24 brain atlas files from https://www.nitrc.org/projects/sri24/.
2
+
3
+ The SRI24 atlas is licensed under the terms of the
4
+
5
+ Creative Commons Attribution-ShareAlike 3.0 Unported (CC BY-SA 3.0)
6
+
7
+ license (https://creativecommons.org/licenses/by-...).
8
+
9
+ In publications using the SRI24 atlas, please cite the following paper:
10
+
11
+ T. Rohlfing, N.M. Zahr, E.V. Sullivan, A. Pfefferbaum, "The SRI24
12
+ Multichannel Atlas of Normal Adult Human Brain Structure," Human
13
+ Brain Mapping, vol. 31, no. 5, pp. 798-819, 2010.
14
+
15
+ http://dx.doi.org/10.1002/hbm.20906
16
+
17
+ """
18
+ import os
19
+ import shutil
20
+ import tempfile
21
+ from pathlib import Path
22
+ from typing import Literal
23
+
24
+ __all__ = [
25
+ "get_sri24",
26
+ ]
27
+
28
+ _ROOT = Path(os.path.dirname(__file__))
29
+
30
+ _sri24_url = "https://www.nitrc.org/frs/download.php/4841/sri24_spm8.zip//?i_agree=1&download_now=1"
31
+
32
+ def _download_sri24() -> None:
33
+ import requests
34
+
35
+ response = requests.get(_sri24_url, stream=True, timeout=30)
36
+ response.raise_for_status()
37
+
38
+ with tempfile.TemporaryDirectory() as tmpdir:
39
+ tmpdir = Path(tmpdir)
40
+ with open(tmpdir / "sri24_spm8.zip", 'wb') as file:
41
+ shutil.copyfileobj(response.raw, file) # type:ignore
42
+
43
+ shutil.unpack_archive(tmpdir / "sri24_spm8.zip", tmpdir / "sri24_spm8")
44
+
45
+ for file in os.listdir(tmpdir / "sri24_spm8" / "templates"):
46
+ shutil.copyfile(tmpdir / "sri24_spm8" / "templates" / file, _ROOT / file)
47
+
48
+
49
+ def get_sri24(type: Literal["EPI", "EPI_brain", "PD", "PD_brain", "T1", "T1_brain", "T2", "T2_brain"]) -> str:
50
+ """Returns path to .nii file of specified SRI-24 template. Templates are downloaded if they haven't been downloaded already.
51
+
52
+ The following templates are available:
53
+ - `"T1"`: post-contrast T1-weighted MRI with skull;
54
+ - `"T1_brain"`: post-contrast T1-weighted MRI without skull;
55
+ - `"T2"`: T2-weighted MRI with skull;
56
+ - `"T2_brain"`: T2-weighted MRI without skull;
57
+ - `"EPI"`: echo-planar imaging MRI with skull;
58
+ - `"EPI_brain"`: echo-planar imaging MRU without skull;
59
+ - `"PD"`: proton density weighted spin-echo imaging MRI with skull;
60
+ - `"PD_brain"`: proton density weighted spin-echo imaging MRI without skull;
61
+
62
+ """
63
+ filename = f"{type}.nii"
64
+ if filename in os.listdir(_ROOT):
65
+ return str(_ROOT / filename)
66
+
67
+ print("SRI24 will be downloaded from https://www.nitrc.org/projects/sri24/, this may take a few minutes.")
68
+ _download_sri24()
69
+
70
+ if filename not in os.listdir(_ROOT):
71
+ raise RuntimeError(
72
+ f"Failed to download {filename}; try downloading it manually from https://www.nitrc.org/projects/sri24/, "
73
+ "then unpack the zip file, open it, open `templates` folder, you will see files such as `EPI.nii`. "
74
+ f"Copy all of those files to {_ROOT}."
75
+ )
76
+
77
+ return str(_ROOT / filename)
mrid/atlas/__init__.py ADDED
@@ -0,0 +1,7 @@
1
+ from . MNI152 import get_mni152
2
+ from .SRI24 import get_sri24
3
+
4
+ __all__ = [
5
+ "get_mni152",
6
+ "get_sri24",
7
+ ]
@@ -0,0 +1 @@
1
+ from .convert import tonumpy, tositk, totensor, ImageLike
@@ -0,0 +1,68 @@
1
+ import importlib.util
2
+ import os
3
+ from typing import TYPE_CHECKING, TypeAlias
4
+
5
+ import numpy as np
6
+ import SimpleITK as sitk
7
+
8
+ from ..utils.torch_utils import TORCH_INSTALLED
9
+
10
+ if TYPE_CHECKING:
11
+ import torch
12
+
13
+ PREFER_DCM2NIIX = False
14
+ ImageLike: TypeAlias = "np.ndarray | sitk.Image | torch.Tensor | str | os.PathLike"
15
+
16
+ def read_dicoms(dir: str | os.PathLike) -> sitk.Image:
17
+ """reads a directory of DICOM files and returns a ``sitk.Image``"""
18
+ # load with dcm2niix
19
+ if PREFER_DCM2NIIX and importlib.util.find_spec("dcm2niix") is not None:
20
+ from ..utils.dcm2niix import dcm2sitk
21
+ return dcm2sitk(dir)
22
+
23
+ # load with SimpleITK
24
+ reader = sitk.ImageSeriesReader()
25
+ dicom_names = reader.GetGDCMSeriesFileNames(str(dir))
26
+
27
+ if not dicom_names:
28
+ raise FileNotFoundError(f"No DICOM series found in directory: {dir}")
29
+
30
+ reader.SetFileNames(dicom_names)
31
+ return reader.Execute()
32
+
33
+ def _read_sitk(path: str | os.PathLike) -> sitk.Image:
34
+ if os.path.isfile(path): return sitk.ReadImage(str(path))
35
+ if os.path.isdir(path): return read_dicoms(str(path))
36
+ raise FileNotFoundError(f"{path} doesn't exist")
37
+
38
+ def tositk(x: ImageLike) -> sitk.Image:
39
+ """Load an image into an ``sitk.Image`` object.
40
+ ``x`` can be a numpy array, a ``sitk.Image``, a ``torch.Tensor`` or a string (path to an image file)."""
41
+ if isinstance(x, np.ndarray): return sitk.GetImageFromArray(x)
42
+ if isinstance(x, sitk.Image): return x
43
+ if isinstance(x, (str, os.PathLike)): return _read_sitk(x)
44
+ if TORCH_INSTALLED:
45
+ import torch
46
+ if isinstance(x, torch.Tensor): return sitk.GetImageFromArray(x.numpy())
47
+ raise TypeError(f"Unsupported type {type(x)}")
48
+
49
+ def tonumpy(x: ImageLike) -> np.ndarray:
50
+ """Load an image into a numpy.ndarray.
51
+ ``x`` can be a numpy array, a ``sitk.Image``, a ``torch.Tensor`` or a string (path to an image file)."""
52
+ if isinstance(x, np.ndarray): return x
53
+ if isinstance(x, sitk.Image): return sitk.GetArrayFromImage(x)
54
+ if isinstance(x, (str, os.PathLike)): return sitk.GetArrayFromImage(_read_sitk(x))
55
+ if TORCH_INSTALLED:
56
+ import torch
57
+ if isinstance(x, torch.Tensor): return x.numpy()
58
+ raise TypeError(f"Unsupported type {type(x)}")
59
+
60
+ def totensor(x: ImageLike) -> "torch.Tensor":
61
+ """Load an image into a torch.Tensor.
62
+ ``x`` can be a numpy array, a ``sitk.Image``, a ``torch.Tensor`` or a string (path to an image file)."""
63
+ import torch
64
+ if isinstance(x, np.ndarray): return torch.from_numpy(x)
65
+ if isinstance(x, sitk.Image): return torch.from_numpy(sitk.GetArrayFromImage(x))
66
+ if isinstance(x, (str, os.PathLike)): return torch.from_numpy(sitk.GetArrayFromImage(_read_sitk(x)))
67
+ if isinstance(x, torch.Tensor): return x
68
+ raise TypeError(f"Unsupported type {type(x)}")
@@ -0,0 +1,12 @@
1
+ from .bias_field_correction import n4_bias_field_correction
2
+ from .cropping import crop_bg, crop_bg_D
3
+ from .registration import resample_to, register, register_D, register_each, resize, downsample, Registration
4
+ from .skullstripping import skullstrip, skullstrip_D, run_hd_bet, predict_brain_mask
5
+
6
+ __all__ = [
7
+ "n4_bias_field_correction",
8
+ "crop_bg", "crop_bg_D",
9
+ "resample_to", "register", "register_D", "register_each", "resize", "downsample",
10
+ "skullstrip", "skullstrip_D", "run_hd_bet", "predict_brain_mask"
11
+
12
+ ]
@@ -0,0 +1,28 @@
1
+ import SimpleITK as sitk
2
+ from ..loading.convert import tositk, ImageLike
3
+
4
+ def n4_bias_field_correction(image: ImageLike, shrink: int = 4) -> sitk.Image:
5
+ """Perform N4 Bias Field Correction to correct low frequency intensity non-uniformity present in MRI image.
6
+
7
+ Args:
8
+ image (ImageLike): Input MRI image to be corrected. Can be any format supported by tositk conversion.
9
+ shrink (int, optional): Shrink factor for reducing image size before correction to speed up computation.
10
+ Default is 4. If set to 1 or less, no shrinking is performed.
11
+
12
+ """
13
+ image = tositk(image)
14
+
15
+ norm_image = sitk.RescaleIntensity(image, 0, 255)
16
+ mask = sitk.OtsuThreshold(norm_image, 0, 1)
17
+
18
+ if shrink > 1:
19
+ reduced = sitk.Shrink(image, [shrink] * image.GetDimension())
20
+ mask = sitk.Shrink(mask, [shrink] * mask.GetDimension())
21
+
22
+ else: reduced = image
23
+
24
+ corrector = sitk.N4BiasFieldCorrectionImageFilter()
25
+ corrector.Execute(reduced, mask)
26
+ log_bias_field = corrector.GetLogBiasFieldAsImage(image)
27
+
28
+ return image / sitk.Cast(sitk.Exp(log_bias_field), image.GetPixelID())
@@ -0,0 +1,36 @@
1
+ from collections.abc import Mapping
2
+ from typing import Any
3
+ import SimpleITK as sitk
4
+
5
+ from ..loading.convert import tositk, ImageLike
6
+
7
+ def _get_bbox(image: sitk.Image):
8
+ rescaled = sitk.RescaleIntensity(image, 0, 255)
9
+ filt = sitk.LabelShapeStatisticsImageFilter()
10
+ filt.Execute(sitk.OtsuThreshold(rescaled, 0, 255))
11
+ return filt.GetBoundingBox(255)
12
+
13
+
14
+ def crop_bg(image: ImageLike) -> sitk.Image:
15
+ """Crops black background of a single 3D image via Otsu's thresholding.
16
+
17
+ Args:
18
+ image (ImageLike): Input 3D image to be cropped. Can be any format supported by tositk conversion.
19
+
20
+ Returns:
21
+ sitk.Image: Cropped image with black background removed, maintaining the same pixel type as input.
22
+ """
23
+ image = tositk(image)
24
+ bbox = _get_bbox(image)
25
+ return sitk.RegionOfInterest( image, bbox[int(len(bbox) / 2) :], bbox[0 : int(len(bbox) / 2)],)
26
+
27
+ def crop_bg_D(images: Mapping[str, ImageLike], key: str) -> dict[str, sitk.Image]:
28
+ """Finds the bounding box of ``images[key]`` and crops all images in ``images`` to that bounding box."""
29
+ images = {k: tositk(v) for k,v in images.items()}
30
+ reference = images[key]
31
+
32
+ bbox = _get_bbox(reference)
33
+
34
+ ret = {k: sitk.RegionOfInterest(v, bbox[int(len(bbox) / 2) :], bbox[0 : int(len(bbox) / 2)]) for k,v in images.items()}
35
+ return ret
36
+
@@ -0,0 +1,251 @@
1
+ from collections.abc import Mapping, Sequence
2
+ from typing import TYPE_CHECKING, Any
3
+
4
+ import numpy as np
5
+ import SimpleITK as sitk
6
+
7
+ from ..loading.convert import tositk, ImageLike
8
+
9
+ def resample_to(input: ImageLike, to: ImageLike, interpolation=sitk.sitkNearestNeighbor) -> sitk.Image:
10
+ """Resample ``input`` to ``reference``.
11
+
12
+ Resampling uses spatial information embedded in the sitk.Image - size, origin, spacing and direction.
13
+
14
+ Note that this information is only available when certain imaging formats are loaded, such as DICOM and NIfTI.
15
+
16
+ ``input`` is transformed in such a way that those attributes will match ``reference``.
17
+ """
18
+ return sitk.Resample(tositk(input), tositk(to), sitk.Transform(), interpolation)
19
+
20
+
21
+ def _default_pmap():
22
+ """Default parameter maps for registration"""
23
+ euler = sitk.GetDefaultParameterMap('translation')
24
+ euler['Transform'] = ['EulerTransform']
25
+ pmap = sitk.VectorOfParameterMap()
26
+ pmap.append(sitk.GetDefaultParameterMap("translation"))
27
+ pmap.append(euler)
28
+ pmap.append(sitk.GetDefaultParameterMap("rigid"))
29
+ pmap.append(sitk.GetDefaultParameterMap("affine"))
30
+ return pmap
31
+
32
+ class Registration:
33
+ """Class for image registration.
34
+
35
+ Args:
36
+ pmap (Any, optional): parameter map, if None, uses default parameter map. Defaults to None.
37
+ log_to_console (bool, optional): if False, disables SimpleElastix logging a lot of stuff to your console. Defaults to False.
38
+ """
39
+ def __init__(self, pmap: Any = None, log_to_console=False):
40
+ if pmap is None: pmap = _default_pmap()
41
+ self.pmap: sitk.VectorOfParameterMap = pmap
42
+ self.log_to_console = log_to_console
43
+
44
+ # create elastix filter
45
+ self.elastix = sitk.ElastixImageFilter()
46
+ if log_to_console: self.elastix.LogToConsoleOn()
47
+ else: self.elastix.LogToConsoleOff()
48
+
49
+ self.elastix.SetParameterMap(self.pmap)
50
+
51
+ self._moving = None
52
+ self._transformed = None
53
+ self.inverse: "Registration | None" = None
54
+
55
+ def find_transform(self, input: ImageLike, to: ImageLike) -> sitk.Image:
56
+ """Find a transform that transforms ``input`` to ``to`` and save it to this ``Registration`` object.
57
+ Returns ``input`` registered to ``to``.
58
+
59
+ Args:
60
+ input (ImageLike): Moving image.
61
+ to (ImageLike): Fixed image.
62
+ """
63
+ if self._transformed is not None:
64
+ raise RuntimeError("`find_transform` has already been called on this Registration object.")
65
+ self._moving = tositk(input)
66
+ to = tositk(to)
67
+
68
+ self.elastix.SetFixedImage(to)
69
+ self.elastix.SetMovingImage(self._moving)
70
+ self.elastix.Execute()
71
+
72
+ self._transformed = self.elastix.GetResultImage()
73
+ return self.elastix.GetResultImage() # return copy
74
+
75
+ def apply_transform(self, input: ImageLike, use_nearest_interpolation: bool = False) -> sitk.Image:
76
+ """Applies transform stored in this ``Registration`` object to ``input``.
77
+
78
+ You have to use ``find_transform`` method first to find the transform.
79
+
80
+ Args:
81
+ input (ImageLike): Moving image to apply transform to.
82
+ use_nearest_interpolation (bool, optional):
83
+ whether to use nearest interpolation, enable when transforming segmentations. Defaults to False.
84
+
85
+ Returns:
86
+ sitk.Image: transformed ``input``.
87
+ """
88
+ if self._transformed is None:
89
+ raise RuntimeError("First find transform parameters using `find_transform` method.")
90
+
91
+ transform = sitk.TransformixImageFilter()
92
+ tmap = self.elastix.GetTransformParameterMap()
93
+ if use_nearest_interpolation:
94
+ for t in tmap:
95
+ t["ResampleInterpolator"] = ["FinalNearestNeighborInterpolator"]
96
+
97
+ transform.SetTransformParameterMap(tmap)
98
+ transform.SetMovingImage(input)
99
+ if not self.log_to_console: transform.LogToConsoleOff()
100
+
101
+ return transform.Execute()
102
+
103
+ def apply_inverse_transform(self, input: ImageLike, use_nearest_interpolation: bool = False) -> sitk.Image:
104
+ """Applies inverse of the transform stored in this ``Registration`` object to ``input``.
105
+
106
+ This is done by finding another transform that undoes the current one. Note that this may not be as robust as
107
+ using other tools like freesurfer (because in SimpleElastix transform inverse is not implemented, and
108
+ "DisplacementMagnitudePenalty" metric is not included in python build).
109
+
110
+ Args:
111
+ input (ImageLike): input image to apply inverse transform to.
112
+ use_nearest_interpolation (bool, optional):
113
+ whether to use nearest interpolation, enable when transforming segmentations. Defaults to False.
114
+ """
115
+ if self.inverse is None:
116
+ if (self._transformed is None) or (self._moving is None):
117
+ raise RuntimeError("First find transform parameters using `find_transform` method.")
118
+ inverse_pmap = self.elastix.GetParameterMap() # this returns a copy
119
+ # for p in inverse_pmap:
120
+ # p["Metric"] = "MeanSquaredDifference" # not implemented
121
+ self.inverse = Registration(pmap=inverse_pmap, log_to_console=self.log_to_console)
122
+ self.inverse.find_transform(
123
+ input=self._transformed,
124
+ to=self._moving
125
+ )
126
+
127
+ return self.inverse.apply_transform(input, use_nearest_interpolation=use_nearest_interpolation)
128
+
129
+
130
+ def register(input: ImageLike, to: ImageLike, pmap: Any = None, log_to_console=False):
131
+ """Register ``input`` to ``reference``. Returns ``input`` with the same shape and spatial position as ``reference``
132
+
133
+ Registering means finding a transform which alligns ``input`` to match with ``reference``,
134
+ it will have the same size, orientation, etc. By default this used affine transform.
135
+
136
+ This uses ``SimpleITK-SimpleElastix`` which is very robust.
137
+ Note that if you don't have it installed, you need to uninstall normal SimpleITK
138
+ and install https://pypi.org/project/SimpleITK-SimpleElastix/, don't worry, it's
139
+ the same as SimpleITK but it additionally includes SimpleElastix.
140
+ """
141
+ reg = Registration(pmap=pmap, log_to_console=log_to_console)
142
+ return reg.find_transform(input=input, to=to)
143
+
144
+
145
+ def register_D(
146
+ images: Mapping[str, ImageLike],
147
+ key: str,
148
+ to: ImageLike,
149
+ pmap: Any = None,
150
+ log_to_console=False,
151
+ ) -> dict[str, sitk.Image]:
152
+ """Register ``images[key]`` to ``reference``, then use that transformation
153
+ to transform other values in ``images`` that are assumed to be aligned with ``images[key]`` (e.g. segmentation).
154
+
155
+ Make sure segmentation with hard edges is under a key that starts with ``"seg"``,
156
+ it will use nearest neighbour interpolation, otherwise it will mess up the edges.
157
+ """
158
+ reg = Registration(pmap=pmap, log_to_console=log_to_console)
159
+ registered = {key: reg.find_transform(images[key], to)}
160
+
161
+ # process segs last because it sets resample interpolator to nearest
162
+ for k,v in sorted(list(images.items()), key = lambda x: 1 if x[0].startswith('seg') else 0):
163
+ if k != key:
164
+ use_nearest_interpolation = k.startswith('seg')
165
+ registered[k] = reg.apply_transform(v, use_nearest_interpolation=use_nearest_interpolation)
166
+
167
+ return registered
168
+
169
+ def register_each(
170
+ images: Mapping[str, ImageLike],
171
+ key: str,
172
+ to: ImageLike,
173
+ pmap: Any = None,
174
+ log_to_console=False,
175
+ ) -> dict[str, sitk.Image]:
176
+ """Register ``images[key]`` to ``reference``, then register all other values in ``images`` to registered ``images[key]``.
177
+
178
+ Use this when you have multiple modalities that do not align."""
179
+ images = {k: tositk(v) for k,v in images.items()}
180
+ to = tositk(to)
181
+
182
+ input = images[key]
183
+ input_reg = register(input=input, to=to, pmap=pmap, log_to_console=log_to_console)
184
+
185
+ registered = {key: input_reg}
186
+ for k,v in images.items():
187
+ if k != key:
188
+ registered[k] = register(input=v, to=input_reg, pmap=pmap, log_to_console=log_to_console)
189
+
190
+ return registered
191
+
192
+ def resize(img: ImageLike, new_size: Sequence[int], interpolator=sitk.sitkLinear) -> sitk.Image:
193
+ """Resize ``sitk.Image`` to ``new_size``. Retains correct spatial information.
194
+ source: https://gist.github.com/lixinqi98/1bbd3596492f20b776fed2778f7cd48c"""
195
+ img = tositk(img)
196
+ new_size = list(reversed(new_size))
197
+
198
+ # img = sitk.ReadImage(img)
199
+ dimension = img.GetDimension()
200
+
201
+ # Physical image size corresponds to the largest physical size in the training set, or any other arbitrary size.
202
+ reference_physical_size = np.zeros(dimension)
203
+
204
+ reference_physical_size[:] = [(sz - 1) * spc if sz * spc > mx else mx for sz, spc, mx in
205
+ zip(img.GetSize(), img.GetSpacing(), reference_physical_size)]
206
+
207
+ # Create the reference image with a zero origin, identity direction cosine matrix and dimension
208
+ reference_origin = np.zeros(dimension)
209
+ reference_direction = np.identity(dimension).flatten()
210
+ reference_size = new_size
211
+ reference_spacing = [phys_sz / (sz - 1) for sz, phys_sz in zip(reference_size, reference_physical_size)]
212
+
213
+ reference_image = sitk.Image(reference_size, img.GetPixelIDValue())
214
+ reference_image.SetOrigin(reference_origin)
215
+ reference_image.SetSpacing(reference_spacing)
216
+ reference_image.SetDirection(reference_direction)
217
+
218
+ # Always use the TransformContinuousIndexToPhysicalPoint to compute an indexed point's physical coordinates as
219
+ # this takes into account size, spacing and direction cosines. For the vast majority of images the direction
220
+ # cosines are the identity matrix, but when this isn't the case simply multiplying the central index by the
221
+ # spacing will not yield the correct coordinates resulting in a long debugging session.
222
+ reference_center = np.array(
223
+ reference_image.TransformContinuousIndexToPhysicalPoint(np.array(reference_image.GetSize()) / 2.0))
224
+
225
+ # Transform which maps from the reference_image to the current img with the translation mapping the image
226
+ # origins to each other.
227
+ transform = sitk.AffineTransform(dimension)
228
+ transform.SetMatrix(img.GetDirection())
229
+ transform.SetTranslation(np.array(img.GetOrigin()) - reference_origin)
230
+ # Modify the transformation to align the centers of the original and reference image instead of their origins.
231
+ centering_transform = sitk.TranslationTransform(dimension)
232
+ img_center = np.array(img.TransformContinuousIndexToPhysicalPoint(np.array(img.GetSize()) / 2.0))
233
+ centering_transform.SetOffset(np.array(transform.GetInverse().TransformPoint(img_center) - reference_center))
234
+
235
+ # centered_transform = sitk.Transform(transform)
236
+ # centered_transform.AddTransform(centering_transform)
237
+
238
+ centered_transform = sitk.CompositeTransform([transform, centering_transform])
239
+
240
+ # Using the linear interpolator as these are intensity images, if there is a need to resample a ground truth
241
+ # segmentation then the segmentation image should be resampled using the NearestNeighbor interpolator so that
242
+ # no new labels are introduced.
243
+
244
+ return sitk.Resample(img, reference_image, centered_transform, interpolator, 0.0)
245
+
246
+ def downsample(image:ImageLike, factor:float, dims: Sequence[int] | None, interpolator=sitk.sitkLinear) -> sitk.Image:
247
+ """factor = 2 for 2x downsampling"""
248
+ image = tositk(image)
249
+ size = sitk.GetArrayFromImage(image).shape
250
+ size = [round(s/factor) if (dims is None or i in dims) else s for i,s in enumerate(size)]
251
+ return resize(image, size, interpolator=interpolator)