[ENH]: Simplify Preprocess interface #473

Merged
synchon merged 13 commits from refactor/preprocessor into main 2025-11-07 11:56:59 +00:00
11 changed files with 90 additions and 288 deletions

View file

@ -0,0 +1 @@
Simplify ``Preprocess`` interface and implementations by `Synchon Mandal`_

View file

@ -14,16 +14,11 @@ new ones, you might need something specific and then you can create your
own Preprocessor. own Preprocessor.
While implementing your own Preprocessor, you need to always inherit from While implementing your own Preprocessor, you need to always inherit from
:class:`.BasePreprocessor` and implement a few methods: :class:`.BasePreprocessor` and implement a few methods and class attributes:
#. ``get_valid_inputs``: This method should return a list of strings
representing the valid data types that the Preprocessor can work on.
Check :ref:`data types <data_types>` for reference.
#. ``get_output_type``: This method should just return the input as it
is unused as of now.
#. ``preprocess``: The method that given the data, preprocesses the data.
#. ``__init__``: The initialisation method, where the Preprocessor is #. ``__init__``: The initialisation method, where the Preprocessor is
configured. configured.
#. ``preprocess``: The method that given the data, preprocesses the data.
As an example, we will develop a ``NilearnSmoothing`` Preprocessor, which As an example, we will develop a ``NilearnSmoothing`` Preprocessor, which
smoothens the data using :func:`nilearn.image.smooth_img`. This is often smoothens the data using :func:`nilearn.image.smooth_img`. This is often
@ -32,37 +27,16 @@ desirable in cases where your data is preprocessed using ``fMRIPrep``, as
.. _extending_preprocessors_input_output: .. _extending_preprocessors_input_output:
Step 1: Configure input and output Step 1: Configure input
---------------------------------- -----------------------
In this step, we define the input and output data types of the Preprocessor. In this step, we define the input data types of the Preprocessor.
For input we can accept ``T1w``, ``T2w`` and ``BOLD`` For input we can accept ``T1w``, ``T2w`` and ``BOLD``
:ref:`data types <data_types>`. :ref:`data types <data_types>` and thus declare them in a class attribute:
.. code-block:: python .. code-block:: python
... _VALID_DATA_TYPES = ["T1w", "T2w", "BOLD"]
def get_valid_inputs(self) -> list[str]:
return ["T1w", "T2w", "BOLD"]
...
The output definition of the Preprocessor is unused now but is kept for
completeness.
.. code-block:: python
...
def get_output_type(self, input_type: str) -> str:
return input_type
...
.. _extending_preprocessors_init: .. _extending_preprocessors_init:
@ -77,7 +51,7 @@ you configure it. Our class will have the following arguments:
pass the value to it. pass the value to it.
2. ``on``: The data type we want the Preprocessor to work on. If the user does 2. ``on``: The data type we want the Preprocessor to work on. If the user does
not specify, it will work on all the data types given by the not specify, it will work on all the data types given by the
``get_valid_inputs`` function. ``_VALID_DATA_TYPES`` attribute.
.. attention:: .. attention::
@ -133,17 +107,6 @@ arguments:
useful if you want to use other data (e.g., ``Warp`` can be used to provide useful if you want to use other data (e.g., ``Warp`` can be used to provide
the transformation matrix file for transformation to subject-native space). the transformation matrix file for transformation to subject-native space).
and it has two return values:
* First is the ``input`` dictionary with necessary data modified. Usually, you
want to replace the ``input["data"]`` with the preprocessed data.
* Second is a dictionary just like ``input`` or ``extra_input`` but with only
specific key-value pairs which you would like to pass down to the Markers.
For example, if your Preprocessor computes some mask with the preprocessed
data, you could pass it through this which would be added and available
in the Marker step with the same key you pass here. Usually, you would
want to pass ``None``.
.. code-block:: python .. code-block:: python
from typing import Any from typing import Any
@ -158,9 +121,9 @@ and it has two return values:
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: dict[str, Any] | None = None, extra_input: dict[str, Any] | None = None,
) -> tuple[dict[str, Any], dict[str, Any] | None]: ) -> dict[str, Any]:
input["data"] = nimg.smooth_img(imgs=input["data"], fwhm=self.fwhm) input["data"] = nimg.smooth_img(imgs=input["data"], fwhm=self.fwhm)
return input, None return input
... ...
@ -187,7 +150,8 @@ decorator and our final code should look like this:
.. code-block:: python .. code-block:: python
from typing import Any, Literal from collections.abc import Sequence
from typing import Any, ClassVar, Literal
from junifer.api.decorators import register_preprocessor from junifer.api.decorators import register_preprocessor
from junifer.preprocess import BasePreprocessor from junifer.preprocess import BasePreprocessor
@ -201,6 +165,8 @@ decorator and our final code should look like this:
_DEPENDENCIES = {"nilearn"} _DEPENDENCIES = {"nilearn"}
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["T1w", "T2w", "BOLD"]
def __init__( def __init__(
self, self,
fwhm: int | float | ArrayLike | Literal["fast"] | None, fwhm: int | float | ArrayLike | Literal["fast"] | None,
@ -209,19 +175,13 @@ decorator and our final code should look like this:
self.fwhm = fwhm self.fwhm = fwhm
super().__init__(on=on) super().__init__(on=on)
def get_valid_inputs(self) -> list[str]:
return ["T1w", "T2w", "BOLD"]
def get_output_type(self, input_type: str) -> str:
return input_type
def preprocess( def preprocess(
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: dict[str, Any] | None = None, extra_input: dict[str, Any] | None = None,
) -> tuple[dict[str, Any], dict[str, Any] | None]: ) -> dict[str, Any]:
input["data"] = nimg.smooth_img(imgs=input["data"], fwhm=self.fwhm) input["data"] = nimg.smooth_img(imgs=input["data"], fwhm=self.fwhm)
return input, None return input
.. _extending_preprocessors_template: .. _extending_preprocessors_template:
@ -238,18 +198,16 @@ Template for a custom Preprocessor
@register_preprocessor @register_preprocessor
class TemplatePreprocessor(BasePreprocessor): class TemplatePreprocessor(BasePreprocessor):
# TODO: add the dependencies
_DEPENDENCIES = {}
# TODO: add the inputs
_VALID_DATA_TYPES = []
def __init__(self, on=None): def __init__(self, on=None):
# TODO: add preprocessor-specific parameters # TODO: add preprocessor-specific parameters
super().__init__(on=on) super().__init__(on=on)
def get_valid_inputs(self):
# TODO: Complete with the valid inputs
valid = []
return valid
def get_output_type(self, input_type):
return input_type
def preprocess(self, input, extra_input): def preprocess(self, input, extra_input):
# TODO: add the preprocessor logic # TODO: add the preprocessor logic
return input, None return input

View file

@ -5,6 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Sequence
from typing import ( from typing import (
Any, Any,
ClassVar, ClassVar,
@ -56,6 +57,7 @@ class TemporalFilter(BasePreprocessor):
""" """
_DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"} _DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"}
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD"]
def __init__( def __init__(
self, self,
@ -74,36 +76,7 @@ class TemporalFilter(BasePreprocessor):
self.t_r = t_r self.t_r = t_r
self.masks = masks self.masks = masks
super().__init__(on="BOLD", required_data_types=["BOLD"]) super().__init__()
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 ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The input to the preprocessor.
Returns
-------
str
The data type output by the preprocessor.
"""
# Does not add any new keys
return input_type
def _validate_data( def _validate_data(
self, self,
@ -130,7 +103,7 @@ class TemporalFilter(BasePreprocessor):
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: ) -> dict[str, Any]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -145,9 +118,6 @@ class TemporalFilter(BasePreprocessor):
dict dict
The computed result as dictionary. If `self.masks` is not None, The computed result as dictionary. If `self.masks` is not None,
then the target data computed mask is updated for further steps. then the target data computed mask is updated for further steps.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
""" """
# Validate data # Validate data
@ -237,4 +207,4 @@ class TemporalFilter(BasePreprocessor):
} }
) )
return input, None return input

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Sequence
from typing import Any, ClassVar, Optional from typing import Any, ClassVar, Optional
import nibabel as nib import nibabel as nib
@ -45,6 +46,7 @@ class TemporalSlicer(BasePreprocessor):
""" """
_DEPENDENCIES: ClassVar[Dependencies] = {"nilearn"} _DEPENDENCIES: ClassVar[Dependencies] = {"nilearn"}
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD"]
def __init__( def __init__(
self, self,
@ -61,42 +63,13 @@ class TemporalSlicer(BasePreprocessor):
self.stop = stop self.stop = stop
self.duration = duration self.duration = duration
self.t_r = t_r self.t_r = t_r
super().__init__(on="BOLD", required_data_types=["BOLD"]) super().__init__()
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 ["BOLD"]
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( def preprocess(
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: ) -> dict[str, Any]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -110,9 +83,6 @@ class TemporalSlicer(BasePreprocessor):
------- -------
dict dict
The computed result as dictionary. The computed result as dictionary.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
Raises Raises
------ ------
@ -233,4 +203,4 @@ class TemporalSlicer(BasePreprocessor):
} }
) )
return input, None return input

View file

@ -5,7 +5,8 @@
# License: AGPL # License: AGPL
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import Any, Optional, Union from collections.abc import Sequence
from typing import Any, ClassVar, Optional, Union
from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..pipeline import PipelineStepMixin, UpdateMetaMixin
from ..utils import logger, raise_error from ..utils import logger, raise_error
@ -15,15 +16,15 @@ __all__ = ["BasePreprocessor"]
class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
"""Abstract base class for all preprocessors. """Abstract base class for preprocessor.
For every interface that is required, one needs to provide a concrete For every preprocessor, one needs to provide a concrete
implementation of this abstract class. implementation of this abstract class.
Parameters Parameters
---------- ----------
on : str or list of str or None, optional on : str or list of str or None, optional
The data type to apply the preprocessor on. If None, The data type(s) to apply the preprocessor on. If None,
will work on all available data types (default None). will work on all available data types (default None).
required_data_types : str or list of str, optional required_data_types : str or list of str, optional
The data types needed for computation. If None, The data types needed for computation. If None,
@ -31,17 +32,27 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
Raises Raises
------ ------
AttributeError
If the preprocessor does not have `_VALID_DATA_TYPES` attribute.
ValueError ValueError
If required input data type(s) is(are) not found. If required input data type(s) is(are) not found.
""" """
_VALID_DATA_TYPES: ClassVar[Sequence[str]]
def __init__( def __init__(
self, self,
on: Optional[Union[list[str], str]] = None, on: Optional[Union[list[str], str]] = None,
required_data_types: Optional[Union[list[str], str]] = None, required_data_types: Optional[Union[list[str], str]] = None,
) -> None: ) -> None:
"""Initialize the class.""" """Initialize the class."""
# Check for missing data types attributes
if not hasattr(self, "_VALID_DATA_TYPES"):
raise_error(
msg="Missing `_VALID_DATA_TYPES` for the preprocessor",
klass=AttributeError,
)
# Use all data types if not provided # Use all data types if not provided
if on is None: if on is None:
on = self.get_valid_inputs() on = self.get_valid_inputs()
@ -58,6 +69,9 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
if required_data_types is None: if required_data_types is None:
self._required_data_types = on self._required_data_types = on
else: else:
# Convert data types to list
if not isinstance(required_data_types, list):
required_data_types = [required_data_types]
self._required_data_types = required_data_types self._required_data_types = required_data_types
def validate_input(self, input: list[str]) -> list[str]: def validate_input(self, input: list[str]) -> list[str]:
@ -89,7 +103,6 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
) )
return [x for x in self._on if x in input] return [x for x in self._on if x in input]
@abstractmethod
def get_valid_inputs(self) -> list[str]: def get_valid_inputs(self) -> list[str]:
"""Get valid data types for input. """Get valid data types for input.
@ -100,12 +113,8 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
preprocessor. preprocessor.
""" """
raise_error( return list(self._VALID_DATA_TYPES)
msg="Concrete classes need to implement get_valid_inputs().",
klass=NotImplementedError,
)
@abstractmethod
def get_output_type(self, input_type: str) -> str: def get_output_type(self, input_type: str) -> str:
"""Get output type. """Get output type.
@ -120,17 +129,15 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
The data type output by the preprocessor. The data type output by the preprocessor.
""" """
raise_error( # Does not add any new keys
msg="Concrete classes need to implement get_output_type().", return input_type
klass=NotImplementedError,
)
@abstractmethod @abstractmethod
def preprocess( def preprocess(
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: ) -> dict[str, Any]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -147,10 +154,6 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
------- -------
dict dict
The computed result as dictionary. The computed result as dictionary.
dict or None
Extra "helper" data types as dictionary to add to the Junifer Data
object. If no new "helper" data type(s) is(are) created, None is to
be passed.
""" """
raise_error( raise_error(
@ -192,19 +195,10 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
f"Extra data type for preprocess: {extra_input.keys()}" f"Extra data type for preprocess: {extra_input.keys()}"
) )
# Preprocess data # Preprocess data
t_out, t_extra_input = self.preprocess( t_out = self.preprocess(input=t_input, extra_input=extra_input)
input=t_input, extra_input=extra_input
)
# Set output to the Junifer Data object # Set output to the Junifer Data object
logger.debug(f"Adding {type_} to output") logger.debug(f"Adding {type_} to output")
out[type_] = t_out out[type_] = t_out
# Check if helper data types are to be added
if t_extra_input is not None:
logger.debug(
f"Adding helper data types: {t_extra_input.keys()} "
"to output"
)
out.update(t_extra_input)
# Update metadata for step # Update metadata for step
self.update_meta(out[type_], "preprocess") self.update_meta(out[type_], "preprocess")
return out return out

View file

@ -5,6 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Sequence
from typing import ( from typing import (
Any, Any,
ClassVar, ClassVar,
@ -175,6 +176,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
""" """
_DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"} _DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"}
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD"]
def __init__( def __init__(
self, self,
@ -251,36 +253,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
"include it in the future", "include it in the future",
klass=ValueError, klass=ValueError,
) )
super().__init__(on="BOLD", required_data_types=["BOLD"]) super().__init__()
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 ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The input to the preprocessor.
Returns
-------
str
The data type output by the preprocessor.
"""
# Does not add any new keys
return input_type
def _map_adhoc_to_fmriprep(self, input: dict[str, Any]) -> None: def _map_adhoc_to_fmriprep(self, input: dict[str, Any]) -> None:
"""Map the adhoc format to the fmpriprep format spec. """Map the adhoc format to the fmpriprep format spec.
@ -621,7 +594,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: ) -> dict[str, Any]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -636,9 +609,6 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
dict dict
The computed result as dictionary. If `self.masks` is not None, The computed result as dictionary. If `self.masks` is not None,
then the target data computed mask is updated for further steps. then the target data computed mask is updated for further steps.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
""" """
# Validate data # Validate data
@ -753,4 +723,4 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
} }
) )
return input, None return input

View file

@ -463,7 +463,7 @@ def test_fMRIPrepConfoundRemover_preprocess() -> None:
pre_extra_input = { pre_extra_input = {
"BOLD": {"confounds": element_data["BOLD"]["confounds"]} "BOLD": {"confounds": element_data["BOLD"]["confounds"]}
} }
output, _ = confound_remover.preprocess(pre_input, pre_extra_input) output = confound_remover.preprocess(pre_input, pre_extra_input)
trans_bold = output["data"].get_fdata() trans_bold = output["data"].get_fdata()
# Transformation is in place # Transformation is in place
assert_array_equal( assert_array_equal(
@ -614,7 +614,7 @@ def test_fMRIPrepConfoundRemover_scrubbing() -> None:
pre_extra_input = { pre_extra_input = {
"BOLD": {"confounds": element_data["BOLD"]["confounds"]} "BOLD": {"confounds": element_data["BOLD"]["confounds"]}
} }
output, _ = confound_remover.preprocess(pre_input, pre_extra_input) output = confound_remover.preprocess(pre_input, pre_extra_input)
trans_bold = output["data"].get_fdata() trans_bold = output["data"].get_fdata()
# Transformation is in place # Transformation is in place
assert_array_equal( assert_array_equal(

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Sequence
from typing import Any, ClassVar, Optional, Union from typing import Any, ClassVar, Optional, Union
from ...api.decorators import register_preprocessor from ...api.decorators import register_preprocessor
@ -82,6 +83,7 @@ class Smoothing(BasePreprocessor):
"depends_on": FSLSmoothing, "depends_on": FSLSmoothing,
}, },
] ]
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["T1w", "T2w", "BOLD"]
def __init__( def __init__(
self, self,
@ -102,40 +104,11 @@ class Smoothing(BasePreprocessor):
) )
super().__init__(on=on) super().__init__(on=on)
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", "T2w", "BOLD"]
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( def preprocess(
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: ) -> dict[str, Any]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -149,9 +122,6 @@ class Smoothing(BasePreprocessor):
------- -------
dict dict
The computed result as dictionary. The computed result as dictionary.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
""" """
logger.debug("Smoothing") logger.debug("Smoothing")
@ -169,4 +139,4 @@ class Smoothing(BasePreprocessor):
**self.smoothing_params, **self.smoothing_params,
) )
return input, None return input

View file

@ -4,6 +4,9 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Sequence
from typing import ClassVar
import pytest import pytest
from junifer.preprocess.base import BasePreprocessor from junifer.preprocess.base import BasePreprocessor
@ -20,19 +23,15 @@ def test_base_preprocessor_subclassing() -> None:
# Create concrete class # Create concrete class
class MyBasePreprocessor(BasePreprocessor): class MyBasePreprocessor(BasePreprocessor):
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD", "T1w"]
def __init__(self, on): def __init__(self, on):
self.parameter = 1 self.parameter = 1
super().__init__(on=on) super().__init__(on=on)
def get_valid_inputs(self):
return ["BOLD", "T1w"]
def get_output_type(self, input_type):
return input_type
def preprocess(self, input, extra_input=None): def preprocess(self, input, extra_input=None):
input["data"] = f"modified_{input['data']}" input["data"] = f"modified_{input['data']}"
return input, extra_input return input
with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"):
MyBasePreprocessor(on=["BOLD", "T2w"]) MyBasePreprocessor(on=["BOLD", "T2w"])

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from collections.abc import Sequence
from typing import Any, ClassVar, Optional, Union from typing import Any, ClassVar, Optional, Union
from templateflow import api as tflow from templateflow import api as tflow
@ -62,6 +63,17 @@ class SpaceWarper(BasePreprocessor):
"depends_on": [FSLWarper, ANTsWarper], "depends_on": [FSLWarper, ANTsWarper],
}, },
] ]
_VALID_DATA_TYPES: ClassVar[Sequence[str]] = [
"T1w",
"T2w",
"BOLD",
"VBM_GM",
"VBM_WM",
"VBM_CSF",
"fALFF",
"GCOR",
"LCOR",
]
def __init__( def __init__(
self, using: str, reference: str, on: Union[list[str], str] self, using: str, reference: str, on: Union[list[str], str]
@ -94,50 +106,11 @@ class SpaceWarper(BasePreprocessor):
else: else:
raise_error(f"Unknown reference: {self.reference}") 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",
"T2w",
"BOLD",
"VBM_GM",
"VBM_WM",
"VBM_CSF",
"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( # noqa: C901 def preprocess( # noqa: C901
self, self,
input: dict[str, Any], input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None, extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: ) -> dict[str, Any]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -151,9 +124,6 @@ class SpaceWarper(BasePreprocessor):
------- -------
dict dict
The computed result as dictionary. The computed result as dictionary.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
Raises Raises
------ ------
@ -288,4 +258,4 @@ class SpaceWarper(BasePreprocessor):
reference=self.reference, reference=self.reference,
) )
return input, None return input

View file

@ -113,7 +113,7 @@ def test_SpaceWarper_native(
# Read data # Read data
element_data = DefaultDataReader().fit_transform(dg[element]) element_data = DefaultDataReader().fit_transform(dg[element])
# Preprocess data # Preprocess data
output, _ = SpaceWarper( output = SpaceWarper(
using=using, using=using,
reference="T1w", reference="T1w",
on="BOLD", on="BOLD",
@ -179,7 +179,7 @@ def test_SpaceWarper_multi_mni(
element_data = DefaultDataReader().fit_transform(dg[element]) element_data = DefaultDataReader().fit_transform(dg[element])
pre_xfm_data = element_data["T1w"]["data"].get_fdata().copy() pre_xfm_data = element_data["T1w"]["data"].get_fdata().copy()
# Preprocess data # Preprocess data
output, _ = SpaceWarper( output = SpaceWarper(
using="ants", using="ants",
reference=space, reference=space,
on=["T1w"], on=["T1w"],