diff --git a/docs/changes/newsfragments/481.bugfix b/docs/changes/newsfragments/481.bugfix new file mode 100644 index 000000000..15a8f8f38 --- /dev/null +++ b/docs/changes/newsfragments/481.bugfix @@ -0,0 +1 @@ +Bypass warper check for callable masks on native data by `Synchon Mandal`_ diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 1bd36b2f3..0274c2a1a 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -470,25 +470,48 @@ class MaskRegistry(BasePipelineDataRegistry): target_space = target_data["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 if target_space == "native": - # 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." + # Check for in-built callable masks so that get_native_warper is + # not called in case only callable masks are used + if ( + set(masks) == {nilearn_callable_masks[0]} + or set(masks) == {nilearn_callable_masks[1]} + 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: # Set warper_spec so that compute_brain_mask does not fail when # target space is non-native @@ -500,10 +523,6 @@ class MaskRegistry(BasePipelineDataRegistry): target_img = target_data["data"] 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 invalid_mask_specs = [ 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: mask_params = {} # From nilearn - if mask_name in [ - "compute_epi_mask", - "compute_background_mask", - ]: + if mask_name in nilearn_callable_masks: mask_img = mask_object(target_img, **mask_params) # custom compute_brain_mask elif mask_name == "compute_brain_mask":