[ENH]: Simplify Preprocess interface #473
11 changed files with 90 additions and 288 deletions
1
docs/changes/newsfragments/473.enh
Normal file
1
docs/changes/newsfragments/473.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Simplify ``Preprocess`` interface and implementations by `Synchon Mandal`_
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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(
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"])
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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"],
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue