[ENH]: Improve BasePreprocessor and fMRIPrepConfoundRemover #260

Merged
synchon merged 10 commits from update/preprocessor-base into main 2023-10-18 09:37:06 +00:00
4 changed files with 168 additions and 118 deletions

View file

@ -0,0 +1 @@
Improve :class:`.BasePreprocessor` for easy subclassing and adapt :class:`.fMRIPrepConfoundRemover` to it by `Synchon Mandal`_

View file

@ -19,22 +19,40 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
on : str or list of str, optional
The kind of data to apply the preprocessor to. If None,
will work on all available data (default None).
required_data_types : str or list of str, optional
The kind of data types needed for computation. If None,
will be equal to ``on`` (default None).
Raises
------
ValueError
If required input data type(s) is(are) not found.
"""
def __init__(
self,
on: Optional[Union[List[str], str]] = None,
required_data_types: Optional[Union[List[str], str]] = None,
) -> None:
"""Initialize the class."""
# Use all data types if not provided
if on is None:
on = self.get_valid_inputs()
# Convert data types to list
if not isinstance(on, list):
on = [on]
# Check if required inputs are found
if any(x not in self.get_valid_inputs() for x in on):
name = self.__class__.__name__
wrong_on = [x for x in on if x not in self.get_valid_inputs()]
raise ValueError(f"{name} cannot be computed on {wrong_on}")
raise_error(f"{name} cannot be computed on {wrong_on}")
self._on = on
# Set required data types for validation
if required_data_types is None:
self._required_data_types = on
else:
self._required_data_types = required_data_types
def validate_input(self, input: List[str]) -> List[str]:
"""Validate input.
@ -55,15 +73,32 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
------
ValueError
If the input does not have the required data.
"""
if not any(x in input for x in self._on):
if any(x not in input for x in self._required_data_types):
raise_error(
"Input does not have the required data."
f"\t Input: {input}"
f"\t Required (any of): {self._on}"
f"\t Required (all of): {self._required_data_types}"
)
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.
Returns
-------
list of str
The list of data types that can be used as input for this
preprocessor.
"""
raise_error(
msg="Concrete classes need to implement get_valid_inputs().",
klass=NotImplementedError,
)
@abstractmethod
def get_output_type(self, input: List[str]) -> List[str]:
"""Get output type.
@ -87,17 +122,34 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
)
@abstractmethod
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
def preprocess(
self,
input: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""Preprocess.
Parameters
----------
input : dict
A single input from the Junifer Data object to preprocess.
extra_input : dict, optional
The other fields in the Junifer Data object. Useful for accessing
other data kind that needs to be used in the computation. For
example, the confound removers can make use of the
confounds if available (default None).
Returns
-------
list of str
The list of data types that can be used as input for this
preprocessor.
str
The key to store the output in the Junifer Data object.
dict
The computed result as dictionary. This will be stored in the
Junifer Data object under the key ``data`` of the data type.
"""
raise_error(
msg="Concrete classes need to implement get_valid_inputs().",
msg="Concrete classes need to implement preprocess().",
klass=NotImplementedError,
)
@ -146,35 +198,3 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
self.update_meta(out[key], "preprocess")
return out
@abstractmethod
def preprocess(
self,
input: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""Preprocess.
Parameters
----------
input : dict
A single input from the Junifer Data object to preprocess.
extra_input : dict, optional
The other fields in the Junifer Data object. Useful for accessing
other data kind that needs to be used in the computation. For
example, the confound removers can make use of the
confounds if available (default None).
Returns
-------
key : str
The key to store the output in the Junifer Data object.
object : dict
The computed result as dictionary. This will be stored in the
Junifer Data object under the key 'key'.
"""
raise_error(
msg="Concrete classes need to implement preprocess().",
klass=NotImplementedError,
)

View file

@ -164,7 +164,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
t_r: Optional[float] = None,
masks: Union[str, Dict, List[Union[Dict, str]], None] = None,
) -> None:
"""Initialise the class."""
"""Initialize the class."""
if strategy is None:
strategy = {
"motion": "full",
@ -208,48 +208,30 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
"include it in the future",
klass=ValueError,
)
super().__init__()
super().__init__(
on="BOLD", required_data_types=["BOLD", "BOLD_confounds"]
)
def validate_input(self, input: List[str]) -> List[str]:
"""Validate the input to the pipeline step.
Parameters
----------
input : list of str
The input to the pipeline step. The list must contain the
available Junifer Data object keys.
def get_valid_inputs(self) -> List[str]:
"""Get valid data types for input.
Returns
-------
list of str
The actual elements of the input that will be processed by this
pipeline step.
Raises
------
ValueError
If the input does not have the required data.
The list of data types that can be used as input for this
preprocessor.
"""
_required_inputs = ["BOLD", "BOLD_confounds"]
if any(x not in input for x in _required_inputs):
raise_error(
msg="Input does not have the required data. \n"
f"Input: {input} \n"
f"Required (all off): {_required_inputs} \n",
klass=ValueError,
)
return [x for x in self._on if x in input]
return ["BOLD"]
def get_output_type(self, input: List[str]) -> List[str]:
"""Get the kind of the pipeline step.
"""Get output type.
Parameters
----------
input : list of str
The input to the pipeline step. The list must contain the
available Junifer Data object keys.
The input to the preprocessor. The list must contain the
available Junifer Data dictionary keys.
Returns
-------
@ -261,17 +243,6 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
# Does not add any new keys
return input
def get_valid_inputs(self) -> List[str]:
"""Get the valid inputs for the pipeline step.
Returns
-------
list of str
The valid inputs for the pipeline step.
"""
return ["BOLD"]
def _map_adhoc_to_fmriprep(self, input: Dict[str, Any]) -> None:
"""Map the adhoc format to the fmpriprep format spec.
@ -333,6 +304,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
spike_name : str
Name of the confound to use for spike detection
Raises
------
ValueError
If invalid confounds file is found.
"""
confounds_df = input["data"]
available_vars = confounds_df.columns
@ -347,7 +323,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
if any(x not in available_vars for x in t_basics):
missing = [x for x in t_basics if x not in available_vars]
raise ValueError(
raise_error(
"Invalid confounds file. Missing basic confounds: "
f"{missing}. "
"Check if this file is really an fmriprep confounds file. "
@ -377,7 +353,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
spike_name = "framewise_displacement"
if self.spike is not None:
if spike_name not in available_vars:
raise ValueError(
raise_error(
"Invalid confounds file. Missing framewise_displacement "
"(spike) confound. "
"Check if this file is really an fmriprep confounds file. "
@ -460,17 +436,32 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
Dictionary containing the rest of the Junifer Data object. Must
include the ``BOLD_confounds`` key.
Raises
------
ValueError
If ``extra_input`` is None or
if ``"BOLD_confounds"`` is not found in ``extra_input`` or
if ``"data"`` key is not found in ``"BOLD_confounds"`` or
if ``"data"`` is not pandas.DataFrame or
if image time series and confounds have different lengths or
if ``"format"`` is not found in ``"BOLD_confounds"`` or
if ``format = "adhoc"`` and ``"mappings"`` key or ``"fmriprep"``
key or correct fMRIPrep mappings or required fMRIPrep mappings are
not found or if invalid confounds format is found.
"""
# Bold must be 4D niimg
check_niimg_4d(input["data"])
if extra_input is None:
raise_error("No extra input provided", ValueError)
raise_error(msg="No extra input provided", klass=ValueError)
if "BOLD_confounds" not in extra_input:
raise_error("No BOLD_confounds provided", ValueError)
raise_error(msg="No BOLD_confounds provided", klass=ValueError)
if "data" not in extra_input["BOLD_confounds"]:
raise_error("No BOLD_confounds data provided", ValueError)
raise_error(
msg="No BOLD_confounds data provided", klass=ValueError
)
# Confounds must be a dataframe
if not isinstance(extra_input["BOLD_confounds"]["data"], pd.DataFrame):
raise_error(
@ -528,14 +519,14 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
]
if len(missing) > 0:
raise ValueError(
raise_error(
"Invalid confounds file. Missing columns: "
f"{missing}. "
"Check if this file matches the adhoc specification for "
"this dataset."
)
elif t_format != "fmriprep":
raise ValueError(f"Invalid confounds format {t_format}")
raise_error(f"Invalid confounds format {t_format}")
def _remove_confounds(
self,
@ -614,18 +605,18 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
Parameters
----------
input : dict
A single input from the Junifer Data object in which to preprocess.
A single input from the Junifer Data object to preprocess.
extra_input : dict, optional
The other fields in the Junifer Data object. Must include the
``BOLD_confounds`` key.
Returns
-------
key : str
str
The key to store the output in the Junifer Data object.
object : dict
dict
The computed result as dictionary. This will be stored in the
Junifer Data object under the key ``key``.
Junifer Data object under the key ``data`` of the data type.
"""
self._validate_data(input, extra_input)

View file

@ -5,7 +5,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
import typing
from typing import List, cast
import nibabel as nib
import numpy as np
@ -40,39 +40,77 @@ def test_fMRIPrepConfoundRemover_init() -> None:
fMRIPrepConfoundRemover(strategy={"motion": "wrong"})
def test_fMRIPrepConfoundRemover_validate_input() -> None:
"""Test fMRIPrepConfoundRemover validate_input."""
@pytest.mark.parametrize(
"input_",
[
["T1w"],
["BOLD"],
["T1w", "BOLD"],
],
)
def test_fMRIPrepConfoundRemover_validate_input_errors(
input_: List[str],
) -> None:
"""Test errors for fMRIPrepConfoundRemover validate_input.
Parameters
----------
input_ : list of str
The input data types.
"""
confound_remover = fMRIPrepConfoundRemover()
# Input is valid when both BOLD and BOLD_confounds are present
input = ["T1w"]
with pytest.raises(ValueError, match="not have the required data"):
confound_remover.validate_input(input)
input = ["BOLD"]
with pytest.raises(ValueError, match="not have the required data"):
confound_remover.validate_input(input)
input = ["BOLD", "T1w"]
with pytest.raises(ValueError, match="not have the required data"):
confound_remover.validate_input(input)
input = ["BOLD", "T1w", "BOLD_confounds"]
confound_remover.validate_input(input)
confound_remover.validate_input(input_)
def test_fMRIPrepConfoundRemover_get_output_type() -> None:
"""Test fMRIPrepConfoundRemover validate_input."""
@pytest.mark.parametrize(
"input_",
[
["BOLD", "BOLD_confounds"],
["T1w", "BOLD", "BOLD_confounds"],
],
)
def test_fMRIPrepConfoundRemover_validate_input(input_: List[str]) -> None:
"""Test fMRIPrepConfoundRemover validate_input.
Parameters
----------
input_ : list of str
The input data types.
"""
confound_remover = fMRIPrepConfoundRemover()
inputs = [
confound_remover.validate_input(input_)
def test_fMRIPrepConfoundRemover_get_valid_inputs() -> None:
"""Test fMRIPrepConfoundRemover get_valid_inputs."""
confound_remover = fMRIPrepConfoundRemover()
assert confound_remover.get_valid_inputs() == ["BOLD"]
@pytest.mark.parametrize(
"input_",
[
["BOLD", "T1w", "BOLD_confounds"],
["BOLD", "VBM_GM", "BOLD_confounds"],
["BOLD", "BOLD_confounds"],
]
],
)
def test_fMRIPrepConfoundRemover_get_output_type(input_: List[str]) -> None:
"""Test fMRIPrepConfoundRemover get_output_type.
Parameters
----------
input_ : list of str
The input data types.
"""
confound_remover = fMRIPrepConfoundRemover()
# Confound remover works in place
for input in inputs:
assert confound_remover.get_output_type(input) == input
assert confound_remover.get_output_type(input_) == input_
def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None:
@ -259,7 +297,7 @@ def test_fMRIPrepConfoundRemover__pick_confounds_adhoc() -> None:
assert set(out.columns) == set(fmriprep_all_vars)
def test_FMRIPRepConfoundRemover__pick_confounds_fmriprep() -> None:
def test_fMRIPRepConfoundRemover__pick_confounds_fmriprep() -> None:
"""Test fMRIPrepConfoundRemover pick confounds on fmriprep confounds."""
confound_remover = fMRIPrepConfoundRemover(
strategy={"wm_csf": "full"}, spike=0.2
@ -292,7 +330,7 @@ def test_FMRIPRepConfoundRemover__pick_confounds_fmriprep() -> None:
assert_frame_equal(out1, out2)
def test_FMRIPRepConfoundRemover__pick_confounds_fmriprep_compute() -> None:
def test_fMRIPRepConfoundRemover__pick_confounds_fmriprep_compute() -> None:
"""Test if fmriprep returns the same derivatives/power2 as we compute."""
confound_remover = fMRIPrepConfoundRemover(strategy={"wm_csf": "full"})
@ -457,7 +495,7 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None:
clean_bold = confound_remover._remove_confounds(
input=input["BOLD"], extra_input=extra_input
)
clean_bold = typing.cast(nib.Nifti1Image, clean_bold)
clean_bold = cast(nib.Nifti1Image, clean_bold)
# TODO: Find a better way to test functionality here
assert (
clean_bold.header.get_zooms() # type: ignore