diff --git a/docs/changes/newsfragments/260.enh b/docs/changes/newsfragments/260.enh new file mode 100644 index 000000000..fa0920093 --- /dev/null +++ b/docs/changes/newsfragments/260.enh @@ -0,0 +1 @@ +Improve :class:`.BasePreprocessor` for easy subclassing and adapt :class:`.fMRIPrepConfoundRemover` to it by `Synchon Mandal`_ diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index 5c8900d97..489988739 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -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, - ) diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index dc14e2ce8..4b2603589 100644 --- a/junifer/preprocess/confounds/fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/fmriprep_confound_remover.py @@ -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) diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index b03f009ae..d197dafd7 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -5,7 +5,7 @@ # Synchon Mandal # 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