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
|
@@ -0,0 +1,185 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import subprocess
|
|
3
|
+
import tempfile
|
|
4
|
+
from collections.abc import Mapping
|
|
5
|
+
from typing import Literal
|
|
6
|
+
|
|
7
|
+
import SimpleITK as sitk
|
|
8
|
+
|
|
9
|
+
from ..loading.convert import ImageLike, tositk
|
|
10
|
+
from ..utils.torch_utils import CUDA_IF_AVAILABLE
|
|
11
|
+
from .registration import register, register_D
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def run_hd_bet(
|
|
15
|
+
input: str | os.PathLike,
|
|
16
|
+
output: str | os.PathLike,
|
|
17
|
+
device: Literal['cpu', 'cuda', 'mps'] = CUDA_IF_AVAILABLE,
|
|
18
|
+
disable_tta: bool = False,
|
|
19
|
+
save_bet_mask: bool = True,
|
|
20
|
+
no_bet_image: bool = False,
|
|
21
|
+
verbose: bool = False,
|
|
22
|
+
) -> None:
|
|
23
|
+
"""Loads ``input`` file (ideally T1-w, postcontrast T1-w, T2-w and FLAIR sequences in MNI152 space)
|
|
24
|
+
and runs HD-BET to generate brain mask.
|
|
25
|
+
|
|
26
|
+
This is a simple wrapper around HD-BET command line interface using subprocess.
|
|
27
|
+
|
|
28
|
+
The documentation for hd-bet (copied from ``hd_bet -h``):
|
|
29
|
+
```bash
|
|
30
|
+
hd_bet -h
|
|
31
|
+
|
|
32
|
+
-i INPUT, --input INPUT
|
|
33
|
+
input. Can be either a single file name or an input folder. If file: must be nifti (.nii.gz) and can only be 3D. No support for 4d images, use fslsplit to split 4d sequences
|
|
34
|
+
into 3d images. If folder: all files ending with .nii.gz within that folder will be brain extracted.
|
|
35
|
+
-o OUTPUT, --output OUTPUT
|
|
36
|
+
output. Can be either a filename or a folder. If it does not exist, the folder will be created
|
|
37
|
+
-device DEVICE used to set on which device the prediction will run. Can be 'cuda' (=GPU), 'cpu' or 'mps'. Default: cuda
|
|
38
|
+
--disable_tta Set this flag to disable test time augmentation. This will make prediction faster at a slight decrease in prediction quality. Recommended for device cpu
|
|
39
|
+
--save_bet_mask Set this flag to keep the bet masks. Otherwise they will be removed once HD_BET is done
|
|
40
|
+
--no_bet_image Set this flag to disable generating the skull stripped/brain extracted image. Only makes sense if you also set --save_bet_mask
|
|
41
|
+
--verbose Talk to me.
|
|
42
|
+
```
|
|
43
|
+
"""
|
|
44
|
+
|
|
45
|
+
command = [
|
|
46
|
+
"hd-bet",
|
|
47
|
+
"-i", os.path.normpath(input),
|
|
48
|
+
"-o", os.path.normpath(output),
|
|
49
|
+
"-device", device,
|
|
50
|
+
]
|
|
51
|
+
if disable_tta: command.append("--disable_tta")
|
|
52
|
+
if save_bet_mask: command.append("--save_bet_mask")
|
|
53
|
+
if no_bet_image: command.append("--no_bet_image")
|
|
54
|
+
if verbose: command.append("--verbose")
|
|
55
|
+
|
|
56
|
+
# run dcm2niix
|
|
57
|
+
subprocess.run(command, check=True)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def predict_brain_mask(
|
|
61
|
+
input: ImageLike,
|
|
62
|
+
register_to_mni152: Literal["T1", "T2"] | None = None,
|
|
63
|
+
device: Literal["cpu", "cuda", "mps"] = CUDA_IF_AVAILABLE,
|
|
64
|
+
disable_tta: bool = False,
|
|
65
|
+
verbose: bool = False,
|
|
66
|
+
) -> sitk.Image:
|
|
67
|
+
"""Returns brain mask of ``input`` predicted by HD-BET.
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
input (ImageLike): input to skullstrip. Recommended T1-w, postcontrast T1-w, T2-w or FLAIR sequence in MNI152 space.
|
|
71
|
+
register_to_mni152 (str | None, optional):
|
|
72
|
+
Should be ``"T1"``, ``"T2"`` or ``None``.
|
|
73
|
+
if specified, ``input`` will be registered to specified MNI152 template,
|
|
74
|
+
and brain mask registered back to original ``input``.
|
|
75
|
+
Note that HD-BET expects images to be in MNI152 space. Defaults to None.
|
|
76
|
+
device (str, optional):
|
|
77
|
+
used to set on which device the prediction will run. Can be 'cuda' (=GPU), 'cpu' or 'mps'.
|
|
78
|
+
Defaults to CUDA_IF_AVAILABLE.
|
|
79
|
+
disable_tta (bool, optional):
|
|
80
|
+
Set this flag to disable test time augmentation. This will make prediction faster
|
|
81
|
+
at a slight decrease in prediction quality. Recommended for device cpu. Defaults to False.
|
|
82
|
+
verbose (bool, optional): purpose currently unknown. Defaults to False.
|
|
83
|
+
"""
|
|
84
|
+
input = tositk(input)
|
|
85
|
+
|
|
86
|
+
# ---------------------------- register to mni152 ---------------------------- #
|
|
87
|
+
if register_to_mni152 is not None:
|
|
88
|
+
from ..atlas.MNI152 import get_mni152
|
|
89
|
+
mni152 = get_mni152(f"2009a {register_to_mni152}w symmetric", skullstripped=False) # type:ignore
|
|
90
|
+
input_mni = register(input, mni152)
|
|
91
|
+
|
|
92
|
+
else:
|
|
93
|
+
input_mni = input
|
|
94
|
+
|
|
95
|
+
# ---------------------------- predict brain mask ---------------------------- #
|
|
96
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
97
|
+
sitk.WriteImage(input_mni, os.path.join(tmpdir, "input.nii.gz"))
|
|
98
|
+
|
|
99
|
+
run_hd_bet(
|
|
100
|
+
input = os.path.join(tmpdir, "input.nii.gz"),
|
|
101
|
+
output = os.path.join(tmpdir, "output.nii.gz"),
|
|
102
|
+
device=device, disable_tta=disable_tta, save_bet_mask=True, verbose=verbose,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
brain_mask_mni = tositk(os.path.join(tmpdir, "output_bet.nii.gz"))
|
|
106
|
+
|
|
107
|
+
# ------------------------- unregister mask if needed ------------------------ #
|
|
108
|
+
if register_to_mni152 is not None:
|
|
109
|
+
study_mni = dict(image=input_mni, seg_brain=brain_mask_mni)
|
|
110
|
+
study = register_D(study_mni, key="image", to=input)
|
|
111
|
+
return study["seg_brain"]
|
|
112
|
+
|
|
113
|
+
return brain_mask_mni
|
|
114
|
+
|
|
115
|
+
def skullstrip(
|
|
116
|
+
input: ImageLike,
|
|
117
|
+
register_to_mni152: Literal["T1", "T2"] | None = None,
|
|
118
|
+
device: Literal["cpu", "cuda", "mps"] = CUDA_IF_AVAILABLE,
|
|
119
|
+
disable_tta: bool = False,
|
|
120
|
+
verbose: bool = False,
|
|
121
|
+
) -> sitk.Image:
|
|
122
|
+
"""Skullstrips ``input`` using HD-BET.
|
|
123
|
+
|
|
124
|
+
Args:
|
|
125
|
+
input (ImageLike): input to skullstrip. Recommended T1-w, postcontrast T1-w, T2-w or FLAIR sequence in MNI152 space.
|
|
126
|
+
register_to_mni152 (str | None, optional):
|
|
127
|
+
Should be ``"T1"``, ``"T2"`` or ``None``.
|
|
128
|
+
if specified, ``input`` will be registered to specified MNI152 template,
|
|
129
|
+
and brain mask registered back to original ``input``.
|
|
130
|
+
Note that HD-BET expects images to be in MNI152 space. Defaults to None.
|
|
131
|
+
device (str, optional):
|
|
132
|
+
used to set on which device the prediction will run. Can be 'cuda' (=GPU), 'cpu' or 'mps'.
|
|
133
|
+
Defaults to CUDA_IF_AVAILABLE.
|
|
134
|
+
disable_tta (bool, optional):
|
|
135
|
+
Set this flag to disable test time augmentation. This will make prediction faster
|
|
136
|
+
at a slight decrease in prediction quality. Recommended for device cpu. Defaults to False.
|
|
137
|
+
verbose (bool, optional): purpose currently unknown. Defaults to False.
|
|
138
|
+
"""
|
|
139
|
+
input = tositk(input)
|
|
140
|
+
mask = predict_brain_mask(input=input, register_to_mni152=register_to_mni152,
|
|
141
|
+
device=device, disable_tta=disable_tta, verbose=verbose)
|
|
142
|
+
|
|
143
|
+
mask = sitk.Cast(mask, input.GetPixelID())
|
|
144
|
+
return sitk.Multiply(input, mask)
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
def skullstrip_D(
|
|
148
|
+
images: Mapping[str, ImageLike],
|
|
149
|
+
key: str,
|
|
150
|
+
register_to_mni152: Literal["T1", "T2"] | None = None,
|
|
151
|
+
device: Literal["cpu", "cuda", "mps"] = CUDA_IF_AVAILABLE,
|
|
152
|
+
disable_tta: bool = False,
|
|
153
|
+
verbose: bool = False,
|
|
154
|
+
) -> dict[str, sitk.Image]:
|
|
155
|
+
"""Predicts brain mask of ``images[key]``, then uses this mask to skull strip all values in ``images``.
|
|
156
|
+
|
|
157
|
+
Args:
|
|
158
|
+
images (Mapping[str, ImageLike]): dictionary of images that align with each other.
|
|
159
|
+
key (str): key of the image to pass to HD-BET for brain mask prediction.
|
|
160
|
+
register_to_mni152 (str | None, optional):
|
|
161
|
+
Should be ``"T1"``, ``"T2"`` or ``None``.
|
|
162
|
+
if specified, ``input`` will be registered to specified MNI152 template,
|
|
163
|
+
and brain mask registered back to original ``input``.
|
|
164
|
+
Note that HD-BET expects images to be in MNI152 space. Defaults to None.
|
|
165
|
+
device (str, optional):
|
|
166
|
+
used to set on which device the prediction will run. Can be 'cuda' (=GPU), 'cpu' or 'mps'.
|
|
167
|
+
Defaults to CUDA_IF_AVAILABLE.
|
|
168
|
+
disable_tta (bool, optional):
|
|
169
|
+
Set this flag to disable test time augmentation. This will make prediction faster
|
|
170
|
+
at a slight decrease in prediction quality. Recommended for device cpu. Defaults to False.
|
|
171
|
+
verbose (bool, optional): purpose currently unknown. Defaults to False.
|
|
172
|
+
"""
|
|
173
|
+
images = {k: tositk(v) for k,v in images.items()}
|
|
174
|
+
|
|
175
|
+
mask = predict_brain_mask(input=images[key], register_to_mni152=register_to_mni152,
|
|
176
|
+
device=device, disable_tta=disable_tta, verbose=verbose)
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
skullstripped = {}
|
|
180
|
+
for k,v in images.items():
|
|
181
|
+
mask_v = sitk.Cast(mask, v.GetPixelID())
|
|
182
|
+
skullstripped[k] = sitk.Multiply(v, mask_v)
|
|
183
|
+
|
|
184
|
+
return skullstripped
|
|
185
|
+
|
mrid/study.py
ADDED
|
@@ -0,0 +1,442 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import pickle
|
|
3
|
+
import shutil
|
|
4
|
+
import tempfile
|
|
5
|
+
import warnings
|
|
6
|
+
from collections import UserDict
|
|
7
|
+
from collections.abc import Callable, Mapping, Sequence
|
|
8
|
+
from functools import partial
|
|
9
|
+
from typing import TYPE_CHECKING, Any, Literal, overload
|
|
10
|
+
|
|
11
|
+
import numpy as np
|
|
12
|
+
import SimpleITK as sitk
|
|
13
|
+
|
|
14
|
+
from . import preprocessing
|
|
15
|
+
from .loading.convert import ImageLike, tonumpy, tositk, totensor
|
|
16
|
+
from .utils.torch_utils import CUDA_IF_AVAILABLE
|
|
17
|
+
|
|
18
|
+
if TYPE_CHECKING:
|
|
19
|
+
import torch
|
|
20
|
+
|
|
21
|
+
def _identity(x): return x
|
|
22
|
+
|
|
23
|
+
class Study(UserDict[str, sitk.Image | Any]):
|
|
24
|
+
@overload
|
|
25
|
+
def __init__(self, /, **kwargs): ...
|
|
26
|
+
@overload
|
|
27
|
+
def __init__(self, dict, /): ...
|
|
28
|
+
def __init__(self, dict=None, /, **kwargs):
|
|
29
|
+
if dict is None: dict = kwargs
|
|
30
|
+
|
|
31
|
+
proc = {}
|
|
32
|
+
for k,v in dict.items():
|
|
33
|
+
if k.startswith("info"):
|
|
34
|
+
proc[k] = v
|
|
35
|
+
else:
|
|
36
|
+
proc[k] = tositk(v)
|
|
37
|
+
|
|
38
|
+
super().__init__(proc)
|
|
39
|
+
|
|
40
|
+
def __setitem__(self, key: str, item: "ImageLike | Any") -> None:
|
|
41
|
+
if not key.startswith("info"): item = tositk(item)
|
|
42
|
+
return super().__setitem__(key, item)
|
|
43
|
+
|
|
44
|
+
def add(self, key: str, item: "ImageLike | Any", reference_key: str | None = None):
|
|
45
|
+
"""Returns a new study with an extra item inserted under ``key``.
|
|
46
|
+
|
|
47
|
+
Args:
|
|
48
|
+
key (str): Key to insert new item under.
|
|
49
|
+
item (ImageLike | Any): Item to insert.
|
|
50
|
+
reference_key (str | None, optional):
|
|
51
|
+
if specified, ``item`` will have SimpleITK attributes copied from ``self[reference_key]``. Defaults to None.
|
|
52
|
+
"""
|
|
53
|
+
study = self.copy()
|
|
54
|
+
if key.startswith('info'):
|
|
55
|
+
if reference_key: raise RuntimeError(f"Can't copy sitk attributes for an non-image item {key}")
|
|
56
|
+
study[key] = item
|
|
57
|
+
return study
|
|
58
|
+
|
|
59
|
+
item = tositk(item)
|
|
60
|
+
if reference_key is not None: item.CopyInformation(study[reference_key])
|
|
61
|
+
study[key] = item
|
|
62
|
+
return study
|
|
63
|
+
|
|
64
|
+
def get_scans(self):
|
|
65
|
+
"""Returns a new ``Study`` with segmentations and info removed."""
|
|
66
|
+
return self.__class__({k:v for k,v in self.items() if not k.startswith(("seg", "info"))})
|
|
67
|
+
|
|
68
|
+
def get_images(self):
|
|
69
|
+
"""Returns a new ``Study`` with info removed."""
|
|
70
|
+
return self.__class__({k:v for k,v in self.items() if not k.startswith("info")})
|
|
71
|
+
|
|
72
|
+
def get_segmentations(self):
|
|
73
|
+
"""Returns a new ``Study`` with scans and info removed."""
|
|
74
|
+
return self.__class__({k:v for k,v in self.items() if k.startswith("seg")})
|
|
75
|
+
|
|
76
|
+
def get_info(self):
|
|
77
|
+
"""Returns a new ``Study`` with scans and segmentations removed."""
|
|
78
|
+
return self.__class__({k:v for k,v in self.items() if k.startswith("info")})
|
|
79
|
+
|
|
80
|
+
def apply(self, fn:Callable[[sitk.Image], sitk.Image] | None, seg_fn: Callable[[sitk.Image], sitk.Image] | None,) -> "Study":
|
|
81
|
+
"""Returns a new ``Study`` with ``fn`` applied to scans and ``seg_fn`` applied to segmentations.
|
|
82
|
+
|
|
83
|
+
Args:
|
|
84
|
+
fn: Function to apply to scan images. Must take and return ``sitk.Image``.
|
|
85
|
+
If None, identity function is used.
|
|
86
|
+
seg_fn: Function to apply to segmentation images. Must take and return ``sitk.Image``.
|
|
87
|
+
If None, identity function is used.
|
|
88
|
+
"""
|
|
89
|
+
if fn is None: fn = _identity
|
|
90
|
+
if seg_fn is None: seg_fn = _identity
|
|
91
|
+
|
|
92
|
+
scans = {k: fn(v) for k,v in self.get_scans().items()}
|
|
93
|
+
seg = {k: seg_fn(v) for k,v in self.get_segmentations().items()}
|
|
94
|
+
|
|
95
|
+
return Study(**scans, **seg, **self.get_info())
|
|
96
|
+
|
|
97
|
+
def cast(self, dtype) -> "Study":
|
|
98
|
+
"""Returns a new study with all scans cast to the specified SimpleITK dtype.
|
|
99
|
+
|
|
100
|
+
Note:
|
|
101
|
+
This operation does not affect segmentations.
|
|
102
|
+
"""
|
|
103
|
+
return self.apply(partial(sitk.Cast, pixelID=dtype), seg_fn=None)
|
|
104
|
+
|
|
105
|
+
def cast_float64(self) -> "Study":
|
|
106
|
+
"""Returns a new study with all scans cast to float64.
|
|
107
|
+
|
|
108
|
+
Note:
|
|
109
|
+
This operation does not affect segmentations.
|
|
110
|
+
"""
|
|
111
|
+
return self.cast(sitk.sitkFloat64)
|
|
112
|
+
|
|
113
|
+
def cast_float32(self) -> "Study":
|
|
114
|
+
"""Returns a new study with all scans cast to float32.
|
|
115
|
+
|
|
116
|
+
Note:
|
|
117
|
+
This operation does not affect segmentations.
|
|
118
|
+
"""
|
|
119
|
+
return self.cast(sitk.sitkFloat32)
|
|
120
|
+
|
|
121
|
+
def normalize(self) -> "Study":
|
|
122
|
+
"""Return a new study where all scans are separately z-normalized to 0 mean and 1 variance.
|
|
123
|
+
|
|
124
|
+
Note:
|
|
125
|
+
This operation does not affect segmentations.
|
|
126
|
+
"""
|
|
127
|
+
return self.apply(sitk.Normalize, seg_fn=None)
|
|
128
|
+
|
|
129
|
+
def rescale_intensity(self, min: float, max: float) -> "Study":
|
|
130
|
+
"""Return a new study where all scans are separately rescaled to the specified intensity range.
|
|
131
|
+
|
|
132
|
+
Args:
|
|
133
|
+
min: Minimum value for the output intensity range.
|
|
134
|
+
max: Maximum value for the output intensity range.
|
|
135
|
+
"""
|
|
136
|
+
return self.apply(partial(sitk.RescaleIntensity, outputMinimum = min, outputMaximum = max), seg_fn=None) # type:ignore
|
|
137
|
+
|
|
138
|
+
def crop_bg(self, key: str) -> "Study":
|
|
139
|
+
"""Return a new study with cropped black background. Finds the foreground bounding box of ``study[key]``,
|
|
140
|
+
and uses that bounding box to crop all other images, including segmentations.
|
|
141
|
+
|
|
142
|
+
Args:
|
|
143
|
+
key: The key of the image to use for finding the foreground bounding box.
|
|
144
|
+
"""
|
|
145
|
+
d = preprocessing.cropping.crop_bg_D(self.get_images(), key)
|
|
146
|
+
return Study(**d, **self.get_info())
|
|
147
|
+
|
|
148
|
+
def skullstrip(
|
|
149
|
+
self,
|
|
150
|
+
key: str,
|
|
151
|
+
register_to_mni152: Literal["T1", "T2"] | None = None,
|
|
152
|
+
device: Literal["cpu", "cuda", "mps"] = CUDA_IF_AVAILABLE,
|
|
153
|
+
disable_tta: bool = False,
|
|
154
|
+
verbose: bool = False,
|
|
155
|
+
) -> "Study":
|
|
156
|
+
"""Predicts brain mask of ``study[key]``, then uses this mask to skull strip all scans. Doesn't affect segmentations.
|
|
157
|
+
|
|
158
|
+
Args:
|
|
159
|
+
key: Key of the image to pass to HD-BET for brain mask prediction.
|
|
160
|
+
register_to_mni152: Should be ``"T1"``, ``"T2"`` or ``None``.
|
|
161
|
+
If specified, ``input`` will be registered to specified MNI152 template,
|
|
162
|
+
and brain mask registered back to original ``input``.
|
|
163
|
+
Note that HD-BET expects images to be in MNI152 space. Defaults to None.
|
|
164
|
+
device: Used to set on which device the prediction will run. Can be 'cuda' (=GPU), 'cpu' or 'mps'.
|
|
165
|
+
Defaults to CUDA_IF_AVAILABLE.
|
|
166
|
+
disable_tta: Set this flag to disable test time augmentation. This will make prediction faster
|
|
167
|
+
at a slight decrease in prediction quality. Recommended for device cpu. Defaults to False.
|
|
168
|
+
verbose: Enable verbose output during processing. Defaults to False.
|
|
169
|
+
"""
|
|
170
|
+
d = preprocessing.skullstripping.skullstrip_D(
|
|
171
|
+
images=self.get_scans(),
|
|
172
|
+
key=key,
|
|
173
|
+
register_to_mni152=register_to_mni152,
|
|
174
|
+
device=device,
|
|
175
|
+
disable_tta=disable_tta,
|
|
176
|
+
verbose=verbose,
|
|
177
|
+
)
|
|
178
|
+
return Study(**d, **self.get_segmentations(), **self.get_info())
|
|
179
|
+
|
|
180
|
+
def resize(self, size: Sequence[int], interpolator=sitk.sitkLinear):
|
|
181
|
+
"""Resize all images to ``size``.
|
|
182
|
+
|
|
183
|
+
Args:
|
|
184
|
+
size: Target size as a sequence of integers (e.g., [height, width, depth]).
|
|
185
|
+
interpolator: Interpolation method for regular images (segmentations always use nearest neighbor).
|
|
186
|
+
"""
|
|
187
|
+
return self.apply(
|
|
188
|
+
partial(preprocessing.registration.resize, new_size=size, interpolator=interpolator,),
|
|
189
|
+
partial(preprocessing.registration.resize, new_size=size, interpolator=sitk.sitkNearestNeighbor,),
|
|
190
|
+
)
|
|
191
|
+
|
|
192
|
+
def downsample(self, factor: float, dims = None, interpolator=sitk.sitkLinear):
|
|
193
|
+
"""Downsample all images. Factor = 2 for 2x downsampling. Dims ``None`` for all dims.
|
|
194
|
+
|
|
195
|
+
Args:
|
|
196
|
+
factor: Downsampling factor (e.g., 2 for 2x downsampling).
|
|
197
|
+
dims: Specific dimensions to downsample, or None for all dimensions.
|
|
198
|
+
interpolator: Interpolation method for regular images (segmentations always use nearest neighbor).
|
|
199
|
+
"""
|
|
200
|
+
return self.apply(
|
|
201
|
+
partial(preprocessing.registration.downsample, factor=factor, dims=dims, interpolator=interpolator,),
|
|
202
|
+
partial(preprocessing.registration.downsample, factor=factor, dims=dims, interpolator=sitk.sitkNearestNeighbor,),
|
|
203
|
+
)
|
|
204
|
+
|
|
205
|
+
def register(self, key: str, to: ImageLike, pmap=None, log_to_console=False) -> "Study":
|
|
206
|
+
"""Returns a new Study, registers ``study[key]`` to ``to``,
|
|
207
|
+
and use transformation parameters to register all other images including segmentation.
|
|
208
|
+
This assumes that all images are in the same space, if they are not, use ``register_many`` method.
|
|
209
|
+
|
|
210
|
+
Args:
|
|
211
|
+
key: The key of the image to use as reference for registration.
|
|
212
|
+
to: Target image or path to register to.
|
|
213
|
+
pmap: Parameter map for registration. If None, uses default parameters.
|
|
214
|
+
log_to_console: Whether to log registration progress to console.
|
|
215
|
+
"""
|
|
216
|
+
d = preprocessing.registration.register_D(
|
|
217
|
+
images=self.get_images(),
|
|
218
|
+
key=key,
|
|
219
|
+
to=to,
|
|
220
|
+
pmap=pmap,
|
|
221
|
+
log_to_console=log_to_console,
|
|
222
|
+
)
|
|
223
|
+
return Study(**d, **self.get_info())
|
|
224
|
+
|
|
225
|
+
def register_each(self, key: str, to: "ImageLike | None" = None, pmap=None, log_to_console=False) -> "Study":
|
|
226
|
+
"""Returns a new study. Registers all other images to ``study[key]``.
|
|
227
|
+
If ``to`` is specified, register ``study[key]`` to ``to`` beforehand.
|
|
228
|
+
|
|
229
|
+
Args:
|
|
230
|
+
key: The key of the image to use as reference for registration.
|
|
231
|
+
to: Target image or path to register the reference image to. If None, uses key as reference.
|
|
232
|
+
pmap: Parameter map for registration. If None, uses default parameters.
|
|
233
|
+
log_to_console: Whether to log registration progress to console.
|
|
234
|
+
|
|
235
|
+
Note:
|
|
236
|
+
If called on a study with segmentations, they will be removed from the returned study.
|
|
237
|
+
"""
|
|
238
|
+
if len(self.get_segmentations()) > 0:
|
|
239
|
+
keys = ', '.join(self.get_segmentations().keys())
|
|
240
|
+
warnings.warn(f"`register_many` was called on a study with segmentations ({keys}), "
|
|
241
|
+
"they will be removed from the returned study", stacklevel=3)
|
|
242
|
+
|
|
243
|
+
d = self.get_scans()
|
|
244
|
+
if to is not None:
|
|
245
|
+
d[key] = preprocessing.registration.register(d[key], to=to, pmap=pmap, log_to_console=log_to_console)
|
|
246
|
+
|
|
247
|
+
for k in d:
|
|
248
|
+
if k != key:
|
|
249
|
+
d[k] = preprocessing.registration.register(d[k], to=d[key], pmap=pmap, log_to_console=log_to_console)
|
|
250
|
+
|
|
251
|
+
return Study(**d, **self.get_info())
|
|
252
|
+
|
|
253
|
+
def resample_to(self, to: "np.ndarray | sitk.Image | torch.Tensor | str", interpolation=sitk.sitkLinear) -> "Study":
|
|
254
|
+
"""Returns a new study, resamples all images including segmentation to `to`.
|
|
255
|
+
Segmentation always uses nearest interpolation"""
|
|
256
|
+
to = tositk(to)
|
|
257
|
+
|
|
258
|
+
return self.apply(
|
|
259
|
+
partial(preprocessing.registration.resample_to, to=to, interpolation=interpolation),
|
|
260
|
+
partial(preprocessing.registration.resample_to, to=to, interpolation=sitk.sitkNearestNeighbor),
|
|
261
|
+
)
|
|
262
|
+
|
|
263
|
+
def n4_bias_field_correction(self, key: str, shrink: int = 4) -> "Study":
|
|
264
|
+
"""Returns a new study with corrected bias field of the image under ``key``. Doesn't affect other images.
|
|
265
|
+
|
|
266
|
+
Args:
|
|
267
|
+
key: The key of the image for which to correct the bias field.
|
|
268
|
+
shrink: By how many times to shrink the size of input image for calculating the bias field.
|
|
269
|
+
The bias field is then applied to original size (unshrunk) image.
|
|
270
|
+
Setting shrink to 1 disables it, but n4 algorithm may take several minutes.
|
|
271
|
+
Setting it to ~4 is good enough in most cases and will be significantly faster (usually few seconds).
|
|
272
|
+
"""
|
|
273
|
+
new = self.copy()
|
|
274
|
+
new[key] = preprocessing.bias_field_correction.n4_bias_field_correction(new[key], shrink=shrink)
|
|
275
|
+
return new
|
|
276
|
+
|
|
277
|
+
def numpy(self, key: str):
|
|
278
|
+
"""returns ``study[key]`` converted to a numpy array."""
|
|
279
|
+
return tonumpy(self[key])
|
|
280
|
+
|
|
281
|
+
def tensor(self, key: str):
|
|
282
|
+
"""returns ``study[key]`` converted to a tensor."""
|
|
283
|
+
return totensor(self[key])
|
|
284
|
+
|
|
285
|
+
def _get_sorted_items(self, scans: bool, seg: bool, order: Sequence[str] | None = None) -> list[tuple[str, sitk.Image]]:
|
|
286
|
+
if not (scans or seg): raise ValueError("At least one of `scans` or `seg` must be True")
|
|
287
|
+
|
|
288
|
+
if order is not None:
|
|
289
|
+
return [(k, self[k]) for k in order]
|
|
290
|
+
|
|
291
|
+
# make sure items are always sorted in the same order
|
|
292
|
+
items = []
|
|
293
|
+
if scans: items = sorted(self.get_scans().items(), key = lambda x: x[0])
|
|
294
|
+
if seg: items.extend(sorted(self.get_segmentations().items(), key = lambda x: x[0]))
|
|
295
|
+
return items
|
|
296
|
+
|
|
297
|
+
|
|
298
|
+
def stack_numpy(self, scans:bool = True, seg: bool = False, dtype=None, order: Sequence[str] | None = None) -> np.ndarray:
|
|
299
|
+
"""Stack images into a numpy array, returns an array of shape ``(n_images, *dims)``.
|
|
300
|
+
|
|
301
|
+
Args:
|
|
302
|
+
scans: Whether to include scan images in the stack.
|
|
303
|
+
seg: Whether to include segmentation images in the stack.
|
|
304
|
+
dtype: Data type for the output array. If None, uses the default type.
|
|
305
|
+
order:
|
|
306
|
+
Specific order for the images in the stack. If specified, ignores ``scans`` and ``seg`` options.
|
|
307
|
+
If None, uses alphabetic sorting.
|
|
308
|
+
"""
|
|
309
|
+
items = self._get_sorted_items(scans=scans, seg=seg, order=order)
|
|
310
|
+
|
|
311
|
+
stacked = np.array([sitk.GetArrayFromImage(v) for k, v in items])
|
|
312
|
+
if dtype is not None: stacked = stacked.astype(dtype, copy=False)
|
|
313
|
+
return stacked
|
|
314
|
+
|
|
315
|
+
def stack_tensor(self, scans:bool = True, seg: bool = False, device=None, dtype=None, order: Sequence[str] | None = None) -> "torch.Tensor":
|
|
316
|
+
"""Stack images into a torch tensor, returns an tensor of shape ``(n_images, *dims)``.
|
|
317
|
+
|
|
318
|
+
Args:
|
|
319
|
+
scans: Whether to include scan images in the stack.
|
|
320
|
+
seg: Whether to include segmentation images in the stack.
|
|
321
|
+
device: Device for the output tensor. If None, uses the default device.
|
|
322
|
+
dtype: Data type for the output tensor. If None, uses the default type.
|
|
323
|
+
order:
|
|
324
|
+
Specific order for the images in the stack. If specified, ignores ``scans`` and ``seg`` options.
|
|
325
|
+
If None, uses alphabetic sorting.
|
|
326
|
+
"""
|
|
327
|
+
import torch
|
|
328
|
+
items = self._get_sorted_items(scans=scans, seg=seg, order=order)
|
|
329
|
+
stacked = torch.stack([torch.from_numpy(sitk.GetArrayFromImage(v)) for _,v in items])
|
|
330
|
+
return stacked.to(device=device, dtype=dtype, memory_format=torch.contiguous_format)
|
|
331
|
+
|
|
332
|
+
def numpy_dict(self) -> dict[str, np.ndarray | Any]:
|
|
333
|
+
"""Returns a dictionary with all images converted to numpy arrays, info is included as is."""
|
|
334
|
+
return {k: (sitk.GetArrayFromImage(v) if isinstance(v, sitk.Image) else v) for k, v in self.items()}
|
|
335
|
+
|
|
336
|
+
def tensor_dict(self) -> "dict[str, torch.Tensor | Any]":
|
|
337
|
+
"""Returns a dictionary with all images converted to tensors, info is included as is."""
|
|
338
|
+
import torch
|
|
339
|
+
return {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v) for k,v in self.numpy_dict()}
|
|
340
|
+
|
|
341
|
+
def plot(self):
|
|
342
|
+
from .utils.plotting import visualize_3d_arrays
|
|
343
|
+
visualize_3d_arrays(self.get_images().numpy_dict())
|
|
344
|
+
|
|
345
|
+
def save(
|
|
346
|
+
self,
|
|
347
|
+
dir: str | os.PathLike,
|
|
348
|
+
prefix: str = "",
|
|
349
|
+
suffix: str = "",
|
|
350
|
+
ext: str = "nii.gz",
|
|
351
|
+
mkdir=True,
|
|
352
|
+
use_compression=True,
|
|
353
|
+
pickle_module = pickle,
|
|
354
|
+
):
|
|
355
|
+
"""Writes this study to a directory, with filenames being ``{path}/{prefix}{key}{suffix}.{ext}``
|
|
356
|
+
|
|
357
|
+
Args:
|
|
358
|
+
dir: Directory to save the study to.
|
|
359
|
+
prefix: Prefix to add to all filenames.
|
|
360
|
+
suffix: Suffix to add to all filenames (before extension).
|
|
361
|
+
ext: File extension for image files. Default is 'nii.gz'.
|
|
362
|
+
mkdir: Whether to create the directory if it doesn't exist. Default is True.
|
|
363
|
+
use_compression: Whether to use compression for image files. Default is True.
|
|
364
|
+
pickle_module: Module to use for pickling info objects. Default is pickle.
|
|
365
|
+
"""
|
|
366
|
+
if ext.startswith('.'): ext = ext[1:]
|
|
367
|
+
|
|
368
|
+
# make directory
|
|
369
|
+
if not os.path.exists(dir):
|
|
370
|
+
if mkdir: os.mkdir(dir)
|
|
371
|
+
else: raise FileNotFoundError(f"Directory {dir} doesn't exist and {mkdir = }")
|
|
372
|
+
|
|
373
|
+
# save
|
|
374
|
+
for k,v in self.items():
|
|
375
|
+
|
|
376
|
+
# save infos
|
|
377
|
+
if k.startswith('info'):
|
|
378
|
+
try:
|
|
379
|
+
with open(os.path.join(dir, f"{prefix}{k}{suffix}.pkl"), "wb") as file:
|
|
380
|
+
pickle_module.dump(v, file)
|
|
381
|
+
|
|
382
|
+
except Exception as e:
|
|
383
|
+
print(f"Couldn't save {k}:\n{e!r}")
|
|
384
|
+
|
|
385
|
+
# save images
|
|
386
|
+
else:
|
|
387
|
+
# this handles non ascii chars
|
|
388
|
+
with tempfile.TemporaryDirectory() as temp_path:
|
|
389
|
+
sitk.WriteImage(v, os.path.join(temp_path, f"{prefix}{k}{suffix}.{ext}"), useCompression=use_compression)
|
|
390
|
+
shutil.move(os.path.join(temp_path, f"{prefix}{k}{suffix}.{ext}"), dir)
|
|
391
|
+
|
|
392
|
+
def load(self, dir: str | os.PathLike, prefix: str = '', suffix: str = '', ext: str = 'nii.gz', pickle_module = pickle):
|
|
393
|
+
"""Returns a new study, updated by data loaded from ``dir``, which can be created by calling ``study.save(dir)``.
|
|
394
|
+
|
|
395
|
+
Args:
|
|
396
|
+
dir: Directory to load the study from.
|
|
397
|
+
prefix: Expected prefix of filenames.
|
|
398
|
+
suffix: Expected suffix of filenames.
|
|
399
|
+
ext: Expected file extension for image files. Default is 'nii.gz'.
|
|
400
|
+
pickle_module: Module used for unpickling info objects. Default is pickle.
|
|
401
|
+
"""
|
|
402
|
+
study = self.copy()
|
|
403
|
+
|
|
404
|
+
files = os.listdir(dir)
|
|
405
|
+
|
|
406
|
+
for f in files:
|
|
407
|
+
full = os.path.join(dir, f)
|
|
408
|
+
name:str = f
|
|
409
|
+
|
|
410
|
+
# check prefix
|
|
411
|
+
if prefix == '' or name.startswith(prefix):
|
|
412
|
+
name = name[len(prefix):]
|
|
413
|
+
|
|
414
|
+
# load images
|
|
415
|
+
if name.endswith(f'{suffix}.{ext}'):
|
|
416
|
+
name = name[:-len(f'{suffix}.{ext}')]
|
|
417
|
+
study[name] = tositk(full)
|
|
418
|
+
|
|
419
|
+
# load infos
|
|
420
|
+
elif name.endswith(f'{suffix}.pkl'):
|
|
421
|
+
name = name[:-len(f'{suffix}.pkl')]
|
|
422
|
+
|
|
423
|
+
try:
|
|
424
|
+
with open(full, 'rb') as file:
|
|
425
|
+
study[name] = pickle_module.load(file)
|
|
426
|
+
except Exception as e:
|
|
427
|
+
print(f"Couldn't load {full}:\n{e!r}")
|
|
428
|
+
|
|
429
|
+
return study
|
|
430
|
+
|
|
431
|
+
@classmethod
|
|
432
|
+
def from_dir(cls, dir: str | os.PathLike, prefix: str = '', suffix: str = '', ext: str = 'nii.gz', pickle_module = pickle):
|
|
433
|
+
"""Load a study from a directory.
|
|
434
|
+
|
|
435
|
+
Args:
|
|
436
|
+
dir: Directory to load the study from.
|
|
437
|
+
prefix: Expected prefix of filenames.
|
|
438
|
+
suffix: Expected suffix of filenames.
|
|
439
|
+
ext: Expected file extension for image files. Default is 'nii.gz'.
|
|
440
|
+
pickle_module: Module used for unpickling info objects. Default is pickle.
|
|
441
|
+
"""
|
|
442
|
+
return cls().load(dir=dir, prefix=prefix, suffix=suffix, ext=ext, pickle_module=pickle_module)
|
mrid/utils/__init__.py
ADDED