[ENH]: Rework masking logic #395

Merged
synchon merged 8 commits from refactor/mask-logic into main 2024-12-06 11:08:19 +00:00
9 changed files with 173 additions and 101 deletions

View 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`_

View file

@ -49,4 +49,3 @@ This is the ideal place to include ``junifer`` extensions.
""" """
return ["my_other_file.py"] return ["my_other_file.py"]

View file

@ -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
------- -------

View file

@ -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

View file

@ -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)

View file

@ -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