From 686e0f86ab2df626c5b619bd52f2c1c204246bee Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 7 Mar 2024 12:14:43 +0100 Subject: [PATCH 1/8] chore: improve docstrings and type annotation for preprocessors --- junifer/preprocess/base.py | 25 ++++++++++--------- junifer/preprocess/bold_warper.py | 14 +++++------ .../confounds/fmriprep_confound_remover.py | 14 +++++------ .../tests/test_fmriprep_confound_remover.py | 21 +++------------- junifer/preprocess/tests/test_bold_warper.py | 23 +++-------------- 5 files changed, 32 insertions(+), 65 deletions(-) diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index 489988739..3634b18dd 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -12,15 +12,18 @@ from ..utils import logger, raise_error class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): - """Provide abstract base class for all preprocessors. + """Abstract base class for all preprocessors. + + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. Parameters ---------- - 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). + on : str or list of str or None, optional + The data type 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 kind of data types needed for computation. If None, + The data types needed for computation. If None, will be equal to ``on`` (default None). Raises @@ -100,20 +103,18 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): ) @abstractmethod - def get_output_type(self, input: List[str]) -> List[str]: + def get_output_type(self, input_type: str) -> str: """Get output type. Parameters ---------- - input : list of str - The input to the preprocessor. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the preprocessor. Returns ------- - list of str - The updated list of available Junifer Data object keys after - the pipeline step. + str + The data type output by the preprocessor. """ raise_error( diff --git a/junifer/preprocess/bold_warper.py b/junifer/preprocess/bold_warper.py index ddeeed93e..83768f3d6 100644 --- a/junifer/preprocess/bold_warper.py +++ b/junifer/preprocess/bold_warper.py @@ -82,24 +82,22 @@ class BOLDWarper(BasePreprocessor): """ return ["BOLD"] - def get_output_type(self, input: List[str]) -> List[str]: + def get_output_type(self, input_type: str) -> str: """Get output type. Parameters ---------- - input : list of str - The input to the preprocessor. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The data type input to the preprocessor. Returns ------- - list of str - The updated list of available Junifer Data object keys after - the pipeline step. + str + The data type output by the preprocessor. """ # Does not add any new keys - return input + return input_type def preprocess( self, diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index 30fb51f41..45fd64eff 100644 --- a/junifer/preprocess/confounds/fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/fmriprep_confound_remover.py @@ -224,24 +224,22 @@ class fMRIPrepConfoundRemover(BasePreprocessor): """ return ["BOLD"] - def get_output_type(self, input: List[str]) -> List[str]: + def get_output_type(self, input_type: str) -> str: """Get output type. Parameters ---------- - input : list of str - The input to the preprocessor. The list must contain the - available Junifer Data dictionary keys. + input_type : str + The input to the preprocessor. Returns ------- - list of str - The updated list of available Junifer Data object keys after - the pipeline step. + str + The data type output by the preprocessor. """ # Does not add any new keys - return input + return input_type def _map_adhoc_to_fmriprep(self, input: Dict[str, Any]) -> None: """Map the adhoc format to the fmpriprep format spec. diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index d197dafd7..1a10ca64f 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -91,26 +91,11 @@ def test_fMRIPrepConfoundRemover_get_valid_inputs() -> None: 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. - - """ +def test_fMRIPrepConfoundRemover_get_output_type() -> None: + """Test fMRIPrepConfoundRemover get_output_type.""" confound_remover = fMRIPrepConfoundRemover() # Confound remover works in place - assert confound_remover.get_output_type(input_) == input_ + assert confound_remover.get_output_type("BOLD") == "BOLD" def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None: diff --git a/junifer/preprocess/tests/test_bold_warper.py b/junifer/preprocess/tests/test_bold_warper.py index 58a26db14..a24548c0e 100644 --- a/junifer/preprocess/tests/test_bold_warper.py +++ b/junifer/preprocess/tests/test_bold_warper.py @@ -4,7 +4,7 @@ # License: AGPL import socket -from typing import TYPE_CHECKING, List, Tuple +from typing import TYPE_CHECKING, Tuple import pytest from numpy.testing import assert_array_equal, assert_raises @@ -31,25 +31,10 @@ def test_BOLDWarper_get_valid_inputs() -> None: assert bold_warper.get_valid_inputs() == ["BOLD"] -@pytest.mark.parametrize( - "input_", - [ - ["BOLD", "T1w", "Warp"], - ["BOLD", "T1w"], - ["BOLD"], - ], -) -def test_BOLDWarper_get_output_type(input_: List[str]) -> None: - """Test BOLDWarper get_output_type. - - Parameters - ---------- - input_ : list of str - The input data types. - - """ +def test_BOLDWarper_get_output_type() -> None: + """Test BOLDWarper get_output_type.""" bold_warper = BOLDWarper(reference="T1w") - assert bold_warper.get_output_type(input_) == input_ + assert bold_warper.get_output_type("BOLD") == "BOLD" @pytest.mark.parametrize( -- 2.52.0 From ceb5c23705d9249c3978e690135ac9114d0975b3 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 7 Mar 2024 12:16:59 +0100 Subject: [PATCH 2/8] refactor: disable return of data type from BasePreprocessor.preprocess() --- junifer/preprocess/base.py | 39 ++++++++----------- .../preprocess/tests/test_preprocess_base.py | 2 +- 2 files changed, 17 insertions(+), 24 deletions(-) diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index 3634b18dd..e36c1a89a 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -5,7 +5,7 @@ # License: AGPL from abc import ABC, abstractmethod -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any, Dict, List, Optional, Union from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..utils import logger, raise_error @@ -127,7 +127,7 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): self, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + ) -> Dict[str, Any]: """Preprocess. Parameters @@ -142,11 +142,8 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): Returns ------- - 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. + The computed result as dictionary. """ raise_error( @@ -172,30 +169,26 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): """ out = input + # For each data type, run preprocessing for type_ in self._on: + # Check if data type is available if type_ in input.keys(): logger.info(f"Preprocessing {type_}") + # Get data dict for data type t_input = input[type_] - # Pass the other data types as extra input, removing - # the current type + # the current type; input is not copied so as to allow + # propagation of extra types like masks extra_input = input extra_input.pop(type_) logger.debug( - f"Extra input for preprocess: {extra_input.keys()}" + f"Extra data type for preprocess: {extra_input.keys()}" ) - key, t_out = self.preprocess( - input=t_input, extra_input=extra_input - ) - - # Add the output to the Junifer Data object - logger.debug(f"Adding {key} to output") - out[key] = t_out - - # In case we are creating a new type, re-add the original input - if key != type_: - logger.debug("Adding original input back to output") - out[type_] = t_input - - self.update_meta(out[key], "preprocess") + # Preprocess data + 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 + # Update metadata for step + self.update_meta(out[type_], "preprocess") return out diff --git a/junifer/preprocess/tests/test_preprocess_base.py b/junifer/preprocess/tests/test_preprocess_base.py index 199679517..7ce47880b 100644 --- a/junifer/preprocess/tests/test_preprocess_base.py +++ b/junifer/preprocess/tests/test_preprocess_base.py @@ -32,7 +32,7 @@ def test_base_preprocessor_subclassing() -> None: def preprocess(self, input, extra_input): input["data"] = f"mofidied_{input['data']}" - return "BOLD", input + return input with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): MyBasePreprocessor(on=["BOLD", "T2w"]) -- 2.52.0 From ccc5d76d75bd3d6bf671e6f7f91154a2723b837a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 7 Mar 2024 12:17:58 +0100 Subject: [PATCH 3/8] refactor: disable return of data type from BOLDWarper.preprocess() --- junifer/preprocess/bold_warper.py | 10 +++------- junifer/preprocess/tests/test_bold_warper.py | 6 ++---- 2 files changed, 5 insertions(+), 11 deletions(-) diff --git a/junifer/preprocess/bold_warper.py b/junifer/preprocess/bold_warper.py index 83768f3d6..6aa8c0dc7 100644 --- a/junifer/preprocess/bold_warper.py +++ b/junifer/preprocess/bold_warper.py @@ -9,7 +9,6 @@ from typing import ( Dict, List, Optional, - Tuple, Union, ) @@ -103,7 +102,7 @@ class BOLDWarper(BasePreprocessor): self, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + ) -> Dict[str, Any]: """Preprocess. Parameters @@ -117,11 +116,8 @@ class BOLDWarper(BasePreprocessor): Returns ------- - 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. + The computed result as dictionary. Raises ------ @@ -239,4 +235,4 @@ class BOLDWarper(BasePreprocessor): input["data"] = nib.load(warped_bold_path) input["space"] = self.ref - return "BOLD", input + return input diff --git a/junifer/preprocess/tests/test_bold_warper.py b/junifer/preprocess/tests/test_bold_warper.py index a24548c0e..d9c4028a1 100644 --- a/junifer/preprocess/tests/test_bold_warper.py +++ b/junifer/preprocess/tests/test_bold_warper.py @@ -86,11 +86,10 @@ def test_BOLDWarper_preprocess_to_native( # Read data element_data = DefaultDataReader().fit_transform(dg[element]) # Preprocess data - data_type, data = BOLDWarper(reference="T1w").preprocess( + data = BOLDWarper(reference="T1w").preprocess( input=element_data["BOLD"], extra_input=element_data, ) - assert data_type == "BOLD" assert isinstance(data, dict) @@ -150,11 +149,10 @@ def test_BOLDWarper_preprocess_to_multi_mni( element_data = DefaultDataReader().fit_transform(dg[element]) pre_xfm_data = element_data["BOLD"]["data"].get_fdata().copy() # Preprocess data - data_type, data = BOLDWarper(reference=space).preprocess( + data = BOLDWarper(reference=space).preprocess( input=element_data["BOLD"], extra_input=element_data, ) - assert data_type == "BOLD" assert isinstance(data, dict) assert data["space"] == space with assert_raises(AssertionError): -- 2.52.0 From 7f288de006ac4c7742dd02f0a8c9857c183cffb3 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 7 Mar 2024 12:18:44 +0100 Subject: [PATCH 4/8] refactor: disable return of data type from fMRIPrepConfoundRemover.preprocess() --- .../confounds/fmriprep_confound_remover.py | 9 ++-- .../tests/test_fmriprep_confound_remover.py | 47 ++++++++----------- 2 files changed, 23 insertions(+), 33 deletions(-) diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index 45fd64eff..474500543 100644 --- a/junifer/preprocess/confounds/fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/fmriprep_confound_remover.py @@ -600,7 +600,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor): self, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + ) -> Dict[str, Any]: """Preprocess. Parameters @@ -613,13 +613,10 @@ class fMRIPrepConfoundRemover(BasePreprocessor): Returns ------- - 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. + The computed result as dictionary. """ self._validate_data(input, extra_input) input["data"] = self._remove_confounds(input, extra_input=extra_input) - return "BOLD", input + return input diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index 1a10ca64f..6dcd4efd2 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -492,22 +492,20 @@ def test_fMRIPrepConfoundRemover__remove_confounds() -> None: def test_fMRIPrepConfoundRemover_preprocess() -> None: """Test fMRIPrepConfoundRemover with all confounds present.""" - - # need reader for the data - reader = DefaultDataReader() # All strategies full, no spike confound_remover = fMRIPrepConfoundRemover() with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: - input = dg["sub-01"] - input = reader.fit_transform(input) - orig_bold = input["BOLD"]["data"].get_fdata().copy() - pre_input = input["BOLD"] - pre_extra_input = {"BOLD_confounds": input["BOLD_confounds"]} - key, output = confound_remover.preprocess(pre_input, pre_extra_input) + element_data = DefaultDataReader().fit_transform(dg["sub-01"]) + orig_bold = element_data["BOLD"]["data"].get_fdata().copy() + pre_input = element_data["BOLD"] + pre_extra_input = {"BOLD_confounds": element_data["BOLD_confounds"]} + output = confound_remover.preprocess(pre_input, pre_extra_input) trans_bold = output["data"].get_fdata() # Transformation is in place - assert_array_equal(trans_bold, input["BOLD"]["data"].get_fdata()) + assert_array_equal( + trans_bold, element_data["BOLD"]["data"].get_fdata() + ) # Data should have the same shape assert orig_bold.shape == trans_bold.shape @@ -516,25 +514,22 @@ def test_fMRIPrepConfoundRemover_preprocess() -> None: assert_raises( AssertionError, assert_array_equal, orig_bold, trans_bold ) - assert key == "BOLD" def test_fMRIPrepConfoundRemover_fit_transform() -> None: """Test fMRIPrepConfoundRemover with all confounds present.""" - - # need reader for the data - reader = DefaultDataReader() # All strategies full, no spike confound_remover = fMRIPrepConfoundRemover() with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: - input = dg["sub-01"] - input = reader.fit_transform(input) - orig_bold = input["BOLD"]["data"].get_fdata().copy() - output = confound_remover.fit_transform(input) + element_data = DefaultDataReader().fit_transform(dg["sub-01"]) + orig_bold = element_data["BOLD"]["data"].get_fdata().copy() + output = confound_remover.fit_transform(element_data) trans_bold = output["BOLD"]["data"].get_fdata() # Transformation is in place - assert_array_equal(trans_bold, input["BOLD"]["data"].get_fdata()) + assert_array_equal( + trans_bold, element_data["BOLD"]["data"].get_fdata() + ) # Data should have the same shape assert orig_bold.shape == trans_bold.shape @@ -565,22 +560,20 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None: def test_fMRIPrepConfoundRemover_fit_transform_masks() -> None: """Test fMRIPrepConfoundRemover with all confounds present.""" - - # need reader for the data - reader = DefaultDataReader() # All strategies full, no spike confound_remover = fMRIPrepConfoundRemover( masks={"compute_brain_mask": {"threshold": 0.2}} ) with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: - input = dg["sub-01"] - input = reader.fit_transform(input) - orig_bold = input["BOLD"]["data"].get_fdata().copy() - output = confound_remover.fit_transform(input) + element_data = DefaultDataReader().fit_transform(dg["sub-01"]) + orig_bold = element_data["BOLD"]["data"].get_fdata().copy() + output = confound_remover.fit_transform(element_data) trans_bold = output["BOLD"]["data"].get_fdata() # Transformation is in place - assert_array_equal(trans_bold, input["BOLD"]["data"].get_fdata()) + assert_array_equal( + trans_bold, element_data["BOLD"]["data"].get_fdata() + ) # Data should have the same shape assert orig_bold.shape == trans_bold.shape -- 2.52.0 From f2b433e5303f3f23f496a6e421d5aea409100802 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Mar 2024 11:33:03 +0100 Subject: [PATCH 5/8] refactor: enable return of helper data type from BasePreprocessor.preprocess() --- junifer/preprocess/base.py | 27 ++++++++++++++----- .../preprocess/tests/test_preprocess_base.py | 12 ++++----- 2 files changed, 26 insertions(+), 13 deletions(-) diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index e36c1a89a..d29b5dc40 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -5,7 +5,7 @@ # License: AGPL from abc import ABC, abstractmethod -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Optional, Tuple, Union from ..pipeline import PipelineStepMixin, UpdateMetaMixin from ..utils import logger, raise_error @@ -127,7 +127,7 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): self, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: """Preprocess. Parameters @@ -144,6 +144,10 @@ 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. For example, computed BOLD mask can be passed via this. + If no new "helper" data types is created, None is to be passed. """ raise_error( @@ -168,7 +172,8 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): The processed output as a dictionary. """ - out = input + # Copy input to not modify the original + out = input.copy() # For each data type, run preprocessing for type_ in self._on: # Check if data type is available @@ -177,18 +182,26 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): # Get data dict for data type t_input = input[type_] # Pass the other data types as extra input, removing - # the current type; input is not copied so as to allow - # propagation of extra types like masks - extra_input = input + # the current type + extra_input = input.copy() extra_input.pop(type_) logger.debug( f"Extra data type for preprocess: {extra_input.keys()}" ) # Preprocess data - t_out = self.preprocess(input=t_input, extra_input=extra_input) + t_out, t_extra_input = 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 diff --git a/junifer/preprocess/tests/test_preprocess_base.py b/junifer/preprocess/tests/test_preprocess_base.py index 7ce47880b..c0481d952 100644 --- a/junifer/preprocess/tests/test_preprocess_base.py +++ b/junifer/preprocess/tests/test_preprocess_base.py @@ -27,12 +27,12 @@ def test_base_preprocessor_subclassing() -> None: def get_valid_inputs(self): return ["BOLD", "T1w"] - def get_output_type(self, input): - return ["timeseries"] + def get_output_type(self, input_type): + return input_type - def preprocess(self, input, extra_input): - input["data"] = f"mofidied_{input['data']}" - return input + def preprocess(self, input, extra_input=None): + input["data"] = f"modified_{input['data']}" + return input, extra_input with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): MyBasePreprocessor(on=["BOLD", "T2w"]) @@ -70,7 +70,7 @@ def test_base_preprocessor_subclassing() -> None: # Check output assert "BOLD" in output assert "data" in output["BOLD"] - assert output["BOLD"]["data"] == "mofidied_data" + assert output["BOLD"]["data"] == "modified_data" assert "path" in output["BOLD"] assert "meta" in output["BOLD"] -- 2.52.0 From 124886b75663f61a5cc1bb6c7cb63c6804412b97 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Mar 2024 11:34:58 +0100 Subject: [PATCH 6/8] refactor: enable return of helper data type from BOLDWarper.preprocess() --- junifer/preprocess/bold_warper.py | 8 ++++++-- junifer/preprocess/tests/test_bold_warper.py | 4 ++-- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/junifer/preprocess/bold_warper.py b/junifer/preprocess/bold_warper.py index 6aa8c0dc7..eb030374f 100644 --- a/junifer/preprocess/bold_warper.py +++ b/junifer/preprocess/bold_warper.py @@ -9,6 +9,7 @@ from typing import ( Dict, List, Optional, + Tuple, Union, ) @@ -102,7 +103,7 @@ class BOLDWarper(BasePreprocessor): self, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: """Preprocess. Parameters @@ -118,6 +119,9 @@ class BOLDWarper(BasePreprocessor): ------- dict The computed result as dictionary. + None + Extra "helper" data types as dictionary to add to the Junifer Data + object. Raises ------ @@ -235,4 +239,4 @@ class BOLDWarper(BasePreprocessor): input["data"] = nib.load(warped_bold_path) input["space"] = self.ref - return input + return input, None diff --git a/junifer/preprocess/tests/test_bold_warper.py b/junifer/preprocess/tests/test_bold_warper.py index d9c4028a1..c36c45fe7 100644 --- a/junifer/preprocess/tests/test_bold_warper.py +++ b/junifer/preprocess/tests/test_bold_warper.py @@ -86,7 +86,7 @@ def test_BOLDWarper_preprocess_to_native( # Read data element_data = DefaultDataReader().fit_transform(dg[element]) # Preprocess data - data = BOLDWarper(reference="T1w").preprocess( + data, _ = BOLDWarper(reference="T1w").preprocess( input=element_data["BOLD"], extra_input=element_data, ) @@ -149,7 +149,7 @@ def test_BOLDWarper_preprocess_to_multi_mni( element_data = DefaultDataReader().fit_transform(dg[element]) pre_xfm_data = element_data["BOLD"]["data"].get_fdata().copy() # Preprocess data - data = BOLDWarper(reference=space).preprocess( + data, _ = BOLDWarper(reference=space).preprocess( input=element_data["BOLD"], extra_input=element_data, ) -- 2.52.0 From 5f41a4d10e667edf5fd82b3f467d941bb9780023 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Mar 2024 11:35:50 +0100 Subject: [PATCH 7/8] refactor: enable return of helper data type from fMRIPrepConfoundRemover.preprocess() and clean up --- .../confounds/fmriprep_confound_remover.py | 176 +++++++----------- .../tests/test_fmriprep_confound_remover.py | 101 ++++------ 2 files changed, 112 insertions(+), 165 deletions(-) diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index 474500543..153667a8f 100644 --- a/junifer/preprocess/confounds/fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/fmriprep_confound_remover.py @@ -6,7 +6,6 @@ # License: AGPL from typing import ( - TYPE_CHECKING, Any, ClassVar, Dict, @@ -19,8 +18,8 @@ from typing import ( import numpy as np import pandas as pd +from nilearn import image as nimg from nilearn._utils.niimg_conversions import check_niimg_4d -from nilearn.image import clean_img from ...api.decorators import register_preprocessor from ...data import get_mask @@ -28,10 +27,6 @@ from ...utils import logger, raise_error from ..base import BasePreprocessor -if TYPE_CHECKING: - from nibabel import MGHImage, Nifti1Image, Nifti2Image - - FMRIPREP_BASICS = { "motion": [ "trans_x", @@ -448,23 +443,21 @@ class fMRIPrepConfoundRemover(BasePreprocessor): not found or if invalid confounds format is found. """ - - # Bold must be 4D niimg + # BOLD must be 4D niimg check_niimg_4d(input["data"]) - + # Check for extra inputs if extra_input is None: - raise_error(msg="No extra input provided", klass=ValueError) + raise_error( + "No extra input provided, requires `BOLD_confounds` data type " + "in particular" + ) if "BOLD_confounds" not in extra_input: - raise_error(msg="No BOLD_confounds provided", klass=ValueError) + raise_error("`BOLD_confounds` data type not provided") if "data" not in extra_input["BOLD_confounds"]: - raise_error( - msg="No BOLD_confounds data provided", klass=ValueError - ) - # Confounds must be a dataframe + raise_error("`BOLD_confounds.data` not provided") + # Confounds must be a pandas.DataFrame if not isinstance(extra_input["BOLD_confounds"]["data"], pd.DataFrame): - raise_error( - "Confounds data must be a pandas dataframe", ValueError - ) + raise_error("`BOLD_confounds.data` must be a `pandas.DataFrame`") confound_df = extra_input["BOLD_confounds"]["data"] bold_img = input["data"] @@ -475,26 +468,20 @@ class fMRIPrepConfoundRemover(BasePreprocessor): f"\tConfounds: {len(confound_df)}" ) + # Check format if "format" not in extra_input["BOLD_confounds"]: - raise_error( - "Confounds format must be specified in " - 'input["BOLD_confounds"]' - ) - + raise_error("`BOLD_confounds.format` not provided") t_format = extra_input["BOLD_confounds"]["format"] - if t_format == "adhoc": if "mappings" not in extra_input["BOLD_confounds"]: raise_error( - "When using adhoc format, you must specify " - "the variables names mappings in " - 'input["BOLD_confounds"]["mappings"]' + "`BOLD_confounds.mappings` need to be set when " + "`BOLD_confounds.format == 'adhoc'`" ) if "fmriprep" not in extra_input["BOLD_confounds"]["mappings"]: raise_error( - "When using adhoc format, you must specify " - "the variables names mappings to fmriprep format in " - 'input["BOLD_confounds"]["mappings"]["fmriprep"]' + "`BOLD_confounds.mappings.fmriprep` need to be set when " + "`BOLD_confounds.format == 'adhoc'`" ) fmriprep_mappings = extra_input["BOLD_confounds"]["mappings"][ "fmriprep" @@ -506,7 +493,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor): ] if len(wrong_names) > 0: raise_error( - "The following mapping values are not valid fmriprep " + "The following mapping values are not valid fMRIPrep " f"names: {wrong_names}" ) # Check that all the required columns are present @@ -526,81 +513,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor): elif t_format != "fmriprep": raise_error(f"Invalid confounds format {t_format}") - def _remove_confounds( - self, - input: Dict[str, Any], - extra_input: Optional[Dict[str, Any]] = None, - ) -> Union["Nifti1Image", "Nifti2Image", "MGHImage", List]: - """Remove confounds from the BOLD image. - - Parameters - ---------- - input : dict - Dictionary containing the ``BOLD`` value from the - Junifer Data object. - extra_input : dict, optional - Dictionary containing the rest of the Junifer Data object. Must - include the ``BOLD_confounds`` key. - - Returns - ------- - Niimg-like object - Input image with confounds removed. - - """ - assert extra_input is not None # Not the case, data is validated - confounds_df = self._pick_confounds(extra_input["BOLD_confounds"]) - confounds_array = confounds_df.values - - bold_img = input["data"] - t_r = self.t_r - if t_r is None: - logger.info("No `t_r` specified, using t_r from nifti header") - zooms = bold_img.header.get_zooms() # type: ignore - t_r = zooms[3] - logger.info( - f"Read t_r from nifti header: {t_r}", - ) - - mask_img = None - if self.masks is not None: - logger.debug(f"Masking with {self.masks}") - mask_img = get_mask( - masks=self.masks, target_data=input, extra_input=extra_input - ) - # Save the mask in the extra input and link it to the bold data - # this allows to use "inherit" down the pipeline - if extra_input is not None: - logger.debug("Setting mask_item") - extra_input["BOLD_mask"] = { - "data": mask_img, - "space": input["space"], - } - input["mask_item"] = "BOLD_mask" - - logger.info("Cleaning image") - logger.debug(f"\tdetrend: {self.detrend}") - logger.debug(f"\tstandardize: {self.standardize}") - logger.debug(f"\tlow_pass: {self.low_pass}") - logger.debug(f"\thigh_pass: {self.high_pass}") - clean_bold = clean_img( - imgs=bold_img, - detrend=self.detrend, - standardize=self.standardize, - confounds=confounds_array, - low_pass=self.low_pass, - high_pass=self.high_pass, - t_r=t_r, - mask_img=mask_img, - ) - - return clean_bold - def preprocess( self, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Dict[str, Any]: + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: """Preprocess. Parameters @@ -615,8 +532,59 @@ class fMRIPrepConfoundRemover(BasePreprocessor): ------- dict The computed result as dictionary. + dict or None + If `self.masks` is not None, then the target data computed mask is + returned else None. """ + # Validate data self._validate_data(input, extra_input) - input["data"] = self._remove_confounds(input, extra_input=extra_input) - return input + # Pick confounds + confounds_df = self._pick_confounds(extra_input["BOLD_confounds"]) # type: ignore + # Get BOLD data + bold_img = input["data"] + # Set t_r + t_r = self.t_r + if t_r is None: + logger.info("No `t_r` specified, using t_r from NIfTI header") + t_r = bold_img.header.get_zooms()[3] # type: ignore + logger.info( + f"Read t_r from NIfTI header: {t_r}", + ) + # Set mask data + mask_img = None + bold_mask_dict = None + if self.masks is not None: + logger.debug(f"Masking with {self.masks}") + mask_img = get_mask( + masks=self.masks, target_data=input, extra_input=extra_input + ) + # Return the BOLD mask and link it to the BOLD data type dict; + # this allows to use "inherit" down the pipeline + if extra_input is not None: + logger.debug("Setting `BOLD.mask_item`") + input["mask_item"] = "BOLD_mask" + bold_mask_dict = { + "BOLD_mask": { + "data": mask_img, + "space": input["space"], + } + } + # Clean image + logger.info("Cleaning image using nilearn") + logger.debug(f"\tdetrend: {self.detrend}") + logger.debug(f"\tstandardize: {self.standardize}") + logger.debug(f"\tlow_pass: {self.low_pass}") + logger.debug(f"\thigh_pass: {self.high_pass}") + input["data"] = nimg.clean_img( + imgs=bold_img, + detrend=self.detrend, + standardize=self.standardize, + confounds=confounds_df.values, + low_pass=self.low_pass, + high_pass=self.high_pass, + t_r=t_r, + mask_img=mask_img, + ) + + return input, bold_mask_dict diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index 6dcd4efd2..cb13452ba 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -5,9 +5,8 @@ # Synchon Mandal # License: AGPL -from typing import List, cast +from typing import List -import nibabel as nib import numpy as np import pandas as pd import pytest @@ -349,78 +348,80 @@ def test_fMRIPRepConfoundRemover__pick_confounds_fmriprep_compute() -> None: def test_fMRIPrepConfoundRemover__validate_data() -> None: """Test fMRIPrepConfoundRemover validate data.""" confound_remover = fMRIPrepConfoundRemover(strategy={"wm_csf": "full"}) - reader = DefaultDataReader() + with OasisVBMTestingDataGrabber() as dg: - input = dg["sub-01"] - input = reader.fit_transform(input) - new_input = input["VBM_GM"] + element_data = DefaultDataReader().fit_transform(dg["sub-01"]) + vbm = element_data["VBM_GM"] with pytest.raises( DimensionError, match="incompatible dimensionality" ): - confound_remover._validate_data(new_input, None) + confound_remover._validate_data(vbm, None) with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: - input = dg["sub-01"] - input = reader.fit_transform(input) - new_input = input["BOLD"] + element_data = DefaultDataReader().fit_transform(dg["sub-01"]) + bold = element_data["BOLD"] with pytest.raises(ValueError, match="No extra input"): - confound_remover._validate_data(new_input, None) - with pytest.raises(ValueError, match="No BOLD_confounds provided"): - confound_remover._validate_data(new_input, {}) + confound_remover._validate_data(bold, None) with pytest.raises( - ValueError, match="No BOLD_confounds data provided" + ValueError, match="`BOLD_confounds` data type not provided" ): - confound_remover._validate_data(new_input, {"BOLD_confounds": {}}) + confound_remover._validate_data(bold, {}) + with pytest.raises( + ValueError, match="`BOLD_confounds.data` not provided" + ): + confound_remover._validate_data(bold, {"BOLD_confounds": {}}) extra_input = { "BOLD_confounds": {"data": "wrong"}, } - msg = "must be a pandas dataframe" + msg = "must be a `pandas.DataFrame`" with pytest.raises(ValueError, match=msg): - confound_remover._validate_data(new_input, extra_input) + confound_remover._validate_data(bold, extra_input) extra_input = {"BOLD_confounds": {"data": pd.DataFrame()}} with pytest.raises(ValueError, match="Image time series and"): - confound_remover._validate_data(new_input, extra_input) + confound_remover._validate_data(bold, extra_input) extra_input = { - "BOLD_confounds": {"data": input["BOLD_confounds"]["data"]} + "BOLD_confounds": {"data": element_data["BOLD_confounds"]["data"]} } - with pytest.raises(ValueError, match="format must be specified"): - confound_remover._validate_data(new_input, extra_input) + with pytest.raises( + ValueError, match="`BOLD_confounds.format` not provided" + ): + confound_remover._validate_data(bold, extra_input) extra_input = { "BOLD_confounds": { - "data": input["BOLD_confounds"]["data"], + "data": element_data["BOLD_confounds"]["data"], "format": "wrong", } } - with pytest.raises(ValueError, match="Invalid confounds format wrong"): - confound_remover._validate_data(new_input, extra_input) + with pytest.raises(ValueError, match="Invalid confounds format"): + confound_remover._validate_data(bold, extra_input) extra_input = { "BOLD_confounds": { - "data": input["BOLD_confounds"]["data"], + "data": element_data["BOLD_confounds"]["data"], "format": "adhoc", } } - with pytest.raises(ValueError, match="variables names mappings"): - confound_remover._validate_data(new_input, extra_input) + with pytest.raises(ValueError, match="need to be set"): + confound_remover._validate_data(bold, extra_input) extra_input = { "BOLD_confounds": { - "data": input["BOLD_confounds"]["data"], + "data": element_data["BOLD_confounds"]["data"], "format": "adhoc", "mappings": {}, } } - with pytest.raises(ValueError, match="mappings to fmriprep"): - confound_remover._validate_data(new_input, extra_input) + with pytest.raises(ValueError, match="need to be set"): + confound_remover._validate_data(bold, extra_input) extra_input = { "BOLD_confounds": { - "data": input["BOLD_confounds"]["data"], + "data": element_data["BOLD_confounds"]["data"], "format": "adhoc", "mappings": { "fmriprep": { @@ -432,11 +433,11 @@ def test_fMRIPrepConfoundRemover__validate_data() -> None: } } with pytest.raises(ValueError, match=r"names: \['wrong'\]"): - confound_remover._validate_data(new_input, extra_input) + confound_remover._validate_data(bold, extra_input) extra_input = { "BOLD_confounds": { - "data": input["BOLD_confounds"]["data"], + "data": element_data["BOLD_confounds"]["data"], "format": "adhoc", "mappings": { "fmriprep": { @@ -448,11 +449,11 @@ def test_fMRIPrepConfoundRemover__validate_data() -> None: } } with pytest.raises(ValueError, match=r"Missing columns: \['wrong'\]"): - confound_remover._validate_data(new_input, extra_input) + confound_remover._validate_data(bold, extra_input) extra_input = { "BOLD_confounds": { - "data": input["BOLD_confounds"]["data"], + "data": element_data["BOLD_confounds"]["data"], "format": "adhoc", "mappings": { "fmriprep": { @@ -463,31 +464,7 @@ def test_fMRIPrepConfoundRemover__validate_data() -> None: }, } } - confound_remover._validate_data(new_input, extra_input) - - -def test_fMRIPrepConfoundRemover__remove_confounds() -> None: - """Test fMRIPrepConfoundRemover remove confounds.""" - confound_remover = fMRIPrepConfoundRemover( - strategy={"wm_csf": "full"}, spike=0.2 - ) - reader = DefaultDataReader() - with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: - input = dg["sub-01"] - input = reader.fit_transform(input) - raw_bold = input["BOLD"]["data"] - extra_input = {k: v for k, v in input.items() if k != "BOLD"} - clean_bold = confound_remover._remove_confounds( - input=input["BOLD"], extra_input=extra_input - ) - clean_bold = cast(nib.Nifti1Image, clean_bold) - # TODO: Find a better way to test functionality here - assert ( - clean_bold.header.get_zooms() # type: ignore - == raw_bold.header.get_zooms() # type: ignore - ) - assert clean_bold.get_fdata().shape == raw_bold.get_fdata().shape - # TODO: Test confound remover with mask, needs #79 to be implemented + confound_remover._validate_data(bold, extra_input) def test_fMRIPrepConfoundRemover_preprocess() -> None: @@ -500,7 +477,7 @@ def test_fMRIPrepConfoundRemover_preprocess() -> None: orig_bold = element_data["BOLD"]["data"].get_fdata().copy() pre_input = element_data["BOLD"] 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( @@ -553,6 +530,8 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None: assert t_meta["t_r"] is None assert t_meta["masks"] is None + assert "BOLD_mask" not in output + assert "dependencies" in output["BOLD"]["meta"] dependencies = output["BOLD"]["meta"]["dependencies"] assert dependencies == {"numpy", "nilearn"} -- 2.52.0 From 5668e631d498d7715b2cf8b8cf1b04c397ef288e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 8 Mar 2024 13:16:48 +0100 Subject: [PATCH 8/8] chore: add changelogs 310.{change,enh} --- docs/changes/newsfragments/310.change | 1 + docs/changes/newsfragments/310.enh | 1 + 2 files changed, 2 insertions(+) create mode 100644 docs/changes/newsfragments/310.change create mode 100644 docs/changes/newsfragments/310.enh diff --git a/docs/changes/newsfragments/310.change b/docs/changes/newsfragments/310.change new file mode 100644 index 000000000..1dddd7b93 --- /dev/null +++ b/docs/changes/newsfragments/310.change @@ -0,0 +1 @@ +Change :meth:`.BasePreprocessor.preprocess` return values to preprocessed target data and "helper" data types as a dictionary by `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/310.enh b/docs/changes/newsfragments/310.enh new file mode 100644 index 000000000..1cb4b4a7b --- /dev/null +++ b/docs/changes/newsfragments/310.enh @@ -0,0 +1 @@ +Improve :class:`.BasePreprocessor` by revamping :meth:`.BasePreprocessor.preprocess` and ``BasePreprocessor._fit_transform`` to handle "helper" data types better and make the pipeline explicit where data is being altered by `Synchon Mandal`_ -- 2.52.0