mrid-python 0.1.5__tar.gz → 0.1.7__tar.gz
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_python-0.1.5 → mrid_python-0.1.7}/PKG-INFO +1 -1
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/atlas/MNI152/__init__.py +1 -1
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/CTseg.py +1 -1
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/__init__.py +2 -2
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/cropping.py +37 -1
- mrid_python-0.1.7/mrid/preprocessing/haca3.py +242 -0
- mrid_python-0.1.7/mrid/preprocessing/spatial.py +51 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/study.py +123 -9
- mrid_python-0.1.7/mrid/training/__init__.py +0 -0
- mrid_python-0.1.7/mrid/training/cropping.py +52 -0
- mrid_python-0.1.7/mrid/training/padding.py +97 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/training/slicer.py +6 -1
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/training/transforms.py +1 -50
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/dcm2niix.py +4 -2
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/plotting.py +28 -25
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/python_utils.py +7 -19
- mrid_python-0.1.7/mrid/utils/sitk_utils.py +18 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/PKG-INFO +1 -1
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/SOURCES.txt +5 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/pyproject.toml +1 -1
- {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_preprocessing.py +5 -5
- {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_study.py +1 -1
- mrid_python-0.1.5/mrid/preprocessing/spatial.py +0 -82
- {mrid_python-0.1.5 → mrid_python-0.1.7}/README.md +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/__init__.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/atlas/SRI24/__init__.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/atlas/__init__.py +0 -0
- {mrid_python-0.1.5/mrid/training → mrid_python-0.1.7/mrid/inference}/__init__.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/loading/__init__.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/loading/convert.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/bias_field_correction.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/hd_bet.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/mask.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/simple_elastix.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/synthstrip.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/__init__.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/dicom_uid_fixer.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/stl_utils.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/torch_utils.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/dependency_links.txt +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/requires.txt +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/top_level.txt +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/setup.cfg +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_loading.py +0 -0
- {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: mrid-python
|
|
3
|
-
Version: 0.1.
|
|
3
|
+
Version: 0.1.7
|
|
4
4
|
Summary: Tools for working with 3D medical images and segmentations - registration, brain skull-stripping, etc.
|
|
5
5
|
Author-email: Ivan Nikishev <nkshv2@gmail.com>
|
|
6
6
|
Project-URL: Homepage, https://github.com/inikishev/mrid
|
|
@@ -60,7 +60,7 @@ def get_mni152(
|
|
|
60
60
|
],
|
|
61
61
|
skullstripped: bool = False,
|
|
62
62
|
):
|
|
63
|
-
"""Returns path to .nii.gz file of specified MNI-152 template.
|
|
63
|
+
"""Returns path to .nii.gz file of specified MNI-152 template. All templates are with skull.
|
|
64
64
|
|
|
65
65
|
The following templates are available:
|
|
66
66
|
- ``"2006 T1w symmetric"``
|
|
@@ -1,9 +1,9 @@
|
|
|
1
1
|
from .bias_field_correction import n4_bias_field_correction
|
|
2
|
-
from .cropping import crop_bg, crop_bg_D
|
|
2
|
+
from .cropping import crop_bg, crop_bg_D, center_crop_or_pad
|
|
3
3
|
from .spatial import downsample, resample_to, resize
|
|
4
4
|
|
|
5
5
|
# lib wrappers
|
|
6
|
-
from . import hd_bet, CTseg, simple_elastix, synthstrip, mask
|
|
6
|
+
from . import hd_bet, CTseg, simple_elastix, synthstrip, mask, haca3
|
|
7
7
|
__all__ = [
|
|
8
8
|
"n4_bias_field_correction",
|
|
9
9
|
"crop_bg", "crop_bg_D",
|
|
@@ -1,4 +1,4 @@
|
|
|
1
|
-
from collections.abc import Mapping
|
|
1
|
+
from collections.abc import Mapping, Sequence
|
|
2
2
|
from typing import Any
|
|
3
3
|
import SimpleITK as sitk
|
|
4
4
|
|
|
@@ -34,3 +34,39 @@ def crop_bg_D(images: Mapping[str, ImageLike], key: str) -> dict[str, sitk.Image
|
|
|
34
34
|
ret = {k: sitk.RegionOfInterest(v, bbox[int(len(bbox) / 2) :], bbox[0 : int(len(bbox) / 2)]) for k,v in images.items()}
|
|
35
35
|
return ret
|
|
36
36
|
|
|
37
|
+
def center_crop_or_pad(image, size: Sequence[int]) -> "sitk.Image":#[192, 224, 192]
|
|
38
|
+
"""Crops or pads image from the center to ``size``."""
|
|
39
|
+
image = tositk(image)
|
|
40
|
+
current_size = list(image.GetSize())
|
|
41
|
+
|
|
42
|
+
low_pad = []
|
|
43
|
+
high_pad = []
|
|
44
|
+
low_crop = []
|
|
45
|
+
high_crop = []
|
|
46
|
+
|
|
47
|
+
for cur, tar in zip(current_size, size):
|
|
48
|
+
diff = tar - cur
|
|
49
|
+
if diff >= 0:
|
|
50
|
+
low = diff // 2
|
|
51
|
+
high = diff - low
|
|
52
|
+
low_pad.append(low)
|
|
53
|
+
high_pad.append(high)
|
|
54
|
+
low_crop.append(0)
|
|
55
|
+
high_crop.append(0)
|
|
56
|
+
else:
|
|
57
|
+
diff = abs(diff)
|
|
58
|
+
low = diff // 2
|
|
59
|
+
high = diff - low
|
|
60
|
+
low_pad.append(0)
|
|
61
|
+
high_pad.append(0)
|
|
62
|
+
low_crop.append(low)
|
|
63
|
+
high_crop.append(high)
|
|
64
|
+
|
|
65
|
+
image = sitk.ConstantPad(image, low_pad, high_pad, 0)
|
|
66
|
+
image = sitk.Crop(image, low_crop, high_crop)
|
|
67
|
+
|
|
68
|
+
# Verify size
|
|
69
|
+
if tuple(image.GetSize()) == tuple(size):
|
|
70
|
+
return image
|
|
71
|
+
|
|
72
|
+
raise RuntimeError(f"Final size is {image.GetSize()} instead of {size}, crop_or_pad failed for some reason.")
|
|
@@ -0,0 +1,242 @@
|
|
|
1
|
+
""""
|
|
2
|
+
This requires [HACA3](https://github.com/lianruizuo/haca3).
|
|
3
|
+
|
|
4
|
+
HACA3 can be installed from source or through singularity image. Mrid currently only supports it being installed in a separate environment (because thats how I currently need to use it), let me know if you need other installation methods supported.
|
|
5
|
+
|
|
6
|
+
You also need to download HACA3 weights harmonization.pt and fusion model weights fusion.pt from the ``4. Usage: Inference`` section in HACA3 github readme.
|
|
7
|
+
"""
|
|
8
|
+
import os
|
|
9
|
+
import shlex
|
|
10
|
+
import subprocess
|
|
11
|
+
import tempfile
|
|
12
|
+
from collections.abc import Sequence
|
|
13
|
+
from pathlib import Path
|
|
14
|
+
from typing import Literal
|
|
15
|
+
|
|
16
|
+
import SimpleITK as sitk
|
|
17
|
+
|
|
18
|
+
from ..loading import ImageLike, tositk
|
|
19
|
+
from .cropping import center_crop_or_pad
|
|
20
|
+
from .simple_elastix import register, register_D
|
|
21
|
+
|
|
22
|
+
# - ```--in-path```: file path to input source image. Multiple ```--in-path``` may be provided if there are multiple
|
|
23
|
+
# source images. See the above example for more details.
|
|
24
|
+
# - ```--target-image```: file path to target image. HACA3 will match the contrast of source images to this target image.
|
|
25
|
+
# - ```--target-theta```: In HACA3, ```theta```
|
|
26
|
+
# is a two-dimensional representation of image contrast. Target image contrast can be directly specified by providing
|
|
27
|
+
# a ```theta``` value, e.g., ```--target-theta 0.5 0.5```. Note: either ```--target-image``` or ```--target-image``` must
|
|
28
|
+
# be provided during inference. If both are provided, only ```--target-theta``` will be used.
|
|
29
|
+
# - ```--norm-val```: normalization value.
|
|
30
|
+
# - ```--out-path```: file path to harmonized image.
|
|
31
|
+
# - ```--harmonization-model```: pretrained HACA3 weights. Pretrained model weights on IXI, OASIS and HCP data can
|
|
32
|
+
# be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/harmonization_public.pt).
|
|
33
|
+
# - ```--fusion-model```: pretrained fusion model weights. HACA3 uses a 3D convolutional network to combine multi-orientation
|
|
34
|
+
# 2D slices into a single 3D volume. Pretrained fusion model can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/fusion.pt).
|
|
35
|
+
# - ```--save-intermediate```: if specified, intermediate results will be saved. Default: ```False```. Action: ```store_true```.
|
|
36
|
+
# - ```--intermediate-out-dir```: directory to save intermediate results.
|
|
37
|
+
# - ```--gpu-id```: integer number specifies which GPU to run HACA3.
|
|
38
|
+
# - ```--num-batches```: During inference, HACA3 takes entire 3D MRI volumes as input. This may cause a considerable amount
|
|
39
|
+
# GPU memory. For reduced GPU memory consumption, source images maybe divided into smaller batches.
|
|
40
|
+
# However, this may slightly increase the inference time.
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def run_HACA3(
|
|
44
|
+
conda_path: str | os.PathLike,
|
|
45
|
+
env_name: str,
|
|
46
|
+
harmonization_model: str | os.PathLike,
|
|
47
|
+
fusion_model: str | os.PathLike,
|
|
48
|
+
in_path: str | os.PathLike | Sequence[str | os.PathLike],
|
|
49
|
+
out_path: str | os.PathLike,
|
|
50
|
+
target_image: str | os.PathLike | None,
|
|
51
|
+
target_theta: tuple[float,float] | None,
|
|
52
|
+
norm_val: float | None = None,
|
|
53
|
+
intermediate_out_dir: str | os.PathLike | None = None,
|
|
54
|
+
gpu_id: int | None = None,
|
|
55
|
+
num_batches: int | None = None,
|
|
56
|
+
) -> None:
|
|
57
|
+
"""Runs ``HACA3`` command-line routine via ``subprocess.run``.
|
|
58
|
+
|
|
59
|
+
Args:
|
|
60
|
+
conda_path: path to ``minconda3`` directory.
|
|
61
|
+
env_name: name of the conda env where HACA3 is installed.
|
|
62
|
+
harmonization_model: pretrained HACA3 weights. Pretrained model weights on IXI, OASIS and HCP
|
|
63
|
+
data can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/harmonization_public.pt).
|
|
64
|
+
fusion_model: pretrained fusion model weights. HACA3 uses a 3D convolutional network to
|
|
65
|
+
combine multi-orientation 2D slices into a single 3D volume. Pretrained fusion model
|
|
66
|
+
can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/fusion.pt).
|
|
67
|
+
in_path: file path to input source image. Multiple paths may be provided if there are
|
|
68
|
+
multiple source images (different modalities). Note that modalities must be in
|
|
69
|
+
MNI space (1mm isotropic resolution). HACA3 assumes a spatial dimension of 192x224x192.
|
|
70
|
+
out_path: file path to harmonized image.
|
|
71
|
+
target_image: file path to target image. HACA3 will match the contrast of
|
|
72
|
+
source images to this target image.
|
|
73
|
+
target_theta: In HACA3, ```theta``` is a two-dimensional representation of image contrast.
|
|
74
|
+
Target image contrast can be directly specified by providing a ```theta``` value, e.g.,
|
|
75
|
+
```target_theta = (0.5, 0.5)```. Note: either ```target_image``` or ```target_theta```
|
|
76
|
+
must be provided during inference.
|
|
77
|
+
norm_val: normalization value. Defaults to None.
|
|
78
|
+
intermediate_out_dir: directory to save intermediate results. Defaults to None.
|
|
79
|
+
gpu_id: integer number specifies which GPU to run HACA3. Defaults to None.
|
|
80
|
+
num_batches: During inference, HACA3 takes entire 3D MRI volumes as input.
|
|
81
|
+
This may cause a considerable amount GPU memory. For reduced GPU memory consumption,
|
|
82
|
+
source images maybe divided into smaller batches. However, this may slightly
|
|
83
|
+
increase the inference time. Defaults to None.
|
|
84
|
+
|
|
85
|
+
"""
|
|
86
|
+
# validate env name
|
|
87
|
+
env_name = str(env_name)
|
|
88
|
+
if not all((c.isalnum() or c in ("-_")) for c in env_name):
|
|
89
|
+
raise RuntimeError(f"env_name contains invalid characters: {env_name}")
|
|
90
|
+
|
|
91
|
+
haca3_command = ['haca3-test']
|
|
92
|
+
|
|
93
|
+
if isinstance(in_path, (str, os.PathLike)):
|
|
94
|
+
in_path = [in_path]
|
|
95
|
+
|
|
96
|
+
if target_theta is not None:
|
|
97
|
+
if not isinstance(target_theta, tuple):
|
|
98
|
+
raise RuntimeError(f"target_theta must be tuple of two float values or None, got {type(target_theta)}")
|
|
99
|
+
if not len(target_theta) == 2:
|
|
100
|
+
raise RuntimeError(f"target_theta must be a length 2 tuple, got length {len(target_theta)}")
|
|
101
|
+
if not all(isinstance(t, (int,float)) for t in target_theta):
|
|
102
|
+
raise RuntimeError(
|
|
103
|
+
f"target_theta must be tuple of two float values, got tuple({tuple(type(v) for v in target_theta)})")
|
|
104
|
+
|
|
105
|
+
for f in in_path:
|
|
106
|
+
haca3_command.append(f'--in-path "{os.path.normpath(f)}"',)
|
|
107
|
+
|
|
108
|
+
haca3_command.extend([
|
|
109
|
+
f'--out-path "{os.path.normpath(out_path)}"',
|
|
110
|
+
f'--harmonization-model "{os.path.normpath(harmonization_model)}"',
|
|
111
|
+
f'--fusion-model "{os.path.normpath(fusion_model)}"',
|
|
112
|
+
])
|
|
113
|
+
|
|
114
|
+
if target_image is not None: haca3_command.append(f'--target-image "{os.path.normpath(target_image)}"')
|
|
115
|
+
if target_theta is not None: haca3_command.append(f'--target-theta {float(target_theta[0])} {float(target_theta[1])}')
|
|
116
|
+
if norm_val is not None: haca3_command.append(f'--norm-val {float(norm_val)}')
|
|
117
|
+
if intermediate_out_dir is not None:
|
|
118
|
+
haca3_command.append('--save-intermediate')
|
|
119
|
+
haca3_command.append(f'--intermediate-out-dir "{os.path.normpath(intermediate_out_dir)}"')
|
|
120
|
+
if gpu_id is not None: haca3_command.append(f'--gpu-id {int(gpu_id)}')
|
|
121
|
+
if num_batches is not None: haca3_command.append(f'--num-batches {int(num_batches)}')
|
|
122
|
+
|
|
123
|
+
# print(f". {conda_path}/etc/profile.d/conda.sh && conda activate {env_name} && {' '.join(haca3_command)}")
|
|
124
|
+
|
|
125
|
+
# so
|
|
126
|
+
# shlex.split doesn't work on ., and conda run doesn't work for whatever reason, so we have to do this
|
|
127
|
+
# all args are explicitly validated so should be ok
|
|
128
|
+
command = f". {conda_path}/etc/profile.d/conda.sh && conda activate {env_name} && {' '.join(haca3_command)}"
|
|
129
|
+
|
|
130
|
+
# run
|
|
131
|
+
subprocess.run(command, shell=True, check=True)
|
|
132
|
+
|
|
133
|
+
def harmonize(
|
|
134
|
+
conda_path: str | os.PathLike,
|
|
135
|
+
env_name: str,
|
|
136
|
+
harmonization_model: str | os.PathLike,
|
|
137
|
+
fusion_model: str | os.PathLike,
|
|
138
|
+
inputs: "ImageLike | Sequence[ImageLike]",
|
|
139
|
+
target_image: "ImageLike | None" = None,
|
|
140
|
+
target_theta: tuple[float,float] | None = None,
|
|
141
|
+
norm_val: float | None = None,
|
|
142
|
+
intermediate_out_dir: str | os.PathLike | None = None,
|
|
143
|
+
gpu_id: int | None = None,
|
|
144
|
+
num_batches: int | None = None,
|
|
145
|
+
) -> sitk.Image:
|
|
146
|
+
"""Harmonizes ``inputs`` using HACA3.
|
|
147
|
+
|
|
148
|
+
Important: Some preprocessing steps are needed before running HACA3:
|
|
149
|
+
|
|
150
|
+
- Inhomogeneity correction
|
|
151
|
+
- Super-resolution for 2D acquired scans. This step is optional, but recommended for optimal performance. See [SMORE](https://github.com/volcanofly/SMORE-Super-resolution-for-3D-medical-images-MRI) for more details.
|
|
152
|
+
- Registration to MNI space (1mm isotropic resolution). HACA3 assumes a spatial dimension of 192x224x192.
|
|
153
|
+
|
|
154
|
+
You can inhomogeneity correction via ``mrid.n4_bias_field_correction(...)``,
|
|
155
|
+
register to MNI152 via ``mrid.simple_elastix.register(...)``, and center-pad to required dimension iva
|
|
156
|
+
``mrid.center_crop_or_pad(image, [192, 224, 192])``.
|
|
157
|
+
|
|
158
|
+
Args:
|
|
159
|
+
conda_path: path to ``minconda3`` directory.
|
|
160
|
+
env_name: name of the conda env where HACA3 is installed.
|
|
161
|
+
harmonization_model: pretrained HACA3 weights. Pretrained model weights on IXI, OASIS and HCP
|
|
162
|
+
data can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/harmonization_public.pt).
|
|
163
|
+
fusion_model: pretrained fusion model weights. HACA3 uses a 3D convolutional network to
|
|
164
|
+
combine multi-orientation 2D slices into a single 3D volume. Pretrained fusion model
|
|
165
|
+
can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/fusion.pt).
|
|
166
|
+
inputs: input source images. Multiple images may be provided if there are
|
|
167
|
+
multiple source images (different modalities). Note that modalities must be in
|
|
168
|
+
MNI space (1mm isotropic resolution). HACA3 assumes a spatial dimension of 192x224x192.
|
|
169
|
+
target_image: target image. HACA3 will match the contrast of
|
|
170
|
+
source images to this target image. One of ``target_image`` or ``target_theta`` must be set.
|
|
171
|
+
target_theta: In HACA3, ```theta``` is a two-dimensional representation of image contrast.
|
|
172
|
+
Target image contrast can be directly specified by providing a ```theta``` value, e.g.,
|
|
173
|
+
```target_theta = (0.5, 0.5)```. One of ``target_image`` or ``target_theta`` must be set.
|
|
174
|
+
Some example values: ``(10.0, 20.0)`` or ``(0.0, 30.0)`` represent T1w;
|
|
175
|
+
``(-18.0, -16.0)`` represents T2W; ``(0.0, 0.0)`` represents FLAIR.
|
|
176
|
+
norm_val: normalization value. Defaults to None.
|
|
177
|
+
intermediate_out_dir: directory to save intermediate results. Defaults to None.
|
|
178
|
+
gpu_id: integer number specifies which GPU to run HACA3. Defaults to None.
|
|
179
|
+
num_batches: During inference, HACA3 takes entire 3D MRI volumes as input.
|
|
180
|
+
This may cause a considerable amount GPU memory. For reduced GPU memory consumption,
|
|
181
|
+
source images maybe divided into smaller batches. However, this may slightly
|
|
182
|
+
increase the inference time. Defaults to None.
|
|
183
|
+
|
|
184
|
+
Raises:
|
|
185
|
+
RuntimeError: _description_
|
|
186
|
+
RuntimeError: _description_
|
|
187
|
+
|
|
188
|
+
Returns:
|
|
189
|
+
_description_
|
|
190
|
+
"""
|
|
191
|
+
|
|
192
|
+
if all(i is None for i in [target_image, target_theta]):
|
|
193
|
+
raise RuntimeError("Either target_image or target_theta must be set")
|
|
194
|
+
|
|
195
|
+
if all(i is not None for i in [target_image, target_theta]):
|
|
196
|
+
raise RuntimeError("Only one of target_image or target_theta must be set")
|
|
197
|
+
|
|
198
|
+
if isinstance(inputs, (str, os.PathLike)) or not isinstance(inputs, Sequence):
|
|
199
|
+
inputs = (inputs, )
|
|
200
|
+
|
|
201
|
+
inputs = [tositk(img) for img in inputs]
|
|
202
|
+
if target_image is not None: target_image = tositk(target_image)
|
|
203
|
+
|
|
204
|
+
# --------------------------------- run HACA3 -------------------------------- #
|
|
205
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
206
|
+
for i, img in enumerate(inputs):
|
|
207
|
+
if tuple(img.GetSize()) != (192, 224, 192):
|
|
208
|
+
raise RuntimeError(
|
|
209
|
+
"all inputs to HACA3 must be in MNI152 space and center-padded to size of ``[192, 224, 192]``. "
|
|
210
|
+
f"Got image {i} of size {img.GetSize()}")
|
|
211
|
+
|
|
212
|
+
sitk.WriteImage(img, os.path.join(tmpdir, f"input_{i}.nii.gz"))
|
|
213
|
+
|
|
214
|
+
if target_image is not None:
|
|
215
|
+
target_path = os.path.join(tmpdir, "target_image.nii.gz")
|
|
216
|
+
if tuple(target_image.GetSize()) != (192, 224, 192):
|
|
217
|
+
raise RuntimeError(
|
|
218
|
+
"all inputs to HACA3 must be in MNI152 space and center-padded to size of ``[192, 224, 192]``. "
|
|
219
|
+
f"Got target_image of size {target_image.GetSize()}")
|
|
220
|
+
|
|
221
|
+
sitk.WriteImage(target_image, target_path)
|
|
222
|
+
else:
|
|
223
|
+
target_path = None
|
|
224
|
+
|
|
225
|
+
run_HACA3(
|
|
226
|
+
conda_path=conda_path,
|
|
227
|
+
env_name=env_name,
|
|
228
|
+
harmonization_model=harmonization_model,
|
|
229
|
+
fusion_model=fusion_model,
|
|
230
|
+
in_path = [os.path.join(tmpdir, f"input_{i}.nii.gz") for i in range(len(inputs))],
|
|
231
|
+
out_path = os.path.join(tmpdir, "output.nii.gz"),
|
|
232
|
+
target_image = target_path,
|
|
233
|
+
target_theta = target_theta,
|
|
234
|
+
norm_val = norm_val,
|
|
235
|
+
intermediate_out_dir = intermediate_out_dir,
|
|
236
|
+
gpu_id = gpu_id,
|
|
237
|
+
num_batches = num_batches,
|
|
238
|
+
)
|
|
239
|
+
|
|
240
|
+
harmonized = tositk(os.path.join(tmpdir, "output_harmonized_fusion.nii.gz"))
|
|
241
|
+
|
|
242
|
+
return harmonized
|
|
@@ -0,0 +1,51 @@
|
|
|
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
|
+
|
|
10
|
+
def resample_to(input: ImageLike, to: ImageLike, interpolation=sitk.sitkNearestNeighbor) -> sitk.Image:
|
|
11
|
+
"""Resample ``input`` to ``reference``.
|
|
12
|
+
|
|
13
|
+
Resampling uses spatial information embedded in the sitk.Image - size, origin, spacing and direction.
|
|
14
|
+
|
|
15
|
+
Note that this information is only available when certain imaging formats are loaded, such as DICOM and NIfTI.
|
|
16
|
+
|
|
17
|
+
``input`` is transformed in such a way that those attributes will match ``reference``.
|
|
18
|
+
"""
|
|
19
|
+
return sitk.Resample(tositk(input), tositk(to), sitk.Transform(), interpolation)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def resize(img: ImageLike, new_size: Sequence[int], interpolator=sitk.sitkLinear) -> sitk.Image:
|
|
23
|
+
"""Resize ``sitk.Image`` to ``new_size`` (numpy axis order, i.e. reversed of sitk size).
|
|
24
|
+
Retains correct spatial information: origin and direction are preserved and spacing is
|
|
25
|
+
scaled so the physical field of view stays constant."""
|
|
26
|
+
img = tositk(img)
|
|
27
|
+
new_size = list(reversed(new_size))
|
|
28
|
+
|
|
29
|
+
old_size = img.GetSize()
|
|
30
|
+
old_spacing = img.GetSpacing()
|
|
31
|
+
new_spacing = []
|
|
32
|
+
for n, o, s in zip(new_size, old_size, old_spacing):
|
|
33
|
+
if n > 1:
|
|
34
|
+
new_spacing.append((o - 1) * s / (n - 1))
|
|
35
|
+
else:
|
|
36
|
+
new_spacing.append(s)
|
|
37
|
+
|
|
38
|
+
reference_image = sitk.Image(new_size, img.GetPixelID())
|
|
39
|
+
reference_image.SetOrigin(img.GetOrigin())
|
|
40
|
+
reference_image.SetSpacing(new_spacing)
|
|
41
|
+
reference_image.SetDirection(img.GetDirection())
|
|
42
|
+
|
|
43
|
+
return sitk.Resample(img, reference_image, sitk.Transform(), interpolator, 0.0)
|
|
44
|
+
|
|
45
|
+
def downsample(image:ImageLike, factor:float, dims: int | Sequence[int] | None, interpolator=sitk.sitkLinear) -> sitk.Image:
|
|
46
|
+
"""factor = 2 for 2x downsampling"""
|
|
47
|
+
if isinstance(dims, int): dims = (dims, )
|
|
48
|
+
image = tositk(image)
|
|
49
|
+
size = sitk.GetArrayFromImage(image).shape
|
|
50
|
+
size = [round(s/factor) if (dims is None or i in dims) else s for i,s in enumerate(size)]
|
|
51
|
+
return resize(image, size, interpolator=interpolator)
|
|
@@ -13,6 +13,7 @@ import SimpleITK as sitk
|
|
|
13
13
|
|
|
14
14
|
from . import preprocessing
|
|
15
15
|
from .loading.convert import ImageLike, tonumpy, tositk, totensor
|
|
16
|
+
from .utils.sitk_utils import sitk_apply_numpy
|
|
16
17
|
from .utils.torch_utils import CUDA_IF_AVAILABLE
|
|
17
18
|
|
|
18
19
|
if TYPE_CHECKING:
|
|
@@ -125,6 +126,26 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
125
126
|
|
|
126
127
|
return Study(**scans, **seg, **self.get_info())
|
|
127
128
|
|
|
129
|
+
def apply_numpy(self, fn:Callable[[np.ndarray], np.ndarray] | None, seg_fn: Callable[[np.ndarray], np.ndarray] | None) -> "Study":
|
|
130
|
+
"""Returns a new ``Study`` with ``fn`` applied to scans and ``seg_fn`` applied to segmentations.
|
|
131
|
+
|
|
132
|
+
Args:
|
|
133
|
+
fn: Function to apply to scan images. Must take and return ``sitk.Image``.
|
|
134
|
+
If None, identity function is used.
|
|
135
|
+
seg_fn: Function to apply to segmentation images. Must take and return ``sitk.Image``.
|
|
136
|
+
If None, identity function is used.
|
|
137
|
+
"""
|
|
138
|
+
scans = self.get_scans()
|
|
139
|
+
seg = self.get_segmentations()
|
|
140
|
+
|
|
141
|
+
if fn is not None:
|
|
142
|
+
scans = {k: sitk_apply_numpy(v, fn) for k,v in scans.items()}
|
|
143
|
+
|
|
144
|
+
if seg_fn is not None:
|
|
145
|
+
seg = {k: sitk_apply_numpy(v, seg_fn) for k,v in seg.items()}
|
|
146
|
+
|
|
147
|
+
return Study(**scans, **seg, **self.get_info())
|
|
148
|
+
|
|
128
149
|
def cast(self, dtype) -> "Study":
|
|
129
150
|
"""Returns a new study with all scans cast to the specified SimpleITK dtype.
|
|
130
151
|
|
|
@@ -176,6 +197,20 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
176
197
|
d = preprocessing.cropping.crop_bg_D(self.get_images(), key)
|
|
177
198
|
return Study(**d, **self.get_info())
|
|
178
199
|
|
|
200
|
+
def center_crop_or_pad(self, size: Sequence[int]):
|
|
201
|
+
"""Returns a new study with all images cropped or padded from the center to ``size``.
|
|
202
|
+
|
|
203
|
+
Args:
|
|
204
|
+
size: target image shape in pixels.
|
|
205
|
+
"""
|
|
206
|
+
shapes = {k: tuple(img.GetSize()) for k, img in self.get_images().items()}
|
|
207
|
+
if len(set(shapes.values())) > 1:
|
|
208
|
+
raise RuntimeError(f"center_crop_or_pad can only be applied to a Study where all images have the same shape. "
|
|
209
|
+
f"Current shapes: {shapes}")
|
|
210
|
+
|
|
211
|
+
fn = partial(preprocessing.center_crop_or_pad, size=size)
|
|
212
|
+
return self.apply(fn=fn, seg_fn=fn)
|
|
213
|
+
|
|
179
214
|
def skullstrip_hd_bet(
|
|
180
215
|
self,
|
|
181
216
|
key: str,
|
|
@@ -213,7 +248,7 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
213
248
|
Positive values expand brain mask by this many pixels, meaning inner parts of the skull will be included;
|
|
214
249
|
Negative values dilate brain mask by this many pixels, meaning outer parts of the brain will be excluded.
|
|
215
250
|
include_mask (bool, optional):
|
|
216
|
-
if True, adds ``"seg_hd_bet"`` with brain mask predicted by HD-BET to returned
|
|
251
|
+
if True, adds ``"seg_hd_bet"`` with brain mask predicted by HD-BET to returned study.
|
|
217
252
|
This adds brain mask BEFORE expanding/dilating if ``expand`` argument is specified.
|
|
218
253
|
keep_original (bool, Optional):
|
|
219
254
|
if True, skull-stripped images are added to the returned study with ``"_hd_bet"`` postfix,
|
|
@@ -279,6 +314,85 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
279
314
|
)
|
|
280
315
|
return Study(**d, **self.get_segmentations(), **self.get_info())
|
|
281
316
|
|
|
317
|
+
def harmonize_haca3(
|
|
318
|
+
self,
|
|
319
|
+
conda_path: str | os.PathLike,
|
|
320
|
+
env_name: str,
|
|
321
|
+
harmonization_model: str | os.PathLike,
|
|
322
|
+
fusion_model: str | os.PathLike,
|
|
323
|
+
keys: "str | Sequence[str]",
|
|
324
|
+
target_image: "ImageLike | None" = None,
|
|
325
|
+
target_theta: tuple[float,float] | None = None,
|
|
326
|
+
norm_val: float | None = None,
|
|
327
|
+
intermediate_out_dir: str | os.PathLike | None = None,
|
|
328
|
+
gpu_id: int | None = None,
|
|
329
|
+
num_batches: int | None = None,
|
|
330
|
+
new_key="harmonized_haca3",
|
|
331
|
+
):
|
|
332
|
+
"""Returns a new study, where a harmonized image is added under key ``new_key``
|
|
333
|
+
(by default ``"harmonized_haca3"``).
|
|
334
|
+
|
|
335
|
+
Important: Some preprocessing steps are needed before running HACA3:
|
|
336
|
+
|
|
337
|
+
- Inhomogeneity correction
|
|
338
|
+
- Super-resolution for 2D acquired scans. This step is optional, but recommended for optimal performance. See [SMORE](https://github.com/volcanofly/SMORE-Super-resolution-for-3D-medical-images-MRI) for more details.
|
|
339
|
+
- Registration to MNI space (1mm isotropic resolution). HACA3 assumes a spatial dimension of 192x224x192.
|
|
340
|
+
|
|
341
|
+
For example, you can do the following:
|
|
342
|
+
```python
|
|
343
|
+
study = study.register_SE("t1", mrid.atlas.get_mni152("2006 T1w symmetric"))
|
|
344
|
+
study = study.n4_bias_field_correction("t1")
|
|
345
|
+
study = study.center_crop_or_pad([192, 224, 192])
|
|
346
|
+
study = study.harmonize_haca3(...)
|
|
347
|
+
```
|
|
348
|
+
|
|
349
|
+
Args:
|
|
350
|
+
conda_path: path to ``minconda3`` directory.
|
|
351
|
+
env_name: name of the conda env where HACA3 is installed.
|
|
352
|
+
harmonization_model: pretrained HACA3 weights. Pretrained model weights on IXI, OASIS and HCP
|
|
353
|
+
data can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/harmonization_public.pt).
|
|
354
|
+
fusion_model: pretrained fusion model weights. HACA3 uses a 3D convolutional network to
|
|
355
|
+
combine multi-orientation 2D slices into a single 3D volume. Pretrained fusion model
|
|
356
|
+
can be downloaded [here](https://iacl.ece.jhu.edu/~lianrui/haca3/fusion.pt).
|
|
357
|
+
keys: key of the input source image. Multiple key may be provided if there are
|
|
358
|
+
multiple source images (different modalities). Note that modalities must be in
|
|
359
|
+
MNI space (1mm isotropic resolution). HACA3 assumes a spatial dimension of 192x224x192.
|
|
360
|
+
target_image: target image. HACA3 will match the contrast of
|
|
361
|
+
source images to this target image. One of ``target_image`` or ``target_theta`` must be set.
|
|
362
|
+
target_theta: In HACA3, ``theta`` is a two-dimensional representation of image contrast.
|
|
363
|
+
Target image contrast can be directly specified by providing a ``theta`` value, e.g.,
|
|
364
|
+
``target_theta = (0.5, 0.5)``. One of ``target_image`` or ``target_theta`` must be set.
|
|
365
|
+
Some example values: ``(10.0, 20.0)`` or ``(0.0, 30.0)`` represent T1w;
|
|
366
|
+
``(-18.0, -16.0)`` represents T2W; ``(0.0, 0.0)`` represents FLAIR.
|
|
367
|
+
norm_val: normalization value. Defaults to None.
|
|
368
|
+
intermediate_out_dir: directory to save intermediate results. Defaults to None.
|
|
369
|
+
gpu_id: integer number specifies which GPU to run HACA3. Defaults to None.
|
|
370
|
+
num_batches: During inference, HACA3 takes entire 3D MRI volumes as input.
|
|
371
|
+
This may cause a considerable amount GPU memory. For reduced GPU memory consumption,
|
|
372
|
+
source images maybe divided into smaller batches. However, this may slightly
|
|
373
|
+
increase the inference time. Defaults to None.
|
|
374
|
+
new_key: name of the harmonized image in retuned study. Defaults to "harmonized_haca3".
|
|
375
|
+
"""
|
|
376
|
+
if isinstance(keys, str): keys = (keys, )
|
|
377
|
+
images = [(k, self[k]) for k in keys]
|
|
378
|
+
|
|
379
|
+
harmonized = preprocessing.haca3.harmonize(
|
|
380
|
+
conda_path=conda_path,
|
|
381
|
+
env_name=env_name,
|
|
382
|
+
harmonization_model=harmonization_model,
|
|
383
|
+
fusion_model=fusion_model,
|
|
384
|
+
inputs=[img for k,img in images],
|
|
385
|
+
target_image=target_image,
|
|
386
|
+
target_theta=target_theta,
|
|
387
|
+
norm_val=norm_val,
|
|
388
|
+
intermediate_out_dir=intermediate_out_dir,
|
|
389
|
+
gpu_id=gpu_id,
|
|
390
|
+
num_batches=num_batches,
|
|
391
|
+
)
|
|
392
|
+
|
|
393
|
+
return self.add(new_key, harmonized)
|
|
394
|
+
|
|
395
|
+
|
|
282
396
|
def resize(self, size: Sequence[int], interpolator=sitk.sitkLinear):
|
|
283
397
|
"""Returns a new study with all images resized to to ``size``.
|
|
284
398
|
|
|
@@ -428,7 +542,7 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
428
542
|
if specified, processed mask is added to returned study with specified postfix rather than
|
|
429
543
|
replacing current ``key``.
|
|
430
544
|
"""
|
|
431
|
-
arr = self.
|
|
545
|
+
arr = self.to_numpy(key)
|
|
432
546
|
|
|
433
547
|
if independent_channels:
|
|
434
548
|
from monai.transforms import remove_small_objects
|
|
@@ -473,14 +587,14 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
473
587
|
from monai.transforms import KeepLargestConnectedComponent
|
|
474
588
|
tfm = KeepLargestConnectedComponent(applied_labels=applied_labels, is_onehot=False, independent=independent,
|
|
475
589
|
connectivity=connectivity, num_components=num_components)
|
|
476
|
-
arr = tfm(self.
|
|
590
|
+
arr = tfm(self.to_numpy(key))
|
|
477
591
|
return self.add(f'{key}{postfix}', arr, reference_key=key)
|
|
478
592
|
|
|
479
|
-
def
|
|
593
|
+
def to_numpy(self, key: str):
|
|
480
594
|
"""returns ``study[key]`` converted to a numpy array."""
|
|
481
595
|
return tonumpy(self[key])
|
|
482
596
|
|
|
483
|
-
def
|
|
597
|
+
def to_tensor(self, key: str):
|
|
484
598
|
"""returns ``study[key]`` converted to a tensor."""
|
|
485
599
|
return totensor(self[key])
|
|
486
600
|
|
|
@@ -531,18 +645,18 @@ class Study(UserDict[str, sitk.Image | Any]):
|
|
|
531
645
|
stacked = torch.stack([torch.from_numpy(sitk.GetArrayFromImage(v)) for _,v in items])
|
|
532
646
|
return stacked.to(device=device, dtype=dtype, memory_format=torch.contiguous_format)
|
|
533
647
|
|
|
534
|
-
def
|
|
648
|
+
def to_numpy_dict(self) -> dict[str, np.ndarray | Any]:
|
|
535
649
|
"""Returns a dictionary with all images converted to numpy arrays, info is included as is."""
|
|
536
650
|
return {k: (sitk.GetArrayFromImage(v) if isinstance(v, sitk.Image) else v) for k, v in self.items()}
|
|
537
651
|
|
|
538
|
-
def
|
|
652
|
+
def to_tensor_dict(self) -> "dict[str, torch.Tensor | Any]":
|
|
539
653
|
"""Returns a dictionary with all images converted to tensors, info is included as is."""
|
|
540
654
|
import torch
|
|
541
|
-
return {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v) for k,v in self.
|
|
655
|
+
return {k: (torch.from_numpy(v) if isinstance(v, np.ndarray) else v) for k,v in self.to_numpy_dict()}
|
|
542
656
|
|
|
543
657
|
def plot(self):
|
|
544
658
|
from .utils.plotting import plot_study
|
|
545
|
-
return plot_study(self.get_images().
|
|
659
|
+
return plot_study(self.get_images().to_numpy_dict())
|
|
546
660
|
|
|
547
661
|
def save(
|
|
548
662
|
self,
|
|
File without changes
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
import random
|
|
2
|
+
from collections.abc import Sequence, Callable
|
|
3
|
+
from typing import Any, Literal, TypeVar
|
|
4
|
+
|
|
5
|
+
ARR = TypeVar("ARR", bound=Any)
|
|
6
|
+
|
|
7
|
+
def crop(
|
|
8
|
+
arr: ARR,
|
|
9
|
+
reduction: Sequence[int],
|
|
10
|
+
where: Literal["start", "end", "center", "random"] = "center",
|
|
11
|
+
) -> ARR:
|
|
12
|
+
"""Crop ``arr`` such that ``output.shape[i] = input.shape[i] - reduction[i]``"""
|
|
13
|
+
|
|
14
|
+
shape = arr.shape[-len(reduction):]
|
|
15
|
+
slices = []
|
|
16
|
+
|
|
17
|
+
for r, sh in zip(reduction, shape):
|
|
18
|
+
if r == 0:
|
|
19
|
+
slices.append(slice(None))
|
|
20
|
+
continue
|
|
21
|
+
|
|
22
|
+
if r < 0: raise ValueError(f"Reduction {r} cannot be negative")
|
|
23
|
+
if r > sh: raise ValueError(f"Reduction {r} exceeds dimension size {sh}")
|
|
24
|
+
|
|
25
|
+
if where == 'start': start, end = 0, sh - r
|
|
26
|
+
elif where == 'end': start, end = r, sh
|
|
27
|
+
elif where == 'center':
|
|
28
|
+
start = r // 2
|
|
29
|
+
end = start + (sh - r)
|
|
30
|
+
elif where == 'random':
|
|
31
|
+
start = random.randint(0, r)
|
|
32
|
+
end = start + (sh - r)
|
|
33
|
+
else:
|
|
34
|
+
raise ValueError(f"Invalid where: {where}")
|
|
35
|
+
|
|
36
|
+
slices.append(slice(start, end))
|
|
37
|
+
|
|
38
|
+
# apply with broadcasting
|
|
39
|
+
return arr[(..., *slices)]
|
|
40
|
+
|
|
41
|
+
def crop_to_shape(
|
|
42
|
+
input: ARR,
|
|
43
|
+
shape: Sequence[int],
|
|
44
|
+
where: Literal["start", "end", "center", "random"] = "center",
|
|
45
|
+
) -> ARR:
|
|
46
|
+
"""Crop ``input`` to ``shape``."""
|
|
47
|
+
|
|
48
|
+
# broadcast
|
|
49
|
+
if len(shape) < input.ndim:
|
|
50
|
+
shape = list(input.shape[:input.ndim - len(shape)]) + list(shape)
|
|
51
|
+
|
|
52
|
+
return crop(input, [i - j for i, j in zip(input.shape, shape)], where=where)
|
|
@@ -0,0 +1,97 @@
|
|
|
1
|
+
import random
|
|
2
|
+
from collections.abc import Sequence, Callable
|
|
3
|
+
from typing import Any, Literal, TypeVar
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import torch
|
|
7
|
+
from . import cropping
|
|
8
|
+
from ..utils.python_utils import reduce_dim
|
|
9
|
+
|
|
10
|
+
ARR = TypeVar("ARR", bound=Any)
|
|
11
|
+
|
|
12
|
+
@torch.no_grad
|
|
13
|
+
def pad(
|
|
14
|
+
input: ARR,
|
|
15
|
+
padding: Sequence[int],
|
|
16
|
+
mode: str = "constant",
|
|
17
|
+
value = None,
|
|
18
|
+
where: Literal["center", "start", "end"] = "center",
|
|
19
|
+
crop: bool = False,
|
|
20
|
+
) -> ARR:
|
|
21
|
+
"""
|
|
22
|
+
TODO REFACTOR
|
|
23
|
+
|
|
24
|
+
``output.shape[i] = input.shape[i] + padding[i]``.
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
input (torch.Tensor): input to pad.
|
|
28
|
+
padding (str): how much padding to add per each dimension of ``input``.
|
|
29
|
+
mode (str, optional):
|
|
30
|
+
padding mode (https://pytorch.org/docs/stable/generated/torch.nn.functional.pad.html).
|
|
31
|
+
Defaults to 'constant'.
|
|
32
|
+
value (_type_, optional): padding constant value. Defaults to None.
|
|
33
|
+
where (str, optional): where to pad.
|
|
34
|
+
if ``center``, will pad start and end of each dimension evenly,
|
|
35
|
+
if ``start``, will pad at the start of each dimension,
|
|
36
|
+
if ``end``, will pad at the end. Defaults to 'center'.
|
|
37
|
+
crop (bool, optional): allow cropping if padding is negative. Defaults to False.
|
|
38
|
+
|
|
39
|
+
Returns:
|
|
40
|
+
torch.Tensor: Padded `input`.
|
|
41
|
+
"""
|
|
42
|
+
pad_values = [i if i > 0 else 0 for i in padding]
|
|
43
|
+
|
|
44
|
+
if sum(pad_values) > 0:
|
|
45
|
+
|
|
46
|
+
# create padding sequence for torch.nn.functional.pad
|
|
47
|
+
if where == 'center':
|
|
48
|
+
dims_padding = [(int(i / 2), int(i / 2)) if i % 2 == 0 else (int(i / 2), int(i / 2) + 1) for i in padding]
|
|
49
|
+
elif where == 'start':
|
|
50
|
+
dims_padding = [(i, 0) for i in padding]
|
|
51
|
+
elif where == 'end':
|
|
52
|
+
dims_padding = [(0, i) for i in padding]
|
|
53
|
+
else: raise ValueError(f'Invalid where: {where}')
|
|
54
|
+
|
|
55
|
+
# broadcasting (e.g. if padding 3×128×128 by [4, 4], it will pad by [0, 4, 4])
|
|
56
|
+
if len(dims_padding) < input.ndim:
|
|
57
|
+
dims_padding = [(0, 0)] * (input.ndim - len(dims_padding)) + dims_padding
|
|
58
|
+
|
|
59
|
+
if mode == 'zeros': mode = 'constant'; value = 0
|
|
60
|
+
elif mode == 'min': mode = 'constant'; value = float(input.min())
|
|
61
|
+
elif mode == 'max': mode = 'constant'; value = float(input.max())
|
|
62
|
+
elif mode == 'mean': mode = 'constant'; value = float(input.mean())
|
|
63
|
+
|
|
64
|
+
if isinstance(input, np.ndarray):
|
|
65
|
+
if mode == 'constant': kwargs = {"constant_values": value}
|
|
66
|
+
else: kwargs = {}
|
|
67
|
+
input = np.pad(input, pad_width = dims_padding, mode = mode, **kwargs) # type:ignore
|
|
68
|
+
else:
|
|
69
|
+
input = torch.nn.functional.pad(input, reduce_dim(reversed(dims_padding)), mode=mode, value=value) # type:ignore
|
|
70
|
+
|
|
71
|
+
if crop:
|
|
72
|
+
crop_values = [-i if i < 0 else 0 for i in padding]
|
|
73
|
+
if sum(crop_values) > 0:
|
|
74
|
+
input = cropping.crop(input, crop_values, where=where)
|
|
75
|
+
|
|
76
|
+
return input
|
|
77
|
+
|
|
78
|
+
def pad_to_shape(
|
|
79
|
+
input:ARR,
|
|
80
|
+
shape:Sequence[int],
|
|
81
|
+
mode:str = "constant",
|
|
82
|
+
value=None,
|
|
83
|
+
where:Literal["center", "start", "end"] = "center",
|
|
84
|
+
crop = False,
|
|
85
|
+
) -> ARR:
|
|
86
|
+
# broadcasting
|
|
87
|
+
if len(shape) < input.ndim:
|
|
88
|
+
shape = list(input.shape[:input.ndim - len(shape)]) + list(shape)
|
|
89
|
+
|
|
90
|
+
return pad(
|
|
91
|
+
input=input,
|
|
92
|
+
padding=[shape[i] - input.shape[i] for i in range(input.ndim)],
|
|
93
|
+
mode=mode,
|
|
94
|
+
value=value,
|
|
95
|
+
where=where,
|
|
96
|
+
crop = crop,
|
|
97
|
+
)
|
|
@@ -43,7 +43,7 @@ class SliceSampler:
|
|
|
43
43
|
raise RuntimeError(f"Segmentation must have a shape of of (D, H, W), got {segmentation.shape}")
|
|
44
44
|
if segmentation.is_floating_point():
|
|
45
45
|
raise RuntimeError(f"Segmentation must have integer data type, got {segmentation.dtype}")
|
|
46
|
-
if
|
|
46
|
+
if data.shape[1:] != segmentation.shape:
|
|
47
47
|
raise RuntimeError(f"Shapes of scans and segmentation do not match: {data.shape = }, {segmentation.shape = }")
|
|
48
48
|
if segmentation.min() < 0:
|
|
49
49
|
raise RuntimeError(f"Segmentation background must have a value of 0, got {segmentation.min() = }")
|
|
@@ -193,6 +193,11 @@ class SliceSampler:
|
|
|
193
193
|
|
|
194
194
|
return get_sample
|
|
195
195
|
|
|
196
|
+
def get_all_dim_slices(self, dim: Literal[0,1,2], around: int, flatten: bool = True):
|
|
197
|
+
return [self.get_slice(
|
|
198
|
+
dim=dim, coord=coord, around=around, flatten=flatten
|
|
199
|
+
) for coord in range(around, self.data.shape[dim+1] - around)]
|
|
200
|
+
|
|
196
201
|
|
|
197
202
|
class SliceDataset(torch.utils.data.Dataset):
|
|
198
203
|
"""A dataset of SliceSamplers.
|
|
@@ -2,58 +2,9 @@ import random
|
|
|
2
2
|
from collections.abc import Sequence, Callable
|
|
3
3
|
from typing import Any, Literal, TypeVar
|
|
4
4
|
|
|
5
|
+
import numpy as np
|
|
5
6
|
import torch
|
|
6
7
|
|
|
7
|
-
ARR = TypeVar("ARR", bound=Any)
|
|
8
|
-
|
|
9
|
-
def crop(
|
|
10
|
-
arr: ARR,
|
|
11
|
-
reduction: Sequence[int],
|
|
12
|
-
where: Literal["start", "end", "center", "random"] = "center",
|
|
13
|
-
) -> ARR:
|
|
14
|
-
"""Crop ``arr`` such that ``output.shape[i] = input.shape[i] - reduction[i]``"""
|
|
15
|
-
|
|
16
|
-
shape = arr.shape[-len(reduction):]
|
|
17
|
-
slices = []
|
|
18
|
-
|
|
19
|
-
for r, sh in zip(reduction, shape):
|
|
20
|
-
if r == 0:
|
|
21
|
-
slices.append(slice(None))
|
|
22
|
-
continue
|
|
23
|
-
|
|
24
|
-
if r < 0: raise ValueError("Reduction cannot be negative")
|
|
25
|
-
if r > sh: raise ValueError(f"Reduction {r} exceeds dimension size {sh}")
|
|
26
|
-
|
|
27
|
-
if where == 'start': start, end = 0, sh - r
|
|
28
|
-
elif where == 'end': start, end = r, sh
|
|
29
|
-
elif where == 'center':
|
|
30
|
-
start = r // 2
|
|
31
|
-
end = start + (sh - r)
|
|
32
|
-
elif where == 'random':
|
|
33
|
-
start = random.randint(0, r)
|
|
34
|
-
end = start + (sh - r)
|
|
35
|
-
else:
|
|
36
|
-
raise ValueError(f"Invalid where: {where}")
|
|
37
|
-
|
|
38
|
-
slices.append(slice(start, end))
|
|
39
|
-
|
|
40
|
-
# apply with broadcasting
|
|
41
|
-
return arr[(..., *slices)]
|
|
42
|
-
|
|
43
|
-
def crop_to_shape(
|
|
44
|
-
input: ARR,
|
|
45
|
-
shape: Sequence[int],
|
|
46
|
-
where: Literal["start", "end", "center", "random"] = "center",
|
|
47
|
-
) -> ARR:
|
|
48
|
-
"""Crop ``input`` to ``shape``."""
|
|
49
|
-
|
|
50
|
-
# broadcast
|
|
51
|
-
if len(shape) < input.ndim:
|
|
52
|
-
shape = list(input.shape[:input.ndim - len(shape)]) + list(shape)
|
|
53
|
-
|
|
54
|
-
return crop(input, [i - j for i, j in zip(input.shape, shape)], where=where)
|
|
55
|
-
|
|
56
|
-
|
|
57
8
|
def shuffle_channels(x:torch.Tensor):
|
|
58
9
|
"""Shuffle first axis in a ``(C, *)`` tensor"""
|
|
59
10
|
return x[torch.randperm(x.shape[0])]
|
|
@@ -47,10 +47,12 @@ def run_dcm2niix(
|
|
|
47
47
|
|
|
48
48
|
# create temporary folders
|
|
49
49
|
tmp_input_dir = os.path.join(tmpdir, "mrid_dcm2niix_input")
|
|
50
|
+
if os.path.exists(tmp_input_dir): shutil.rmtree(tmp_input_dir)
|
|
51
|
+
shutil.copytree(inpath, tmp_input_dir)
|
|
52
|
+
|
|
50
53
|
tmp_output_dir = os.path.join(tmpdir, "mrid_dcm2niix_output")
|
|
51
54
|
if os.path.exists(tmp_output_dir): shutil.rmtree(tmp_output_dir)
|
|
52
55
|
os.mkdir(tmp_output_dir)
|
|
53
|
-
shutil.copytree(inpath, tmp_input_dir)
|
|
54
56
|
|
|
55
57
|
# create output dir if it doesn't exist
|
|
56
58
|
if outfolder != '':
|
|
@@ -72,7 +74,7 @@ def run_dcm2niix(
|
|
|
72
74
|
out_files = [i for i in os.listdir(tmp_output_dir) if i.lower().strip().endswith('.nii.gz')]
|
|
73
75
|
|
|
74
76
|
# move them to output folder
|
|
75
|
-
shutil.copytree(tmp_output_dir, outfolder)
|
|
77
|
+
shutil.copytree(tmp_output_dir, outfolder, dirs_exist_ok=True)
|
|
76
78
|
|
|
77
79
|
if len(out_files) > 1:
|
|
78
80
|
warnings.warn(f"More than one NIfTI file was created in {outfolder}, path to the first one will be returned. Something may be wrong.")
|
|
@@ -1,12 +1,16 @@
|
|
|
1
1
|
from collections.abc import Mapping
|
|
2
2
|
|
|
3
3
|
import numpy as np
|
|
4
|
+
|
|
4
5
|
from ..loading import ImageLike, tonumpy
|
|
5
6
|
|
|
6
|
-
def plot_study(data: Mapping[str, ImageLike]):
|
|
7
|
+
def plot_study(data: "ImageLike | Mapping[str, ImageLike]"):
|
|
7
8
|
import matplotlib.gridspec as gridspec
|
|
8
9
|
import matplotlib.pyplot as plt
|
|
9
10
|
|
|
11
|
+
if not isinstance(data, Mapping):
|
|
12
|
+
data = {"image": data}
|
|
13
|
+
|
|
10
14
|
data = {k: tonumpy(v) for k,v in data.items()}
|
|
11
15
|
n_vals = len(data)
|
|
12
16
|
|
|
@@ -78,29 +82,28 @@ def plot_study(data: Mapping[str, ImageLike]):
|
|
|
78
82
|
fig.text(box.x0 + box.width/2, box.y1 + 0.04, modality_name,
|
|
79
83
|
ha='center', va='bottom', fontsize=14, fontweight='bold')
|
|
80
84
|
|
|
81
|
-
# plt.show()
|
|
82
85
|
return fig
|
|
83
86
|
|
|
84
|
-
if __name__ == "__main__":
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
|
|
94
|
-
|
|
95
|
-
|
|
96
|
-
|
|
97
|
-
|
|
98
|
-
|
|
99
|
-
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
106
|
-
|
|
87
|
+
# if __name__ == "__main__":
|
|
88
|
+
|
|
89
|
+
# def dummy_data(shape, type='sphere'):
|
|
90
|
+
# x, y, z = np.indices(shape)
|
|
91
|
+
# cx, cy, cz = shape[0]//2, shape[1]//2, shape[2]//2
|
|
92
|
+
# mask = np.zeros(shape)
|
|
93
|
+
# if type == 'sphere':
|
|
94
|
+
# r = min(shape)//3
|
|
95
|
+
# mask[(x-cx)**2 + (y-cy)**2 + (z-cz)**2 < r**2] = 1
|
|
96
|
+
# elif type == 'cube':
|
|
97
|
+
# r = min(shape)//4
|
|
98
|
+
# mask[cx-r:cx+r, cy-r:cy+r, cz-r:cz+r] = 1
|
|
99
|
+
# elif type == 'noise':
|
|
100
|
+
# mask = np.random.rand(*shape)
|
|
101
|
+
|
|
102
|
+
# return mask
|
|
103
|
+
|
|
104
|
+
# plot_study({
|
|
105
|
+
# "Sphere": dummy_data((60, 60, 60), 'sphere'),
|
|
106
|
+
# "Cube": dummy_data((60, 60, 60), 'cube'),
|
|
107
|
+
# "Noise": dummy_data((60, 60, 60), 'noise'),
|
|
108
|
+
# "Another sphere": dummy_data((60, 60, 60), 'sphere'),
|
|
109
|
+
# })
|
|
@@ -1,7 +1,7 @@
|
|
|
1
1
|
import importlib.util
|
|
2
|
-
from collections.abc import Mapping, Sequence
|
|
3
|
-
from typing import Any
|
|
4
|
-
|
|
2
|
+
from collections.abc import Mapping, Sequence, Iterable
|
|
3
|
+
from typing import Any, TypeVar
|
|
4
|
+
import functools, operator
|
|
5
5
|
|
|
6
6
|
# lazy loader from https://stackoverflow.com/a/78312674/15673832
|
|
7
7
|
class LazyLoader:
|
|
@@ -30,19 +30,7 @@ class LazyLoader:
|
|
|
30
30
|
return getattr (self._mod, attr)
|
|
31
31
|
|
|
32
32
|
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
# self.keys = keys
|
|
38
|
-
|
|
39
|
-
# def pack(self, unpacked: Sequence) -> T:
|
|
40
|
-
# if self.keys is not None: return self.type(dict(zip(self.keys, unpacked)))
|
|
41
|
-
# return self.type(unpacked)
|
|
42
|
-
|
|
43
|
-
# def unpack_struct[T: Sequence | Mapping](struct: T) -> tuple[Any, _Packer[T]]:
|
|
44
|
-
# if isinstance(struct, Sequence):
|
|
45
|
-
# return list(struct), _Packer(type(struct))
|
|
46
|
-
# if isinstance(struct, Mapping):
|
|
47
|
-
# return list(struct.values()), _Packer(type(struct), list(struct.values()))
|
|
48
|
-
# raise TypeError(f"Transformation functions accept lists and dictionaries, but received {type(struct)}")
|
|
33
|
+
T = TypeVar("T")
|
|
34
|
+
def reduce_dim(x:Iterable[Iterable[T]]) -> list[T]:
|
|
35
|
+
"""Reduces one level of nesting. Takes an iterable of iterables of X, and returns an iterable of X."""
|
|
36
|
+
return functools.reduce(operator.iconcat, x, [])
|
|
@@ -0,0 +1,18 @@
|
|
|
1
|
+
from collections.abc import Callable
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import SimpleITK as sitk
|
|
5
|
+
|
|
6
|
+
from ..loading import ImageLike, tositk
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def sitk_apply_numpy(image: ImageLike, func: Callable[[np.ndarray], np.ndarray]) -> sitk.Image:
|
|
10
|
+
image = tositk(image)
|
|
11
|
+
array = sitk.GetArrayFromImage(image)
|
|
12
|
+
shape = array.shape
|
|
13
|
+
array = func(array)
|
|
14
|
+
if shape != array.shape:
|
|
15
|
+
raise RuntimeError(f"Function {func} changed array shape from {shape} to {array.shape}.")
|
|
16
|
+
res = sitk.GetImageFromArray(array)
|
|
17
|
+
res.CopyInformation(image)
|
|
18
|
+
return res
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: mrid-python
|
|
3
|
-
Version: 0.1.
|
|
3
|
+
Version: 0.1.7
|
|
4
4
|
Summary: Tools for working with 3D medical images and segmentations - registration, brain skull-stripping, etc.
|
|
5
5
|
Author-email: Ivan Nikishev <nkshv2@gmail.com>
|
|
6
6
|
Project-URL: Homepage, https://github.com/inikishev/mrid
|
|
@@ -5,18 +5,22 @@ mrid/study.py
|
|
|
5
5
|
mrid/atlas/__init__.py
|
|
6
6
|
mrid/atlas/MNI152/__init__.py
|
|
7
7
|
mrid/atlas/SRI24/__init__.py
|
|
8
|
+
mrid/inference/__init__.py
|
|
8
9
|
mrid/loading/__init__.py
|
|
9
10
|
mrid/loading/convert.py
|
|
10
11
|
mrid/preprocessing/CTseg.py
|
|
11
12
|
mrid/preprocessing/__init__.py
|
|
12
13
|
mrid/preprocessing/bias_field_correction.py
|
|
13
14
|
mrid/preprocessing/cropping.py
|
|
15
|
+
mrid/preprocessing/haca3.py
|
|
14
16
|
mrid/preprocessing/hd_bet.py
|
|
15
17
|
mrid/preprocessing/mask.py
|
|
16
18
|
mrid/preprocessing/simple_elastix.py
|
|
17
19
|
mrid/preprocessing/spatial.py
|
|
18
20
|
mrid/preprocessing/synthstrip.py
|
|
19
21
|
mrid/training/__init__.py
|
|
22
|
+
mrid/training/cropping.py
|
|
23
|
+
mrid/training/padding.py
|
|
20
24
|
mrid/training/slicer.py
|
|
21
25
|
mrid/training/transforms.py
|
|
22
26
|
mrid/utils/__init__.py
|
|
@@ -24,6 +28,7 @@ mrid/utils/dcm2niix.py
|
|
|
24
28
|
mrid/utils/dicom_uid_fixer.py
|
|
25
29
|
mrid/utils/plotting.py
|
|
26
30
|
mrid/utils/python_utils.py
|
|
31
|
+
mrid/utils/sitk_utils.py
|
|
27
32
|
mrid/utils/stl_utils.py
|
|
28
33
|
mrid/utils/torch_utils.py
|
|
29
34
|
mrid_python.egg-info/PKG-INFO
|
|
@@ -13,7 +13,7 @@ build-backend = "setuptools.build_meta"
|
|
|
13
13
|
name = "mrid-python"
|
|
14
14
|
description = "Tools for working with 3D medical images and segmentations - registration, brain skull-stripping, etc."
|
|
15
15
|
|
|
16
|
-
version = "0.1.
|
|
16
|
+
version = "0.1.7"
|
|
17
17
|
dependencies = [
|
|
18
18
|
"numpy",
|
|
19
19
|
"SimpleITK",
|
|
@@ -10,8 +10,8 @@ def test_resize():
|
|
|
10
10
|
study = Study(t1=np.random.rand(10, 20, 30).astype(np.float32))
|
|
11
11
|
resized = study.resize([5, 10, 15])
|
|
12
12
|
assert 't1' in resized
|
|
13
|
-
assert study.
|
|
14
|
-
assert resized.
|
|
13
|
+
assert study.to_numpy("t1").shape == (10, 20, 30)
|
|
14
|
+
assert resized.to_numpy("t1").shape == (5, 10, 15)
|
|
15
15
|
|
|
16
16
|
|
|
17
17
|
def test_downsample():
|
|
@@ -19,8 +19,8 @@ def test_downsample():
|
|
|
19
19
|
study = Study(t1=np.random.rand(10, 20, 30).astype(np.float32))
|
|
20
20
|
downsampled = study.downsample(factor=2.0) # 2x downsampling
|
|
21
21
|
assert 't1' in downsampled
|
|
22
|
-
assert study.
|
|
23
|
-
assert downsampled.
|
|
22
|
+
assert study.to_numpy("t1").shape == (10, 20, 30)
|
|
23
|
+
assert downsampled.to_numpy("t1").shape == (5, 10, 15)
|
|
24
24
|
|
|
25
25
|
|
|
26
26
|
def test_bias_field_correction():
|
|
@@ -40,4 +40,4 @@ def test_crop_bg():
|
|
|
40
40
|
assert 't1' in cropped
|
|
41
41
|
assert 't2' in cropped
|
|
42
42
|
|
|
43
|
-
assert study.
|
|
43
|
+
assert study.to_numpy("t1").shape == study.to_numpy("t2").shape
|
|
@@ -90,7 +90,7 @@ def test_numpy_method():
|
|
|
90
90
|
data = np.random.rand(10, 20, 30).astype(np.float32)
|
|
91
91
|
study = Study(t1=data)
|
|
92
92
|
|
|
93
|
-
numpy_array = study.
|
|
93
|
+
numpy_array = study.to_numpy('t1')
|
|
94
94
|
assert isinstance(numpy_array, np.ndarray)
|
|
95
95
|
assert numpy_array.shape == (10, 20, 30)
|
|
96
96
|
|
|
@@ -1,82 +0,0 @@
|
|
|
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
|
-
|
|
10
|
-
def resample_to(input: ImageLike, to: ImageLike, interpolation=sitk.sitkNearestNeighbor) -> sitk.Image:
|
|
11
|
-
"""Resample ``input`` to ``reference``.
|
|
12
|
-
|
|
13
|
-
Resampling uses spatial information embedded in the sitk.Image - size, origin, spacing and direction.
|
|
14
|
-
|
|
15
|
-
Note that this information is only available when certain imaging formats are loaded, such as DICOM and NIfTI.
|
|
16
|
-
|
|
17
|
-
``input`` is transformed in such a way that those attributes will match ``reference``.
|
|
18
|
-
"""
|
|
19
|
-
return sitk.Resample(tositk(input), tositk(to), sitk.Transform(), interpolation)
|
|
20
|
-
|
|
21
|
-
|
|
22
|
-
def resize(img: ImageLike, new_size: Sequence[int], interpolator=sitk.sitkLinear) -> sitk.Image:
|
|
23
|
-
"""Resize ``sitk.Image`` to ``new_size``. Retains correct spatial information.
|
|
24
|
-
source: https://gist.github.com/lixinqi98/1bbd3596492f20b776fed2778f7cd48c"""
|
|
25
|
-
img = tositk(img)
|
|
26
|
-
new_size = list(reversed(new_size))
|
|
27
|
-
|
|
28
|
-
# img = sitk.ReadImage(img)
|
|
29
|
-
dimension = img.GetDimension()
|
|
30
|
-
|
|
31
|
-
# Physical image size corresponds to the largest physical size in the training set, or any other arbitrary size.
|
|
32
|
-
reference_physical_size = np.zeros(dimension)
|
|
33
|
-
|
|
34
|
-
reference_physical_size[:] = [(sz - 1) * spc if sz * spc > mx else mx for sz, spc, mx in
|
|
35
|
-
zip(img.GetSize(), img.GetSpacing(), reference_physical_size)]
|
|
36
|
-
|
|
37
|
-
# Create the reference image with a zero origin, identity direction cosine matrix and dimension
|
|
38
|
-
reference_origin = np.zeros(dimension)
|
|
39
|
-
reference_direction = np.identity(dimension).flatten()
|
|
40
|
-
reference_size = new_size
|
|
41
|
-
reference_spacing = [phys_sz / (sz - 1) for sz, phys_sz in zip(reference_size, reference_physical_size)]
|
|
42
|
-
|
|
43
|
-
reference_image = sitk.Image(reference_size, img.GetPixelIDValue())
|
|
44
|
-
reference_image.SetOrigin(reference_origin)
|
|
45
|
-
reference_image.SetSpacing(reference_spacing)
|
|
46
|
-
reference_image.SetDirection(reference_direction)
|
|
47
|
-
|
|
48
|
-
# Always use the TransformContinuousIndexToPhysicalPoint to compute an indexed point's physical coordinates as
|
|
49
|
-
# this takes into account size, spacing and direction cosines. For the vast majority of images the direction
|
|
50
|
-
# cosines are the identity matrix, but when this isn't the case simply multiplying the central index by the
|
|
51
|
-
# spacing will not yield the correct coordinates resulting in a long debugging session.
|
|
52
|
-
reference_center = np.array(
|
|
53
|
-
reference_image.TransformContinuousIndexToPhysicalPoint(np.array(reference_image.GetSize()) / 2.0))
|
|
54
|
-
|
|
55
|
-
# Transform which maps from the reference_image to the current img with the translation mapping the image
|
|
56
|
-
# origins to each other.
|
|
57
|
-
transform = sitk.AffineTransform(dimension)
|
|
58
|
-
transform.SetMatrix(img.GetDirection())
|
|
59
|
-
transform.SetTranslation(np.array(img.GetOrigin()) - reference_origin)
|
|
60
|
-
# Modify the transformation to align the centers of the original and reference image instead of their origins.
|
|
61
|
-
centering_transform = sitk.TranslationTransform(dimension)
|
|
62
|
-
img_center = np.array(img.TransformContinuousIndexToPhysicalPoint(np.array(img.GetSize()) / 2.0))
|
|
63
|
-
centering_transform.SetOffset(np.array(transform.GetInverse().TransformPoint(img_center) - reference_center))
|
|
64
|
-
|
|
65
|
-
# centered_transform = sitk.Transform(transform)
|
|
66
|
-
# centered_transform.AddTransform(centering_transform)
|
|
67
|
-
|
|
68
|
-
centered_transform = sitk.CompositeTransform([transform, centering_transform])
|
|
69
|
-
|
|
70
|
-
# Using the linear interpolator as these are intensity images, if there is a need to resample a ground truth
|
|
71
|
-
# segmentation then the segmentation image should be resampled using the NearestNeighbor interpolator so that
|
|
72
|
-
# no new labels are introduced.
|
|
73
|
-
|
|
74
|
-
return sitk.Resample(img, reference_image, centered_transform, interpolator, 0.0)
|
|
75
|
-
|
|
76
|
-
def downsample(image:ImageLike, factor:float, dims: int | Sequence[int] | None, interpolator=sitk.sitkLinear) -> sitk.Image:
|
|
77
|
-
"""factor = 2 for 2x downsampling"""
|
|
78
|
-
if isinstance(dims, int): dims = (dims, )
|
|
79
|
-
image = tositk(image)
|
|
80
|
-
size = sitk.GetArrayFromImage(image).shape
|
|
81
|
-
size = [round(s/factor) if (dims is None or i in dims) else s for i,s in enumerate(size)]
|
|
82
|
-
return resize(image, size, interpolator=interpolator)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|