diff --git a/docs/changes/newsfragments/415.change b/docs/changes/newsfragments/415.change new file mode 100644 index 000000000..2e52fdf1f --- /dev/null +++ b/docs/changes/newsfragments/415.change @@ -0,0 +1 @@ +Add ``template_space`` parameter for ``compute_brain_mask`` and ``resolution`` parameter for :func:`.get_template` by `Fede Raimondo`_ diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index 0206f1f6d..150e13bd4 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -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 diff --git a/junifer/data/template_spaces.py b/junifer/data/template_spaces.py index 42b826727..714d6e29c 100644 --- a/junifer/data/template_spaces.py +++ b/junifer/data/template_spaces.py @@ -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