[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.
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
Loading…
Reference in a new issue