From 979b22cb1601d25fcf43272e188c85e09dd63f93 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 6 Dec 2024 14:48:57 +0100 Subject: [PATCH 1/6] update: get correct interpolation method for masks --- junifer/data/masks/_masks.py | 26 +++++++++++++++++++++++++- 1 file changed, 25 insertions(+), 1 deletion(-) diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 150e13bd4..4dd3dbb76 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -203,7 +203,9 @@ def compute_brain_mask( else: # Resample template to target image 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 @@ -561,6 +563,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): mask_img = nimg.resample_to_img( source_img=mask_img, target_img=target_data["data"], + interpolation=_get_interpolation_method(mask_img), ) # Starting with new mask else: @@ -632,6 +635,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): mask_img = nimg.resample_to_img( source_img=mask_img, target_img=target_img, + interpolation=_get_interpolation_method(mask_img), ) else: # Warp mask if target space is native as @@ -761,3 +765,23 @@ def _load_ukb_mask(name: str) -> Path: mask_fname = _masks_path / "ukb" / 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" -- 2.52.0 From fb8330ebf2cc81ccc2001d0d4b52b46109246810 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 6 Dec 2024 14:49:49 +0100 Subject: [PATCH 2/6] chore: improve log message in _masks.py --- junifer/data/masks/_masks.py | 14 +++++++++++++- 1 file changed, 13 insertions(+), 1 deletion(-) diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 4dd3dbb76..54b1f032a 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -199,8 +199,20 @@ def compute_brain_mask( target_data=target_data, warp_data=warp_data, ) - # Resample template to target image else: + # Warp template to correct space + if template_space != target_std_space: + logger.debug( + f"Warping template to {target_std_space} space using ANTs." + ) + template = ANTsMaskWarper().warp( + mask_name=f"template_{target_std_space}_for_compute_brain_mask", + mask_img=template, + src=template_space, + dst=target_std_space, + target_data=target_data, + warp_data=None, + ) # Resample template to target image resampled_template = nimg.resample_to_img( source_img=template, -- 2.52.0 From 9e8ee8c2ff3984b7e29e8c7cf0c8cc8620522c05 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 6 Dec 2024 14:50:53 +0100 Subject: [PATCH 3/6] style: simplify mask_name in _masks.py --- junifer/data/masks/_masks.py | 21 +++------------------ 1 file changed, 3 insertions(+), 18 deletions(-) diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 54b1f032a..ae72f2810 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -165,33 +165,18 @@ def compute_brain_mask( ) mask_name = f"template_{target_std_space}_for_compute_brain_mask" - - # Warp template to correct space (MNI to MNI) - if template_space != "native" and template_space != target_std_space: - logger.debug( - f"Warping template to {target_std_space} space using ANTs." - ) - template = ANTsMaskWarper().warp( - mask_name=mask_name, - mask_img=template, - src=template_space, - dst=target_std_space, - target_data=target_data, - 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_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=f"template_{target_std_space}_for_compute_brain_mask", + mask_name=mask_name, # use template here mask_img=template, src=target_std_space, @@ -206,7 +191,7 @@ def compute_brain_mask( f"Warping template to {target_std_space} space using ANTs." ) template = ANTsMaskWarper().warp( - mask_name=f"template_{target_std_space}_for_compute_brain_mask", + mask_name=mask_name, mask_img=template, src=template_space, dst=target_std_space, -- 2.52.0 From 49ae13e8d707f4d2b965697a9312b12b5d06f827 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 6 Dec 2024 15:42:21 +0100 Subject: [PATCH 4/6] Use the right interpolation for ants masks warping --- junifer/data/masks/_ants_mask_warper.py | 24 +++++++++++++++++++++++- 1 file changed, 23 insertions(+), 1 deletion(-) diff --git a/junifer/data/masks/_ants_mask_warper.py b/junifer/data/masks/_ants_mask_warper.py index e68d78873..cbac12325 100644 --- a/junifer/data/masks/_ants_mask_warper.py +++ b/junifer/data/masks/_ants_mask_warper.py @@ -7,6 +7,7 @@ import uuid from typing import TYPE_CHECKING, Any, Optional import nibabel as nib +import numpy as np from ...pipeline import WorkDirManager from ...utils import logger, raise_error, run_ext_cmd @@ -20,6 +21,27 @@ if TYPE_CHECKING: __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 for mask space warping via ANTs. @@ -143,7 +165,7 @@ class ANTsMaskWarper: "antsApplyTransforms", "-d 3", "-e 3", - "-n 'GenericLabel[NearestNeighbor]'", + f"-n {_get_interpolation_method(mask_img)}", f"-i {prewarp_mask_path.resolve()}", f"-r {template_space_img_path.resolve()}", f"-t {xfm_file_path.resolve()}", -- 2.52.0 From 13604de71d0084577832071cadef07a575027c62 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 6 Dec 2024 15:44:22 +0100 Subject: [PATCH 5/6] Use right interpolation for FSL mask warper --- junifer/data/masks/_fsl_mask_warper.py | 23 ++++++++++++++++++++++- 1 file changed, 22 insertions(+), 1 deletion(-) diff --git a/junifer/data/masks/_fsl_mask_warper.py b/junifer/data/masks/_fsl_mask_warper.py index 9e0b3961a..8510318e3 100644 --- a/junifer/data/masks/_fsl_mask_warper.py +++ b/junifer/data/masks/_fsl_mask_warper.py @@ -7,6 +7,7 @@ import uuid from typing import TYPE_CHECKING, Any import nibabel as nib +import numpy as np from ...pipeline import WorkDirManager from ...utils import logger, run_ext_cmd @@ -19,6 +20,26 @@ if TYPE_CHECKING: __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 for mask space warping via FSL FLIRT. @@ -71,7 +92,7 @@ class FSLMaskWarper: # Set applywarp command applywarp_cmd = [ "applywarp", - "--interp=nn", + f"--interp={_get_interpolation_method(mask_img)}", f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference']['path'].resolve()}", -- 2.52.0 From ea9f435171ae1629eab82a24dea3699de5ddc3b7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 9 Dec 2024 12:24:21 +0100 Subject: [PATCH 6/6] chore: lint --- junifer/data/masks/_ants_mask_warper.py | 1 - 1 file changed, 1 deletion(-) diff --git a/junifer/data/masks/_ants_mask_warper.py b/junifer/data/masks/_ants_mask_warper.py index cbac12325..5b07a6675 100644 --- a/junifer/data/masks/_ants_mask_warper.py +++ b/junifer/data/masks/_ants_mask_warper.py @@ -41,7 +41,6 @@ def _get_interpolation_method(img: "Nifti1Image") -> str: return "LanczosWindowedSinc" - class ANTsMaskWarper: """Class for mask space warping via ANTs. -- 2.52.0