diff --git a/docs/changes/newsfragments/293.enh b/docs/changes/newsfragments/293.enh new file mode 100644 index 000000000..a9084c0cf --- /dev/null +++ b/docs/changes/newsfragments/293.enh @@ -0,0 +1 @@ +Adapt :class:`.BOLDWarper` to use FSL or ANTs depending on warp file extension by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/293.feature b/docs/changes/newsfragments/293.feature new file mode 100644 index 000000000..7d6f01a28 --- /dev/null +++ b/docs/changes/newsfragments/293.feature @@ -0,0 +1 @@ +Introduce ``junifer.preprocess.ants.ants_apply_transforms_warper._AntsApplyTransformsWarper`` to wrap ANTs' ``antsApplyTransforms`` by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/293.misc b/docs/changes/newsfragments/293.misc new file mode 100644 index 000000000..5196cadcf --- /dev/null +++ b/docs/changes/newsfragments/293.misc @@ -0,0 +1 @@ +Add support for accessing ANTs' ``ResampleImage`` via Docker wrapper by `Synchon Mandal`_ diff --git a/ignore_words.txt b/ignore_words.txt index d333dfe52..0760787ff 100644 --- a/ignore_words.txt +++ b/ignore_words.txt @@ -2,4 +2,5 @@ master nin chang sepulcre -arange \ No newline at end of file +arange +sinc diff --git a/junifer/api/res/ants/ResampleImage b/junifer/api/res/ants/ResampleImage new file mode 100644 index 000000000..42accbcae --- /dev/null +++ b/junifer/api/res/ants/ResampleImage @@ -0,0 +1,3 @@ +#!/bin/bash + +run_ants_docker.sh ResampleImage "$@" diff --git a/junifer/data/coordinates.py b/junifer/data/coordinates.py index b702d0a69..752f3c4a0 100644 --- a/junifer/data/coordinates.py +++ b/junifer/data/coordinates.py @@ -251,12 +251,6 @@ def get_coordinates( # Create component-scoped tempdir tempdir = WorkDirManager().get_tempdir(prefix="coordinates") - # Save existing coordinates to a component-scoped tempfile - pretransform_coordinates_path = ( - tempdir / "pretransform_coordinates.txt" - ) - np.savetxt(pretransform_coordinates_path, seeds) - # Create element-scoped tempdir so that transformed coordinates is # available later as numpy stores file path reference for # loading on computation @@ -264,47 +258,140 @@ def get_coordinates( prefix="coordinates" ) - # Create an element-scoped tempfile for transformed coordinates output - img2imgcoord_out_path = element_tempdir / "coordinates_transformed.txt" - # Set img2imgcoord command - img2imgcoord_cmd = [ - "cat", - f"{pretransform_coordinates_path.resolve()}", - "| img2imgcoord -mm", - f"-src {target_data['path'].resolve()}", - f"-dest {target_data['reference_path'].resolve()}", - f"-warp {extra_input['Warp']['path'].resolve()}", - f"> {img2imgcoord_out_path.resolve()};", - f"sed -i 1d {img2imgcoord_out_path.resolve()}", - ] - # Call img2imgcoord - img2imgcoord_cmd_str = " ".join(img2imgcoord_cmd) - logger.info( - f"img2imgcoord command to be executed: {img2imgcoord_cmd_str}" - ) - img2imgcoord_process = subprocess.run( - img2imgcoord_cmd_str, # string needed with shell=True - stdin=subprocess.DEVNULL, - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - shell=True, # needed for respecting $PATH - check=False, - ) - # Check for success or failure - if img2imgcoord_process.returncode == 0: - logger.info( - "img2imgcoord succeeded with the following output: " - f"{img2imgcoord_process.stdout}" + # Check for warp file type to use correct tool + warp_file_ext = extra_input["Warp"]["path"].suffix + if warp_file_ext == ".mat": + # Save existing coordinates to a component-scoped tempfile + pretransform_coordinates_path = ( + tempdir / "pretransform_coordinates.txt" ) - else: - raise_error( - msg="img2imgcoord failed with the following error: " - f"{img2imgcoord_process.stdout}", - klass=RuntimeError, + np.savetxt(pretransform_coordinates_path, seeds) + + # Create an element-scoped tempfile for transformed coordinates + # output + transformed_coords_path = ( + element_tempdir / "coordinates_transformed.txt" ) - # Load coordinates - seeds = np.loadtxt(img2imgcoord_out_path) + logger.debug("Using FSL for coordinates transformation") + # Set img2imgcoord command + img2imgcoord_cmd = [ + "cat", + f"{pretransform_coordinates_path.resolve()}", + "| img2imgcoord -mm", + f"-src {target_data['path'].resolve()}", + f"-dest {target_data['reference_path'].resolve()}", + f"-warp {extra_input['Warp']['path'].resolve()}", + f"> {transformed_coords_path.resolve()};", + f"sed -i 1d {transformed_coords_path.resolve()}", + ] + # Call img2imgcoord + img2imgcoord_cmd_str = " ".join(img2imgcoord_cmd) + logger.info( + f"img2imgcoord command to be executed: {img2imgcoord_cmd_str}" + ) + img2imgcoord_process = subprocess.run( + img2imgcoord_cmd_str, # string needed with shell=True + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + # Check for success or failure + if img2imgcoord_process.returncode == 0: + logger.info( + "img2imgcoord succeeded with the following output: " + f"{img2imgcoord_process.stdout}" + ) + else: + raise_error( + msg="img2imgcoord failed with the following error: " + f"{img2imgcoord_process.stdout}", + klass=RuntimeError, + ) + + # Load coordinates + seeds = np.loadtxt(transformed_coords_path) + + elif warp_file_ext == ".h5": + # Save existing coordinates to a component-scoped tempfile + pretransform_coordinates_path = ( + tempdir / "pretransform_coordinates.csv" + ) + np.savetxt( + pretransform_coordinates_path, + seeds, + delimiter=",", + # Add header while saving to make ANTs work + header="x,y,z", + ) + + # Create an element-scoped tempfile for transformed coordinates + # output + transformed_coords_path = ( + element_tempdir / "coordinates_transformed.csv" + ) + + logger.debug("Using ANTs for coordinates transformation") + # Set antsApplyTransformsToPoints command + apply_transforms_to_points_cmd = [ + "antsApplyTransformsToPoints", + "-d 3", + "-p 1", + "-f 0", + f"-i {pretransform_coordinates_path.resolve()}", + f"-o {transformed_coords_path.resolve()}", + f"-t {extra_input['Warp']['path'].resolve()};", + ] + # Call antsApplyTransformsToPoints + apply_transforms_to_points_cmd_str = " ".join( + apply_transforms_to_points_cmd + ) + logger.info( + "antsApplyTransformsToPoints command to be executed: " + f"{apply_transforms_to_points_cmd_str}" + ) + apply_transforms_to_points_process = subprocess.run( + # string needed with shell=True + apply_transforms_to_points_cmd_str, + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + if apply_transforms_to_points_process.returncode == 0: + logger.info( + "antsApplyTransformsToPoints succeeded with the following " + f"output: {apply_transforms_to_points_process.stdout}" + ) + else: + raise_error( + msg=( + "antsApplyTransformsToPoints failed with the " + "following error: " + f"{apply_transforms_to_points_process.stdout}" + ), + klass=RuntimeError, + ) + + # Load coordinates + seeds = np.loadtxt( + # Skip header when reading + transformed_coords_path, + delimiter=",", + skiprows=1, + ) + + else: + raise_error( + msg=( + "Unknown warp / transformation file extension: " + f"{warp_file_ext}" + ), + klass=RuntimeError, + ) # Delete tempdir WorkDirManager().delete_tempdir(tempdir) diff --git a/junifer/data/masks.py b/junifer/data/masks.py index e3cabc6e5..24d88ee57 100644 --- a/junifer/data/masks.py +++ b/junifer/data/masks.py @@ -1,6 +1,7 @@ """Provide functions for masks.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL import subprocess @@ -199,7 +200,9 @@ def get_mask( # noqa: C901 ------ RuntimeError If masks are in different spaces and they need to be intersected / - unionized. + unionized or + if warp / transformation file extension is not ".mat" or ".h5" or + if external tool execution failed. ValueError If extra key is provided in addition to mask name in ``masks`` or if no mask is provided or @@ -364,50 +367,104 @@ def get_mask( # noqa: C901 element_tempdir = WorkDirManager().get_element_tempdir(prefix="masks") # Create an element-scoped tempfile for warped output - applywarp_out_path = element_tempdir / "mask_warped.nii.gz" - # Set applywarp command - applywarp_cmd = [ - "applywarp", - "--interp=nn", - f"-i {prewarp_mask_path.resolve()}", - # use resampled reference - f"-r {target_data['reference_path'].resolve()}", - f"-w {extra_input['Warp']['path'].resolve()}", - f"-o {applywarp_out_path.resolve()}", - ] - # Call applywarp - applywarp_cmd_str = " ".join(applywarp_cmd) - logger.info(f"applywarp command to be executed: {applywarp_cmd_str}") - applywarp_process = subprocess.run( - applywarp_cmd_str, # string needed with shell=True - stdin=subprocess.DEVNULL, - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - shell=True, # needed for respecting $PATH - check=False, - ) - # Check for success or failure - if applywarp_process.returncode == 0: + warped_mask_path = element_tempdir / "mask_warped.nii.gz" + + # Check for warp file type to use correct tool + warp_file_ext = extra_input["Warp"]["path"].suffix + if warp_file_ext == ".mat": + logger.debug("Using FSL for mask warping") + # Set applywarp command + applywarp_cmd = [ + "applywarp", + "--interp=nn", + f"-i {prewarp_mask_path.resolve()}", + # use resampled reference + f"-r {target_data['reference_path'].resolve()}", + f"-w {extra_input['Warp']['path'].resolve()}", + f"-o {warped_mask_path.resolve()}", + ] + # Call applywarp + applywarp_cmd_str = " ".join(applywarp_cmd) logger.info( - "applywarp succeeded with the following output: " - f"{applywarp_process.stdout}" + f"applywarp command to be executed: {applywarp_cmd_str}" ) + applywarp_process = subprocess.run( + applywarp_cmd_str, # string needed with shell=True + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + # Check for success or failure + if applywarp_process.returncode == 0: + logger.info( + "applywarp succeeded with the following output: " + f"{applywarp_process.stdout}" + ) + else: + raise_error( + msg="applywarp failed with the following error: " + f"{applywarp_process.stdout}", + klass=RuntimeError, + ) + elif warp_file_ext == ".h5": + logger.debug("Using ANTs for mask warping") + # Set antsApplyTransforms command + apply_transforms_cmd = [ + "antsApplyTransforms", + "-d 3", + "-e 3", + "-n 'GenericLabel[NearestNeighbor]'", + f"-i {prewarp_mask_path.resolve()}", + # use resampled reference + f"-r {target_data['reference_path'].resolve()}", + f"-t {extra_input['Warp']['path'].resolve()}", + f"-o {warped_mask_path.resolve()}", + ] + # Call antsApplyTransforms + apply_transforms_cmd_str = " ".join(apply_transforms_cmd) + logger.info( + "antsApplyTransforms command to be executed: " + f"{apply_transforms_cmd_str}" + ) + apply_transforms_process = subprocess.run( + apply_transforms_cmd_str, # string needed with shell=True + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + if apply_transforms_process.returncode == 0: + logger.info( + "antsApplyTransforms succeeded with the following output: " + f"{apply_transforms_process.stdout}" + ) + else: + raise_error( + msg=( + "antsApplyTransforms failed with the following error: " + f"{apply_transforms_process.stdout}" + ), + klass=RuntimeError, + ) else: raise_error( - msg="applywarp failed with the following error: " - f"{applywarp_process.stdout}", + msg=( + "Unknown warp / transformation file extension: " + f"{warp_file_ext}" + ), klass=RuntimeError, ) # Load nifti - mask_img = nib.load(applywarp_out_path) + mask_img = nib.load(warped_mask_path) # Delete tempdir WorkDirManager().delete_tempdir(tempdir) - # Type-cast to remove errors - mask_img = typing.cast("Nifti1Image", mask_img) - return mask_img + return mask_img # type: ignore def load_mask( diff --git a/junifer/data/parcellations.py b/junifer/data/parcellations.py index 3442efd37..e33c36340 100644 --- a/junifer/data/parcellations.py +++ b/junifer/data/parcellations.py @@ -231,7 +231,9 @@ def get_parcellation( Raises ------ RuntimeError - If parcellations are in different spaces and they need to be merged. + If parcellations are in different spaces and they need to be merged or + if warp / transformation file extension is not ".mat" or ".h5" or + if external tool execution failed. ValueError If ``extra_input`` is None when ``target_data``'s space is native. @@ -302,55 +304,107 @@ def get_parcellation( element_tempdir = WorkDirManager().get_element_tempdir( prefix="parcellations" ) - # Create an element-scoped tempfile for warped output - applywarp_out_path = element_tempdir / "parcellation_warped.nii.gz" - # Set applywarp command - applywarp_cmd = [ - "applywarp", - "--interp=nn", - f"-i {prewarp_parcellation_path.resolve()}", - # use resampled reference - f"-r {target_data['reference_path'].resolve()}", - f"-w {extra_input['Warp']['path'].resolve()}", - f"-o {applywarp_out_path.resolve()}", - ] - # Call applywarp - applywarp_cmd_str = " ".join(applywarp_cmd) - logger.info(f"applywarp command to be executed: {applywarp_cmd_str}") - applywarp_process = subprocess.run( - applywarp_cmd_str, # string needed with shell=True - stdin=subprocess.DEVNULL, - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - shell=True, # needed for respecting $PATH - check=False, + warped_parcellation_path = ( + element_tempdir / "parcellation_warped.nii.gz" ) - # Check for success or failure - if applywarp_process.returncode == 0: + + # Check for warp file type to use correct tool + warp_file_ext = extra_input["Warp"]["path"].suffix + if warp_file_ext == ".mat": + logger.debug("Using FSL for parcellation warping") + # Set applywarp command + applywarp_cmd = [ + "applywarp", + "--interp=nn", + f"-i {prewarp_parcellation_path.resolve()}", + # use resampled reference + f"-r {target_data['reference_path'].resolve()}", + f"-w {extra_input['Warp']['path'].resolve()}", + f"-o {warped_parcellation_path.resolve()}", + ] + # Call applywarp + applywarp_cmd_str = " ".join(applywarp_cmd) logger.info( - "applywarp succeeded with the following output: " - f"{applywarp_process.stdout}" + f"applywarp command to be executed: {applywarp_cmd_str}" ) + applywarp_process = subprocess.run( + applywarp_cmd_str, # string needed with shell=True + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + # Check for success or failure + if applywarp_process.returncode == 0: + logger.info( + "applywarp succeeded with the following output: " + f"{applywarp_process.stdout}" + ) + else: + raise_error( + msg="applywarp failed with the following error: " + f"{applywarp_process.stdout}", + klass=RuntimeError, + ) + elif warp_file_ext == ".h5": + logger.debug("Using ANTs for parcellation warping") + # Set antsApplyTransforms command + apply_transforms_cmd = [ + "antsApplyTransforms", + "-d 3", + "-e 3", + "-n 'GenericLabel[NearestNeighbor]'", + f"-i {prewarp_parcellation_path.resolve()}", + # use resampled reference + f"-r {target_data['reference_path'].resolve()}", + f"-t {extra_input['Warp']['path'].resolve()}", + f"-o {warped_parcellation_path.resolve()}", + ] + # Call antsApplyTransforms + apply_transforms_cmd_str = " ".join(apply_transforms_cmd) + logger.info( + "antsApplyTransforms command to be executed: " + f"{apply_transforms_cmd_str}" + ) + apply_transforms_process = subprocess.run( + apply_transforms_cmd_str, # string needed with shell=True + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + if apply_transforms_process.returncode == 0: + logger.info( + "antsApplyTransforms succeeded with the following output: " + f"{apply_transforms_process.stdout}" + ) + else: + raise_error( + msg=( + "antsApplyTransforms failed with the following error: " + f"{apply_transforms_process.stdout}" + ), + klass=RuntimeError, + ) else: raise_error( - msg="applywarp failed with the following error: " - f"{applywarp_process.stdout}", + msg=( + "Unknown warp / transformation file extension: " + f"{warp_file_ext}" + ), klass=RuntimeError, ) # Load nifti - resampled_parcellation_img = nib.load(applywarp_out_path) + resampled_parcellation_img = nib.load(warped_parcellation_path) # Delete tempdir WorkDirManager().delete_tempdir(tempdir) - # Stupid casting - resampled_parcellation_img = typing.cast( - "Nifti1Image", resampled_parcellation_img - ) - - return resampled_parcellation_img, labels + return resampled_parcellation_img, labels # type: ignore def load_parcellation( diff --git a/junifer/preprocess/ants/__init__.py b/junifer/preprocess/ants/__init__.py new file mode 100644 index 000000000..4ffb8cbc3 --- /dev/null +++ b/junifer/preprocess/ants/__init__.py @@ -0,0 +1,4 @@ +"""Provide imports for ants sub-package.""" + +# Authors: Synchon Mandal +# License: AGPL diff --git a/junifer/preprocess/ants/ants_apply_transforms_warper.py b/junifer/preprocess/ants/ants_apply_transforms_warper.py new file mode 100644 index 000000000..8bf2266df --- /dev/null +++ b/junifer/preprocess/ants/ants_apply_transforms_warper.py @@ -0,0 +1,280 @@ +"""Provide class for warping via ANTs antsApplyTransforms.""" + +# Authors: Synchon Mandal +# License: AGPL + +import subprocess +from pathlib import Path +from typing import ( + TYPE_CHECKING, + Any, + ClassVar, + Dict, + List, + Optional, + Tuple, + Union, + cast, +) + +import nibabel as nib +import numpy as np + +from ...pipeline import WorkDirManager +from ...utils import logger, raise_error +from ..base import BasePreprocessor + + +if TYPE_CHECKING: + from nibabel import Nifti1Image + + +class _AntsApplyTransformsWarper(BasePreprocessor): + """Class for warping NIfTI images via ANTs antsApplyTransforms. + + Warps ANTs ``antsApplyTransforms``. + + Parameters + ---------- + reference : str + The data type to use as reference for warping. + on : str + The data type to use for warping. + + Raises + ------ + ValueError + If a list was passed for ``on``. + + """ + + _EXT_DEPENDENCIES: ClassVar[ + List[Dict[str, Union[str, bool, List[str]]]] + ] = [ + { + "name": "ants", + "optional": False, + "commands": ["ResampleImage", "antsApplyTransforms"], + }, + ] + + def __init__(self, reference: str, on: str) -> None: + """Initialize the class.""" + self.ref = reference + # Check only single data type is passed + if isinstance(on, list): + raise_error("Can only work on single data type, list was passed.") + self.on = on # needed for the base validation to work + super().__init__( + on=self.on, required_data_types=[self.on, self.ref, "Warp"] + ) + + def get_valid_inputs(self) -> List[str]: + """Get valid data types for input. + + Returns + ------- + list of str + The list of data types that can be used as input for this + preprocessor. + + """ + # Constructed dynamically + return [self.on] + + def get_output_type(self, input: List[str]) -> List[str]: + """Get output type. + + Parameters + ---------- + input : list of str + The input to the preprocessor. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + list of str + The updated list of available Junifer Data object keys after + the pipeline step. + + """ + # Does not add any new keys + return input + + def _run_apply_transforms( + self, + input_data: Dict, + ref_path: Path, + warp_path: Path, + ) -> Tuple["Nifti1Image", Path]: + """Run ``antsApplyTransforms``. + + Parameters + ---------- + input_data : dict + The input data. + ref_path : pathlib.Path + The path to the reference file. + warp_path : pathlib.Path + The path to the warp file. + + Returns + ------- + Niimg-like object + The warped input image. + pathlib.Path + The path to the resampled reference image. + + Raises + ------ + RuntimeError + If ANTs command fails. + + """ + # Get the min of the voxel sizes from input and use it as the + # resolution + resolution = np.min(input_data["data"].header.get_zooms()[:3]) + + # Create element-specific tempdir for storing post-warping assets + tempdir = WorkDirManager().get_element_tempdir( + prefix="applytransforms" + ) + + # Create a tempfile for resampled reference output + resample_image_out_path = tempdir / "reference_resampled.nii.gz" + # Set ResampleImage command + resample_image_cmd = [ + "ResampleImage", + "3", # image dimension + f"{ref_path.resolve()}", + f"{resample_image_out_path.resolve()}", + f"{resolution}x{resolution}x{resolution}", + "0", # option for spacing and not size + "3 3", # Lanczos windowed sinc + ] + # Call ResampleImage + resample_image_cmd_str = " ".join(resample_image_cmd) + logger.info( + f"ResampleImage command to be executed: {resample_image_cmd_str}" + ) + resample_image_process = subprocess.run( + resample_image_cmd_str, + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + if resample_image_process.returncode == 0: + logger.info( + "ResampleImage succeeded with the following output: " + f"{resample_image_process.stdout}" + ) + else: + raise_error( + msg="ResampleImage failed with the following error: " + f"{resample_image_process.stdout}", + klass=RuntimeError, + ) + + # Create a tempfile for warped output + apply_transforms_out_path = tempdir / "input_warped.nii.gz" + # Set antsApplyTransforms command + apply_transforms_cmd = [ + "antsApplyTransforms", + "-d 3", + "-e 3", + "-n LanczosWindowedSinc", + f"-i {input_data['path'].resolve()}", + # use resampled reference + f"-r {resample_image_out_path.resolve()}", + f"-t {warp_path.resolve()}", + f"-o {apply_transforms_out_path.resolve()}", + ] + # Call antsApplyTransforms + apply_transforms_cmd_str = " ".join(apply_transforms_cmd) + logger.info( + "antsApplyTransforms command to be executed: " + f"{apply_transforms_cmd_str}" + ) + apply_transforms_process = subprocess.run( + apply_transforms_cmd_str, # string needed with shell=True + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + shell=True, # needed for respecting $PATH + check=False, + ) + if apply_transforms_process.returncode == 0: + logger.info( + "antsApplyTransforms succeeded with the following output: " + f"{apply_transforms_process.stdout}" + ) + else: + raise_error( + msg=( + "antsApplyTransforms failed with the following error: " + f"{apply_transforms_process.stdout}" + ), + klass=RuntimeError, + ) + + # Load nifti + output_img = nib.load(apply_transforms_out_path) + + # Stupid casting + output_img = cast("Nifti1Image", output_img) + return output_img, resample_image_out_path + + def preprocess( + self, + input: Dict[str, Any], + extra_input: Optional[Dict[str, Any]] = None, + ) -> Tuple[str, Dict[str, Any]]: + """Preprocess. + + Parameters + ---------- + input : dict + A single input from the Junifer Data object in which to preprocess. + extra_input : dict, optional + The other fields in the Junifer Data object. Must include the + ``Warp`` and ``ref`` value's keys. + + Returns + ------- + str + The key to store the output in the Junifer Data object. + dict + The computed result as dictionary. This will be stored in the + Junifer Data object under the key ``data`` of the data type. + + Raises + ------ + ValueError + If ``extra_input`` is None. + + """ + logger.debug("Warping via ANTs using antsApplyTransforms") + # Check for extra inputs + if extra_input is None: + raise_error( + f"No extra input provided, requires `Warp` and `{self.ref}` " + "data types in particular." + ) + # Retrieve data type info to warp + to_warp_input = input + # Retrieve data type info to use as reference + ref_input = extra_input[self.ref] + # Retrieve Warp data + warp = extra_input["Warp"] + # Replace original data with warped data and add resampled reference + # path + input["data"], input["reference_path"] = self._run_apply_transforms( + input_data=to_warp_input, + ref_path=ref_input["path"], + warp_path=warp["path"], + ) + # Use reference input's space as warped input's space + input["space"] = ref_input["space"] + return self.on, input diff --git a/junifer/preprocess/ants/tests/test_ants_apply_transforms_warper.py b/junifer/preprocess/ants/tests/test_ants_apply_transforms_warper.py new file mode 100644 index 000000000..68b43c648 --- /dev/null +++ b/junifer/preprocess/ants/tests/test_ants_apply_transforms_warper.py @@ -0,0 +1,124 @@ +"""Provide tests for AntsApplyTransformsWarper.""" + +# Authors: Synchon Mandal +# License: AGPL + +import socket +from pathlib import Path +from typing import List + +import nibabel as nib +import pytest + +from junifer.datagrabber import DMCC13Benchmark +from junifer.datareader import DefaultDataReader +from junifer.pipeline.utils import _check_ants +from junifer.preprocess.ants.ants_apply_transforms_warper import ( + _AntsApplyTransformsWarper, +) + + +def test_AntsApplyTransformsWarper_init() -> None: + """Test AntsApplyTransformsWarper init.""" + ants_apply_transforms_warper = _AntsApplyTransformsWarper( + reference="T1w", on="BOLD" + ) + assert ants_apply_transforms_warper.ref == "T1w" + assert ants_apply_transforms_warper.on == "BOLD" + assert ants_apply_transforms_warper._on == ["BOLD"] + + +def test_AntsApplyTransformsWarper_get_valid_inputs() -> None: + """Test AntsApplyTransformsWarper get_valid_inputs.""" + ants_apply_transforms_warper = _AntsApplyTransformsWarper( + reference="T1w", on="BOLD" + ) + assert ants_apply_transforms_warper.get_valid_inputs() == ["BOLD"] + + +@pytest.mark.parametrize( + "input_", + [ + ["BOLD", "T1w", "Warp"], + ["BOLD", "T1w"], + ["BOLD"], + ], +) +def test_AntsApplyTransformsWarper_get_output_type(input_: List[str]) -> None: + """Test AntsApplyTransformsWarper get_output_type. + + Parameters + ---------- + input_ : list of str + The input data types. + + """ + ants_apply_transforms_warper = _AntsApplyTransformsWarper( + reference="T1w", on="BOLD" + ) + assert ants_apply_transforms_warper.get_output_type(input_) == input_ + + +@pytest.mark.skipif( + _check_ants() is False, reason="requires ANTs to be in PATH" +) +@pytest.mark.skipif( + socket.gethostname() != "juseless", + reason="only for juseless", +) +def test_AntsApplyTransformsWarper__run_apply_transform() -> None: + """Test AntsApplyTransformsWarper _run_apply_transform.""" + with DMCC13Benchmark( + types=["BOLD", "T1w", "Warp"], + sessions=["wave1bas"], + tasks=["Rest"], + phase_encodings=["AP"], + runs=["1"], + native_t1w=True, + ) as dg: + # Read data + element_data = DefaultDataReader().fit_transform( + dg[("f9057kp", "wave1bas", "Rest", "AP", "1")] + ) + # Preprocess data + warped_data, resampled_ref_path = _AntsApplyTransformsWarper( + reference="T1w", on="BOLD" + )._run_apply_transforms( + input_data=element_data["BOLD"], + ref_path=element_data["T1w"]["path"], + warp_path=element_data["Warp"]["path"], + ) + assert isinstance(warped_data, nib.Nifti1Image) + assert isinstance(resampled_ref_path, Path) + + +@pytest.mark.skipif( + _check_ants() is False, reason="requires ANTs to be in PATH" +) +@pytest.mark.skipif( + socket.gethostname() != "juseless", + reason="only for juseless", +) +def test_AntsApplyTransformsWarper_preprocess() -> None: + """Test AntsApplyTransformsWarper preprocess.""" + with DMCC13Benchmark( + types=["BOLD", "T1w", "Warp"], + sessions=["wave1bas"], + tasks=["Rest"], + phase_encodings=["AP"], + runs=["1"], + native_t1w=True, + ) as dg: + # Read data + element_data = DefaultDataReader().fit_transform( + dg[("f9057kp", "wave1bas", "Rest", "AP", "1")] + ) + # Preprocess data + data_type, data = _AntsApplyTransformsWarper( + reference="T1w", on="BOLD" + ).preprocess( + input=element_data["BOLD"], + extra_input=element_data, + ) + assert isinstance(data_type, str) + assert isinstance(data, dict) diff --git a/junifer/preprocess/bold_warper.py b/junifer/preprocess/bold_warper.py index ee8638fe8..432fb6ff1 100644 --- a/junifer/preprocess/bold_warper.py +++ b/junifer/preprocess/bold_warper.py @@ -15,6 +15,7 @@ from typing import ( from ..api.decorators import register_preprocessor from ..utils import logger, raise_error +from .ants.ants_apply_transforms_warper import _AntsApplyTransformsWarper from .base import BasePreprocessor from .fsl.apply_warper import _ApplyWarper @@ -35,8 +36,13 @@ class BOLDWarper(BasePreprocessor): ] = [ { "name": "fsl", - "optional": False, - "commands": ["applywarp"], + "optional": True, + "commands": ["flirt", "applywarp"], + }, + { + "name": "ants", + "optional": True, + "commands": ["ResampleImage", "antsApplyTransforms"], }, ] @@ -105,6 +111,8 @@ class BOLDWarper(BasePreprocessor): ------ ValueError If ``extra_input`` is None. + RuntimeError + If warp / transformation file extension is not ".mat" or ".h5". """ logger.debug("Warping BOLD using BOLDWarper") @@ -114,11 +122,34 @@ class BOLDWarper(BasePreprocessor): f"No extra input provided, requires `Warp` and `{self.ref}` " "data types in particular." ) - # Initialize ApplyWarper for computation - apply_warper = _ApplyWarper(reference=self.ref, on="BOLD") - # Replace original BOLD data with warped BOLD data - _, input = apply_warper.preprocess( - input=input, - extra_input=extra_input, - ) + # Check for warp file type to use correct tool + warp_file_ext = extra_input["Warp"]["path"].suffix + if warp_file_ext == ".mat": + logger.debug("Using FSL with BOLDWarper") + # Initialize ApplyWarper for computation + apply_warper = _ApplyWarper(reference=self.ref, on="BOLD") + # Replace original BOLD data with warped BOLD data + _, input = apply_warper.preprocess( + input=input, + extra_input=extra_input, + ) + elif warp_file_ext == ".h5": + logger.debug("Using ANTs with BOLDWarper") + # Initialize AntsApplyTransformsWarper for computation + ants_apply_transforms_warper = _AntsApplyTransformsWarper( + reference=self.ref, on="BOLD" + ) + # Replace original BOLD data with warped BOLD data + _, input = ants_apply_transforms_warper.preprocess( + input=input, + extra_input=extra_input, + ) + else: + raise_error( + msg=( + "Unknown warp / transformation file extension: " + f"{warp_file_ext}" + ), + klass=RuntimeError, + ) return "BOLD", input diff --git a/junifer/preprocess/fsl/apply_warper.py b/junifer/preprocess/fsl/apply_warper.py index 59a464f21..1b3cc2633 100644 --- a/junifer/preprocess/fsl/apply_warper.py +++ b/junifer/preprocess/fsl/apply_warper.py @@ -54,7 +54,7 @@ class _ApplyWarper(BasePreprocessor): { "name": "fsl", "optional": False, - "commands": ["applywarp"], + "commands": ["flirt", "applywarp"], }, ] diff --git a/junifer/preprocess/fsl/tests/test_apply_warper.py b/junifer/preprocess/fsl/tests/test_apply_warper.py index 24aee85cd..af2177628 100644 --- a/junifer/preprocess/fsl/tests/test_apply_warper.py +++ b/junifer/preprocess/fsl/tests/test_apply_warper.py @@ -3,12 +3,16 @@ # Authors: Synchon Mandal # License: AGPL +import socket +from pathlib import Path from typing import List +import nibabel as nib import pytest -# from junifer.datareader import DefaultDataReader -# from junifer.pipeline.utils import _check_fsl +from junifer.datagrabber import DataladHCP1200 +from junifer.datareader import DefaultDataReader +from junifer.pipeline.utils import _check_fsl from junifer.preprocess.fsl.apply_warper import _ApplyWarper @@ -47,27 +51,54 @@ def test_ApplyWarper_get_output_type(input_: List[str]) -> None: assert apply_warper.get_output_type(input_) == input_ -@pytest.mark.skip(reason="requires testing dataset") -# @pytest.mark.skipif( -# _check_fsl() is False, reason="requires fsl to be in PATH" -# ) +@pytest.mark.skipif(_check_fsl() is False, reason="requires FSL to be in PATH") +@pytest.mark.skipif( + socket.gethostname() != "juseless", + reason="only for juseless", +) def test_ApplyWarper__run_applywarp() -> None: """Test ApplyWarper _run_applywarp.""" - # Initialize datareader - # reader = DefaultDataReader() - # Initialize preprocessor - # bold_warper = _ApplyWarper(reference="T1w", on="BOLD") - # TODO(synchon): setup datagrabber and run pipeline + with DataladHCP1200( + tasks=["REST1"], + phase_encodings=["LR"], + ica_fix=True, + ) as dg: + # Read data + element_data = DefaultDataReader().fit_transform( + dg[("100206", "REST1", "LR")] + ) + # Preprocess data + warped_data, resampled_ref_path = _ApplyWarper( + reference="T1w", on="BOLD" + )._run_applywarp( + input_data=element_data["BOLD"], + ref_path=element_data["T1w"]["path"], + warp_path=element_data["Warp"]["path"], + ) + assert isinstance(warped_data, nib.Nifti1Image) + assert isinstance(resampled_ref_path, Path) -@pytest.mark.skip(reason="requires testing dataset") -# @pytest.mark.skipif( -# _check_fsl() is False, reason="requires fsl to be in PATH" -# ) +@pytest.mark.skipif(_check_fsl() is False, reason="requires FSL to be in PATH") +@pytest.mark.skipif( + socket.gethostname() != "juseless", + reason="only for juseless", +) def test_ApplyWarper_preprocess() -> None: """Test ApplyWarper preprocess.""" - # Initialize datareader - # reader = DefaultDataReader() - # Initialize preprocessor - # bold_warper = _ApplyWarper(reference="T1w", on="BOLD") - # TODO(synchon): setup datagrabber and run pipeline + with DataladHCP1200( + tasks=["REST1"], + phase_encodings=["LR"], + ica_fix=True, + ) as dg: + # Read data + element_data = DefaultDataReader().fit_transform( + dg[("100206", "REST1", "LR")] + ) + # Preprocess data + data_type, data = _ApplyWarper(reference="T1w", on="BOLD").preprocess( + input=element_data["BOLD"], + extra_input=element_data, + ) + assert isinstance(data_type, str) + assert isinstance(data, dict) diff --git a/junifer/preprocess/tests/test_bold_warper.py b/junifer/preprocess/tests/test_bold_warper.py index a10f29a69..ac2d00720 100644 --- a/junifer/preprocess/tests/test_bold_warper.py +++ b/junifer/preprocess/tests/test_bold_warper.py @@ -3,15 +3,21 @@ # Authors: Synchon Mandal # License: AGPL -from typing import List +import socket +from typing import TYPE_CHECKING, List, Tuple import pytest -# from junifer.datareader import DefaultDataReader -# from junifer.pipeline.utils import _check_fsl +from junifer.datagrabber import DataladHCP1200, DMCC13Benchmark +from junifer.datareader import DefaultDataReader +from junifer.pipeline.utils import _check_ants, _check_fsl from junifer.preprocess import BOLDWarper +if TYPE_CHECKING: + from junifer.datagrabber import BaseDataGrabber + + def test_BOLDWarper_init() -> None: """Test BOLDWarper init.""" bold_warper = BOLDWarper(reference="T1w") @@ -45,14 +51,58 @@ def test_BOLDWarper_get_output_type(input_: List[str]) -> None: assert bold_warper.get_output_type(input_) == input_ -@pytest.mark.skip(reason="requires testing dataset") -# @pytest.mark.skipif( -# _check_fsl() is False, reason="requires fsl to be in PATH" -# ) -def test_BOLDWarper_preprocess() -> None: - """Test BOLDWarper preprocess.""" - # Initialize datareader - # reader = DefaultDataReader() - # Initialize preprocessor - # bold_warper = BOLDWarper(reference="T1w") - # TODO(synchon): setup datagrabber and run pipeline +@pytest.mark.parametrize( + "datagrabber, element", + [ + [ + DMCC13Benchmark( + types=["BOLD", "T1w", "Warp"], + sessions=["wave1bas"], + tasks=["Rest"], + phase_encodings=["AP"], + runs=["1"], + native_t1w=True, + ), + ("f9057kp", "wave1bas", "Rest", "AP", "1"), + ], + [ + DataladHCP1200( + tasks=["REST1"], + phase_encodings=["LR"], + ica_fix=True, + ), + ("100206", "REST1", "LR"), + ], + ], +) +@pytest.mark.skipif(_check_fsl() is False, reason="requires FSL to be in PATH") +@pytest.mark.skipif( + _check_ants() is False, reason="requires ANTs to be in PATH" +) +@pytest.mark.skipif( + socket.gethostname() != "juseless", + reason="only for juseless", +) +def test_BOLDWarper_preprocess( + datagrabber: "BaseDataGrabber", element: Tuple[str, ...] +) -> None: + """Test BOLDWarper preprocess. + + Parameters + ---------- + datagrabber : DataGrabber-like object + The parametrized DataGrabber objects. + element : tuple of str + The parametrized elements. + + """ + with datagrabber as dg: + # Read data + element_data = DefaultDataReader().fit_transform(dg[element]) + # Preprocess data + data_type, data = BOLDWarper(reference="T1w").preprocess( + input=element_data["BOLD"], + extra_input=element_data, + ) + assert data_type == "BOLD" + assert isinstance(data, dict)