[ENH]: Improve BasePreprocessor and fMRIPrepConfoundRemover #260
4 changed files with 168 additions and 118 deletions
1
docs/changes/newsfragments/260.enh
Normal file
1
docs/changes/newsfragments/260.enh
Normal file
|
|
@ -0,0 +1 @@
|
|||
Improve :class:`.BasePreprocessor` for easy subclassing and adapt :class:`.fMRIPrepConfoundRemover` to it by `Synchon Mandal`_
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in a new issue