[BUG]: compute_epi_mask fails for data in native space #481

Merged
synchon merged 2 commits from fix/native-data-callable-mask-apply into main 2026-01-28 14:53:46 +00:00
2 changed files with 41 additions and 24 deletions

View file

@ -0,0 +1 @@
Bypass warper check for callable masks on native data by `Synchon Mandal`_

View file

@ -470,25 +470,48 @@ class MaskRegistry(BasePipelineDataRegistry):
target_space = target_data["space"] target_space = target_data["space"]
logger.debug(f"Getting masks: {masks} in {target_space} space") logger.debug(f"Getting masks: {masks} in {target_space} space")
# Set once for future use
nilearn_callable_masks = [
"compute_epi_mask",
"compute_background_mask",
]
# Convert masks to list if not already
if not isinstance(masks, list):
masks = [masks]
# 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 in-built callable masks so that get_native_warper is
if extra_input is None: # not called in case only callable masks are used
raise_error( if (
"No extra input provided, requires `Warp` and `T1w` " set(masks) == {nilearn_callable_masks[0]}
"data types in particular for transformation to " or set(masks) == {nilearn_callable_masks[1]}
f"{target_data['space']} space for further computation." or set(masks) == set(nilearn_callable_masks)
):
logger.debug(
"Target space is native. "
f"No warping will be done for: {masks}"
)
else:
# Check for extra inputs
if extra_input is None:
raise_error(
"No extra input provided, requires `Warp` and `T1w` "
"data types in particular for transformation to "
f"{target_data['space']} space for further "
"computation."
)
# Get native space warper spec
warper_spec = get_native_warper(
target_data=target_data,
other_data=extra_input,
)
# Set target standard space to warp file space source
target_std_space = warper_spec["src"]
logger.debug(
"Target space is native. "
f"Will warp from {target_std_space}"
) )
# Get native space warper spec
warper_spec = get_native_warper(
target_data=target_data,
other_data=extra_input,
)
# Set target standard space to warp file space source
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 # Set warper_spec so that compute_brain_mask does not fail when
# target space is non-native # target space is non-native
@ -500,10 +523,6 @@ class MaskRegistry(BasePipelineDataRegistry):
target_img = target_data["data"] target_img = target_data["data"]
resolution = np.min(target_img.header.get_zooms()[:3]) resolution = np.min(target_img.header.get_zooms()[:3])
# Convert masks to list if not already
if not isinstance(masks, list):
masks = [masks]
# Check that masks passed as dicts have only one key # Check that masks passed as dicts have only one key
invalid_mask_specs = [ 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
@ -588,10 +607,7 @@ class MaskRegistry(BasePipelineDataRegistry):
if mask_params is None: if mask_params is None:
mask_params = {} mask_params = {}
# From nilearn # From nilearn
if mask_name in [ if mask_name in nilearn_callable_masks:
"compute_epi_mask",
"compute_background_mask",
]:
mask_img = mask_object(target_img, **mask_params) mask_img = mask_object(target_img, **mask_params)
# custom compute_brain_mask # custom compute_brain_mask
elif mask_name == "compute_brain_mask": elif mask_name == "compute_brain_mask":