diff --git a/docs/changes/newsfragments/301.feature b/docs/changes/newsfragments/301.feature new file mode 100644 index 000000000..5ad08b386 --- /dev/null +++ b/docs/changes/newsfragments/301.feature @@ -0,0 +1 @@ +Introduce :class:`.SpaceWarper` for warping ``T1w``, ``BOLD``, ``VBM_GM``, ``VBM_WM``, ``fALFF``, ``GCOR`` and ``LCOR`` data to other spaces by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/301.removal b/docs/changes/newsfragments/301.removal new file mode 100644 index 000000000..9b2c343ed --- /dev/null +++ b/docs/changes/newsfragments/301.removal @@ -0,0 +1 @@ +Deprecate :class:`.BOLDWarper` and mark for removal in v0.0.4 by `Synchon Mandal`_ diff --git a/junifer/preprocess/__init__.py b/junifer/preprocess/__init__.py index 804433ecc..15042f0ee 100644 --- a/junifer/preprocess/__init__.py +++ b/junifer/preprocess/__init__.py @@ -8,3 +8,4 @@ from .base import BasePreprocessor from .confounds import fMRIPrepConfoundRemover from .bold_warper import BOLDWarper +from .warping import SpaceWarper diff --git a/junifer/preprocess/bold_warper.py b/junifer/preprocess/bold_warper.py index 22aeaceb1..304f59bbc 100644 --- a/junifer/preprocess/bold_warper.py +++ b/junifer/preprocess/bold_warper.py @@ -30,6 +30,10 @@ from .fsl.apply_warper import _ApplyWarper class BOLDWarper(BasePreprocessor): """Class for warping BOLD NIfTI images. + .. deprecated:: 0.0.3 + `BOLDWarper` will be removed in v0.0.4, it is replaced by + `SpaceWarper` because the latter works also with T1w data. + Parameters ---------- using : {"fsl", "ants"} diff --git a/junifer/preprocess/warping/__init__.py b/junifer/preprocess/warping/__init__.py new file mode 100644 index 000000000..1b7a3bc2b --- /dev/null +++ b/junifer/preprocess/warping/__init__.py @@ -0,0 +1,6 @@ +"""Provide imports for warping sub-package.""" + +# Authors: Synchon Mandal +# License: AGPL + +from .space_warper import SpaceWarper diff --git a/junifer/preprocess/warping/_ants_warper.py b/junifer/preprocess/warping/_ants_warper.py new file mode 100644 index 000000000..a78d67702 --- /dev/null +++ b/junifer/preprocess/warping/_ants_warper.py @@ -0,0 +1,167 @@ +"""Provide class for space warping via ANTs antsApplyTransforms.""" + +# Authors: Synchon Mandal +# License: AGPL + +from typing import ( + Any, + ClassVar, + Dict, + List, + Set, + Union, +) + +import nibabel as nib +import numpy as np + +from ...data import get_template, get_xfm +from ...pipeline import WorkDirManager +from ...utils import logger, run_ext_cmd + + +class ANTsWarper: + """Class for space warping via ANTs antsApplyTransforms. + + This class uses ANTs' ``ResampleImage`` for resampling (if required) and + ``antsApplyTransforms`` for transformation. + + """ + + _EXT_DEPENDENCIES: ClassVar[List[Dict[str, Union[str, List[str]]]]] = [ + { + "name": "ants", + "commands": ["ResampleImage", "antsApplyTransforms"], + }, + ] + + _DEPENDENCIES: ClassVar[Set[str]] = {"numpy", "nibabel"} + + def preprocess( + self, + input: Dict[str, Any], + extra_input: Dict[str, Any], + reference: str, + ) -> Dict[str, Any]: + """Preprocess using ANTs. + + Parameters + ---------- + input : dict + A single input from the Junifer Data object in which to preprocess. + extra_input : dict + The other fields in the Junifer Data object. Should have ``T1w`` + and ``Warp`` data types. + reference : str + The data type or template space to use as reference for warping. + + Returns + ------- + dict + The ``input`` dictionary with modified ``data`` and ``space`` key + values and new ``reference_path`` key whose value points to the + reference file used for warping. + + """ + # Create element-specific tempdir for storing post-warping assets + element_tempdir = WorkDirManager().get_element_tempdir( + prefix="ants_warper" + ) + + # Native space warping + if reference == "T1w": + logger.debug("Using ANTs for space warping") + + # Get the min of the voxel sizes from input and use it as the + # resolution + resolution = np.min(input["data"].header.get_zooms()[:3]) + + # Create a tempfile for resampled reference output + resample_image_out_path = ( + element_tempdir / "resampled_reference.nii.gz" + ) + # Set ResampleImage command + resample_image_cmd = [ + "ResampleImage", + "3", # image dimension + f"{extra_input['T1w']['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 + run_ext_cmd(name="ResampleImage", cmd=resample_image_cmd) + + # Create a tempfile for warped output + apply_transforms_out_path = element_tempdir / "output.nii.gz" + # Set antsApplyTransforms command + apply_transforms_cmd = [ + "antsApplyTransforms", + "-d 3", + "-e 3", + "-n LanczosWindowedSinc", + f"-i {input['path'].resolve()}", + # use resampled reference + f"-r {resample_image_out_path.resolve()}", + f"-t {extra_input['Warp']['path'].resolve()}", + f"-o {apply_transforms_out_path.resolve()}", + ] + # Call antsApplyTransforms + run_ext_cmd(name="antsApplyTransforms", cmd=apply_transforms_cmd) + + # Load nifti + input["data"] = nib.load(apply_transforms_out_path) + # Save resampled reference path + input["reference_path"] = resample_image_out_path + # Use reference input's space as warped input's space + input["space"] = extra_input["T1w"]["space"] + + # Template space warping + else: + logger.debug( + f"Using ANTs to warp data from {input['space']} to {reference}" + ) + + # Get xfm file + xfm_file_path = get_xfm(src=input["space"], dst=reference) + # Get template space image + template_space_img = get_template( + space=reference, + target_data=input, + extra_input=None, + ) + + # Create component-scoped tempdir + tempdir = WorkDirManager().get_tempdir(prefix="ants_warper") + # Save template + template_space_img_path = tempdir / f"{reference}_T1w.nii.gz" + nib.save(template_space_img, template_space_img_path) + + # Create a tempfile for warped output + warped_output_path = element_tempdir / ( + f"data_warped_from_{input['space']}_to_" f"{reference}.nii.gz" + ) + + # Set antsApplyTransforms command + apply_transforms_cmd = [ + "antsApplyTransforms", + "-d 3", + "-e 3", + "-n LanczosWindowedSinc", + f"-i {input['path'].resolve()}", + f"-r {template_space_img_path.resolve()}", + f"-t {xfm_file_path.resolve()}", + f"-o {warped_output_path.resolve()}", + ] + # Call antsApplyTransforms + run_ext_cmd(name="antsApplyTransforms", cmd=apply_transforms_cmd) + + # Delete tempdir + WorkDirManager().delete_tempdir(tempdir) + + # Modify target data + input["data"] = nib.load(warped_output_path) + input["space"] = reference + + return input diff --git a/junifer/preprocess/warping/_fsl_warper.py b/junifer/preprocess/warping/_fsl_warper.py new file mode 100644 index 000000000..32ded0e60 --- /dev/null +++ b/junifer/preprocess/warping/_fsl_warper.py @@ -0,0 +1,109 @@ +"""Provide class for space warping via FSL FLIRT.""" + +# Authors: Synchon Mandal +# License: AGPL + +from typing import ( + Any, + ClassVar, + Dict, + List, + Set, + Union, +) + +import nibabel as nib +import numpy as np + +from ...pipeline import WorkDirManager +from ...utils import logger, run_ext_cmd + + +class FSLWarper: + """Class for space warping via FSL FLIRT. + + This class uses FSL FLIRT's ``flirt`` for resampling and ``applywarp`` for + transformation. + + """ + + _EXT_DEPENDENCIES: ClassVar[List[Dict[str, Union[str, List[str]]]]] = [ + { + "name": "fsl", + "commands": ["flirt", "applywarp"], + }, + ] + + _DEPENDENCIES: ClassVar[Set[str]] = {"numpy", "nibabel"} + + def preprocess( + self, + input: Dict[str, Any], + extra_input: Dict[str, Any], + ) -> Dict[str, Any]: + """Preprocess using FSL. + + Parameters + ---------- + input : dict + A single input from the Junifer Data object in which to preprocess. + extra_input : dict + The other fields in the Junifer Data object. Should have ``T1w`` + and ``Warp`` data types. + + Returns + ------- + dict + The ``input`` dictionary with modified ``data`` and ``space`` key + values and new ``reference_path`` key whose value points to the + reference file used for warping. + + """ + logger.debug("Using FSL for space warping") + + # Get the min of the voxel sizes from input and use it as the + # resolution + resolution = np.min(input["data"].header.get_zooms()[:3]) + + # Create element-specific tempdir for storing post-warping assets + element_tempdir = WorkDirManager().get_element_tempdir( + prefix="fsl_warper" + ) + + # Create a tempfile for resampled reference output + flirt_out_path = element_tempdir / "resampled_reference.nii.gz" + # Set flirt command + flirt_cmd = [ + "flirt", + "-interp spline", + f"-in {extra_input['T1w']['path'].resolve()}", + f"-ref {extra_input['T1w']['path'].resolve()}", + f"-applyisoxfm {resolution}", + f"-out {flirt_out_path.resolve()}", + ] + # Call flirt + run_ext_cmd(name="flirt", cmd=flirt_cmd) + + # Create a tempfile for warped output + applywarp_out_path = element_tempdir / "output.nii.gz" + # Set applywarp command + applywarp_cmd = [ + "applywarp", + "--interp=spline", + f"-i {input['path'].resolve()}", + f"-r {flirt_out_path.resolve()}", # use resampled reference + f"-w {extra_input['Warp']['path'].resolve()}", + f"-o {applywarp_out_path.resolve()}", + ] + # Call applywarp + run_ext_cmd(name="applywarp", cmd=applywarp_cmd) + + # Load nifti + input["data"] = nib.load(applywarp_out_path) + # Save resampled reference path + input["reference_path"] = flirt_out_path + + # Use reference input's space as warped input's space + input["space"] = extra_input["T1w"]["space"] + + return input diff --git a/junifer/preprocess/warping/space_warper.py b/junifer/preprocess/warping/space_warper.py new file mode 100644 index 000000000..18b9e9cca --- /dev/null +++ b/junifer/preprocess/warping/space_warper.py @@ -0,0 +1,203 @@ +"""Provide class for warping data to other template spaces.""" + +# Authors: Synchon Mandal +# License: AGPL + +from typing import Any, ClassVar, Dict, List, Optional, Tuple, Type, Union + +from templateflow import api as tflow + +from ...api.decorators import register_preprocessor +from ...utils import logger, raise_error +from ..base import BasePreprocessor +from ._ants_warper import ANTsWarper +from ._fsl_warper import FSLWarper + + +__all__ = ["SpaceWarper"] + + +@register_preprocessor +class SpaceWarper(BasePreprocessor): + """Class for warping data to other template spaces. + + Parameters + ---------- + using : {"fsl", "ants"} + Implementation to use for warping: + + * "fsl" : Use FSL's ``applywarp`` + * "ants" : Use ANTs' ``antsApplyTransforms`` + + reference : str + The data type to use as reference for warping, can be either a data + type like ``"T1w"`` or a template space like ``"MNI152NLin2009cAsym"``. + Use ``"T1w"`` for native space warping and named templates for + template space warping. + on : {"T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"} or list \ + of the options + The data type to warp. + + Raises + ------ + ValueError + If ``using`` is invalid or + if ``reference`` is invalid. + + """ + + _CONDITIONAL_DEPENDENCIES: ClassVar[List[Dict[str, Union[str, Type]]]] = [ + { + "using": "fsl", + "depends_on": FSLWarper, + }, + { + "using": "ants", + "depends_on": ANTsWarper, + }, + ] + + def __init__( + self, using: str, reference: str, on: Union[List[str], str] + ) -> None: + """Initialize the class.""" + # Validate `using` parameter + valid_using = [dep["using"] for dep in self._CONDITIONAL_DEPENDENCIES] + if using not in valid_using: + raise_error( + f"Invalid value for `using`, should be one of: {valid_using}" + ) + self.using = using + self.reference = reference + # Set required data types based on reference and + # initialize superclass + if self.reference == "T1w": + required_data_types = [self.reference, "Warp"] + # Listify on + if not isinstance(on, list): + on = [on] + # Extend required data types + required_data_types.extend(on) + + super().__init__( + on=on, + required_data_types=required_data_types, + ) + elif self.reference in tflow.templates(): + super().__init__(on=on) + else: + raise_error(f"Unknown reference: {self.reference}") + + 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. + + """ + return ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] + + def get_output_type(self, input_type: str) -> str: + """Get output type. + + Parameters + ---------- + input_type : str + The data type input to the preprocessor. + + Returns + ------- + str + The data type output by the preprocessor. + + """ + # Does not add any new keys + return input_type + + def preprocess( + self, + input: Dict[str, Any], + extra_input: Optional[Dict[str, Any]] = None, + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: + """Preprocess. + + Parameters + ---------- + input : dict + The input from the Junifer Data object. + extra_input : dict, optional + The other fields in the Junifer Data object. + + Returns + ------- + dict + The computed result as dictionary. + None + Extra "helper" data types as dictionary to add to the Junifer Data + object. + + Raises + ------ + ValueError + If ``extra_input`` is None when transforming to native space + i.e., using ``"T1w"`` as reference. + RuntimeError + If the data is in the correct space and does not require + warping or + if FSL is used for template space warping. + + """ + logger.info(f"Warping to {self.reference} space using SpaceWarper") + # Transform to native space + if self.using in ["fsl", "ants"] and self.reference == "T1w": + # Check for extra inputs + if extra_input is None: + raise_error( + "No extra input provided, requires `Warp` and " + f"`{self.reference}` data types in particular." + ) + # Conditional preprocessor + if self.using == "fsl": + input = FSLWarper().preprocess( + input=input, + extra_input=extra_input, + ) + elif self.using == "ants": + input = ANTsWarper().preprocess( + input=input, + extra_input=extra_input, + reference=self.reference, + ) + # Transform to template space with ANTs possible + elif self.using == "ants" and self.reference != "T1w": + # Check pre-requirements for space manipulation + if self.reference == input["space"]: + raise_error( + ( + f"The target data is in {self.reference} space " + "and thus warping will not be performed, hence you " + "should remove the SpaceWarper from the preprocess " + "step." + ), + klass=RuntimeError, + ) + + input = ANTsWarper().preprocess( + input=input, + extra_input={}, + reference=self.reference, + ) + # Transform to template space with FSL not possible + elif self.using == "fsl" and self.reference != "T1w": + raise_error( + ( + f"Warping to {self.reference} space not possible with " + "FSL, use ANTs instead." + ), + klass=RuntimeError, + ) + + return input, None diff --git a/junifer/preprocess/warping/tests/test_space_warper.py b/junifer/preprocess/warping/tests/test_space_warper.py new file mode 100644 index 000000000..1d21d9fae --- /dev/null +++ b/junifer/preprocess/warping/tests/test_space_warper.py @@ -0,0 +1,198 @@ +"""Provide tests for SpaceWarper.""" + +# Authors: Synchon Mandal +# License: AGPL + +import socket +from typing import TYPE_CHECKING, Tuple, Type + +import pytest +from numpy.testing import assert_array_equal, assert_raises + +from junifer.datagrabber import DataladHCP1200, DMCC13Benchmark +from junifer.datareader import DefaultDataReader +from junifer.pipeline.utils import _check_ants, _check_fsl +from junifer.preprocess import SpaceWarper +from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber + + +if TYPE_CHECKING: + from junifer.datagrabber import BaseDataGrabber + + +@pytest.mark.parametrize( + "using, reference, error_type, error_msg", + [ + ("jam", "T1w", ValueError, "`using`"), + ("ants", "juice", ValueError, "reference"), + ("ants", "MNI152NLin2009cAsym", RuntimeError, "remove"), + ("fsl", "MNI152NLin2009cAsym", RuntimeError, "ANTs"), + ], +) +def test_SpaceWarper_errors( + using: str, + reference: str, + error_type: Type[Exception], + error_msg: str, +) -> None: + """Test SpaceWarper errors. + + Parameters + ---------- + using : str + The parametrized implementation method. + reference : str + The parametrized reference to use. + error_type : Exception-like object + The parametrized exception to check. + error_msg : str + The parametrized exception message to check. + + """ + with PartlyCloudyTestingDataGrabber() as dg: + # Read data + element_data = DefaultDataReader().fit_transform(dg["sub-01"]) + # Preprocess data + with pytest.raises(error_type, match=error_msg): + SpaceWarper( + using=using, + reference=reference, + on="BOLD", + ).preprocess( + input=element_data["BOLD"], + extra_input=element_data, + ) + + +@pytest.mark.parametrize( + "datagrabber, element, using", + [ + [ + DMCC13Benchmark( + types=["BOLD", "T1w", "Warp"], + sessions=["wave1bas"], + tasks=["Rest"], + phase_encodings=["AP"], + runs=["1"], + native_t1w=True, + ), + ("f9057kp", "wave1bas", "Rest", "AP", "1"), + "ants", + ], + [ + DataladHCP1200( + tasks=["REST1"], + phase_encodings=["LR"], + ica_fix=True, + ), + ("100206", "REST1", "LR"), + "fsl", + ], + ], +) +@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_SpaceWarper_native( + datagrabber: "BaseDataGrabber", element: Tuple[str, ...], using: str +) -> None: + """Test SpaceWarper for native space warping. + + Parameters + ---------- + datagrabber : DataGrabber-like object + The parametrized DataGrabber objects. + element : tuple of str + The parametrized elements. + using : str + The parametrized implementation method. + + """ + with datagrabber as dg: + # Read data + element_data = DefaultDataReader().fit_transform(dg[element]) + # Preprocess data + output, _ = SpaceWarper( + using=using, + reference="T1w", + on="BOLD", + ).preprocess( + input=element_data["BOLD"], + extra_input=element_data, + ) + # Check + assert isinstance(output, dict) + + +@pytest.mark.parametrize( + "datagrabber, element, space", + [ + [ + DMCC13Benchmark( + types=["T1w"], + sessions=["wave1bas"], + tasks=["Rest"], + phase_encodings=["AP"], + runs=["1"], + native_t1w=False, + ), + ("f9057kp", "wave1bas", "Rest", "AP", "1"), + "MNI152NLin2009aAsym", + ], + [ + DMCC13Benchmark( + types=["T1w"], + sessions=["wave1bas"], + tasks=["Rest"], + phase_encodings=["AP"], + runs=["1"], + native_t1w=False, + ), + ("f9057kp", "wave1bas", "Rest", "AP", "1"), + "MNI152NLin6Asym", + ], + ], +) +@pytest.mark.skipif( + _check_ants() is False, reason="requires ANTs to be in PATH" +) +def test_SpaceWarper_multi_mni( + datagrabber: "BaseDataGrabber", + element: Tuple[str, ...], + space: str, +) -> None: + """Test SpaceWarper for MNI space warping. + + Parameters + ---------- + datagrabber : DataGrabber-like object + The parametrized DataGrabber objects. + element : tuple of str + The parametrized elements. + space : str + The parametrized template space to transform to. + + """ + with datagrabber as dg: + # Read data + element_data = DefaultDataReader().fit_transform(dg[element]) + pre_xfm_data = element_data["T1w"]["data"].get_fdata().copy() + # Preprocess data + output, _ = SpaceWarper( + using="ants", + reference=space, + on=["T1w"], + ).preprocess( + input=element_data["T1w"], + extra_input=element_data, + ) + # Checks + assert isinstance(output, dict) + assert output["space"] == space + with assert_raises(AssertionError): + assert_array_equal(pre_xfm_data, output["data"])