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

View file

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

View file

@ -86,10 +86,6 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None:
@pytest.mark.skipif(
_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:
"""Test ReHoParcels implementation comparison.