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 +5 -0
- mrid/atlas/MNI152/__init__.py +80 -0
- mrid/atlas/SRI24/__init__.py +77 -0
- mrid/atlas/__init__.py +7 -0
- mrid/loading/__init__.py +1 -0
- mrid/loading/convert.py +68 -0
- mrid/preprocessing/__init__.py +12 -0
- mrid/preprocessing/bias_field_correction.py +28 -0
- mrid/preprocessing/cropping.py +36 -0
- mrid/preprocessing/registration.py +251 -0
- mrid/preprocessing/skullstripping.py +185 -0
- mrid/study.py +442 -0
- mrid/utils/__init__.py +3 -0
- mrid/utils/dcm2niix.py +90 -0
- mrid/utils/dicom_uid_fixer.py +86 -0
- mrid/utils/plotting.py +102 -0
- mrid/utils/python_utils.py +48 -0
- mrid/utils/stl_utils.py +70 -0
- mrid/utils/torch_utils.py +16 -0
- mrid_python-0.1.3.dist-info/METADATA +140 -0
- mrid_python-0.1.3.dist-info/RECORD +27 -0
- mrid_python-0.1.3.dist-info/WHEEL +5 -0
- mrid_python-0.1.3.dist-info/top_level.txt +2 -0
- tests/test_loading.py +82 -0
- tests/test_preprocessing.py +43 -0
- tests/test_study.py +136 -0
- tests/test_utils.py +16 -0
mrid/__init__.py
ADDED
|
@@ -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
mrid/loading/__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
1
|
+
from .convert import tonumpy, tositk, totensor, ImageLike
|
mrid/loading/convert.py
ADDED
|
@@ -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)
|