diff --git a/docs/changes/newsfragments/390.change b/docs/changes/newsfragments/390.change new file mode 100644 index 000000000..6356751c2 --- /dev/null +++ b/docs/changes/newsfragments/390.change @@ -0,0 +1 @@ +:class:`.SpaceWarper`'s ``using`` now supports ``"auto"`` allowing auto tool selection based on DataGrabber specification by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/390.enh b/docs/changes/newsfragments/390.enh new file mode 100644 index 000000000..96545d4b6 --- /dev/null +++ b/docs/changes/newsfragments/390.enh @@ -0,0 +1 @@ +Update ``Warp`` data type to be a list of dictionaries and adapt to allow multiple transforms to be specified by `Synchon Mandal`_ diff --git a/docs/extending/dependencies.rst b/docs/extending/dependencies.rst index 441ba5414..48132f990 100644 --- a/docs/extending/dependencies.rst +++ b/docs/extending/dependencies.rst @@ -39,7 +39,7 @@ and others who use it will thank you. Handling external dependencies from toolboxes --------------------------------------------- -You can also specify dependencies of external toolboxes like AFIN, FSL and ANTs, +You can also specify dependencies of external toolboxes like AFNI, FSL and ANTs, by having a class attribute like so: .. code-block:: python @@ -87,6 +87,10 @@ that it shows the problem a bit better and how we solve it: "using": "ants", "depends_on": ANTsWarper, }, + { + "using": "auto", + "depends_on": [FSLWarper, ANTsWarper], + }, ] def __init__( @@ -100,13 +104,20 @@ Here, you see a new class attribute ``_CONDITIONAL_DEPENDENCIES`` which is a list of dictionaries with two keys: * ``using`` (str) : lowercased name of the toolbox -* ``depends_on`` (object) : a class which implements the particular tool's use +* ``depends_on`` (object or list of objects) : a class or list of classes which \ + implements the particular tool's use It is mandatory to have the ``using`` positional argument in the constructor in this case as the validation starts with this and moves further. It is also mandatory to only allow the value of ``using`` argument to be one of them specified in the ``using`` key of ``_CONDITIONAL_DEPENDENCIES`` entries. +For some cases, like we have here, it might be worth to have ``"auto"`` for +``"using"`` which eases the choice of a particular tool by the user and +instead lets ``junifer`` do it automatically. In that case, ``"depends_on"`` +needs to be a list of the tool implementation classes. This also requires the +user to have all the tools in the ``PATH``. + For brevity, we only show the ``FSLWarper`` here but ``ANTsWarper`` looks very similar. ``FSLWarper`` looks like this (only the relevant part is shown here): diff --git a/docs/understanding/preprocess.rst b/docs/understanding/preprocess.rst index c5701aa18..08cd73ddb 100644 --- a/docs/understanding/preprocess.rst +++ b/docs/understanding/preprocess.rst @@ -146,14 +146,23 @@ Warping to subject's native space To warp to subject's native space, the dataset needs to provide ``T1w`` and ``Warp`` data types and the DataGrabber needs to at least have ``["BOLD", "T1w", "Warp"]`` (if you are warping ``BOLD``) as the ``types`` -parameter's value. The :class:`.SpaceWarper`'s ``reference`` parameter needs +parameter's value. + +The :class:`.SpaceWarper`'s ``reference`` parameter needs to be set to ``T1w``, which means that the ``BOLD`` data will be transformed using the ``T1w`` as reference (it's resampled internally to match the -resolution of the ``BOLD``). The ``Warp`` data type is new and it's only purpose -is to provide the warp or transformation file (can be linear, non-linear or -linear + non-linear transform) for the purpose. For ``using`` parameter, you can -pass either ``"fsl"`` or ``"ants"`` depending on the warp or transformation file -format. +resolution of the ``BOLD``). + +The ``Warp`` data type provides the warp or transformation file (can be linear, +non-linear or linear + non-linear transform) for the purpose. For ``using`` +parameter, you can pass either ``fsl`` or ``ants`` depending on the warp or +transformation file format. You can also provide ``auto`` to ``using`` in which +case either ``FSL`` or ``ANTs`` will be used based on the file format provided +by the DataGrabber. This also requires that both the tools are in the ``PATH``. + +And finally, you would need to set the ``on`` parameter to ``BOLD`` to make it +clear which data type you intend to warp, as the :class:`.SpaceWarper` is also +capable of warping ``T1w``. An example YAML might look like this: @@ -163,6 +172,7 @@ An example YAML might look like this: - kind: SpaceWarper using: fsl reference: T1w + on: BOLD .. _preprocess_warping_template: @@ -174,8 +184,8 @@ data type that you want to work on) in ``MNI152NLin6Asym`` template space but you would like to compute features in ``MNI152NLin2009cAsym`` template space, you can also use the :class:`.SpaceWarper` by setting the ``reference`` parameter to the template space's name, in this case, -``reference="MNI152NLin2009cAsym"``. The ``using`` parameter needs to be set -to ``"ants"`` as we need it to warp the data. +``reference: MNI152NLin2009cAsym``. The ``using`` parameter needs to be set +to ``ants`` as we need it to warp the data. .. note:: @@ -190,3 +200,4 @@ For an YAML example: - kind: SpaceWarper using: ants reference: MNI152NLin2009cAsym + on: BOLD diff --git a/junifer/data/coordinates/_ants_coordinates_warper.py b/junifer/data/coordinates/_ants_coordinates_warper.py index e1d7f35db..d9b4ccf09 100644 --- a/junifer/data/coordinates/_ants_coordinates_warper.py +++ b/junifer/data/coordinates/_ants_coordinates_warper.py @@ -26,7 +26,7 @@ class ANTsCoordinatesWarper: self, seeds: ArrayLike, target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> ArrayLike: """Warp ``seeds`` to correct space. @@ -37,10 +37,8 @@ class ANTsCoordinatesWarper: target_data : dict The corresponding item of the data object to which the coordinates will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of coordinates - (default None). + warp_data : dict or None + The warp data item of the data object. Returns ------- @@ -79,7 +77,7 @@ class ANTsCoordinatesWarper: "-f 0", f"-i {pretransform_coordinates_path.resolve()}", f"-o {transformed_coords_path.resolve()}", - f"-t {extra_input['Warp']['path'].resolve()};", + f"-t {warp_data['path'].resolve()}", ] # Call antsApplyTransformsToPoints run_ext_cmd( diff --git a/junifer/data/coordinates/_coordinates.py b/junifer/data/coordinates/_coordinates.py index 01073de9b..e51b9fb9d 100644 --- a/junifer/data/coordinates/_coordinates.py +++ b/junifer/data/coordinates/_coordinates.py @@ -14,6 +14,7 @@ from numpy.typing import ArrayLike from ...utils import logger, raise_error from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry +from ..utils import get_native_warper from ._ants_coordinates_warper import ANTsCoordinatesWarper from ._fsl_coordinates_warper import FSLCoordinatesWarper @@ -311,7 +312,8 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ RuntimeError - If warp / transformation file extension is not ".mat" or ".h5". + If warping specification required for warping using ANTs, is not + found. ValueError If ``extra_input`` is None when ``target_data``'s space is native. @@ -329,27 +331,38 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): f"{target_data['space']} space for further computation." ) - # Check for warp file type to use correct tool - warp_file_ext = extra_input["Warp"]["path"].suffix - if warp_file_ext == ".mat": + # Get native space warper spec + warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + ) + # Conditional for warping tool implementation + if warper_spec["warper"] == "fsl": seeds = FSLCoordinatesWarper().warp( seeds=seeds, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) - elif warp_file_ext == ".h5": + elif warper_spec["warper"] == "ants": + # Requires the inverse warp + inverse_warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + inverse=True, + ) + # Check warper + if inverse_warper_spec["warper"] != "ants": + raise_error( + klass=RuntimeError, + msg=( + "Warping specification mismatch for native space " + "warping of coordinates using ANTs." + ), + ) seeds = ANTsCoordinatesWarper().warp( seeds=seeds, target_data=target_data, - extra_input=extra_input, - ) - else: - raise_error( - msg=( - "Unknown warp / transformation file extension: " - f"{warp_file_ext}" - ), - klass=RuntimeError, + warp_data=warper_spec, ) return seeds, labels diff --git a/junifer/data/coordinates/_fsl_coordinates_warper.py b/junifer/data/coordinates/_fsl_coordinates_warper.py index 72948e0f9..afe65280c 100644 --- a/junifer/data/coordinates/_fsl_coordinates_warper.py +++ b/junifer/data/coordinates/_fsl_coordinates_warper.py @@ -26,7 +26,7 @@ class FSLCoordinatesWarper: self, seeds: ArrayLike, target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> ArrayLike: """Warp ``seeds`` to correct space. @@ -37,10 +37,8 @@ class FSLCoordinatesWarper: target_data : dict The corresponding item of the data object to which the coordinates will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of coordinates - (default None). + warp_data : dict + The warp data item of the data object. Returns ------- @@ -72,7 +70,7 @@ class FSLCoordinatesWarper: "| img2imgcoord -mm", f"-src {target_data['path'].resolve()}", f"-dest {target_data['reference_path'].resolve()}", - f"-warp {extra_input['Warp']['path'].resolve()}", + f"-warp {warp_data['path'].resolve()}", f"> {transformed_coords_path.resolve()};", f"sed -i 1d {transformed_coords_path.resolve()}", ] diff --git a/junifer/data/masks/_ants_mask_warper.py b/junifer/data/masks/_ants_mask_warper.py index 572c63097..90329eb57 100644 --- a/junifer/data/masks/_ants_mask_warper.py +++ b/junifer/data/masks/_ants_mask_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd from ..template_spaces import get_template, get_xfm @@ -34,7 +34,7 @@ class ANTsMaskWarper: src: str, dst: str, target_data: Dict[str, Any], - extra_input: Optional[Dict[str, Any]] = None, + warp_data: Optional[Dict[str, Any]], ) -> "Nifti1Image": """Warp ``mask_img`` to correct space. @@ -55,17 +55,20 @@ class ANTsMaskWarper: target_data : dict The corresponding item of the data object to which the mask will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of mask - (default None). - + warp_data : dict or None + The warp data item of the data object. The value is unused if + ``dst!="T1w"``. Returns ------- nibabel.nifti1.Nifti1Image The transformed mask image. + Raises + ------ + RuntimeError + If ``warp_data`` is None when ``dst="T1w"``. + """ # Create element-scoped tempdir so that warped mask is # available later as nibabel stores file path reference for @@ -80,7 +83,11 @@ class ANTsMaskWarper: ) # Native space warping - if dst == "T1w": + if dst == "native": + # Warp data check + if warp_data is None: + raise_error("No `warp_data` provided") + logger.debug("Using ANTs for mask transformation") # Save existing mask image to a tempfile @@ -98,7 +105,7 @@ class ANTsMaskWarper: f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-t {extra_input['Warp']['path'].resolve()}", + f"-t {warp_data['path'].resolve()}", f"-o {warped_mask_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/data/masks/_fsl_mask_warper.py b/junifer/data/masks/_fsl_mask_warper.py index 933d49b90..4aab3ed02 100644 --- a/junifer/data/masks/_fsl_mask_warper.py +++ b/junifer/data/masks/_fsl_mask_warper.py @@ -31,7 +31,7 @@ class FSLMaskWarper: mask_name: str, mask_img: "Nifti1Image", target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> "Nifti1Image": """Warp ``mask_img`` to correct space. @@ -44,10 +44,8 @@ class FSLMaskWarper: target_data : dict The corresponding item of the data object to which the mask will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of mask - (default None). + warp_data : dict + The warp data item of the data object. Returns ------- @@ -77,7 +75,7 @@ class FSLMaskWarper: f"-i {prewarp_mask_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-w {extra_input['Warp']['path'].resolve()}", + f"-w {warp_data['path'].resolve()}", f"-o {warped_mask_path.resolve()}", ] # Call applywarp diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index be7a1b688..a868ccf30 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -29,7 +29,7 @@ from ...utils import logger, raise_error from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..template_spaces import get_template -from ..utils import closest_resolution +from ..utils import closest_resolution, get_native_warper from ._ants_mask_warper import ANTsMaskWarper from ._fsl_mask_warper import FSLMaskWarper @@ -105,14 +105,16 @@ def compute_brain_mask( "data type to infer target template space." ) # Set target standard space to warp file space source - target_std_space = extra_input["Warp"]["src"] + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + target_std_space = entry["src"] # Fetch template in closest resolution template = get_template( space=target_std_space, target_data=target_data, extra_input=extra_input, - template_type=mask_type if mask_type in ["gm", "wm"] else "T1w", + template_type=mask_type, ) # Resample template to target image target_img = target_data["data"] @@ -357,8 +359,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ - RuntimeError - If warp / transformation file extension is not ".mat" or ".h5". ValueError If extra key is provided in addition to mask name in ``masks`` or if no mask is provided or @@ -372,8 +372,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): """ # Check pre-requirements for space manipulation target_space = target_data["space"] - # Set target standard space to target space - target_std_space = target_space # Extra data type requirement check if target space is native if target_space == "native": # Check for extra inputs @@ -383,8 +381,16 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): "data types in particular for transformation to " f"{target_data['space']} space for further computation." ) + # Get native space warper spec + warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + ) # Set target standard space to warp file space source - target_std_space = extra_input["Warp"]["src"] + target_std_space = warper_spec["src"] + else: + # Set target standard space to target space + target_std_space = target_space # Get the min of the voxels sizes and use it as the resolution target_img = target_data["data"] @@ -489,7 +495,7 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): src=mask_space, dst=target_std_space, target_data=target_data, - extra_input=None, + warp_data=None, ) all_masks.append(mask_img) @@ -510,32 +516,22 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Warp mask if target data is native if target_space == "native": - # extra_input check done earlier - # Check for warp file type to use correct tool - warp_file_ext = extra_input["Warp"]["path"].suffix - if warp_file_ext == ".mat": + # extra_input check done earlier and warper_spec exists + if warper_spec["warper"] == "fsl": mask_img = FSLMaskWarper().warp( mask_name="native", mask_img=mask_img, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) - elif warp_file_ext == ".h5": + elif warper_spec["warper"] == "ants": mask_img = ANTsMaskWarper().warp( mask_name="native", mask_img=mask_img, src="", - dst="T1w", + dst="native", target_data=target_data, - extra_input=extra_input, - ) - else: - raise_error( - msg=( - "Unknown warp / transformation file extension: " - f"{warp_file_ext}" - ), - klass=RuntimeError, + warp_data=warper_spec, ) return mask_img diff --git a/junifer/data/parcellations/_ants_parcellation_warper.py b/junifer/data/parcellations/_ants_parcellation_warper.py index 64de58ef0..8eade0ce1 100644 --- a/junifer/data/parcellations/_ants_parcellation_warper.py +++ b/junifer/data/parcellations/_ants_parcellation_warper.py @@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, Optional import nibabel as nib from ...pipeline import WorkDirManager -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd from ..template_spaces import get_template, get_xfm @@ -34,7 +34,7 @@ class ANTsParcellationWarper: src: str, dst: str, target_data: Dict[str, Any], - extra_input: Optional[Dict[str, Any]] = None, + warp_data: Optional[Dict[str, Any]], ) -> "Nifti1Image": """Warp ``parcellation_img`` to correct space. @@ -55,17 +55,20 @@ class ANTsParcellationWarper: target_data : dict The corresponding item of the data object to which the parcellation will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of parcellation - (default None). - + warp_data : dict or None + The warp data item of the data object. The value is unused if + ``dst!="T1w"``. Returns ------- nibabel.nifti1.Nifti1Image The transformed parcellation image. + Raises + ------ + ValueError + If ``warp_data`` is None when ``dst="T1w"``. + """ # Create element-scoped tempdir so that warped parcellation is # available later as nibabel stores file path reference for @@ -80,7 +83,11 @@ class ANTsParcellationWarper: ) # Native space warping - if dst == "T1w": + if dst == "native": + # Warp data check + if warp_data is None: + raise_error("No `warp_data` provided") + logger.debug("Using ANTs for parcellation transformation") # Save existing parcellation image to a tempfile @@ -102,7 +109,7 @@ class ANTsParcellationWarper: f"-i {prewarp_parcellation_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-t {extra_input['Warp']['path'].resolve()}", + f"-t {warp_data['path'].resolve()}", f"-o {warped_parcellation_path.resolve()}", ] # Call antsApplyTransforms diff --git a/junifer/data/parcellations/_fsl_parcellation_warper.py b/junifer/data/parcellations/_fsl_parcellation_warper.py index 7fc9717ca..f9af85504 100644 --- a/junifer/data/parcellations/_fsl_parcellation_warper.py +++ b/junifer/data/parcellations/_fsl_parcellation_warper.py @@ -31,7 +31,7 @@ class FSLParcellationWarper: parcellation_name: str, parcellation_img: "Nifti1Image", target_data: Dict[str, Any], - extra_input: Dict[str, Any], + warp_data: Dict[str, Any], ) -> "Nifti1Image": """Warp ``parcellation_img`` to correct space. @@ -44,10 +44,8 @@ class FSLParcellationWarper: target_data : dict The corresponding item of the data object to which the parcellation will be applied. - extra_input : dict, optional - The other fields in the data object. Useful for accessing other - data kinds that needs to be used in the computation of parcellation - (default None). + warp_data : dict + The warp data item of the data object. Returns ------- @@ -81,7 +79,7 @@ class FSLParcellationWarper: f"-i {prewarp_parcellation_path.resolve()}", # use resampled reference f"-r {target_data['reference_path'].resolve()}", - f"-w {extra_input['Warp']['path'].resolve()}", + f"-w {warp_data['path'].resolve()}", f"-o {warped_parcellation_path.resolve()}", ] # Call applywarp diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 3e9cb4195..129b13226 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -23,7 +23,7 @@ from nilearn import datasets, image from ...utils import logger, raise_error, warn_with_log from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry -from ..utils import closest_resolution +from ..utils import closest_resolution, get_native_warper from ._ants_parcellation_warper import ANTsParcellationWarper from ._fsl_parcellation_warper import FSLParcellationWarper @@ -394,16 +394,12 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): Raises ------ - RuntimeError - If warp / transformation file extension is not ".mat" or ".h5". ValueError If ``extra_input`` is None when ``target_data``'s space is native. """ # Check pre-requirements for space manipulation target_space = target_data["space"] - # Set target standard space to target space - target_std_space = target_space # Extra data type requirement check if target space is native if target_space == "native": # Check for extra inputs @@ -413,8 +409,16 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): "data types in particular for transformation to " f"{target_data['space']} space for further computation." ) + # Get native space warper spec + warper_spec = get_native_warper( + target_data=target_data, + other_data=extra_input, + ) # Set target standard space to warp file space source - target_std_space = extra_input["Warp"]["src"] + target_std_space = warper_spec["src"] + else: + # Set target standard space to target space + target_std_space = target_space # Get the min of the voxels sizes and use it as the resolution target_img = target_data["data"] @@ -428,12 +432,14 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): all_parcellations = [] all_labels = [] for name in parcellations: + # Load parcellation img, labels, _, space = self.load( name=name, resolution=resolution, ) - # Convert parcellation spaces if required + # Convert parcellation spaces if required; + # cannot be "native" due to earlier check if space != target_std_space: raw_img = ANTsParcellationWarper().warp( parcellation_name=name, @@ -441,7 +447,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): src=space, dst=target_std_space, target_data=target_data, - extra_input=None, + warp_data=None, ) # Remove extra dimension added by ANTs img = image.math_img("np.squeeze(img)", img=raw_img) @@ -471,32 +477,22 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Warp parcellation if target space is native if target_space == "native": - # extra_input check done earlier - # Check for warp file type to use correct tool - warp_file_ext = extra_input["Warp"]["path"].suffix - if warp_file_ext == ".mat": + # extra_input check done earlier and warper_spec exists + if warper_spec["warper"] == "fsl": resampled_parcellation_img = FSLParcellationWarper().warp( parcellation_name="native", parcellation_img=resampled_parcellation_img, target_data=target_data, - extra_input=extra_input, + warp_data=warper_spec, ) - elif warp_file_ext == ".h5": + elif warper_spec["warper"] == "ants": resampled_parcellation_img = ANTsParcellationWarper().warp( parcellation_name="native", parcellation_img=resampled_parcellation_img, src="", - dst="T1w", + dst="native", target_data=target_data, - extra_input=extra_input, - ) - else: - raise_error( - msg=( - "Unknown warp / transformation file extension: " - f"{warp_file_ext}" - ), - klass=RuntimeError, + warp_data=warper_spec, ) return resampled_parcellation_img, labels diff --git a/junifer/data/utils.py b/junifer/data/utils.py index dd73a5388..a03171417 100644 --- a/junifer/data/utils.py +++ b/junifer/data/utils.py @@ -4,14 +4,14 @@ # Synchon Mandal # License: AGPL -from typing import List, Optional, Union +from typing import Dict, List, MutableMapping, Optional, Union import numpy as np -from ..utils.logging import logger +from ..utils import logger, raise_error -__all__ = ["closest_resolution"] +__all__ = ["closest_resolution", "get_native_warper"] def closest_resolution( @@ -49,3 +49,67 @@ def closest_resolution( closest = np.min(valid_resolution) return closest + + +def get_native_warper( + target_data: MutableMapping, + other_data: MutableMapping, + inverse: bool = False, +) -> Dict: + """Get correct warping specification for native space. + + Parameters + ---------- + target_data : dict + The target data from the pipeline data object. + other_data : dict + The other data in the pipeline data object. + inverse : bool, optional + Whether to get the inverse warping specification (default False). + + Returns + ------- + dict + The correct warping specification. + + Raises + ------ + RuntimeError + If no warper or multiple possible warpers are found. + + """ + # Get possible warpers + possible_warpers = [] + for entry in other_data["Warp"]: + if not inverse: + if ( + entry["src"] == target_data["prewarp_space"] + and entry["dst"] == "native" + ): + possible_warpers.append(entry) + else: + if ( + entry["dst"] == target_data["prewarp_space"] + and entry["src"] == "native" + ): + possible_warpers.append(entry) + + # Check for no warper + if not possible_warpers: + raise_error( + klass=RuntimeError, + msg="Could not find correct warping specification", + ) + + # Check for multiple possible warpers + if len(possible_warpers) > 1: + raise_error( + klass=RuntimeError, + msg=( + "Cannot proceed as multiple warping specification found, " + "adjust either the DataGrabber or the working space: " + f"{possible_warpers}" + ), + ) + + return possible_warpers[0] diff --git a/junifer/datagrabber/aomic/id1000.py b/junifer/datagrabber/aomic/id1000.py index 7d0586255..cf5fc90e4 100644 --- a/junifer/datagrabber/aomic/id1000.py +++ b/junifer/datagrabber/aomic/id1000.py @@ -169,15 +169,28 @@ class DataladAOMICID1000(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + "warper": "ants", + }, + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + "warper": "ants", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/aomic/piop1.py b/junifer/datagrabber/aomic/piop1.py index d2763bb18..29751e842 100644 --- a/junifer/datagrabber/aomic/piop1.py +++ b/junifer/datagrabber/aomic/piop1.py @@ -204,15 +204,28 @@ class DataladAOMICPIOP1(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + "warper": "ants", + }, + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + "warper": "ants", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/aomic/piop2.py b/junifer/datagrabber/aomic/piop2.py index a9852ab75..53d922d13 100644 --- a/junifer/datagrabber/aomic/piop2.py +++ b/junifer/datagrabber/aomic/piop2.py @@ -202,15 +202,28 @@ class DataladAOMICPIOP2(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + "warper": "ants", + }, + { + "pattern": ( + "derivatives/fmriprep/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + "warper": "ants", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 9e6d0246f..00b64cb69 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -96,7 +96,12 @@ class BaseDataGrabber(ABC, UpdateMetaMixin): # Update metadata for _, t_val in out.items(): self.update_meta(t_val, "datagrabber") - t_val["meta"]["element"] = named_element + # Conditional for list dtype vals like Warp + if isinstance(t_val, list): + for entry in t_val: + entry["meta"]["element"] = named_element + else: + t_val["meta"]["element"] = named_element return out diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 2d4cb163b..55f3bb1fa 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -181,14 +181,21 @@ class DataladDataGrabber(BaseDataGrabber): """ to_get = [] for type_val in out.values(): - # Iterate to check for nested "types" like mask - for k, v in type_val.items(): - # Add base data type path - if k == "path": - to_get.append(v) - # Add nested data type path - if isinstance(v, dict) and "path" in v: - to_get.append(v["path"]) + # Conditional for list dtype vals like Warp + if isinstance(type_val, list): + for entry in type_val: + for k, v in entry.items(): + if k == "path": + to_get.append(v) + else: + # Iterate to check for nested "types" like mask + for k, v in type_val.items(): + # Add base data type path + if k == "path": + to_get.append(v) + # Add nested data type path + if isinstance(v, dict) and "path" in v: + to_get.append(v["path"]) if len(to_get) > 0: logger.debug(f"Getting {len(to_get)} files using datalad:") diff --git a/junifer/datagrabber/dmcc13_benchmark.py b/junifer/datagrabber/dmcc13_benchmark.py index af4d32f96..7111624b1 100644 --- a/junifer/datagrabber/dmcc13_benchmark.py +++ b/junifer/datagrabber/dmcc13_benchmark.py @@ -150,7 +150,7 @@ class DMCC13Benchmark(PatternDataladDataGrabber): "mask": { "pattern": ( "derivatives/fmriprep-1.3.2/{subject}/{session}/" - "/func/{subject}_{session}_task-{task}_acq-mb4" + "func/{subject}_{session}_task-{task}_acq-mb4" "{phase_encoding}_run-{run}_" "space-MNI152NLin2009cAsym_desc-brain_mask.nii.gz" ), @@ -221,15 +221,28 @@ class DMCC13Benchmark(PatternDataladDataGrabber): "space": "native", }, }, - "Warp": { - "pattern": ( - "derivatives/fmriprep-1.3.2/{subject}/anat/" - "{subject}_from-MNI152NLin2009cAsym_to-T1w_" - "mode-image_xfm.h5" - ), - "src": "MNI152NLin2009cAsym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "derivatives/fmriprep-1.3.2/{subject}/anat/" + "{subject}_from-MNI152NLin2009cAsym_to-T1w_" + "mode-image_xfm.h5" + ), + "src": "MNI152NLin2009cAsym", + "dst": "native", + "warper": "ants", + }, + { + "pattern": ( + "derivatives/fmriprep-1.3.2/{subject}/anat/" + "{subject}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), + "src": "native", + "dst": "MNI152NLin2009cAsym", + "warper": "ants", + }, + ], } ) # Set default types diff --git a/junifer/datagrabber/hcp1200/hcp1200.py b/junifer/datagrabber/hcp1200/hcp1200.py index 0fa8a8018..c2455d16e 100644 --- a/junifer/datagrabber/hcp1200/hcp1200.py +++ b/junifer/datagrabber/hcp1200/hcp1200.py @@ -122,13 +122,24 @@ class HCP1200(PatternDataGrabber): "pattern": "{subject}/T1w/T1w_acpc_dc_restore.nii.gz", "space": "native", }, - "Warp": { - "pattern": ( - "{subject}/MNINonLinear/xfms/standard2acpc_dc.nii.gz" - ), - "src": "MNI152NLin6Asym", - "dst": "native", - }, + "Warp": [ + { + "pattern": ( + "{subject}/MNINonLinear/xfms/standard2acpc_dc.nii.gz" + ), + "src": "MNI152NLin6Asym", + "dst": "native", + "warper": "fsl", + }, + { + "pattern": ( + "{subject}/MNINonLinear/xfms/acpc_dc2standard.nii.gz" + ), + "src": "native", + "dst": "MNI152NLin6Asym", + "warper": "fsl", + }, + ], } # The replacements replacements = ["subject", "task", "phase_encoding"] diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index 9790e7dba..25012b074 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -89,7 +89,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): .. code-block:: none { - "mandatory": ["pattern", "src", "dst"], + "mandatory": ["pattern", "src", "dst", "warper"], "optional": [] } @@ -127,11 +127,22 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): "pattern": "...", "space": "...", }, - "Warp": { - "pattern": "...", - "src": "...", - "dst": "...", - } + } + + except ``Warp``, which needs to be a list of dictionaries as there can + be multiple spaces to warp (for example, with fMRIPrep): + + .. code-block:: none + + { + "Warp": [ + { + "pattern": "...", + "src": "...", + "dst": "...", + "warper": "...", + }, + ], } taken from :class:`.HCP1200`. @@ -380,42 +391,61 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin): t_pattern = self.patterns[t_type] # Copy data type dictionary in output out[t_type] = deepcopy(t_pattern) - # Iterate to check for nested "types" like mask - for k, v in t_pattern.items(): - # Resolve pattern for base data type - if k == "pattern": - logger.info(f"Resolving path from pattern for {t_type}") - # Resolve pattern - base_data_type_pattern_path = self._get_path_from_patterns( - element=element, - pattern=v, - data_type=t_type, - ) - # Remove pattern key - out[t_type].pop("pattern") - # Add path key - out[t_type].update({"path": base_data_type_pattern_path}) - # Resolve pattern for nested data type - if isinstance(v, dict) and "pattern" in v: - # Set nested type key for easier access - t_nested_type = f"{t_type}.{k}" + # Conditional for list dtype vals like Warp + if isinstance(t_pattern, list): + for idx, entry in enumerate(t_pattern): logger.info( - f"Resolving path from pattern for {t_nested_type}" + f"Resolving path from pattern for {t_type}.{idx}" ) # Resolve pattern - nested_data_type_pattern_path = ( - self._get_path_from_patterns( - element=element, - pattern=v["pattern"], - data_type=t_nested_type, - ) + dtype_pattern_path = self._get_path_from_patterns( + element=element, + pattern=entry["pattern"], + data_type=f"{t_type}.{idx}", ) # Remove pattern key - out[t_type][k].pop("pattern") + out[t_type][idx].pop("pattern") # Add path key - out[t_type][k].update( - {"path": nested_data_type_pattern_path} - ) + out[t_type][idx].update({"path": dtype_pattern_path}) + else: + # Iterate to check for nested "types" like mask + for k, v in t_pattern.items(): + # Resolve pattern for base data type + if k == "pattern": + logger.info( + f"Resolving path from pattern for {t_type}" + ) + # Resolve pattern + base_dtype_pattern_path = self._get_path_from_patterns( + element=element, + pattern=v, + data_type=t_type, + ) + # Remove pattern key + out[t_type].pop("pattern") + # Add path key + out[t_type].update({"path": base_dtype_pattern_path}) + # Resolve pattern for nested data type + if isinstance(v, dict) and "pattern" in v: + # Set nested type key for easier access + t_nested_type = f"{t_type}.{k}" + logger.info( + f"Resolving path from pattern for {t_nested_type}" + ) + # Resolve pattern + nested_dtype_pattern_path = ( + self._get_path_from_patterns( + element=element, + pattern=v["pattern"], + data_type=t_nested_type, + ) + ) + # Remove pattern key + out[t_type][k].pop("pattern") + # Add path key + out[t_type][k].update( + {"path": nested_dtype_pattern_path} + ) return out diff --git a/junifer/datagrabber/pattern_validation_mixin.py b/junifer/datagrabber/pattern_validation_mixin.py index 9c130b902..e9e1a8602 100644 --- a/junifer/datagrabber/pattern_validation_mixin.py +++ b/junifer/datagrabber/pattern_validation_mixin.py @@ -3,7 +3,7 @@ # Authors: Synchon Mandal # License: AGPL -from typing import Dict, List +from typing import Dict, List, Union from ..utils import logger, raise_error, warn_with_log @@ -36,7 +36,7 @@ PATTERNS_SCHEMA = { }, }, "Warp": { - "mandatory": ["pattern", "src", "dst"], + "mandatory": ["pattern", "src", "dst", "warper"], "optional": {}, }, "VBM_GM": { @@ -96,7 +96,7 @@ class PatternValidationMixin: def _validate_replacements( self, replacements: List[str], - patterns: Dict[str, Dict[str, str]], + patterns: Dict[str, Union[Dict[str, str], List[Dict[str, str]]]], partial_pattern_ok: bool, ) -> None: """Validate the replacements. @@ -132,39 +132,51 @@ class PatternValidationMixin: if any(not isinstance(x, str) for x in replacements): raise_error( - msg="`replacements` must be a list of strings.", + msg="`replacements` must be a list of strings", klass=TypeError, ) + # Make a list of all patterns recursively + all_patterns = [] + for dtype_val in patterns.values(): + # Conditional for list dtype vals like Warp + if isinstance(dtype_val, list): + for entry in dtype_val: + all_patterns.append(entry.get("pattern", "")) + else: + all_patterns.append(dtype_val.get("pattern", "")) + # Check for stray replacements for x in replacements: - if all( - x not in y - for y in [ - data_type_val.get("pattern", "") - for data_type_val in patterns.values() - ] - ): + if all(x not in y for y in all_patterns): if partial_pattern_ok: warn_with_log( f"Replacement: `{x}` is not part of any pattern, " "things might not work as expected if you are unsure " - "of what you are doing" + "of what you are doing." ) else: raise_error( - msg=f"Replacement: {x} is not part of any pattern." + msg=f"Replacement: `{x}` is not part of any pattern" ) # Check that at least one pattern has all the replacements at_least_one = False - for data_type_val in patterns.values(): - if all( - x in data_type_val.get("pattern", "") for x in replacements - ): - at_least_one = True + for dtype_val in patterns.values(): + # Conditional for list dtype vals like Warp + if isinstance(dtype_val, list): + for entry in dtype_val: + if all( + x in entry.get("pattern", "") for x in replacements + ): + at_least_one = True + else: + if all( + x in dtype_val.get("pattern", "") for x in replacements + ): + at_least_one = True if not at_least_one and not partial_pattern_ok: raise_error( - msg="At least one pattern must contain all replacements." + msg="At least one pattern must contain all replacements" ) def _validate_mandatory_keys( @@ -207,7 +219,7 @@ class PatternValidationMixin: warn_with_log( f"Mandatory key: `{key}` not found for {data_type}, " "things might not work as expected if you are unsure " - "of what you are doing" + "of what you are doing." ) else: raise_error( @@ -215,7 +227,7 @@ class PatternValidationMixin: klass=KeyError, ) else: - logger.debug(f"Mandatory key: `{key}` found for {data_type}") + logger.debug(f"Mandatory key: `{key}` found for {data_type}.") def _identify_stray_keys( self, keys: List[str], schema: List[str], data_type: str @@ -251,7 +263,7 @@ class PatternValidationMixin: self, types: List[str], replacements: List[str], - patterns: Dict[str, Dict[str, str]], + patterns: Dict[str, Union[Dict[str, str], List[Dict[str, str]]]], partial_pattern_ok: bool = False, ) -> None: """Validate the patterns. @@ -298,87 +310,185 @@ class PatternValidationMixin: msg="`patterns` must contain all `types`", klass=ValueError ) # Check against schema - for data_type_key, data_type_val in patterns.items(): + for dtype_key, dtype_val in patterns.items(): # Check if valid data type is provided - if data_type_key not in PATTERNS_SCHEMA: + if dtype_key not in PATTERNS_SCHEMA: raise_error( - f"Unknown data type: {data_type_key}, " + f"Unknown data type: {dtype_key}, " f"should be one of: {list(PATTERNS_SCHEMA.keys())}" ) - # Check mandatory keys for data type - self._validate_mandatory_keys( - keys=list(data_type_val), - schema=PATTERNS_SCHEMA[data_type_key]["mandatory"], - data_type=data_type_key, - partial_pattern_ok=partial_pattern_ok, - ) - # Check optional keys for data type - for optional_key, optional_val in PATTERNS_SCHEMA[data_type_key][ - "optional" - ].items(): - if optional_key not in data_type_val: - logger.debug( - f"Optional key: `{optional_key}` missing for " - f"{data_type_key}" - ) - else: - logger.debug( - f"Optional key: `{optional_key}` found for " - f"{data_type_key}" - ) - # Set nested type name for easier access - nested_data_type = f"{data_type_key}.{optional_key}" - nested_mandatory_keys_schema = PATTERNS_SCHEMA[ - data_type_key - ]["optional"][optional_key]["mandatory"] - nested_optional_keys_schema = PATTERNS_SCHEMA[ - data_type_key - ]["optional"][optional_key]["optional"] - # Check mandatory keys for nested type + # Conditional for list dtype vals like Warp + if isinstance(dtype_val, list): + for idx, entry in enumerate(dtype_val): + # Check mandatory keys for data type self._validate_mandatory_keys( - keys=list(optional_val["mandatory"]), - schema=nested_mandatory_keys_schema, - data_type=nested_data_type, + keys=list(entry), + schema=PATTERNS_SCHEMA[dtype_key]["mandatory"], + data_type=f"{dtype_key}.{idx}", partial_pattern_ok=partial_pattern_ok, ) - # Check optional keys for nested type - for nested_optional_key in nested_optional_keys_schema: - if nested_optional_key not in optional_val["optional"]: + # Check optional keys for data type + for optional_key, optional_val in PATTERNS_SCHEMA[ + dtype_key + ]["optional"].items(): + if optional_key not in entry: logger.debug( - f"Optional key: `{nested_optional_key}` " - f"missing for {nested_data_type}" + f"Optional key: `{optional_key}` missing for " + f"{dtype_key}.{idx}" ) else: logger.debug( - f"Optional key: `{nested_optional_key}` found " - f"for {nested_data_type}" + f"Optional key: `{optional_key}` found for " + f"{dtype_key}.{idx}" ) - # Check stray key for nested data type + # Set nested type name for easier access + nested_dtype = f"{dtype_key}.{idx}.{optional_key}" + nested_mandatory_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["mandatory"] + nested_optional_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["optional"] + # Check mandatory keys for nested type + self._validate_mandatory_keys( + keys=list(optional_val["mandatory"]), + schema=nested_mandatory_keys_schema, + data_type=nested_dtype, + partial_pattern_ok=partial_pattern_ok, + ) + # Check optional keys for nested type + for ( + nested_optional_key + ) in nested_optional_keys_schema: + if ( + nested_optional_key + not in optional_val["optional"] + ): + logger.debug( + f"Optional key: " + f"`{nested_optional_key}` missing for " + f"{nested_dtype}" + ) + else: + logger.debug( + f"Optional key: " + f"`{nested_optional_key}` found for " + f"{nested_dtype}" + ) + # Check stray key for nested data type + self._identify_stray_keys( + keys=( + optional_val["mandatory"] + + optional_val["optional"] + ), + schema=( + nested_mandatory_keys_schema + + nested_optional_keys_schema + ), + data_type=nested_dtype, + ) + # Check stray key for data type self._identify_stray_keys( - keys=optional_val["mandatory"] - + optional_val["optional"], - schema=nested_mandatory_keys_schema - + nested_optional_keys_schema, - data_type=nested_data_type, + keys=list(entry.keys()), + schema=( + PATTERNS_SCHEMA[dtype_key]["mandatory"] + + list( + PATTERNS_SCHEMA[dtype_key]["optional"].keys() + ) + ), + data_type=dtype_key, ) - # Check stray key for data type - self._identify_stray_keys( - keys=list(data_type_val.keys()), - schema=( - PATTERNS_SCHEMA[data_type_key]["mandatory"] - + list(PATTERNS_SCHEMA[data_type_key]["optional"].keys()) - ), - data_type=data_type_key, - ) - # Wildcard check in patterns - if "}*" in data_type_val.get("pattern", ""): - raise_error( - msg=( - f"`{data_type_key}.pattern` must not contain `*` " - "following a replacement" - ), - klass=ValueError, + # Wildcard check in patterns + if "}*" in entry.get("pattern", ""): + raise_error( + msg=( + f"`{dtype_key}.pattern` must not contain `*` " + "following a replacement" + ), + klass=ValueError, + ) + else: + # Check mandatory keys for data type + self._validate_mandatory_keys( + keys=list(dtype_val), + schema=PATTERNS_SCHEMA[dtype_key]["mandatory"], + data_type=dtype_key, + partial_pattern_ok=partial_pattern_ok, ) + # Check optional keys for data type + for optional_key, optional_val in PATTERNS_SCHEMA[dtype_key][ + "optional" + ].items(): + if optional_key not in dtype_val: + logger.debug( + f"Optional key: `{optional_key}` missing for " + f"{dtype_key}." + ) + else: + logger.debug( + f"Optional key: `{optional_key}` found for " + f"{dtype_key}." + ) + # Set nested type name for easier access + nested_dtype = f"{dtype_key}.{optional_key}" + nested_mandatory_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["mandatory"] + nested_optional_keys_schema = PATTERNS_SCHEMA[ + dtype_key + ]["optional"][optional_key]["optional"] + # Check mandatory keys for nested type + self._validate_mandatory_keys( + keys=list(optional_val["mandatory"]), + schema=nested_mandatory_keys_schema, + data_type=nested_dtype, + partial_pattern_ok=partial_pattern_ok, + ) + # Check optional keys for nested type + for nested_optional_key in nested_optional_keys_schema: + if ( + nested_optional_key + not in optional_val["optional"] + ): + logger.debug( + f"Optional key: `{nested_optional_key}` " + f"missing for {nested_dtype}" + ) + else: + logger.debug( + f"Optional key: `{nested_optional_key}` " + f"found for {nested_dtype}" + ) + # Check stray key for nested data type + self._identify_stray_keys( + keys=( + optional_val["mandatory"] + + optional_val["optional"] + ), + schema=( + nested_mandatory_keys_schema + + nested_optional_keys_schema + ), + data_type=nested_dtype, + ) + # Check stray key for data type + self._identify_stray_keys( + keys=list(dtype_val.keys()), + schema=( + PATTERNS_SCHEMA[dtype_key]["mandatory"] + + list(PATTERNS_SCHEMA[dtype_key]["optional"].keys()) + ), + data_type=dtype_key, + ) + # Wildcard check in patterns + if "}*" in dtype_val.get("pattern", ""): + raise_error( + msg=( + f"`{dtype_key}.pattern` must not contain `*` " + "following a replacement" + ), + klass=ValueError, + ) # Validate replacements self._validate_replacements( diff --git a/junifer/datagrabber/tests/test_dmcc13_benchmark.py b/junifer/datagrabber/tests/test_dmcc13_benchmark.py index 2b2566de1..ca41950bb 100644 --- a/junifer/datagrabber/tests/test_dmcc13_benchmark.py +++ b/junifer/datagrabber/tests/test_dmcc13_benchmark.py @@ -116,7 +116,12 @@ def test_DMCC13Benchmark( data_file_names.extend( [ "sub-01_desc-preproc_T1w.nii.gz", - "sub-01_from-MNI152NLin2009cAsym_to-T1w_mode-image_xfm.h5", + [ + "sub-01_from-MNI152NLin2009cAsym_to-T1w" + "_mode-image_xfm.h5", + "sub-01_from-T1w_to-MNI152NLin2009cAsym" + "_mode-image_xfm.h5", + ], ] ) else: @@ -127,14 +132,26 @@ def test_DMCC13Benchmark( for data_type, data_file_name in zip(data_types, data_file_names): # Assert data type assert data_type in out - # Assert data file path exists - assert out[data_type]["path"].exists() - # Assert data file path is a file - assert out[data_type]["path"].is_file() - # Assert data file name - assert out[data_type]["path"].name == data_file_name - # Assert metadata - assert "meta" in out[data_type] + # Conditional for Warp + if data_type == "Warp": + for idx, fname in enumerate(data_file_name): + # Assert data file path exists + assert out[data_type][idx]["path"].exists() + # Assert data file path is a file + assert out[data_type][idx]["path"].is_file() + # Assert data file name + assert out[data_type][idx]["path"].name == fname + # Assert metadata + assert "meta" in out[data_type][idx] + else: + # Assert data file path exists + assert out[data_type]["path"].exists() + # Assert data file path is a file + assert out[data_type]["path"].is_file() + # Assert data file name + assert out[data_type]["path"].name == data_file_name + # Assert metadata + assert "meta" in out[data_type] # Check BOLD nested data types for type_, file_name in zip( diff --git a/junifer/pipeline/pipeline_step_mixin.py b/junifer/pipeline/pipeline_step_mixin.py index b4f45594f..84590dc2e 100644 --- a/junifer/pipeline/pipeline_step_mixin.py +++ b/junifer/pipeline/pipeline_step_mixin.py @@ -196,10 +196,14 @@ class PipelineStepMixin: for dependency in obj._CONDITIONAL_DEPENDENCIES: if dependency["using"] == obj.using: depends_on = dependency["depends_on"] - # Check dependencies - _check_dependencies(depends_on) - # Check external dependencies - _check_ext_dependencies(depends_on) + # Conditional to make `using="auto"` work + if not isinstance(depends_on, list): + depends_on = [depends_on] + for entry in depends_on: + # Check dependencies + _check_dependencies(entry) + # Check external dependencies + _check_ext_dependencies(entry) # Check dependencies _check_dependencies(self) diff --git a/junifer/pipeline/update_meta_mixin.py b/junifer/pipeline/update_meta_mixin.py index c88978244..abbbc07d0 100644 --- a/junifer/pipeline/update_meta_mixin.py +++ b/junifer/pipeline/update_meta_mixin.py @@ -4,7 +4,7 @@ # Synchon Mandal # License: AGPL -from typing import Dict +from typing import Dict, List, Union __all__ = ["UpdateMetaMixin"] @@ -15,14 +15,14 @@ class UpdateMetaMixin: def update_meta( self, - input: Dict, + input: Union[Dict, List[Dict]], step_name: str, ) -> None: """Update metadata. Parameters ---------- - input : dict + input : dict or list of dict The data object to update. step_name : str The name of the pipeline step. @@ -36,17 +36,21 @@ class UpdateMetaMixin: for k, v in vars(self).items(): if not k.startswith("_"): t_meta[k] = v - # Add "meta" to the step's local context dict - if "meta" not in input: - input["meta"] = {} - # Add step name - input["meta"][step_name] = t_meta - # Add step dependencies - if "dependencies" not in input["meta"]: - input["meta"]["dependencies"] = set() - # Update step dependencies - dependencies = getattr(self, "_DEPENDENCIES", set()) - if dependencies is not None: - if not isinstance(dependencies, (set, list)): - dependencies = {dependencies} - input["meta"]["dependencies"].update(dependencies) + # Conditional for list dtype vals like Warp + if not isinstance(input, list): + input = [input] + for entry in input: + # Add "meta" to the step's entry's local context dict + if "meta" not in entry: + entry["meta"] = {} + # Add step name + entry["meta"][step_name] = t_meta + # Add step dependencies + if "dependencies" not in entry["meta"]: + entry["meta"]["dependencies"] = set() + # Update step dependencies + dependencies = getattr(self, "_DEPENDENCIES", set()) + if dependencies is not None: + if not isinstance(dependencies, (set, list)): + dependencies = {dependencies} + entry["meta"]["dependencies"].update(dependencies) diff --git a/junifer/preprocess/warping/_ants_warper.py b/junifer/preprocess/warping/_ants_warper.py index dfc4e7bf8..9fdcaa0e6 100644 --- a/junifer/preprocess/warping/_ants_warper.py +++ b/junifer/preprocess/warping/_ants_warper.py @@ -15,7 +15,7 @@ import numpy as np from ...data import get_template, get_xfm from ...pipeline import WorkDirManager from ...typing import Dependencies, ExternalDependencies -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd __all__ = ["ANTsWarper"] @@ -63,6 +63,11 @@ class ANTsWarper: values and new ``reference_path`` key whose value points to the reference file used for warping. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ # Create element-specific tempdir for storing post-warping assets element_tempdir = WorkDirManager().get_element_tempdir( @@ -77,6 +82,17 @@ class ANTsWarper: # resolution resolution = np.min(input["data"].header.get_zooms()[:3]) + # Get warp file path + warp_file_path = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + warp_file_path = entry["path"] + if warp_file_path is None: + raise_error( + klass=RuntimeError, + msg="Could not find correct warp file path", + ) + # Create a tempfile for resampled reference output resample_image_out_path = ( element_tempdir / "resampled_reference.nii.gz" @@ -105,7 +121,7 @@ class ANTsWarper: f"-i {input['path'].resolve()}", # use resampled reference f"-r {resample_image_out_path.resolve()}", - f"-t {extra_input['Warp']['path'].resolve()}", + f"-t {warp_file_path.resolve()}", f"-o {apply_transforms_out_path.resolve()}", ] # Call antsApplyTransforms @@ -115,6 +131,8 @@ class ANTsWarper: input["data"] = nib.load(apply_transforms_out_path) # Save resampled reference path input["reference_path"] = resample_image_out_path + # Keep pre-warp space for further operations + input["prewarp_space"] = input["space"] # Use reference input's space as warped input's space input["space"] = extra_input["T1w"]["space"] @@ -163,6 +181,9 @@ class ANTsWarper: # Modify target data input["data"] = nib.load(warped_output_path) + # Keep pre-warp space for further operations + input["prewarp_space"] = input["space"] + # Update warped input's space input["space"] = reference return input diff --git a/junifer/preprocess/warping/_fsl_warper.py b/junifer/preprocess/warping/_fsl_warper.py index c04ae6582..1261ba4a6 100644 --- a/junifer/preprocess/warping/_fsl_warper.py +++ b/junifer/preprocess/warping/_fsl_warper.py @@ -14,7 +14,7 @@ import numpy as np from ...pipeline import WorkDirManager from ...typing import Dependencies, ExternalDependencies -from ...utils import logger, run_ext_cmd +from ...utils import logger, raise_error, run_ext_cmd __all__ = ["FSLWarper"] @@ -59,6 +59,11 @@ class FSLWarper: values and new ``reference_path`` key whose value points to the reference file used for warping. + Raises + ------ + RuntimeError + If warp file path could not be found in ``extra_input``. + """ logger.debug("Using FSL for space warping") @@ -66,6 +71,16 @@ class FSLWarper: # resolution resolution = np.min(input["data"].header.get_zooms()[:3]) + # Get warp file path + warp_file_path = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + warp_file_path = entry["path"] + if warp_file_path is None: + raise_error( + klass=RuntimeError, msg="Could not find correct warp file path" + ) + # Create element-specific tempdir for storing post-warping assets element_tempdir = WorkDirManager().get_element_tempdir( prefix="fsl_warper" @@ -93,7 +108,7 @@ class FSLWarper: "--interp=spline", f"-i {input['path'].resolve()}", f"-r {flirt_out_path.resolve()}", # use resampled reference - f"-w {extra_input['Warp']['path'].resolve()}", + f"-w {warp_file_path.resolve()}", f"-o {applywarp_out_path.resolve()}", ] # Call applywarp @@ -103,7 +118,8 @@ class FSLWarper: input["data"] = nib.load(applywarp_out_path) # Save resampled reference path input["reference_path"] = flirt_out_path - + # Keep pre-warp space for further operations + input["prewarp_space"] = input["space"] # Use reference input's space as warped input's space input["space"] = extra_input["T1w"]["space"] diff --git a/junifer/preprocess/warping/space_warper.py b/junifer/preprocess/warping/space_warper.py index 445bc4c5c..fdab703c3 100644 --- a/junifer/preprocess/warping/space_warper.py +++ b/junifer/preprocess/warping/space_warper.py @@ -24,11 +24,12 @@ class SpaceWarper(BasePreprocessor): Parameters ---------- - using : {"fsl", "ants"} + using : {"fsl", "ants", "auto"} Implementation to use for warping: * "fsl" : Use FSL's ``applywarp`` * "ants" : Use ANTs' ``antsApplyTransforms`` + * "auto" : Auto-select tool when ``reference="T1w"`` reference : str The data type to use as reference for warping, can be either a data @@ -56,6 +57,10 @@ class SpaceWarper(BasePreprocessor): "using": "ants", "depends_on": ANTsWarper, }, + { + "using": "auto", + "depends_on": [FSLWarper, ANTsWarper], + }, ] def __init__( @@ -156,14 +161,16 @@ class SpaceWarper(BasePreprocessor): 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 + If warper could not be found in ``extra_input`` when + ``using="auto"`` or + if the data is in the correct space and does not require warping or - if FSL is used for template space warping. + if FSL is used when ``reference="T1w"``. """ logger.info(f"Warping to {self.reference} space using SpaceWarper") # Transform to native space - if self.using in ["fsl", "ants"] and self.reference == "T1w": + if self.using in ["fsl", "ants", "auto"] and self.reference == "T1w": # Check for extra inputs if extra_input is None: raise_error( @@ -182,6 +189,26 @@ class SpaceWarper(BasePreprocessor): extra_input=extra_input, reference=self.reference, ) + elif self.using == "auto": + warper = None + for entry in extra_input["Warp"]: + if entry["dst"] == "native": + warper = entry["warper"] + if warper is None: + raise_error( + klass=RuntimeError, msg="Could not find correct warper" + ) + if warper == "fsl": + input = FSLWarper().preprocess( + input=input, + extra_input=extra_input, + ) + elif warper == "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 diff --git a/tools/create_dmcc13_benchmark_example_dataset.py b/tools/create_dmcc13_benchmark_example_dataset.py index 8be13edb3..996a897c3 100644 --- a/tools/create_dmcc13_benchmark_example_dataset.py +++ b/tools/create_dmcc13_benchmark_example_dataset.py @@ -78,7 +78,11 @@ if __name__ == "__main__": ( f"anat/sub-{sub:02d}_from-MNI152NLin2009cAsym_to-T1w_" "mode-image_xfm.h5" - ), # Warp + ), # Warp to native + ( + f"anat/sub-{sub:02d}_from-T1w_to-MNI152NLin2009cAsym_" + "mode-image_xfm.h5" + ), # Warp from native ] # Tasks for functional data diff --git a/tools/create_hcp1200_example_dataset.py b/tools/create_hcp1200_example_dataset.py index 0759cbb9b..9929bbd9f 100644 --- a/tools/create_hcp1200_example_dataset.py +++ b/tools/create_hcp1200_example_dataset.py @@ -44,11 +44,15 @@ if __name__ == "__main__": sub_warp_datadir = subdir / "MNINonLinear" / "xfms" # Create subject data directory sub_warp_datadir.mkdir(parents=True) - # Set subject data file - sub_warp_datafile = sub_warp_datadir / "standard2acpc_dc.nii.gz" - # Write subject data file - with open(sub_warp_datafile, "w") as f: - f.write("placeholder") + # Set subject data files + sub_warp_datafiles = [ + sub_warp_datadir / "standard2acpc_dc.nii.gz", + sub_warp_datadir / "acpc_dc2standard.nii.gz", + ] + # Write subject data files + for datafile in sub_warp_datafiles: + with open(datafile, "w") as f: + f.write("placeholder") # BOLD data for task, phase_encoding in product(