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.
Files changed (45) hide show
  1. {mrid_python-0.1.5 → mrid_python-0.1.7}/PKG-INFO +1 -1
  2. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/atlas/MNI152/__init__.py +1 -1
  3. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/CTseg.py +1 -1
  4. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/__init__.py +2 -2
  5. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/cropping.py +37 -1
  6. mrid_python-0.1.7/mrid/preprocessing/haca3.py +242 -0
  7. mrid_python-0.1.7/mrid/preprocessing/spatial.py +51 -0
  8. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/study.py +123 -9
  9. mrid_python-0.1.7/mrid/training/__init__.py +0 -0
  10. mrid_python-0.1.7/mrid/training/cropping.py +52 -0
  11. mrid_python-0.1.7/mrid/training/padding.py +97 -0
  12. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/training/slicer.py +6 -1
  13. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/training/transforms.py +1 -50
  14. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/dcm2niix.py +4 -2
  15. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/plotting.py +28 -25
  16. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/python_utils.py +7 -19
  17. mrid_python-0.1.7/mrid/utils/sitk_utils.py +18 -0
  18. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/PKG-INFO +1 -1
  19. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/SOURCES.txt +5 -0
  20. {mrid_python-0.1.5 → mrid_python-0.1.7}/pyproject.toml +1 -1
  21. {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_preprocessing.py +5 -5
  22. {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_study.py +1 -1
  23. mrid_python-0.1.5/mrid/preprocessing/spatial.py +0 -82
  24. {mrid_python-0.1.5 → mrid_python-0.1.7}/README.md +0 -0
  25. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/__init__.py +0 -0
  26. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/atlas/SRI24/__init__.py +0 -0
  27. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/atlas/__init__.py +0 -0
  28. {mrid_python-0.1.5/mrid/training → mrid_python-0.1.7/mrid/inference}/__init__.py +0 -0
  29. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/loading/__init__.py +0 -0
  30. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/loading/convert.py +0 -0
  31. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/bias_field_correction.py +0 -0
  32. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/hd_bet.py +0 -0
  33. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/mask.py +0 -0
  34. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/simple_elastix.py +0 -0
  35. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/preprocessing/synthstrip.py +0 -0
  36. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/__init__.py +0 -0
  37. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/dicom_uid_fixer.py +0 -0
  38. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/stl_utils.py +0 -0
  39. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid/utils/torch_utils.py +0 -0
  40. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/dependency_links.txt +0 -0
  41. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/requires.txt +0 -0
  42. {mrid_python-0.1.5 → mrid_python-0.1.7}/mrid_python.egg-info/top_level.txt +0 -0
  43. {mrid_python-0.1.5 → mrid_python-0.1.7}/setup.cfg +0 -0
  44. {mrid_python-0.1.5 → mrid_python-0.1.7}/tests/test_loading.py +0 -0
  45. {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.5
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"``
@@ -73,7 +73,7 @@ def run_CTseg(
73
73
  f"spm_CTseg('/data/{pth_ct.name}', '{dir_out}', true, true, true, true, 1.0)",
74
74
  ]
75
75
 
76
- # run dcm2niix
76
+ # run
77
77
  subprocess.run(command, check=True)
78
78
 
79
79
  # this creates
@@ -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 dictionary.
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.numpy(key)
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.numpy(key))
590
+ arr = tfm(self.to_numpy(key))
477
591
  return self.add(f'{key}{postfix}', arr, reference_key=key)
478
592
 
479
- def numpy(self, key: str):
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 tensor(self, key: str):
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 numpy_dict(self) -> dict[str, np.ndarray | Any]:
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 tensor_dict(self) -> "dict[str, torch.Tensor | Any]":
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.numpy_dict()}
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().numpy_dict())
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 segmentation.shape[1:] != data.shape:
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
- def dummy_data(shape, type='sphere'):
87
- x, y, z = np.indices(shape)
88
- cx, cy, cz = shape[0]//2, shape[1]//2, shape[2]//2
89
- mask = np.zeros(shape)
90
- if type == 'sphere':
91
- r = min(shape)//3
92
- mask[(x-cx)**2 + (y-cy)**2 + (z-cz)**2 < r**2] = 1
93
- elif type == 'cube':
94
- r = min(shape)//4
95
- mask[cx-r:cx+r, cy-r:cy+r, cz-r:cz+r] = 1
96
- elif type == 'noise':
97
- mask = np.random.rand(*shape)
98
-
99
- return mask
100
-
101
- plot_study({
102
- "Sphere": dummy_data((60, 60, 60), 'sphere'),
103
- "Cube": dummy_data((60, 60, 60), 'cube'),
104
- "Noise": dummy_data((60, 60, 60), 'noise'),
105
- "Another sphere": dummy_data((60, 60, 60), 'sphere'),
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
- # this allows transforms to support any kind of container
34
- # class _Packer[T: Sequence | Mapping]:
35
- # def __init__(self, type: type[T], keys: list | None = None):
36
- # self.type: Any = type
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.5
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.5"
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.numpy("t1").shape == (10, 20, 30)
14
- assert resized.numpy("t1").shape == (5, 10, 15)
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.numpy("t1").shape == (10, 20, 30)
23
- assert downsampled.numpy("t1").shape == (5, 10, 15)
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.numpy("t1").shape == study.numpy("t2").shape
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.numpy('t1')
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