[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",
threshold: float = 0.5,
source: str = "template",
template_space: Optional[str] = None,
extra_input: Optional[dict[str, Any]] = None,
) -> "Nifti1Image":
"""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
subject's data (``VBM_GM`` or ``VBM_WM``). If "template", the mask is
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
The other fields in the data object. Useful for accessing other data
types (default None).
@ -93,6 +97,7 @@ def compute_brain_mask(
If ``mask_type`` is invalid or
if ``source`` 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 ``extra_input`` is None when ``source="subject"`` or
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"]:
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
if target_data["space"] == "native":
# Warp data check
@ -138,15 +146,40 @@ def compute_brain_mask(
)
template = extra_input[key]["data"]
template_space = extra_input[key]["space"]
logger.debug(f"Using {key} in {template_space} for mask computation.")
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
template = get_template(
space=target_std_space,
space=template_space,
target_img=target_data["data"],
extra_input=None,
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
if target_data["space"] == "native" and template_space != "native":
if warp_data["warper"] == "fsl":
@ -168,13 +201,15 @@ def compute_brain_mask(
)
# Resample template to target image
else:
# Resample template to target image
resampled_template = nimg.resample_to_img(
source_img=template, target_img=target_data["data"]
)
# Threshold resampled template and get mask
logger.debug("Thresholding template to get mask.")
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

View file

@ -125,6 +125,7 @@ def get_template(
target_img: nib.Nifti1Image,
extra_input: Optional[dict[str, Any]] = None,
template_type: str = "T1w",
resolution: Optional[Union[int, "str"]] = None,
) -> nib.Nifti1Image:
"""Get template for the space, tailored for the target image.
@ -140,6 +141,10 @@ def get_template(
types (default None).
template_type : {"T1w", "brain", "gm", "wm", "csf"}, optional
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
-------
@ -149,7 +154,8 @@ def get_template(
Raises
------
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
If required template is not found.
@ -162,18 +168,30 @@ def get_template(
if template_type not in ["T1w", "brain", "gm", "wm", "csf"]:
raise_error(f"Unknown template type: {template_type}")
# Get the min of the voxels sizes and use it as the resolution
resolution = np.min(target_img.header.get_zooms()[:3]).astype(int)
if isinstance(resolution, str) and resolution != "highest":
raise_error(
"Invalid resolution value. Must be an integer or 'highest'"
)
# Fetch available resolutions for the template
available_resolutions = [
int(min(val["zooms"]))
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
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
try:
suffix = None