mrid-python 0.1.4__py3-none-any.whl → 0.1.6__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -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"``
File without changes
@@ -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
@@ -235,7 +235,7 @@ def skullstrip_D(
235
235
 
236
236
  # optionally add with skullstripped postfix
237
237
  if keep_original:
238
- skullstripped = {f"{k}_hd_bet": v for k,v in skullstripped}
238
+ skullstripped = {(f"{k}_hd_bet" if k != "seg_hd_bet" else k): v for k,v in skullstripped.items()}
239
239
  skullstripped.update(images.copy())
240
240
 
241
241
  return skullstripped
@@ -1,6 +1,5 @@
1
1
  import SimpleITK as sitk
2
2
  import numpy as np
3
- import SimpleITK as sitk
4
3
 
5
4
  from ..loading.convert import ImageLike, tositk, tonumpy
6
5
 
@@ -15,12 +14,12 @@ def expand_binary_mask(binary_mask: ImageLike, expand: int) -> sitk.Image:
15
14
  Negative values dilate the mask by this many pixels.
16
15
  """
17
16
  binary_mask = tositk(binary_mask)
18
- if expand > 0:
17
+ if expand < 0:
19
18
  inverted_mask = 1 - binary_mask
20
- return 1 - sitk.BinaryDilate(inverted_mask, (expand, expand, expand))
19
+ return 1 - sitk.BinaryDilate(inverted_mask, (-expand, -expand, -expand))
21
20
 
22
- if expand < 0:
23
- return sitk.BinaryDilate(binary_mask, (-expand, -expand, -expand))
21
+ if expand > 0:
22
+ return sitk.BinaryDilate(binary_mask, (expand, expand, expand))
24
23
 
25
24
  return binary_mask
26
25
 
@@ -20,61 +20,31 @@ def resample_to(input: ImageLike, to: ImageLike, interpolation=sitk.sitkNearestN
20
20
 
21
21
 
22
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"""
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."""
25
26
  img = tositk(img)
26
27
  new_size = list(reversed(new_size))
27
28
 
28
- # img = sitk.ReadImage(img)
29
- dimension = img.GetDimension()
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)
30
37
 
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)
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())
33
42
 
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)]
43
+ return sitk.Resample(img, reference_image, sitk.Transform(), interpolator, 0.0)
36
44
 
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: Sequence[int] | None, interpolator=sitk.sitkLinear) -> sitk.Image:
45
+ def downsample(image:ImageLike, factor:float, dims: int | Sequence[int] | None, interpolator=sitk.sitkLinear) -> sitk.Image:
77
46
  """factor = 2 for 2x downsampling"""
47
+ if isinstance(dims, int): dims = (dims, )
78
48
  image = tositk(image)
79
49
  size = sitk.GetArrayFromImage(image).shape
80
50
  size = [round(s/factor) if (dims is None or i in dims) else s for i,s in enumerate(size)]
@@ -70,6 +70,7 @@ def run_synthstrip(
70
70
  fill: int | None = None,
71
71
  no_csf: bool | None = None,
72
72
  model: str | os.PathLike | None = None,
73
+ verbose: bool = True,
73
74
  ):
74
75
  """Runs ``synthstrip`` command-line routine via ``subprocess.run``.
75
76
 
@@ -108,7 +109,10 @@ def run_synthstrip(
108
109
  if model is not None: command.extend(["--model", os.path.normpath(model)])
109
110
 
110
111
  # run
111
- subprocess.run(command, check=True)
112
+ if verbose:
113
+ subprocess.run(command, check=True)
114
+ else:
115
+ subprocess.run(command, check=True, stdout=subprocess.DEVNULL, stderr=subprocess.STDOUT)
112
116
 
113
117
  def predict_brain_mask(
114
118
  synthstrip_script_path: str | os.PathLike,
@@ -117,6 +121,7 @@ def predict_brain_mask(
117
121
  border: int | None = None,
118
122
  threads: int | None = None,
119
123
  model: str | os.PathLike | None = None,
124
+ verbose: bool = True,
120
125
  ):
121
126
  """Returns brain mask of ``input`` predicted by ``synthstrip``.
122
127
 
@@ -141,6 +146,7 @@ def predict_brain_mask(
141
146
  border=border,
142
147
  threads=threads,
143
148
  model=model,
149
+ verbose=verbose,
144
150
  )
145
151
 
146
152
  brain_mask = tositk(os.path.join(tmpdir, "synthstrip_mask.nii.gz"))
@@ -155,6 +161,7 @@ def skullstrip(
155
161
  threads: int | None = None,
156
162
  model: str | os.PathLike | None = None,
157
163
  expand: int = 0,
164
+ verbose: bool = True,
158
165
  ):
159
166
  """Skullstrips ``input`` using synthstrip.
160
167
 
@@ -181,6 +188,7 @@ def skullstrip(
181
188
  border=border,
182
189
  threads=threads,
183
190
  model=model,
191
+ verbose=verbose
184
192
  )
185
193
  if expand != 0:
186
194
  mask = expand_binary_mask(mask, expand=expand)
@@ -202,6 +210,8 @@ def skullstrip_D(
202
210
  include_mask: bool = False,
203
211
  keep_original: bool = False,
204
212
 
213
+ verbose: bool = True,
214
+
205
215
  ) -> dict[str, sitk.Image]:
206
216
  """Predicts brain mask of ``images[key]`` using synthstrip, then uses this mask to skull strip all values in ``images``.
207
217
 
@@ -233,6 +243,7 @@ def skullstrip_D(
233
243
  border=border,
234
244
  threads=threads,
235
245
  model=model,
246
+ verbose=verbose,
236
247
  )
237
248
  skullstripped = {}
238
249
 
@@ -252,7 +263,7 @@ def skullstrip_D(
252
263
 
253
264
  # optionally add with skullstripped postfix
254
265
  if keep_original:
255
- skullstripped = {f"{k}_synthstrip": v for k,v in skullstripped}
266
+ skullstripped = {(f"{k}_synthstrip" if k != "seg_synthstrip" else k): v for k,v in skullstripped.items()}
256
267
  skullstripped.update(images.copy())
257
268
 
258
269
  return skullstripped