[ENH]: Rework masking logic #395
9 changed files with 173 additions and 101 deletions
1
docs/changes/newsfragments/395.enh
Normal file
1
docs/changes/newsfragments/395.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Refactor the masking logic in the pipeline to account for and optimise space transformations and merges by `Synchon Mandal`_ and `Fede Raimondo`_
|
||||||
|
|
@ -49,4 +49,3 @@ This is the ideal place to include ``junifer`` extensions.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
return ["my_other_file.py"]
|
return ["my_other_file.py"]
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -46,7 +46,7 @@ class ANTsMaskWarper:
|
||||||
The mask image to transform.
|
The mask image to transform.
|
||||||
src : str
|
src : str
|
||||||
The data type or template space to warp from.
|
The data type or template space to warp from.
|
||||||
It should be empty string if ``dst="T1w"``.
|
It should be empty string if ``dst="native"``.
|
||||||
dst : str
|
dst : str
|
||||||
The data type or template space to warp to.
|
The data type or template space to warp to.
|
||||||
`"native"` is the only allowed data type and it uses the resampled
|
`"native"` is the only allowed data type and it uses the resampled
|
||||||
|
|
@ -58,7 +58,7 @@ class ANTsMaskWarper:
|
||||||
will be applied.
|
will be applied.
|
||||||
warp_data : dict or None
|
warp_data : dict or None
|
||||||
The warp data item of the data object. The value is unused if
|
The warp data item of the data object. The value is unused if
|
||||||
``dst!="T1w"``.
|
``dst!="native"``.
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
|
|
|
||||||
|
|
@ -44,23 +44,24 @@ _masks_path = Path(__file__).parent
|
||||||
|
|
||||||
def compute_brain_mask(
|
def compute_brain_mask(
|
||||||
target_data: dict[str, Any],
|
target_data: dict[str, Any],
|
||||||
extra_input: Optional[dict[str, Any]] = None,
|
warp_data: Optional[dict[str, Any]] = None,
|
||||||
mask_type: str = "brain",
|
mask_type: str = "brain",
|
||||||
threshold: float = 0.5,
|
threshold: float = 0.5,
|
||||||
) -> "Nifti1Image":
|
) -> "Nifti1Image":
|
||||||
"""Compute the whole-brain, grey-matter or white-matter mask.
|
"""Compute the whole-brain, grey-matter or white-matter mask.
|
||||||
|
|
||||||
This mask is calculated using the template space and resolution as found
|
This mask is calculated using the template space and resolution as found
|
||||||
in the ``target_data``.
|
in the ``target_data``. If target space is native, then the template is
|
||||||
|
warped to native and then thresholded.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
target_data : dict
|
target_data : dict
|
||||||
The corresponding item of the data object for which mask will be
|
The corresponding item of the data object for which mask will be
|
||||||
loaded.
|
loaded.
|
||||||
extra_input : dict, optional
|
warp_data : dict or None, optional
|
||||||
The other fields in the data object. Useful for accessing other data
|
The warp data item of the data object. Needs to be provided if
|
||||||
types (default None).
|
``target_data`` is in native space (default None).
|
||||||
mask_type : {"brain", "gm", "wm"}, optional
|
mask_type : {"brain", "gm", "wm"}, optional
|
||||||
Type of mask to be computed:
|
Type of mask to be computed:
|
||||||
|
|
||||||
|
|
@ -81,7 +82,7 @@ def compute_brain_mask(
|
||||||
------
|
------
|
||||||
ValueError
|
ValueError
|
||||||
If ``mask_type`` is invalid or
|
If ``mask_type`` is invalid or
|
||||||
if ``extra_input`` is None when ``target_data``'s space is native.
|
if ``warp_data`` is None when ``target_data``'s space is native.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.debug(f"Computing {mask_type} mask")
|
logger.debug(f"Computing {mask_type} mask")
|
||||||
|
|
@ -90,39 +91,45 @@ def compute_brain_mask(
|
||||||
raise_error(f"Unknown mask type: {mask_type}")
|
raise_error(f"Unknown mask type: {mask_type}")
|
||||||
|
|
||||||
# Check pre-requirements for space manipulation
|
# Check pre-requirements for space manipulation
|
||||||
target_space = target_data["space"]
|
if target_data["space"] == "native":
|
||||||
# Set target standard space to target space
|
# Warp data check
|
||||||
target_std_space = target_space
|
if warp_data is None:
|
||||||
# Extra data type requirement check if target space is native
|
raise_error("No `warp_data` provided")
|
||||||
if target_space == "native":
|
# Set space to fetch template using
|
||||||
# Check for extra inputs
|
target_std_space = warp_data["src"]
|
||||||
if extra_input is None:
|
else:
|
||||||
raise_error(
|
# Set space to fetch template using
|
||||||
"No extra input provided, requires `Warp` "
|
target_std_space = target_data["space"]
|
||||||
"data type to infer target template space."
|
|
||||||
)
|
|
||||||
# Set target standard space to warp file space source
|
|
||||||
for entry in extra_input["Warp"]:
|
|
||||||
if entry["dst"] == "native":
|
|
||||||
target_std_space = entry["src"]
|
|
||||||
|
|
||||||
target_img = target_data["data"]
|
|
||||||
# Fetch template in closest resolution
|
# Fetch template in closest resolution
|
||||||
template = get_template(
|
template = get_template(
|
||||||
space=target_std_space,
|
space=target_std_space,
|
||||||
target_img=target_img,
|
target_img=target_data["data"],
|
||||||
extra_input=extra_input,
|
extra_input=None,
|
||||||
template_type=mask_type,
|
template_type=mask_type,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Resample and warp template if target space is native
|
||||||
|
if target_data["space"] == "native":
|
||||||
|
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
|
# Resample template to target image
|
||||||
|
else:
|
||||||
resampled_template = resample_to_img(
|
resampled_template = resample_to_img(
|
||||||
source_img=template, target_img=target_img
|
source_img=template, target_img=target_data["data"]
|
||||||
)
|
)
|
||||||
|
|
||||||
# Threshold and get mask
|
# Threshold resampled template and get mask
|
||||||
mask = (get_data(resampled_template) >= threshold).astype("int8")
|
mask = (get_data(resampled_template) >= threshold).astype("int8")
|
||||||
|
|
||||||
return new_img_like(target_img, mask) # type: ignore
|
return new_img_like(target_data["data"], mask) # type: ignore
|
||||||
|
|
||||||
|
|
||||||
class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
|
|
@ -369,6 +376,8 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
"""
|
"""
|
||||||
# Check pre-requirements for space manipulation
|
# Check pre-requirements for space manipulation
|
||||||
target_space = target_data["space"]
|
target_space = target_data["space"]
|
||||||
|
logger.debug(f"Getting masks: {masks} in {target_space} space")
|
||||||
|
|
||||||
# Extra data type requirement check if target space is native
|
# Extra data type requirement check if target space is native
|
||||||
if target_space == "native":
|
if target_space == "native":
|
||||||
# Check for extra inputs
|
# Check for extra inputs
|
||||||
|
|
@ -385,7 +394,13 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
)
|
)
|
||||||
# Set target standard space to warp file space source
|
# Set target standard space to warp file space source
|
||||||
target_std_space = warper_spec["src"]
|
target_std_space = warper_spec["src"]
|
||||||
|
logger.debug(
|
||||||
|
f"Target space is native. Will warp from {target_std_space}"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
|
# Set warper_spec so that compute_brain_mask does not fail when
|
||||||
|
# target space is non-native
|
||||||
|
warper_spec = None
|
||||||
# Set target standard space to target space
|
# Set target standard space to target space
|
||||||
target_std_space = target_space
|
target_std_space = target_space
|
||||||
|
|
||||||
|
|
@ -398,31 +413,33 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
masks = [masks]
|
masks = [masks]
|
||||||
|
|
||||||
# Check that masks passed as dicts have only one key
|
# Check that masks passed as dicts have only one key
|
||||||
invalid_elements = [
|
invalid_mask_specs = [
|
||||||
x for x in masks if isinstance(x, dict) and len(x) != 1
|
x for x in masks if isinstance(x, dict) and len(x) != 1
|
||||||
]
|
]
|
||||||
if len(invalid_elements) > 0:
|
if invalid_mask_specs:
|
||||||
raise_error(
|
raise_error(
|
||||||
"Each of the masks dictionary must have only one key, "
|
"Each of the masks dictionary must have only one key, "
|
||||||
"the name of the mask. The following dictionaries are "
|
"the name of the mask. The following dictionaries are "
|
||||||
f"invalid: {invalid_elements}"
|
f"invalid: {invalid_mask_specs}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Check params for the intersection function
|
# Store params for nilearn.masking.intersect_mask()
|
||||||
intersect_params = {}
|
intersect_params = {}
|
||||||
true_masks = []
|
# Store all mask specs for further operations
|
||||||
|
mask_specs = []
|
||||||
for t_mask in masks:
|
for t_mask in masks:
|
||||||
if isinstance(t_mask, dict):
|
if isinstance(t_mask, dict):
|
||||||
|
# Get params to pass to nilearn.masking.intersect_mask()
|
||||||
if "threshold" in t_mask:
|
if "threshold" in t_mask:
|
||||||
intersect_params["threshold"] = t_mask["threshold"]
|
intersect_params["threshold"] = t_mask["threshold"]
|
||||||
continue
|
continue
|
||||||
elif "connected" in t_mask:
|
if "connected" in t_mask:
|
||||||
intersect_params["connected"] = t_mask["connected"]
|
intersect_params["connected"] = t_mask["connected"]
|
||||||
continue
|
continue
|
||||||
# All the other elements are masks
|
# Add mask spec
|
||||||
true_masks.append(t_mask)
|
mask_specs.append(t_mask)
|
||||||
|
|
||||||
if len(true_masks) == 0:
|
if not mask_specs:
|
||||||
raise_error("No mask was passed. At least one mask is required.")
|
raise_error("No mask was passed. At least one mask is required.")
|
||||||
|
|
||||||
# Get the nested mask data type for the input data type
|
# Get the nested mask data type for the input data type
|
||||||
|
|
@ -430,7 +447,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
|
|
||||||
# Get all the masks
|
# Get all the masks
|
||||||
all_masks = []
|
all_masks = []
|
||||||
for t_mask in true_masks:
|
for t_mask in mask_specs:
|
||||||
if isinstance(t_mask, dict):
|
if isinstance(t_mask, dict):
|
||||||
mask_name = next(iter(t_mask.keys()))
|
mask_name = next(iter(t_mask.keys()))
|
||||||
mask_params = t_mask[mask_name]
|
mask_params = t_mask[mask_name]
|
||||||
|
|
@ -441,33 +458,57 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
# If mask is being inherited from the datagrabber or a
|
# If mask is being inherited from the datagrabber or a
|
||||||
# preprocessor, check that it's accessible
|
# preprocessor, check that it's accessible
|
||||||
if mask_name == "inherit":
|
if mask_name == "inherit":
|
||||||
|
logger.debug("Using inherited mask.")
|
||||||
if inherited_mask_item is None:
|
if inherited_mask_item is None:
|
||||||
raise_error(
|
raise_error(
|
||||||
"Cannot inherit mask from the target data. Either the "
|
"Cannot inherit mask from the target data. Either the "
|
||||||
"DataGrabber or a Preprocessor does not provide "
|
"DataGrabber or a Preprocessor does not provide "
|
||||||
"`mask` for the target data type."
|
"`mask` for the target data type."
|
||||||
)
|
)
|
||||||
|
logger.debug(
|
||||||
|
f"Inherited mask is in {inherited_mask_item['space']} "
|
||||||
|
"space."
|
||||||
|
)
|
||||||
mask_img = inherited_mask_item["data"]
|
mask_img = inherited_mask_item["data"]
|
||||||
|
|
||||||
|
if inherited_mask_item["space"] != target_space:
|
||||||
|
raise_error(
|
||||||
|
"Inherited mask space does not match target space."
|
||||||
|
)
|
||||||
|
logger.debug("Resampling inherited mask to target image.")
|
||||||
|
# Resample inherited mask to target image
|
||||||
|
mask_img = resample_to_img(
|
||||||
|
source_img=mask_img,
|
||||||
|
target_img=target_data["data"],
|
||||||
|
)
|
||||||
# Starting with new mask
|
# Starting with new mask
|
||||||
else:
|
else:
|
||||||
# Load mask
|
# Load mask
|
||||||
|
logger.debug(f"Loading mask {t_mask}.")
|
||||||
mask_object, _, mask_space = self.load(
|
mask_object, _, mask_space = self.load(
|
||||||
mask_name, path_only=False, resolution=resolution
|
mask_name, path_only=False, resolution=resolution
|
||||||
)
|
)
|
||||||
# Replace mask space with target space if mask's space is
|
# If mask is callable like from nilearn; space will be inherit
|
||||||
# inherit
|
# so no check for that
|
||||||
if mask_space == "inherit":
|
|
||||||
mask_space = target_std_space
|
|
||||||
# If mask is callable like from nilearn
|
|
||||||
if callable(mask_object):
|
if callable(mask_object):
|
||||||
|
logger.debug("Computing mask (callable).")
|
||||||
if mask_params is None:
|
if mask_params is None:
|
||||||
mask_params = {}
|
mask_params = {}
|
||||||
# From nilearn
|
# From nilearn
|
||||||
if mask_name != "compute_brain_mask":
|
if mask_name in [
|
||||||
|
"compute_epi_mask",
|
||||||
|
"compute_background_mask",
|
||||||
|
]:
|
||||||
mask_img = mask_object(target_img, **mask_params)
|
mask_img = mask_object(target_img, **mask_params)
|
||||||
# Not from nilearn
|
# custom compute_brain_mask
|
||||||
|
elif mask_name == "compute_brain_mask":
|
||||||
|
mask_img = mask_object(
|
||||||
|
target_data, warper_spec, **mask_params
|
||||||
|
)
|
||||||
|
# custom registered; arm kept for clarity
|
||||||
else:
|
else:
|
||||||
mask_img = mask_object(target_data, **mask_params)
|
mask_img = mask_object(target_img, **mask_params)
|
||||||
|
|
||||||
# Mask is a Nifti1Image
|
# Mask is a Nifti1Image
|
||||||
else:
|
else:
|
||||||
# Mask params provided
|
# Mask params provided
|
||||||
|
|
@ -477,22 +518,56 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
"Cannot pass callable params to a non-callable "
|
"Cannot pass callable params to a non-callable "
|
||||||
"mask."
|
"mask."
|
||||||
)
|
)
|
||||||
# Resample mask to target image
|
|
||||||
mask_img = resample_to_img(
|
# Resample and warp mask to standard space
|
||||||
source_img=mask_object,
|
|
||||||
target_img=target_img,
|
|
||||||
interpolation="nearest",
|
|
||||||
copy=True,
|
|
||||||
)
|
|
||||||
# Convert mask space if required
|
|
||||||
if mask_space != target_std_space:
|
if mask_space != target_std_space:
|
||||||
|
logger.debug(
|
||||||
|
f"Warping {t_mask} to {target_std_space} space "
|
||||||
|
"using ants."
|
||||||
|
)
|
||||||
mask_img = ANTsMaskWarper().warp(
|
mask_img = ANTsMaskWarper().warp(
|
||||||
mask_name=mask_name,
|
mask_name=mask_name,
|
||||||
mask_img=mask_img,
|
mask_img=mask_object,
|
||||||
src=mask_space,
|
src=mask_space,
|
||||||
dst=target_std_space,
|
dst=target_std_space,
|
||||||
target_data=target_data,
|
target_data=target_data,
|
||||||
warp_data=None,
|
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"],
|
||||||
|
)
|
||||||
|
# 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":
|
||||||
|
logger.debug(
|
||||||
|
"Warping mask to native space using "
|
||||||
|
f"{warper_spec['warper']}."
|
||||||
|
)
|
||||||
|
mask_name = f"{mask_name}_to_native"
|
||||||
|
# extra_input check done earlier and warper_spec exists
|
||||||
|
if warper_spec["warper"] == "fsl":
|
||||||
|
mask_img = FSLMaskWarper().warp(
|
||||||
|
mask_name=mask_name,
|
||||||
|
mask_img=mask_img,
|
||||||
|
target_data=target_data,
|
||||||
|
warp_data=warper_spec,
|
||||||
|
)
|
||||||
|
elif warper_spec["warper"] == "ants":
|
||||||
|
mask_img = ANTsMaskWarper().warp(
|
||||||
|
mask_name=mask_name,
|
||||||
|
mask_img=mask_img,
|
||||||
|
src="",
|
||||||
|
dst="native",
|
||||||
|
target_data=target_data,
|
||||||
|
warp_data=warper_spec,
|
||||||
)
|
)
|
||||||
|
|
||||||
all_masks.append(mask_img)
|
all_masks.append(mask_img)
|
||||||
|
|
@ -500,10 +575,11 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
# Multiple masks, need intersection / union
|
# Multiple masks, need intersection / union
|
||||||
if len(all_masks) > 1:
|
if len(all_masks) > 1:
|
||||||
# Intersect / union of masks
|
# Intersect / union of masks
|
||||||
|
logger.debug("Intersecting masks.")
|
||||||
mask_img = intersect_masks(all_masks, **intersect_params)
|
mask_img = intersect_masks(all_masks, **intersect_params)
|
||||||
# Single mask
|
# Single mask
|
||||||
else:
|
else:
|
||||||
if len(intersect_params) > 0:
|
if intersect_params:
|
||||||
# Yes, I'm this strict!
|
# Yes, I'm this strict!
|
||||||
raise_error(
|
raise_error(
|
||||||
"Cannot pass parameters to the intersection function "
|
"Cannot pass parameters to the intersection function "
|
||||||
|
|
@ -511,26 +587,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
)
|
)
|
||||||
mask_img = all_masks[0]
|
mask_img = all_masks[0]
|
||||||
|
|
||||||
# Warp mask if target data is native
|
|
||||||
if target_space == "native":
|
|
||||||
# extra_input check done earlier and warper_spec exists
|
|
||||||
if warper_spec["warper"] == "fsl":
|
|
||||||
mask_img = FSLMaskWarper().warp(
|
|
||||||
mask_name="native",
|
|
||||||
mask_img=mask_img,
|
|
||||||
target_data=target_data,
|
|
||||||
warp_data=warper_spec,
|
|
||||||
)
|
|
||||||
elif warper_spec["warper"] == "ants":
|
|
||||||
mask_img = ANTsMaskWarper().warp(
|
|
||||||
mask_name="native",
|
|
||||||
mask_img=mask_img,
|
|
||||||
src="",
|
|
||||||
dst="native",
|
|
||||||
target_data=target_data,
|
|
||||||
warp_data=warper_spec,
|
|
||||||
)
|
|
||||||
|
|
||||||
return mask_img
|
return mask_img
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -64,7 +64,6 @@ def test_compute_brain_mask(mask_type: str, threshold: float) -> None:
|
||||||
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
||||||
mask = compute_brain_mask(
|
mask = compute_brain_mask(
|
||||||
target_data=element_data["BOLD"],
|
target_data=element_data["BOLD"],
|
||||||
extra_input=None,
|
|
||||||
mask_type=mask_type,
|
mask_type=mask_type,
|
||||||
)
|
)
|
||||||
assert isinstance(mask, nib.nifti1.Nifti1Image)
|
assert isinstance(mask, nib.nifti1.Nifti1Image)
|
||||||
|
|
|
||||||
|
|
@ -400,6 +400,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
"""
|
"""
|
||||||
# Check pre-requirements for space manipulation
|
# Check pre-requirements for space manipulation
|
||||||
target_space = target_data["space"]
|
target_space = target_data["space"]
|
||||||
|
logger.debug(f"Getting {parcellations} in {target_space} space.")
|
||||||
# Extra data type requirement check if target space is native
|
# Extra data type requirement check if target space is native
|
||||||
if target_space == "native":
|
if target_space == "native":
|
||||||
# Check for extra inputs
|
# Check for extra inputs
|
||||||
|
|
@ -416,6 +417,9 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
)
|
)
|
||||||
# Set target standard space to warp file space source
|
# Set target standard space to warp file space source
|
||||||
target_std_space = warper_spec["src"]
|
target_std_space = warper_spec["src"]
|
||||||
|
logger.debug(
|
||||||
|
f"Target space is native. Will warp from {target_std_space}"
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
# Set target standard space to target space
|
# Set target standard space to target space
|
||||||
target_std_space = target_space
|
target_std_space = target_space
|
||||||
|
|
@ -433,6 +437,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
all_labels = []
|
all_labels = []
|
||||||
for name in parcellations:
|
for name in parcellations:
|
||||||
# Load parcellation
|
# Load parcellation
|
||||||
|
logger.debug(f"Loading parcellation {name}")
|
||||||
img, labels, _, space = self.load(
|
img, labels, _, space = self.load(
|
||||||
name=name,
|
name=name,
|
||||||
resolution=resolution,
|
resolution=resolution,
|
||||||
|
|
@ -441,6 +446,9 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
# Convert parcellation spaces if required;
|
# Convert parcellation spaces if required;
|
||||||
# cannot be "native" due to earlier check
|
# cannot be "native" due to earlier check
|
||||||
if space != target_std_space:
|
if space != target_std_space:
|
||||||
|
logger.debug(
|
||||||
|
f"Warping {name} to {target_std_space} space using ants."
|
||||||
|
)
|
||||||
raw_img = ANTsParcellationWarper().warp(
|
raw_img = ANTsParcellationWarper().warp(
|
||||||
parcellation_name=name,
|
parcellation_name=name,
|
||||||
parcellation_img=img,
|
parcellation_img=img,
|
||||||
|
|
@ -452,6 +460,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
# Remove extra dimension added by ANTs
|
# Remove extra dimension added by ANTs
|
||||||
img = image.math_img("np.squeeze(img)", img=raw_img)
|
img = image.math_img("np.squeeze(img)", img=raw_img)
|
||||||
|
|
||||||
|
logger.debug(f"Resampling {name} to target image.")
|
||||||
# Resample parcellation to target image
|
# Resample parcellation to target image
|
||||||
img_to_merge = image.resample_to_img(
|
img_to_merge = image.resample_to_img(
|
||||||
source_img=img,
|
source_img=img,
|
||||||
|
|
@ -469,6 +478,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
labels = all_labels[0]
|
labels = all_labels[0]
|
||||||
# Parcellations are already transformed to target standard space
|
# Parcellations are already transformed to target standard space
|
||||||
else:
|
else:
|
||||||
|
logger.debug("Merging parcellations.")
|
||||||
resampled_parcellation_img, labels = merge_parcellations(
|
resampled_parcellation_img, labels = merge_parcellations(
|
||||||
parcellations_list=all_parcellations,
|
parcellations_list=all_parcellations,
|
||||||
parcellations_names=parcellations,
|
parcellations_names=parcellations,
|
||||||
|
|
@ -477,6 +487,10 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
||||||
|
|
||||||
# Warp parcellation if target space is native
|
# Warp parcellation if target space is native
|
||||||
if target_space == "native":
|
if target_space == "native":
|
||||||
|
logger.debug(
|
||||||
|
"Warping parcellation to native space using "
|
||||||
|
f"{warper_spec['warper']}."
|
||||||
|
)
|
||||||
# extra_input check done earlier and warper_spec exists
|
# extra_input check done earlier and warper_spec exists
|
||||||
if warper_spec["warper"] == "fsl":
|
if warper_spec["warper"] == "fsl":
|
||||||
resampled_parcellation_img = FSLParcellationWarper().warp(
|
resampled_parcellation_img = FSLParcellationWarper().warp(
|
||||||
|
|
@ -1194,7 +1208,10 @@ def _retrieve_aicha(
|
||||||
|
|
||||||
# Load labels
|
# Load labels
|
||||||
labels = pd.read_csv(
|
labels = pd.read_csv(
|
||||||
parcellation_lname, sep="\t", header=None, skiprows=[0] # type: ignore
|
parcellation_lname,
|
||||||
|
sep="\t",
|
||||||
|
header=None,
|
||||||
|
skiprows=[0], # type: ignore
|
||||||
)[0].to_list()
|
)[0].to_list()
|
||||||
|
|
||||||
return parcellation_fname, labels
|
return parcellation_fname, labels
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue