[ENH]: Enable compute_brain_mask to use subject-specific probseg files #404

Merged
fraimondo merged 4 commits from enh/compute_subject_gm_mask into main 2024-12-06 11:39:34 +00:00
3 changed files with 71 additions and 20 deletions

View file

@ -0,0 +1 @@
Add ``source`` parameter and resurrect ``extra_input`` parameter for ``compute_brain_mask`` by `Fede Raimondo`_

View file

@ -0,0 +1 @@
Enable ``compute_brain_mask`` to use subject-specific probseg files by `Fede Raimondo`_

View file

@ -47,6 +47,8 @@ def compute_brain_mask(
warp_data: 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,
source: str = "template",
extra_input: Optional[dict[str, Any]] = None,
) -> "Nifti1Image": ) -> "Nifti1Image":
"""Compute the whole-brain, grey-matter or white-matter mask. """Compute the whole-brain, grey-matter or white-matter mask.
@ -72,6 +74,13 @@ def compute_brain_mask(
(default "brain"). (default "brain").
threshold : float, optional threshold : float, optional
The value under which the template is cut off (default 0.5). The value under which the template is cut off (default 0.5).
source : {"subject", "template"}, optional
The source of the mask. If "subject", the mask is computed from the
subject's data (``VBM_GM`` or ``VBM_WM``). If "template", the mask is
computed from the template data (default "template").
extra_input : dict, optional
The other fields in the data object. Useful for accessing other data
types (default None).
Returns Returns
------- -------
@ -82,7 +91,13 @@ def compute_brain_mask(
------ ------
ValueError ValueError
If ``mask_type`` is invalid or If ``mask_type`` is invalid or
if ``warp_data`` is None when ``target_data``'s space is native. if ``source`` is invalid or
if ``source="subject"`` and ``mask_type`` is invalid or
if ``warp_data`` is None when ``target_data``'s space is native or
if ``extra_input`` is None when ``source="subject"`` or
if ``VBM_GM`` or ``VBM_WM`` data types are not in ``extra_input``
when ``source="subject"`` and ``mask_type`` is ``"gm"`` or ``"wm"``
respectively.
""" """
logger.debug(f"Computing {mask_type} mask") logger.debug(f"Computing {mask_type} mask")
@ -90,6 +105,12 @@ def compute_brain_mask(
if mask_type not in ["brain", "gm", "wm"]: if mask_type not in ["brain", "gm", "wm"]:
raise_error(f"Unknown mask type: {mask_type}") raise_error(f"Unknown mask type: {mask_type}")
if source not in ["subject", "template"]:
raise_error(f"Unknown mask source: {source}")
if source == "subject" and mask_type not in ["gm", "wm"]:
raise_error(f"Unknown mask type: {mask_type} for subject space")
# Check pre-requirements for space manipulation # Check pre-requirements for space manipulation
if target_data["space"] == "native": if target_data["space"] == "native":
# Warp data check # Warp data check
@ -101,25 +122,50 @@ def compute_brain_mask(
# Set space to fetch template using # Set space to fetch template using
target_std_space = target_data["space"] target_std_space = target_data["space"]
# Fetch template in closest resolution if source == "subject":
template = get_template( key = f"VBM_{mask_type.upper()}"
space=target_std_space, # Check for extra inputs
target_img=target_data["data"], if extra_input is None:
extra_input=None, raise_error(
template_type=mask_type, f"No extra input provided, requires `{key}` "
) "data type to infer target template data and space."
)
# Resample and warp template if target space is native # Check for missing data type
if target_data["space"] == "native": if key not in extra_input:
resampled_template = ANTsMaskWarper().warp( raise_error(
mask_name=f"template_{target_std_space}_for_compute_brain_mask", f"Cannot compute {mask_type} from subject's data. "
# use template here f"Missing {key} in extra input."
mask_img=template, )
src=target_std_space, template = extra_input[key]["data"]
dst="native", template_space = extra_input[key]["space"]
target_data=target_data, else:
warp_data=warp_data, # Fetch template in closest resolution
template = get_template(
space=target_std_space,
target_img=target_data["data"],
extra_input=None,
template_type=mask_type,
) )
template_space = target_std_space
# Resample and warp template if target space is native
if target_data["space"] == "native" and template_space != "native":
if warp_data["warper"] == "fsl":
resampled_template = FSLMaskWarper().warp(
mask_name=f"template_{target_std_space}_for_compute_brain_mask",
mask_img=template,
target_data=target_data,
warp_data=warp_data,
)
elif warp_data["warper"] == "ants":
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: else:
resampled_template = resample_to_img( resampled_template = resample_to_img(
@ -503,7 +549,10 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
# custom compute_brain_mask # custom compute_brain_mask
elif mask_name == "compute_brain_mask": elif mask_name == "compute_brain_mask":
mask_img = mask_object( mask_img = mask_object(
target_data, warper_spec, **mask_params target_data=target_data,
warp_data=warper_spec,
extra_input=extra_input,
**mask_params,
) )
# custom registered; arm kept for clarity # custom registered; arm kept for clarity
else: else: