[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.
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
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 <data_types>`.
:ref:`data types <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

View file

@ -5,6 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -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

View file

@ -5,6 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -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(

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -4,6 +4,9 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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"])

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# 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

View file

@ -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"],