[ENH]: Add helper function for getting correct interpolator for masks #416
3 changed files with 84 additions and 21 deletions
|
|
@ -7,6 +7,7 @@ import uuid
|
||||||
from typing import TYPE_CHECKING, Any, Optional
|
from typing import TYPE_CHECKING, Any, Optional
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
from ...pipeline import WorkDirManager
|
from ...pipeline import WorkDirManager
|
||||||
from ...utils import logger, raise_error, run_ext_cmd
|
from ...utils import logger, raise_error, run_ext_cmd
|
||||||
|
|
@ -20,6 +21,26 @@ if TYPE_CHECKING:
|
||||||
__all__ = ["ANTsMaskWarper"]
|
__all__ = ["ANTsMaskWarper"]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_interpolation_method(img: "Nifti1Image") -> str:
|
||||||
|
"""Get correct interpolation method for `img`.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
img : nibabel.nifti1.Nifti1Image
|
||||||
|
The image.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
str
|
||||||
|
The interpolation method.
|
||||||
|
|
||||||
|
"""
|
||||||
|
if np.array_equal(np.unique(img.get_fdata()), [0, 1]):
|
||||||
|
return "'GenericLabel[NearestNeighbor]'"
|
||||||
|
else:
|
||||||
|
return "LanczosWindowedSinc"
|
||||||
|
|
||||||
|
|
||||||
class ANTsMaskWarper:
|
class ANTsMaskWarper:
|
||||||
"""Class for mask space warping via ANTs.
|
"""Class for mask space warping via ANTs.
|
||||||
|
|
||||||
|
|
@ -143,7 +164,7 @@ class ANTsMaskWarper:
|
||||||
"antsApplyTransforms",
|
"antsApplyTransforms",
|
||||||
"-d 3",
|
"-d 3",
|
||||||
"-e 3",
|
"-e 3",
|
||||||
"-n 'GenericLabel[NearestNeighbor]'",
|
f"-n {_get_interpolation_method(mask_img)}",
|
||||||
f"-i {prewarp_mask_path.resolve()}",
|
f"-i {prewarp_mask_path.resolve()}",
|
||||||
f"-r {template_space_img_path.resolve()}",
|
f"-r {template_space_img_path.resolve()}",
|
||||||
f"-t {xfm_file_path.resolve()}",
|
f"-t {xfm_file_path.resolve()}",
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@ import uuid
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
from ...pipeline import WorkDirManager
|
from ...pipeline import WorkDirManager
|
||||||
from ...utils import logger, run_ext_cmd
|
from ...utils import logger, run_ext_cmd
|
||||||
|
|
@ -19,6 +20,26 @@ if TYPE_CHECKING:
|
||||||
__all__ = ["FSLMaskWarper"]
|
__all__ = ["FSLMaskWarper"]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_interpolation_method(img: "Nifti1Image") -> str:
|
||||||
|
"""Get correct interpolation method for `img`.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
img : nibabel.nifti1.Nifti1Image
|
||||||
|
The image.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
str
|
||||||
|
The interpolation method.
|
||||||
|
|
||||||
|
"""
|
||||||
|
if np.array_equal(np.unique(img.get_fdata()), [0, 1]):
|
||||||
|
return "nn"
|
||||||
|
else:
|
||||||
|
return "spline"
|
||||||
|
|
||||||
|
|
||||||
class FSLMaskWarper:
|
class FSLMaskWarper:
|
||||||
"""Class for mask space warping via FSL FLIRT.
|
"""Class for mask space warping via FSL FLIRT.
|
||||||
|
|
||||||
|
|
@ -71,7 +92,7 @@ class FSLMaskWarper:
|
||||||
# Set applywarp command
|
# Set applywarp command
|
||||||
applywarp_cmd = [
|
applywarp_cmd = [
|
||||||
"applywarp",
|
"applywarp",
|
||||||
"--interp=nn",
|
f"--interp={_get_interpolation_method(mask_img)}",
|
||||||
f"-i {prewarp_mask_path.resolve()}",
|
f"-i {prewarp_mask_path.resolve()}",
|
||||||
# use resampled reference
|
# use resampled reference
|
||||||
f"-r {target_data['reference']['path'].resolve()}",
|
f"-r {target_data['reference']['path'].resolve()}",
|
||||||
|
|
|
||||||
|
|
@ -165,9 +165,28 @@ def compute_brain_mask(
|
||||||
)
|
)
|
||||||
|
|
||||||
mask_name = f"template_{target_std_space}_for_compute_brain_mask"
|
mask_name = f"template_{target_std_space}_for_compute_brain_mask"
|
||||||
|
# Resample and warp template if target space is native
|
||||||
# Warp template to correct space (MNI to MNI)
|
if target_data["space"] == "native" and template_space != "native":
|
||||||
if template_space != "native" and template_space != target_std_space:
|
if warp_data["warper"] == "fsl":
|
||||||
|
resampled_template = FSLMaskWarper().warp(
|
||||||
|
mask_name=mask_name,
|
||||||
|
mask_img=template,
|
||||||
|
target_data=target_data,
|
||||||
|
warp_data=warp_data,
|
||||||
|
)
|
||||||
|
elif warp_data["warper"] == "ants":
|
||||||
|
resampled_template = ANTsMaskWarper().warp(
|
||||||
|
mask_name=mask_name,
|
||||||
|
# use template here
|
||||||
|
mask_img=template,
|
||||||
|
src=target_std_space,
|
||||||
|
dst="native",
|
||||||
|
target_data=target_data,
|
||||||
|
warp_data=warp_data,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Warp template to correct space
|
||||||
|
if template_space != target_std_space:
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"Warping template to {target_std_space} space using ANTs."
|
f"Warping template to {target_std_space} space using ANTs."
|
||||||
)
|
)
|
||||||
|
|
@ -179,31 +198,11 @@ def compute_brain_mask(
|
||||||
target_data=target_data,
|
target_data=target_data,
|
||||||
warp_data=None,
|
warp_data=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Resample and warp template if target space is native
|
|
||||||
if target_data["space"] == "native" and template_space != "native":
|
|
||||||
if warp_data["warper"] == "fsl":
|
|
||||||
resampled_template = FSLMaskWarper().warp(
|
|
||||||
mask_name=f"template_{target_std_space}_for_compute_brain_mask",
|
|
||||||
mask_img=template,
|
|
||||||
target_data=target_data,
|
|
||||||
warp_data=warp_data,
|
|
||||||
)
|
|
||||||
elif warp_data["warper"] == "ants":
|
|
||||||
resampled_template = ANTsMaskWarper().warp(
|
|
||||||
mask_name=f"template_{target_std_space}_for_compute_brain_mask",
|
|
||||||
# use template here
|
|
||||||
mask_img=template,
|
|
||||||
src=target_std_space,
|
|
||||||
dst="native",
|
|
||||||
target_data=target_data,
|
|
||||||
warp_data=warp_data,
|
|
||||||
)
|
|
||||||
# Resample template to target image
|
|
||||||
else:
|
|
||||||
# Resample template to target image
|
# Resample template to target image
|
||||||
resampled_template = nimg.resample_to_img(
|
resampled_template = nimg.resample_to_img(
|
||||||
source_img=template, target_img=target_data["data"]
|
source_img=template,
|
||||||
|
target_img=target_data["data"],
|
||||||
|
interpolation=_get_interpolation_method(template),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Threshold resampled template and get mask
|
# Threshold resampled template and get mask
|
||||||
|
|
@ -561,6 +560,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
mask_img = nimg.resample_to_img(
|
mask_img = nimg.resample_to_img(
|
||||||
source_img=mask_img,
|
source_img=mask_img,
|
||||||
target_img=target_data["data"],
|
target_img=target_data["data"],
|
||||||
|
interpolation=_get_interpolation_method(mask_img),
|
||||||
)
|
)
|
||||||
# Starting with new mask
|
# Starting with new mask
|
||||||
else:
|
else:
|
||||||
|
|
@ -632,6 +632,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
mask_img = nimg.resample_to_img(
|
mask_img = nimg.resample_to_img(
|
||||||
source_img=mask_img,
|
source_img=mask_img,
|
||||||
target_img=target_img,
|
target_img=target_img,
|
||||||
|
interpolation=_get_interpolation_method(mask_img),
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# Warp mask if target space is native as
|
# Warp mask if target space is native as
|
||||||
|
|
@ -761,3 +762,23 @@ def _load_ukb_mask(name: str) -> Path:
|
||||||
mask_fname = _masks_path / "ukb" / mask_fname
|
mask_fname = _masks_path / "ukb" / mask_fname
|
||||||
|
|
||||||
return mask_fname
|
return mask_fname
|
||||||
|
|
||||||
|
|
||||||
|
def _get_interpolation_method(img: "Nifti1Image") -> str:
|
||||||
|
"""Get correct interpolation method for `img`.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
img : nibabel.nifti1.Nifti1Image
|
||||||
|
The image.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
str
|
||||||
|
The interpolation method.
|
||||||
|
|
||||||
|
"""
|
||||||
|
if np.array_equal(np.unique(img.get_fdata()), [0, 1]):
|
||||||
|
return "nearest"
|
||||||
|
else:
|
||||||
|
return "continuous"
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue