[ENH]: Simplify space warping for parcellations and masks #414

Merged
synchon merged 7 commits from fix/non-native-data-warping into main 2024-12-09 10:18:03 +00:00
4 changed files with 45 additions and 45 deletions

View file

@ -0,0 +1 @@
Simplify logic for parcellation and mask space warping by `Synchon Mandal`_

View file

@ -14,8 +14,8 @@ from typing import (
) )
import nibabel as nib import nibabel as nib
import nilearn.image as nimg
import numpy as np import numpy as np
from nilearn.image import get_data, new_img_like, resample_to_img
from nilearn.masking import ( from nilearn.masking import (
compute_background_mask, compute_background_mask,
compute_epi_mask, compute_epi_mask,
@ -168,14 +168,14 @@ def compute_brain_mask(
) )
# Resample template to target image # Resample template to target image
else: else:
resampled_template = 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"]
) )
# Threshold resampled template and get mask # Threshold resampled template and get mask
mask = (get_data(resampled_template) >= threshold).astype("int8") mask = (nimg.get_data(resampled_template) >= threshold).astype("int8")
return new_img_like(target_data["data"], mask) # type: ignore return nimg.new_img_like(target_data["data"], mask) # type: ignore
class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
@ -523,7 +523,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
) )
logger.debug("Resampling inherited mask to target image.") logger.debug("Resampling inherited mask to target image.")
# Resample inherited mask to target image # Resample inherited mask to target image
mask_img = 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"],
) )
@ -568,34 +568,41 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
"mask." "mask."
) )
# Set here to simplify things later
mask_img: nib.nifti1.Nifti1Image = mask_object
# Resample and warp mask to standard space # Resample and warp mask to standard space
if mask_space != target_std_space: if mask_space != target_std_space:
logger.debug( logger.debug(
f"Warping {t_mask} to {target_std_space} space " f"Warping {t_mask} to {target_std_space} space "
"using ants." "using ANTs."
) )
mask_img = ANTsMaskWarper().warp( mask_img = ANTsMaskWarper().warp(
mask_name=mask_name, mask_name=mask_name,
mask_img=mask_object, mask_img=mask_img,
src=mask_space, src=mask_space,
dst=target_std_space, dst=target_std_space,
target_data=target_data, target_data=target_data,
warp_data=warper_spec, warp_data=warper_spec,
) )
# Remove extra dimension added by ANTs
mask_img = nimg.math_img(
"np.squeeze(img)", img=mask_img
)
else: if target_space != "native":
# Resample mask to target image; no further warping # No warping is going to happen, just resampling,
# because we are in the correct space
logger.debug(f"Resampling {t_mask} to target image.") logger.debug(f"Resampling {t_mask} to target image.")
if target_space != "native": mask_img = nimg.resample_to_img(
mask_img = resample_to_img( source_img=mask_img,
source_img=mask_object, target_img=target_img,
target_img=target_data["data"], )
) else:
# Set mask_img in case no warping happens before this # Warp mask if target space is native as
else: # either the image is in the right non-native space or
mask_img = mask_object # it's warped from one non-native space to another
# Resample and warp mask if target data is native # non-native space
if target_space == "native":
logger.debug( logger.debug(
"Warping mask to native space using " "Warping mask to native space using "
f"{warper_spec['warper']}." f"{warper_spec['warper']}."

View file

@ -16,9 +16,10 @@ from typing import TYPE_CHECKING, Any, Optional, Union
import httpx import httpx
import nibabel as nib import nibabel as nib
import nilearn.image as nimg
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from nilearn import datasets, image from nilearn import datasets
from ...utils import logger, raise_error, warn_with_log from ...utils import logger, raise_error, warn_with_log
from ...utils.singleton import Singleton from ...utils.singleton import Singleton
@ -470,28 +471,23 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
warp_data=None, warp_data=None,
) )
# Remove extra dimension added by ANTs # Remove extra dimension added by ANTs
img = image.math_img("np.squeeze(img)", img=raw_img) img = nimg.math_img("np.squeeze(img)", img=raw_img)
# Set correct affine as resolution won't be correct
img = image.resample_img( if target_space != "native":
img=img, # No warping is going to happen, just resampling, because
target_affine=target_img.affine, # we are in the correct space
logger.debug(f"Resampling {name} to target image.")
# Resample parcellation to target image
img = nimg.resample_to_img(
source_img=img,
target_img=target_img,
interpolation="nearest", interpolation="nearest",
copy=True,
) )
else: else:
if target_space != "native": # Warp parcellation if target space is native as either
# No warping is going to happen, just resampling, because # the image is in the right non-native space or it's
# we are in the correct space # warped from one non-native space to another non-native space
logger.debug(f"Resampling {name} to target image.")
# Resample parcellation to target image
img = image.resample_to_img(
source_img=img,
target_img=target_img,
interpolation="nearest",
copy=True,
)
# Warp parcellation if target space is native
if target_space == "native":
logger.debug( logger.debug(
"Warping parcellation to native space using " "Warping parcellation to native space using "
f"{warper_spec['warper']}." f"{warper_spec['warper']}."
@ -1807,7 +1803,7 @@ def merge_parcellations(
"The parcellations have different resolutions!" "The parcellations have different resolutions!"
"Resampling all parcellations to the first one in the list." "Resampling all parcellations to the first one in the list."
) )
t_parc = image.resample_to_img( t_parc = nimg.resample_to_img(
t_parc, ref_parc, interpolation="nearest", copy=True t_parc, ref_parc, interpolation="nearest", copy=True
) )
@ -1833,6 +1829,6 @@ def merge_parcellations(
"parcellation that was first in the list." "parcellation that was first in the list."
) )
parcellation_img_res = image.new_img_like(parcellations_list[0], parc_data) parcellation_img_res = nimg.new_img_like(parcellations_list[0], parc_data)
return parcellation_img_res, labels return parcellation_img_res, labels

View file

@ -86,10 +86,6 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@pytest.mark.skipif( @pytest.mark.skipif(
_check_afni() is False, reason="requires AFNI to be in PATH" _check_afni() is False, reason="requires AFNI to be in PATH"
) )
@pytest.mark.xfail(
reason="junifer ReHo needs to use the correct mask",
raises=AssertionError,
)
def test_ReHoParcels_comparison(tmp_path: Path) -> None: def test_ReHoParcels_comparison(tmp_path: Path) -> None:
"""Test ReHoParcels implementation comparison. """Test ReHoParcels implementation comparison.