[ENH]: Allow template space request for compute_brain_mask #415
3 changed files with 61 additions and 7 deletions
1
docs/changes/newsfragments/415.change
Normal file
1
docs/changes/newsfragments/415.change
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Add ``template_space`` parameter for ``compute_brain_mask`` and ``resolution`` parameter for :func:`.get_template` by `Fede Raimondo`_
|
||||||
|
|
@ -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
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue