[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

@ -1 +1 @@
Allow for external python files in the ``with`` section of the yaml to import other files by `Fede Raimondo`_
Allow for external python files in the ``with`` section of the yaml to import other files by `Fede Raimondo`_

View file

@ -1 +1 @@
Add a ``cleanup`` parameter to the :class:`junifer.pipeline.WorkDirManager`, which allows to disable cleaning up temporary directories by `Fede Raimondo`_
Add a ``cleanup`` parameter to the :class:`junifer.pipeline.WorkDirManager`, which allows to disable cleaning up temporary directories by `Fede Raimondo`_

View file

@ -1 +1 @@
Refactor ``Singleton`` class to use a metaclass instead of a class decorator by `Fede Raimondo`_
Refactor ``Singleton`` class to use a metaclass instead of a class decorator by `Fede Raimondo`_

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

@ -31,7 +31,7 @@ This is the ideal place to include ``junifer`` extensions.
Some ``junifer`` commands will not consider files imported from files
included in the ``with`` statement, unless this is known to junifer. If
``my_file.py`` imports ``my_other_file.py``, the ``run`` command will work,
but ``queue`` will not create a proper job. This is because we need to
but ``queue`` will not create a proper job. This is because we need to
let junifer know that ``my_other_file.py`` is also part of the code. To do
so, we need to include a special function in ``my_file.py`` which tells
``junifer`` about the dependencies of the module:
@ -39,14 +39,13 @@ This is the ideal place to include ``junifer`` extensions.
.. code-block:: python
def junifer_module_deps() -> List[str]:
"""Return the dependencies of the module.
"""Return the dependencies of the module.
Returns
-------
list of str
The list of dependencies.
Returns
-------
list of str
The list of dependencies.
"""
return ["my_other_file.py"]
"""
return ["my_other_file.py"]

View file

@ -46,7 +46,7 @@ class ANTsMaskWarper:
The mask image to transform.
src : str
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
The data type or template space to warp to.
`"native"` is the only allowed data type and it uses the resampled
@ -58,7 +58,7 @@ class ANTsMaskWarper:
will be applied.
warp_data : dict or None
The warp data item of the data object. The value is unused if
``dst!="T1w"``.
``dst!="native"``.
Returns
-------

View file

@ -44,23 +44,24 @@ _masks_path = Path(__file__).parent
def compute_brain_mask(
target_data: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None,
warp_data: Optional[dict[str, Any]] = None,
mask_type: str = "brain",
threshold: float = 0.5,
) -> "Nifti1Image":
"""Compute the whole-brain, grey-matter or white-matter mask.
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
----------
target_data : dict
The corresponding item of the data object for which mask will be
loaded.
extra_input : dict, optional
The other fields in the data object. Useful for accessing other data
types (default None).
warp_data : dict or None, optional
The warp data item of the data object. Needs to be provided if
``target_data`` is in native space (default None).
mask_type : {"brain", "gm", "wm"}, optional
Type of mask to be computed:
@ -81,7 +82,7 @@ def compute_brain_mask(
------
ValueError
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")
@ -90,39 +91,45 @@ def compute_brain_mask(
raise_error(f"Unknown mask type: {mask_type}")
# Check pre-requirements for space manipulation
target_space = target_data["space"]
# Set target standard space to target space
target_std_space = target_space
# 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` "
"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"]
if target_data["space"] == "native":
# Warp data check
if warp_data is None:
raise_error("No `warp_data` provided")
# Set space to fetch template using
target_std_space = warp_data["src"]
else:
# Set space to fetch template using
target_std_space = target_data["space"]
target_img = target_data["data"]
# Fetch template in closest resolution
template = get_template(
space=target_std_space,
target_img=target_img,
extra_input=extra_input,
target_img=target_data["data"],
extra_input=None,
template_type=mask_type,
)
# Resample template to target image
resampled_template = resample_to_img(
source_img=template, target_img=target_img
)
# Threshold and get mask
# 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
else:
resampled_template = resample_to_img(
source_img=template, target_img=target_data["data"]
)
# Threshold resampled template and get mask
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):
@ -369,6 +376,8 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
"""
# Check pre-requirements for space manipulation
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
if target_space == "native":
# Check for extra inputs
@ -385,7 +394,13 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
)
# 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
warper_spec = None
# Set target standard space to target space
target_std_space = target_space
@ -398,31 +413,33 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
masks = [masks]
# 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
]
if len(invalid_elements) > 0:
if invalid_mask_specs:
raise_error(
"Each of the masks dictionary must have only one key, "
"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 = {}
true_masks = []
# Store all mask specs for further operations
mask_specs = []
for t_mask in masks:
if isinstance(t_mask, dict):
# Get params to pass to nilearn.masking.intersect_mask()
if "threshold" in t_mask:
intersect_params["threshold"] = t_mask["threshold"]
continue
elif "connected" in t_mask:
if "connected" in t_mask:
intersect_params["connected"] = t_mask["connected"]
continue
# All the other elements are masks
true_masks.append(t_mask)
# Add mask spec
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.")
# Get the nested mask data type for the input data type
@ -430,7 +447,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
# Get all the masks
all_masks = []
for t_mask in true_masks:
for t_mask in mask_specs:
if isinstance(t_mask, dict):
mask_name = next(iter(t_mask.keys()))
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
# preprocessor, check that it's accessible
if mask_name == "inherit":
logger.debug("Using inherited mask.")
if inherited_mask_item is None:
raise_error(
"Cannot inherit mask from the target data. Either the "
"DataGrabber or a Preprocessor does not provide "
"`mask` for the target data type."
)
logger.debug(
f"Inherited mask is in {inherited_mask_item['space']} "
"space."
)
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
else:
# Load mask
logger.debug(f"Loading mask {t_mask}.")
mask_object, _, mask_space = self.load(
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_std_space
# If mask is callable like from nilearn
# If mask is callable like from nilearn; space will be inherit
# so no check for that
if callable(mask_object):
logger.debug("Computing mask (callable).")
if mask_params is None:
mask_params = {}
# 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)
# 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:
mask_img = mask_object(target_data, **mask_params)
mask_img = mask_object(target_img, **mask_params)
# Mask is a Nifti1Image
else:
# Mask params provided
@ -477,33 +518,68 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
"Cannot pass callable params to a non-callable "
"mask."
)
# Resample mask to target image
mask_img = resample_to_img(
source_img=mask_object,
target_img=target_img,
interpolation="nearest",
copy=True,
)
# Convert mask space if required
if mask_space != target_std_space:
mask_img = ANTsMaskWarper().warp(
mask_name=mask_name,
mask_img=mask_img,
src=mask_space,
dst=target_std_space,
target_data=target_data,
warp_data=None,
)
# Resample and warp mask to standard 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_name=mask_name,
mask_img=mask_object,
src=mask_space,
dst=target_std_space,
target_data=target_data,
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)
# Multiple masks, need intersection / union
if len(all_masks) > 1:
# Intersect / union of masks
logger.debug("Intersecting masks.")
mask_img = intersect_masks(all_masks, **intersect_params)
# Single mask
else:
if len(intersect_params) > 0:
if intersect_params:
# Yes, I'm this strict!
raise_error(
"Cannot pass parameters to the intersection function "
@ -511,26 +587,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
)
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

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"])
mask = compute_brain_mask(
target_data=element_data["BOLD"],
extra_input=None,
mask_type=mask_type,
)
assert isinstance(mask, nib.nifti1.Nifti1Image)

View file

@ -400,6 +400,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
"""
# Check pre-requirements for space manipulation
target_space = target_data["space"]
logger.debug(f"Getting {parcellations} in {target_space} space.")
# Extra data type requirement check if target space is native
if target_space == "native":
# Check for extra inputs
@ -416,6 +417,9 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
)
# 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 target standard space to target space
target_std_space = target_space
@ -433,6 +437,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
all_labels = []
for name in parcellations:
# Load parcellation
logger.debug(f"Loading parcellation {name}")
img, labels, _, space = self.load(
name=name,
resolution=resolution,
@ -441,6 +446,9 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
# Convert parcellation spaces if required;
# cannot be "native" due to earlier check
if space != target_std_space:
logger.debug(
f"Warping {name} to {target_std_space} space using ants."
)
raw_img = ANTsParcellationWarper().warp(
parcellation_name=name,
parcellation_img=img,
@ -452,6 +460,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
# Remove extra dimension added by ANTs
img = image.math_img("np.squeeze(img)", img=raw_img)
logger.debug(f"Resampling {name} to target image.")
# Resample parcellation to target image
img_to_merge = image.resample_to_img(
source_img=img,
@ -469,6 +478,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
labels = all_labels[0]
# Parcellations are already transformed to target standard space
else:
logger.debug("Merging parcellations.")
resampled_parcellation_img, labels = merge_parcellations(
parcellations_list=all_parcellations,
parcellations_names=parcellations,
@ -477,6 +487,10 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
# Warp parcellation if target space is 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
if warper_spec["warper"] == "fsl":
resampled_parcellation_img = FSLParcellationWarper().warp(
@ -1194,7 +1208,10 @@ def _retrieve_aicha(
# Load labels
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()
return parcellation_fname, labels