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`_ diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index 489988739..d29b5dc40 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( @@ -126,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]]: + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: """Preprocess. Parameters @@ -141,11 +142,12 @@ 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. + 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( @@ -170,31 +172,36 @@ 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 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 - extra_input = input + extra_input = input.copy() 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( + # Preprocess data + t_out, t_extra_input = 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") + # 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/bold_warper.py b/junifer/preprocess/bold_warper.py index ddeeed93e..eb030374f 100644 --- a/junifer/preprocess/bold_warper.py +++ b/junifer/preprocess/bold_warper.py @@ -82,30 +82,28 @@ 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, input: Dict[str, Any], extra_input: Optional[Dict[str, Any]] = None, - ) -> Tuple[str, Dict[str, Any]]: + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: """Preprocess. Parameters @@ -119,11 +117,11 @@ 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. + None + Extra "helper" data types as dictionary to add to the Junifer Data + object. Raises ------ @@ -241,4 +239,4 @@ class BOLDWarper(BasePreprocessor): input["data"] = nib.load(warped_bold_path) input["space"] = self.ref - return "BOLD", input + return input, None diff --git a/junifer/preprocess/confounds/fmriprep_confound_remover.py b/junifer/preprocess/confounds/fmriprep_confound_remover.py index 30fb51f41..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", @@ -224,24 +219,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. @@ -450,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"] @@ -477,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" @@ -508,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 @@ -528,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, - ) -> Tuple[str, Dict[str, Any]]: + ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]: """Preprocess. Parameters @@ -615,13 +530,61 @@ 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. + 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 "BOLD", 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 d197dafd7..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 @@ -91,26 +90,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: @@ -364,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": { @@ -447,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": { @@ -463,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": { @@ -478,51 +464,25 @@ 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: """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 @@ -531,25 +491,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 @@ -573,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"} @@ -580,22 +539,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 diff --git a/junifer/preprocess/tests/test_bold_warper.py b/junifer/preprocess/tests/test_bold_warper.py index 58a26db14..c36c45fe7 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( @@ -101,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) @@ -165,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): diff --git a/junifer/preprocess/tests/test_preprocess_base.py b/junifer/preprocess/tests/test_preprocess_base.py index 199679517..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 "BOLD", 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"]