diff --git a/docs/changes/newsfragments/473.enh b/docs/changes/newsfragments/473.enh new file mode 100644 index 000000000..b9a2706ac --- /dev/null +++ b/docs/changes/newsfragments/473.enh @@ -0,0 +1 @@ +Simplify ``Preprocess`` interface and implementations by `Synchon Mandal`_ diff --git a/docs/extending/preprocessor.rst b/docs/extending/preprocessor.rst index 3fb26828a..7dea78df7 100644 --- a/docs/extending/preprocessor.rst +++ b/docs/extending/preprocessor.rst @@ -14,16 +14,11 @@ new ones, you might need something specific and then you can create your own Preprocessor. 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 ` 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 configured. +#. ``preprocess``: The method that given the data, preprocesses the data. As an example, we will develop a ``NilearnSmoothing`` Preprocessor, which 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: -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`` -:ref:`data types `. +:ref:`data types ` and thus declare them in a class attribute: .. code-block:: python - ... - - - 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 - - - ... + _VALID_DATA_TYPES = ["T1w", "T2w", "BOLD"] .. _extending_preprocessors_init: @@ -77,7 +51,7 @@ you configure it. Our class will have the following arguments: pass the value to it. 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 - ``get_valid_inputs`` function. + ``_VALID_DATA_TYPES`` attribute. .. attention:: @@ -133,17 +107,6 @@ arguments: 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). -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 from typing import Any @@ -158,9 +121,9 @@ and it has two return values: self, input: dict[str, Any], 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) - return input, None + return input ... @@ -187,7 +150,8 @@ decorator and our final code should look like this: .. 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.preprocess import BasePreprocessor @@ -201,6 +165,8 @@ decorator and our final code should look like this: _DEPENDENCIES = {"nilearn"} + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["T1w", "T2w", "BOLD"] + def __init__( self, fwhm: int | float | ArrayLike | Literal["fast"] | None, @@ -209,19 +175,13 @@ decorator and our final code should look like this: self.fwhm = fwhm 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( self, input: dict[str, Any], 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) - return input, None + return input .. _extending_preprocessors_template: @@ -238,18 +198,16 @@ Template for a custom Preprocessor @register_preprocessor class TemplatePreprocessor(BasePreprocessor): + # TODO: add the dependencies + _DEPENDENCIES = {} + + # TODO: add the inputs + _VALID_DATA_TYPES = [] + def __init__(self, on=None): # TODO: add preprocessor-specific parameters 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): # TODO: add the preprocessor logic - return input, None + return input diff --git a/junifer/preprocess/_temporal_filter.py b/junifer/preprocess/_temporal_filter.py index 517bfd6ca..17270c682 100644 --- a/junifer/preprocess/_temporal_filter.py +++ b/junifer/preprocess/_temporal_filter.py @@ -5,6 +5,7 @@ # Synchon Mandal # License: AGPL +from collections.abc import Sequence from typing import ( Any, ClassVar, @@ -56,6 +57,7 @@ class TemporalFilter(BasePreprocessor): """ _DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"} + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD"] def __init__( self, @@ -74,36 +76,7 @@ class TemporalFilter(BasePreprocessor): self.t_r = t_r self.masks = masks - super().__init__(on="BOLD", required_data_types=["BOLD"]) - - 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 + super().__init__() def _validate_data( self, @@ -130,7 +103,7 @@ class TemporalFilter(BasePreprocessor): self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, - ) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: + ) -> dict[str, Any]: """Preprocess. Parameters @@ -145,9 +118,6 @@ class TemporalFilter(BasePreprocessor): dict The computed result as dictionary. If `self.masks` is not None, 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 @@ -237,4 +207,4 @@ class TemporalFilter(BasePreprocessor): } ) - return input, None + return input diff --git a/junifer/preprocess/_temporal_slicer.py b/junifer/preprocess/_temporal_slicer.py index bd82d3640..90dfcc922 100644 --- a/junifer/preprocess/_temporal_slicer.py +++ b/junifer/preprocess/_temporal_slicer.py @@ -3,6 +3,7 @@ # Authors: Synchon Mandal # License: AGPL +from collections.abc import Sequence from typing import Any, ClassVar, Optional import nibabel as nib @@ -45,6 +46,7 @@ class TemporalSlicer(BasePreprocessor): """ _DEPENDENCIES: ClassVar[Dependencies] = {"nilearn"} + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD"] def __init__( self, @@ -61,42 +63,13 @@ class TemporalSlicer(BasePreprocessor): self.stop = stop self.duration = duration self.t_r = t_r - super().__init__(on="BOLD", required_data_types=["BOLD"]) - - 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 + super().__init__() def preprocess( self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, - ) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: + ) -> dict[str, Any]: """Preprocess. Parameters @@ -110,9 +83,6 @@ class TemporalSlicer(BasePreprocessor): ------- dict The computed result as dictionary. - None - Extra "helper" data types as dictionary to add to the Junifer Data - object. Raises ------ @@ -233,4 +203,4 @@ class TemporalSlicer(BasePreprocessor): } ) - return input, None + return input diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index 875590548..f2200b212 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -5,7 +5,8 @@ # License: AGPL 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 ..utils import logger, raise_error @@ -15,15 +16,15 @@ __all__ = ["BasePreprocessor"] 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. Parameters ---------- 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). required_data_types : str or list of str, optional The data types needed for computation. If None, @@ -31,17 +32,27 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): Raises ------ + AttributeError + If the preprocessor does not have `_VALID_DATA_TYPES` attribute. ValueError If required input data type(s) is(are) not found. """ + _VALID_DATA_TYPES: ClassVar[Sequence[str]] + def __init__( self, on: Optional[Union[list[str], str]] = None, required_data_types: Optional[Union[list[str], str]] = None, ) -> None: """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 if on is None: on = self.get_valid_inputs() @@ -58,6 +69,9 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): if required_data_types is None: self._required_data_types = on 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 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] - @abstractmethod def get_valid_inputs(self) -> list[str]: """Get valid data types for input. @@ -100,12 +113,8 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): preprocessor. """ - raise_error( - msg="Concrete classes need to implement get_valid_inputs().", - klass=NotImplementedError, - ) + return list(self._VALID_DATA_TYPES) - @abstractmethod def get_output_type(self, input_type: str) -> str: """Get output type. @@ -120,17 +129,15 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): The data type output by the preprocessor. """ - raise_error( - msg="Concrete classes need to implement get_output_type().", - klass=NotImplementedError, - ) + # Does not add any new keys + return input_type @abstractmethod def preprocess( self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, - ) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: + ) -> dict[str, Any]: """Preprocess. Parameters @@ -147,10 +154,6 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): ------- dict 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( @@ -192,19 +195,10 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): f"Extra data type for preprocess: {extra_input.keys()}" ) # Preprocess data - t_out, t_extra_input = self.preprocess( - input=t_input, extra_input=extra_input - ) + t_out = self.preprocess(input=t_input, extra_input=extra_input) # Set output to the Junifer Data object logger.debug(f"Adding {type_} to output") 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 self.update_meta(out[type_], "preprocess") return out diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index 9bdf66533..753c21a38 100644 --- a/junifer/preprocess/confounds/fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/fmriprep_confound_remover.py @@ -5,6 +5,7 @@ # Synchon Mandal # License: AGPL +from collections.abc import Sequence from typing import ( Any, ClassVar, @@ -175,6 +176,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor): """ _DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"} + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD"] def __init__( self, @@ -251,36 +253,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor): "include it in the future", klass=ValueError, ) - super().__init__(on="BOLD", required_data_types=["BOLD"]) - - 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 + super().__init__() def _map_adhoc_to_fmriprep(self, input: dict[str, Any]) -> None: """Map the adhoc format to the fmpriprep format spec. @@ -621,7 +594,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor): self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, - ) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: + ) -> dict[str, Any]: """Preprocess. Parameters @@ -636,9 +609,6 @@ class fMRIPrepConfoundRemover(BasePreprocessor): dict The computed result as dictionary. If `self.masks` is not None, 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 @@ -753,4 +723,4 @@ class fMRIPrepConfoundRemover(BasePreprocessor): } ) - return input, None + return input diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index 2ed820fb1..ab026701e 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -463,7 +463,7 @@ def test_fMRIPrepConfoundRemover_preprocess() -> None: pre_extra_input = { "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() # Transformation is in place assert_array_equal( @@ -614,7 +614,7 @@ def test_fMRIPrepConfoundRemover_scrubbing() -> None: pre_extra_input = { "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() # Transformation is in place assert_array_equal( diff --git a/junifer/preprocess/smoothing/smoothing.py b/junifer/preprocess/smoothing/smoothing.py index 810dcbb25..99160b5c2 100644 --- a/junifer/preprocess/smoothing/smoothing.py +++ b/junifer/preprocess/smoothing/smoothing.py @@ -3,6 +3,7 @@ # Authors: Synchon Mandal # License: AGPL +from collections.abc import Sequence from typing import Any, ClassVar, Optional, Union from ...api.decorators import register_preprocessor @@ -82,6 +83,7 @@ class Smoothing(BasePreprocessor): "depends_on": FSLSmoothing, }, ] + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["T1w", "T2w", "BOLD"] def __init__( self, @@ -102,40 +104,11 @@ class Smoothing(BasePreprocessor): ) 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( self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, - ) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: + ) -> dict[str, Any]: """Preprocess. Parameters @@ -149,9 +122,6 @@ class Smoothing(BasePreprocessor): ------- dict The computed result as dictionary. - None - Extra "helper" data types as dictionary to add to the Junifer Data - object. """ logger.debug("Smoothing") @@ -169,4 +139,4 @@ class Smoothing(BasePreprocessor): **self.smoothing_params, ) - return input, None + return input diff --git a/junifer/preprocess/tests/test_preprocess_base.py b/junifer/preprocess/tests/test_preprocess_base.py index c0481d952..24322ffc6 100644 --- a/junifer/preprocess/tests/test_preprocess_base.py +++ b/junifer/preprocess/tests/test_preprocess_base.py @@ -4,6 +4,9 @@ # Synchon Mandal # License: AGPL +from collections.abc import Sequence +from typing import ClassVar + import pytest from junifer.preprocess.base import BasePreprocessor @@ -20,19 +23,15 @@ def test_base_preprocessor_subclassing() -> None: # Create concrete class class MyBasePreprocessor(BasePreprocessor): + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = ["BOLD", "T1w"] + def __init__(self, on): self.parameter = 1 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): input["data"] = f"modified_{input['data']}" - return input, extra_input + return input with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): MyBasePreprocessor(on=["BOLD", "T2w"]) diff --git a/junifer/preprocess/warping/space_warper.py b/junifer/preprocess/warping/space_warper.py index a85e6955b..fd7ebc848 100644 --- a/junifer/preprocess/warping/space_warper.py +++ b/junifer/preprocess/warping/space_warper.py @@ -3,6 +3,7 @@ # Authors: Synchon Mandal # License: AGPL +from collections.abc import Sequence from typing import Any, ClassVar, Optional, Union from templateflow import api as tflow @@ -62,6 +63,17 @@ class SpaceWarper(BasePreprocessor): "depends_on": [FSLWarper, ANTsWarper], }, ] + _VALID_DATA_TYPES: ClassVar[Sequence[str]] = [ + "T1w", + "T2w", + "BOLD", + "VBM_GM", + "VBM_WM", + "VBM_CSF", + "fALFF", + "GCOR", + "LCOR", + ] def __init__( self, using: str, reference: str, on: Union[list[str], str] @@ -94,50 +106,11 @@ class SpaceWarper(BasePreprocessor): 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", - "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 self, input: dict[str, Any], extra_input: Optional[dict[str, Any]] = None, - ) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]: + ) -> dict[str, Any]: """Preprocess. Parameters @@ -151,9 +124,6 @@ class SpaceWarper(BasePreprocessor): ------- dict The computed result as dictionary. - None - Extra "helper" data types as dictionary to add to the Junifer Data - object. Raises ------ @@ -288,4 +258,4 @@ class SpaceWarper(BasePreprocessor): reference=self.reference, ) - return input, None + return input diff --git a/junifer/preprocess/warping/tests/test_space_warper.py b/junifer/preprocess/warping/tests/test_space_warper.py index e79b32a01..d301a091a 100644 --- a/junifer/preprocess/warping/tests/test_space_warper.py +++ b/junifer/preprocess/warping/tests/test_space_warper.py @@ -113,7 +113,7 @@ def test_SpaceWarper_native( # Read data element_data = DefaultDataReader().fit_transform(dg[element]) # Preprocess data - output, _ = SpaceWarper( + output = SpaceWarper( using=using, reference="T1w", on="BOLD", @@ -179,7 +179,7 @@ def test_SpaceWarper_multi_mni( element_data = DefaultDataReader().fit_transform(dg[element]) pre_xfm_data = element_data["T1w"]["data"].get_fdata().copy() # Preprocess data - output, _ = SpaceWarper( + output = SpaceWarper( using="ants", reference=space, on=["T1w"],