[ENH]: Allow template space request for compute_brain_mask #415

Merged
synchon merged 4 commits from enh/gm_template_space into main 2024-12-09 11:15:17 +00:00
3 changed files with 61 additions and 7 deletions

View file

@ -0,0 +1 @@
Add ``template_space`` parameter for ``compute_brain_mask`` and ``resolution`` parameter for :func:`.get_template` by `Fede Raimondo`_

View file

@ -48,6 +48,7 @@ def compute_brain_mask(
mask_type: str = "brain", mask_type: str = "brain",
threshold: float = 0.5, threshold: float = 0.5,
source: str = "template", source: str = "template",
template_space: Optional[str] = None,
extra_input: Optional[dict[str, Any]] = None, 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.
@ -78,6 +79,9 @@ def compute_brain_mask(
The source of the mask. If "subject", the mask is computed from the 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 subject's data (``VBM_GM`` or ``VBM_WM``). If "template", the mask is
computed from the template data (default "template"). computed from the template data (default "template").
template_space : str, optional
The space of the template. If not provided, the space is inferred from
the ``target_data`` (default None).
extra_input : dict, optional extra_input : dict, optional
The other fields in the data object. Useful for accessing other data The other fields in the data object. Useful for accessing other data
types (default None). types (default None).
@ -93,6 +97,7 @@ def compute_brain_mask(
If ``mask_type`` is invalid or If ``mask_type`` is invalid or
if ``source`` is invalid or if ``source`` is invalid or
if ``source="subject"`` and ``mask_type`` is invalid or if ``source="subject"`` and ``mask_type`` is invalid or
if ``template_space`` is provided when ``source="subject"`` or
if ``warp_data`` is None when ``target_data``'s space is native or if ``warp_data`` is None when ``target_data``'s space is native or
if ``extra_input`` is None when ``source="subject"`` or if ``extra_input`` is None when ``source="subject"`` or
if ``VBM_GM`` or ``VBM_WM`` data types are not in ``extra_input`` if ``VBM_GM`` or ``VBM_WM`` data types are not in ``extra_input``
@ -111,6 +116,9 @@ def compute_brain_mask(
if source == "subject" and mask_type not in ["gm", "wm"]: if source == "subject" and mask_type not in ["gm", "wm"]:
raise_error(f"Unknown mask type: {mask_type} for subject space") raise_error(f"Unknown mask type: {mask_type} for subject space")
if source == "subject" and template_space is not None:
raise_error("Cannot provide `template_space` when source is `subject`")
# 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
@ -138,15 +146,40 @@ def compute_brain_mask(
) )
template = extra_input[key]["data"] template = extra_input[key]["data"]
template_space = extra_input[key]["space"] template_space = extra_input[key]["space"]
logger.debug(f"Using {key} in {template_space} for mask computation.")
else: else:
template_resolution = None
if template_space is None:
template_space = target_std_space
elif template_space != target_std_space:
# We re going to warp, so get the highest resolution
template_resolution = "highest"
# Fetch template in closest resolution # Fetch template in closest resolution
template = get_template( template = get_template(
space=target_std_space, space=template_space,
target_img=target_data["data"], target_img=target_data["data"],
extra_input=None, extra_input=None,
template_type=mask_type, template_type=mask_type,
resolution=template_resolution,
) )
template_space = target_std_space
mask_name = f"template_{target_std_space}_for_compute_brain_mask"
# Warp template to correct space (MNI to MNI)
if template_space != "native" and template_space != target_std_space:
logger.debug(
f"Warping template to {target_std_space} space using ANTs."
)
template = ANTsMaskWarper().warp(
mask_name=mask_name,
mask_img=template,
src=template_space,
dst=target_std_space,
target_data=target_data,
warp_data=None,
)
# Resample and warp template if target space is native # Resample and warp template if target space is native
if target_data["space"] == "native" and template_space != "native": if target_data["space"] == "native" and template_space != "native":
if warp_data["warper"] == "fsl": if warp_data["warper"] == "fsl":
@ -168,13 +201,15 @@ def compute_brain_mask(
) )
# Resample template to target image # Resample template to target image
else: else:
# Resample template to target image
resampled_template = nimg.resample_to_img( resampled_template = nimg.resample_to_img(
source_img=template, target_img=target_data["data"] source_img=template, target_img=target_data["data"]
) )
# Threshold resampled template and get mask # Threshold resampled template and get mask
logger.debug("Thresholding template to get mask.")
mask = (nimg.get_data(resampled_template) >= threshold).astype("int8") mask = (nimg.get_data(resampled_template) >= threshold).astype("int8")
logger.debug("Mask computation from brain template complete.")
return nimg.new_img_like(target_data["data"], mask) # type: ignore return nimg.new_img_like(target_data["data"], mask) # type: ignore

View file

@ -125,6 +125,7 @@ def get_template(
target_img: nib.Nifti1Image, target_img: nib.Nifti1Image,
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
template_type: str = "T1w", template_type: str = "T1w",
resolution: Optional[Union[int, "str"]] = None,
) -> nib.Nifti1Image: ) -> nib.Nifti1Image:
"""Get template for the space, tailored for the target image. """Get template for the space, tailored for the target image.
@ -140,6 +141,10 @@ def get_template(
types (default None). types (default None).
template_type : {"T1w", "brain", "gm", "wm", "csf"}, optional template_type : {"T1w", "brain", "gm", "wm", "csf"}, optional
The template type to retrieve (default "T1w"). The template type to retrieve (default "T1w").
resolution : int or "highest", optional
The resolution of the template to fetch. If None, the closest
resolution to the target image is used (default None). If "highest",
the highest resolution is used.
Returns Returns
------- -------
@ -149,7 +154,8 @@ def get_template(
Raises Raises
------ ------
ValueError ValueError
If ``space`` or ``template_type`` is invalid. If ``space`` or ``template_type`` is invalid or
if ``resolution`` is not at int or "highest".
RuntimeError RuntimeError
If required template is not found. If required template is not found.
@ -162,18 +168,30 @@ def get_template(
if template_type not in ["T1w", "brain", "gm", "wm", "csf"]: if template_type not in ["T1w", "brain", "gm", "wm", "csf"]:
raise_error(f"Unknown template type: {template_type}") raise_error(f"Unknown template type: {template_type}")
# Get the min of the voxels sizes and use it as the resolution if isinstance(resolution, str) and resolution != "highest":
resolution = np.min(target_img.header.get_zooms()[:3]).astype(int) raise_error(
"Invalid resolution value. Must be an integer or 'highest'"
)
# Fetch available resolutions for the template # Fetch available resolutions for the template
available_resolutions = [ available_resolutions = [
int(min(val["zooms"])) int(min(val["zooms"]))
for val in tflow.get_metadata(space)["res"].values() for val in tflow.get_metadata(space)["res"].values()
] ]
# Get the min of the voxels sizes and use it as the resolution
if resolution is None:
resolution = np.min(target_img.header.get_zooms()[:3]).astype(int)
elif resolution == "highest":
resolution = 0
# Use the closest resolution if desired resolution is not found # Use the closest resolution if desired resolution is not found
resolution = closest_resolution(resolution, available_resolutions) resolution = closest_resolution(resolution, available_resolutions)
logger.info(f"Downloading template {space} in resolution {resolution}") logger.info(
f"Downloading template {space} ({template_type} in "
f"resolution {resolution}"
)
# Retrieve template # Retrieve template
try: try:
suffix = None suffix = None