[BUG]: Mask "inherit" will be warped twice if working in native space. #284

Merged
synchon merged 6 commits from fix/inherit-mask-single-warp into main 2024-01-15 08:45:34 +00:00
4 changed files with 25 additions and 14 deletions

View file

@ -0,0 +1 @@
Avoid warping mask preprocessed with :class:`.fMRIPrepConfoundRemover` and used by markers with ``mask="inherit"`` in subject-native template space by `Fede Raimondo`_ and `Synchon Mandal`_

View file

@ -280,12 +280,16 @@ def get_mask( # noqa: C901
f"because the item ({inherited_mask_item}) does not exist." f"because the item ({inherited_mask_item}) does not exist."
) )
mask_img = extra_input[inherited_mask_item]["data"] mask_img = extra_input[inherited_mask_item]["data"]
mask_space = target_data["space"]
# Starting with new mask # Starting with new mask
else: else:
# Load mask # Load mask
mask_object, _, mask_space = load_mask( mask_object, _, mask_space = load_mask(
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 inherit
if mask_space == "inherit":
mask_space = target_data["space"]
# If mask is callable like from nilearn # If mask is callable like from nilearn
if callable(mask_object): if callable(mask_object):
if mask_params is None: if mask_params is None:
@ -311,15 +315,17 @@ def get_mask( # noqa: C901
# Multiple masks, need intersection / union # Multiple masks, need intersection / union
if len(all_masks) > 1: if len(all_masks) > 1:
# Filter out "inherit" and make a set for spaces # Make a set of unique spaces
filtered_spaces = set(filter(lambda x: x != "inherit", all_spaces)) unique_spaces = set(all_spaces)
# Intersect / union of masks only if all masks are in the same space # Intersect / union of masks only if all masks are in the same space
if len(filtered_spaces) == 1: if len(unique_spaces) == 1:
mask_img = intersect_masks(all_masks, **intersect_params) mask_img = intersect_masks(all_masks, **intersect_params)
# Store the mask space for further checks
mask_space = next(iter(unique_spaces))
else: else:
raise_error( raise_error(
msg=( msg=(
f"Masks are in different spaces: {filtered_spaces}, " f"Masks are in different spaces: {unique_spaces}, "
"unable to merge." "unable to merge."
), ),
klass=RuntimeError, klass=RuntimeError,
fraimondo commented 2023-12-14 11:55:42 +00:00 (Migrated from github.com)

I do not follow here the logic. Why we do not consider "inherit" in "filtered spaces"? Indeed inherit should be replaced by the space of the "inherited" mask, no?

I do not follow here the logic. Why we do not consider "inherit" in "filtered spaces"? Indeed inherit should be replaced by the space of the "inherited" mask, no?
synchon commented 2024-01-11 11:41:40 +00:00 (Migrated from github.com)

We need to check actual space and not "inherit" coz we get the unique ones by making a set out of the spaces and inherit doesn't make sense there. That's a fair point that the space of the inherited mask should be replaced, will take a look.

We need to check actual space and not "inherit" coz we get the unique ones by making a set out of the spaces and inherit doesn't make sense there. That's a fair point that the space of the inherited mask should be replaced, will take a look.
synchon commented 2024-01-12 10:19:53 +00:00 (Migrated from github.com)

So I've updated the logic to check for correct target space and do away with "inherit".

So I've updated the logic to check for correct target space and do away with "inherit".
@ -333,9 +339,10 @@ def get_mask( # noqa: C901
"when there is only one mask." "when there is only one mask."
) )
mask_img = all_masks[0] mask_img = all_masks[0]
mask_space = all_spaces[0]
# Warp mask if target data is native # Warp mask if target data is native and mask space is not native
if target_data["space"] == "native": if target_data["space"] == "native" and target_data["space"] != mask_space:
# Check for extra inputs # Check for extra inputs
if extra_input is None: if extra_input is None:
raise_error( raise_error(

View file

@ -392,7 +392,9 @@ def test_get_mask_inherit() -> None:
# Now get the mask using the inherit functionality, passing the # Now get the mask using the inherit functionality, passing the
# computed mask as extra data # computed mask as extra data
extra_input = {"BOLD_MASK": {"data": gm_mask}} extra_input = {
"BOLD_MASK": {"data": gm_mask, "space": input["BOLD"]["space"]}
}
input["BOLD"]["mask_item"] = "BOLD_MASK" input["BOLD"]["mask_item"] = "BOLD_MASK"
mask2 = get_mask( mask2 = get_mask(
masks="inherit", target_data=input["BOLD"], extra_input=extra_input masks="inherit", target_data=input["BOLD"], extra_input=extra_input
@ -405,11 +407,9 @@ def test_get_mask_inherit() -> None:
@pytest.mark.parametrize( @pytest.mark.parametrize(
"masks,params", "masks,params",
[ [
(["GM_prob0.2", "compute_brain_mask"], {}), (["GM_prob0.2", "GM_prob0.2_cortex"], {}),
( (["compute_brain_mask", "compute_background_mask"], {}),
["GM_prob0.2", "compute_brain_mask"], (["compute_brain_mask", "compute_epi_mask"], {}),
{"threshold": 0.2},
),
], ],
) )
def test_get_mask_multiple( def test_get_mask_multiple(

View file

@ -574,7 +574,10 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
# this allows to use "inherit" down the pipeline # this allows to use "inherit" down the pipeline
if extra_input is not None: if extra_input is not None:
logger.debug("Setting mask_item") logger.debug("Setting mask_item")
extra_input["BOLD_mask"] = {"data": mask_img} extra_input["BOLD_mask"] = {
"data": mask_img,
"space": input["space"],
}
input["mask_item"] = "BOLD_mask" input["mask_item"] = "BOLD_mask"
logger.info("Cleaning image") logger.info("Cleaning image")