[ENH]: Improve BasePreprocessor #310

Merged
synchon merged 8 commits from refactor/preprocessor-fit-transform into main 2024-03-13 08:45:38 +00:00
8 changed files with 187 additions and 277 deletions

View file

@ -0,0 +1 @@
Change :meth:`.BasePreprocessor.preprocess` return values to preprocessed target data and "helper" data types as a dictionary by `Synchon Mandal`_

View file

@ -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`_

View file

@ -12,15 +12,18 @@ from ..utils import logger, raise_error
class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin): 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 Parameters
---------- ----------
on : str or list of str, optional on : str or list of str or None, optional
The kind of data to apply the preprocessor to. If None, The data type to apply the preprocessor on. If None,
will work on all available data (default None). will work on all available data types (default None).
required_data_types : str or list of str, optional 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). will be equal to ``on`` (default None).
Raises Raises
@ -100,20 +103,18 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
) )
@abstractmethod @abstractmethod
def get_output_type(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output type. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the preprocessor. The list must contain the The data type input to the preprocessor.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of available Junifer Data object keys after The data type output by the preprocessor.
the pipeline step.
""" """
raise_error( raise_error(
@ -126,7 +127,7 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
self, self,
input: Dict[str, Any], input: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None, extra_input: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]: ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -141,11 +142,12 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
Returns Returns
------- -------
str
The key to store the output in the Junifer Data object.
dict dict
The computed result as dictionary. This will be stored in the The computed result as dictionary.
Junifer Data object under the key ``data`` of the data type. 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( raise_error(
@ -170,31 +172,36 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
The processed output as a dictionary. 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: for type_ in self._on:
# Check if data type is available
if type_ in input.keys(): if type_ in input.keys():
logger.info(f"Preprocessing {type_}") logger.info(f"Preprocessing {type_}")
# Get data dict for data type
t_input = input[type_] t_input = input[type_]
# Pass the other data types as extra input, removing # Pass the other data types as extra input, removing
# the current type # the current type
extra_input = input extra_input = input.copy()
extra_input.pop(type_) extra_input.pop(type_)
logger.debug( 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 input=t_input, extra_input=extra_input
) )
# Set output to the Junifer Data object
# Add the output to the Junifer Data object logger.debug(f"Adding {type_} to output")
logger.debug(f"Adding {key} to output") out[type_] = t_out
out[key] = t_out # Check if helper data types are to be added
if t_extra_input is not None:
# In case we are creating a new type, re-add the original input logger.debug(
if key != type_: f"Adding helper data types: {t_extra_input.keys()} "
logger.debug("Adding original input back to output") "to output"
out[type_] = t_input )
out.update(t_extra_input)
self.update_meta(out[key], "preprocess") # Update metadata for step
self.update_meta(out[type_], "preprocess")
return out return out

View file

@ -82,30 +82,28 @@ class BOLDWarper(BasePreprocessor):
""" """
return ["BOLD"] return ["BOLD"]
def get_output_type(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output type. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the preprocessor. The list must contain the The data type input to the preprocessor.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of available Junifer Data object keys after The data type output by the preprocessor.
the pipeline step.
""" """
# Does not add any new keys # Does not add any new keys
return input return input_type
def preprocess( def preprocess(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None, extra_input: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]: ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -119,11 +117,11 @@ class BOLDWarper(BasePreprocessor):
Returns Returns
------- -------
str
The key to store the output in the Junifer Data object.
dict dict
The computed result as dictionary. This will be stored in the The computed result as dictionary.
Junifer Data object under the key ``data`` of the data type. None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
Raises Raises
------ ------
@ -241,4 +239,4 @@ class BOLDWarper(BasePreprocessor):
input["data"] = nib.load(warped_bold_path) input["data"] = nib.load(warped_bold_path)
input["space"] = self.ref input["space"] = self.ref
return "BOLD", input return input, None

View file

@ -6,7 +6,6 @@
# License: AGPL # License: AGPL
from typing import ( from typing import (
TYPE_CHECKING,
Any, Any,
ClassVar, ClassVar,
Dict, Dict,
@ -19,8 +18,8 @@ from typing import (
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from nilearn import image as nimg
from nilearn._utils.niimg_conversions import check_niimg_4d from nilearn._utils.niimg_conversions import check_niimg_4d
from nilearn.image import clean_img
from ...api.decorators import register_preprocessor from ...api.decorators import register_preprocessor
from ...data import get_mask from ...data import get_mask
@ -28,10 +27,6 @@ from ...utils import logger, raise_error
from ..base import BasePreprocessor from ..base import BasePreprocessor
if TYPE_CHECKING:
from nibabel import MGHImage, Nifti1Image, Nifti2Image
FMRIPREP_BASICS = { FMRIPREP_BASICS = {
"motion": [ "motion": [
"trans_x", "trans_x",
@ -224,24 +219,22 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
""" """
return ["BOLD"] return ["BOLD"]
def get_output_type(self, input: List[str]) -> List[str]: def get_output_type(self, input_type: str) -> str:
"""Get output type. """Get output type.
Parameters Parameters
---------- ----------
input : list of str input_type : str
The input to the preprocessor. The list must contain the The input to the preprocessor.
available Junifer Data dictionary keys.
Returns Returns
------- -------
list of str str
The updated list of available Junifer Data object keys after The data type output by the preprocessor.
the pipeline step.
""" """
# Does not add any new keys # Does not add any new keys
return input return input_type
def _map_adhoc_to_fmriprep(self, input: Dict[str, Any]) -> None: def _map_adhoc_to_fmriprep(self, input: Dict[str, Any]) -> None:
"""Map the adhoc format to the fmpriprep format spec. """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. 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_niimg_4d(input["data"])
# Check for extra inputs
if extra_input is None: 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: 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"]: if "data" not in extra_input["BOLD_confounds"]:
raise_error( raise_error("`BOLD_confounds.data` not provided")
msg="No BOLD_confounds data provided", klass=ValueError # Confounds must be a pandas.DataFrame
)
# Confounds must be a dataframe
if not isinstance(extra_input["BOLD_confounds"]["data"], pd.DataFrame): if not isinstance(extra_input["BOLD_confounds"]["data"], pd.DataFrame):
raise_error( raise_error("`BOLD_confounds.data` must be a `pandas.DataFrame`")
"Confounds data must be a pandas dataframe", ValueError
)
confound_df = extra_input["BOLD_confounds"]["data"] confound_df = extra_input["BOLD_confounds"]["data"]
bold_img = input["data"] bold_img = input["data"]
@ -477,26 +468,20 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
f"\tConfounds: {len(confound_df)}" f"\tConfounds: {len(confound_df)}"
) )
# Check format
if "format" not in extra_input["BOLD_confounds"]: if "format" not in extra_input["BOLD_confounds"]:
raise_error( raise_error("`BOLD_confounds.format` not provided")
"Confounds format must be specified in "
'input["BOLD_confounds"]'
)
t_format = extra_input["BOLD_confounds"]["format"] t_format = extra_input["BOLD_confounds"]["format"]
if t_format == "adhoc": if t_format == "adhoc":
if "mappings" not in extra_input["BOLD_confounds"]: if "mappings" not in extra_input["BOLD_confounds"]:
raise_error( raise_error(
"When using adhoc format, you must specify " "`BOLD_confounds.mappings` need to be set when "
"the variables names mappings in " "`BOLD_confounds.format == 'adhoc'`"
'input["BOLD_confounds"]["mappings"]'
) )
if "fmriprep" not in extra_input["BOLD_confounds"]["mappings"]: if "fmriprep" not in extra_input["BOLD_confounds"]["mappings"]:
raise_error( raise_error(
"When using adhoc format, you must specify " "`BOLD_confounds.mappings.fmriprep` need to be set when "
"the variables names mappings to fmriprep format in " "`BOLD_confounds.format == 'adhoc'`"
'input["BOLD_confounds"]["mappings"]["fmriprep"]'
) )
fmriprep_mappings = extra_input["BOLD_confounds"]["mappings"][ fmriprep_mappings = extra_input["BOLD_confounds"]["mappings"][
"fmriprep" "fmriprep"
@ -508,7 +493,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
] ]
if len(wrong_names) > 0: if len(wrong_names) > 0:
raise_error( raise_error(
"The following mapping values are not valid fmriprep " "The following mapping values are not valid fMRIPrep "
f"names: {wrong_names}" f"names: {wrong_names}"
) )
# Check that all the required columns are present # Check that all the required columns are present
@ -528,81 +513,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
elif t_format != "fmriprep": elif t_format != "fmriprep":
raise_error(f"Invalid confounds format {t_format}") 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( def preprocess(
self, self,
input: Dict[str, Any], input: Dict[str, Any],
extra_input: Optional[Dict[str, Any]] = None, extra_input: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]: ) -> Tuple[Dict[str, Any], Optional[Dict[str, Dict[str, Any]]]]:
"""Preprocess. """Preprocess.
Parameters Parameters
@ -615,13 +530,61 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
Returns Returns
------- -------
str
The key to store the output in the Junifer Data object.
dict dict
The computed result as dictionary. This will be stored in the The computed result as dictionary.
Junifer Data object under the key ``data`` of the data type. 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) self._validate_data(input, extra_input)
input["data"] = self._remove_confounds(input, extra_input=extra_input) # Pick confounds
return "BOLD", input 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

View file

@ -5,9 +5,8 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import List, cast from typing import List
import nibabel as nib
import numpy as np import numpy as np
import pandas as pd import pandas as pd
import pytest import pytest
@ -91,26 +90,11 @@ def test_fMRIPrepConfoundRemover_get_valid_inputs() -> None:
assert confound_remover.get_valid_inputs() == ["BOLD"] assert confound_remover.get_valid_inputs() == ["BOLD"]
@pytest.mark.parametrize( def test_fMRIPrepConfoundRemover_get_output_type() -> None:
"input_", """Test fMRIPrepConfoundRemover get_output_type."""
[
["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 = fMRIPrepConfoundRemover()
# Confound remover works in place # 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: 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: def test_fMRIPrepConfoundRemover__validate_data() -> None:
"""Test fMRIPrepConfoundRemover validate data.""" """Test fMRIPrepConfoundRemover validate data."""
confound_remover = fMRIPrepConfoundRemover(strategy={"wm_csf": "full"}) confound_remover = fMRIPrepConfoundRemover(strategy={"wm_csf": "full"})
reader = DefaultDataReader()
with OasisVBMTestingDataGrabber() as dg: with OasisVBMTestingDataGrabber() as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) vbm = element_data["VBM_GM"]
new_input = input["VBM_GM"]
with pytest.raises( with pytest.raises(
DimensionError, match="incompatible dimensionality" DimensionError, match="incompatible dimensionality"
): ):
confound_remover._validate_data(new_input, None) confound_remover._validate_data(vbm, None)
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) bold = element_data["BOLD"]
new_input = input["BOLD"]
with pytest.raises(ValueError, match="No extra input"): with pytest.raises(ValueError, match="No extra input"):
confound_remover._validate_data(new_input, None) confound_remover._validate_data(bold, None)
with pytest.raises(ValueError, match="No BOLD_confounds provided"):
confound_remover._validate_data(new_input, {})
with pytest.raises( 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 = { extra_input = {
"BOLD_confounds": {"data": "wrong"}, "BOLD_confounds": {"data": "wrong"},
} }
msg = "must be a pandas dataframe" msg = "must be a `pandas.DataFrame`"
with pytest.raises(ValueError, match=msg): 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()}} extra_input = {"BOLD_confounds": {"data": pd.DataFrame()}}
with pytest.raises(ValueError, match="Image time series and"): 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 = { 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"): with pytest.raises(
confound_remover._validate_data(new_input, extra_input) ValueError, match="`BOLD_confounds.format` not provided"
):
confound_remover._validate_data(bold, extra_input)
extra_input = { extra_input = {
"BOLD_confounds": { "BOLD_confounds": {
"data": input["BOLD_confounds"]["data"], "data": element_data["BOLD_confounds"]["data"],
"format": "wrong", "format": "wrong",
} }
} }
with pytest.raises(ValueError, match="Invalid confounds format wrong"): with pytest.raises(ValueError, match="Invalid confounds format"):
confound_remover._validate_data(new_input, extra_input) confound_remover._validate_data(bold, extra_input)
extra_input = { extra_input = {
"BOLD_confounds": { "BOLD_confounds": {
"data": input["BOLD_confounds"]["data"], "data": element_data["BOLD_confounds"]["data"],
"format": "adhoc", "format": "adhoc",
} }
} }
with pytest.raises(ValueError, match="variables names mappings"): with pytest.raises(ValueError, match="need to be set"):
confound_remover._validate_data(new_input, extra_input) confound_remover._validate_data(bold, extra_input)
extra_input = { extra_input = {
"BOLD_confounds": { "BOLD_confounds": {
"data": input["BOLD_confounds"]["data"], "data": element_data["BOLD_confounds"]["data"],
"format": "adhoc", "format": "adhoc",
"mappings": {}, "mappings": {},
} }
} }
with pytest.raises(ValueError, match="mappings to fmriprep"): with pytest.raises(ValueError, match="need to be set"):
confound_remover._validate_data(new_input, extra_input) confound_remover._validate_data(bold, extra_input)
extra_input = { extra_input = {
"BOLD_confounds": { "BOLD_confounds": {
"data": input["BOLD_confounds"]["data"], "data": element_data["BOLD_confounds"]["data"],
"format": "adhoc", "format": "adhoc",
"mappings": { "mappings": {
"fmriprep": { "fmriprep": {
@ -447,11 +433,11 @@ def test_fMRIPrepConfoundRemover__validate_data() -> None:
} }
} }
with pytest.raises(ValueError, match=r"names: \['wrong'\]"): 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 = { extra_input = {
"BOLD_confounds": { "BOLD_confounds": {
"data": input["BOLD_confounds"]["data"], "data": element_data["BOLD_confounds"]["data"],
"format": "adhoc", "format": "adhoc",
"mappings": { "mappings": {
"fmriprep": { "fmriprep": {
@ -463,11 +449,11 @@ def test_fMRIPrepConfoundRemover__validate_data() -> None:
} }
} }
with pytest.raises(ValueError, match=r"Missing columns: \['wrong'\]"): 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 = { extra_input = {
"BOLD_confounds": { "BOLD_confounds": {
"data": input["BOLD_confounds"]["data"], "data": element_data["BOLD_confounds"]["data"],
"format": "adhoc", "format": "adhoc",
"mappings": { "mappings": {
"fmriprep": { "fmriprep": {
@ -478,51 +464,25 @@ def test_fMRIPrepConfoundRemover__validate_data() -> None:
}, },
} }
} }
confound_remover._validate_data(new_input, extra_input) confound_remover._validate_data(bold, 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
def test_fMRIPrepConfoundRemover_preprocess() -> None: def test_fMRIPrepConfoundRemover_preprocess() -> None:
"""Test fMRIPrepConfoundRemover with all confounds present.""" """Test fMRIPrepConfoundRemover with all confounds present."""
# need reader for the data
reader = DefaultDataReader()
# All strategies full, no spike # All strategies full, no spike
confound_remover = fMRIPrepConfoundRemover() confound_remover = fMRIPrepConfoundRemover()
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) orig_bold = element_data["BOLD"]["data"].get_fdata().copy()
orig_bold = input["BOLD"]["data"].get_fdata().copy() pre_input = element_data["BOLD"]
pre_input = input["BOLD"] pre_extra_input = {"BOLD_confounds": element_data["BOLD_confounds"]}
pre_extra_input = {"BOLD_confounds": input["BOLD_confounds"]} output, _ = confound_remover.preprocess(pre_input, pre_extra_input)
key, output = confound_remover.preprocess(pre_input, pre_extra_input)
trans_bold = output["data"].get_fdata() trans_bold = output["data"].get_fdata()
# Transformation is in place # 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 # Data should have the same shape
assert orig_bold.shape == trans_bold.shape assert orig_bold.shape == trans_bold.shape
@ -531,25 +491,22 @@ def test_fMRIPrepConfoundRemover_preprocess() -> None:
assert_raises( assert_raises(
AssertionError, assert_array_equal, orig_bold, trans_bold AssertionError, assert_array_equal, orig_bold, trans_bold
) )
assert key == "BOLD"
def test_fMRIPrepConfoundRemover_fit_transform() -> None: def test_fMRIPrepConfoundRemover_fit_transform() -> None:
"""Test fMRIPrepConfoundRemover with all confounds present.""" """Test fMRIPrepConfoundRemover with all confounds present."""
# need reader for the data
reader = DefaultDataReader()
# All strategies full, no spike # All strategies full, no spike
confound_remover = fMRIPrepConfoundRemover() confound_remover = fMRIPrepConfoundRemover()
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) orig_bold = element_data["BOLD"]["data"].get_fdata().copy()
orig_bold = input["BOLD"]["data"].get_fdata().copy() output = confound_remover.fit_transform(element_data)
output = confound_remover.fit_transform(input)
trans_bold = output["BOLD"]["data"].get_fdata() trans_bold = output["BOLD"]["data"].get_fdata()
# Transformation is in place # 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 # Data should have the same shape
assert orig_bold.shape == trans_bold.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["t_r"] is None
assert t_meta["masks"] is None assert t_meta["masks"] is None
assert "BOLD_mask" not in output
assert "dependencies" in output["BOLD"]["meta"] assert "dependencies" in output["BOLD"]["meta"]
dependencies = output["BOLD"]["meta"]["dependencies"] dependencies = output["BOLD"]["meta"]["dependencies"]
assert dependencies == {"numpy", "nilearn"} assert dependencies == {"numpy", "nilearn"}
@ -580,22 +539,20 @@ def test_fMRIPrepConfoundRemover_fit_transform() -> None:
def test_fMRIPrepConfoundRemover_fit_transform_masks() -> None: def test_fMRIPrepConfoundRemover_fit_transform_masks() -> None:
"""Test fMRIPrepConfoundRemover with all confounds present.""" """Test fMRIPrepConfoundRemover with all confounds present."""
# need reader for the data
reader = DefaultDataReader()
# All strategies full, no spike # All strategies full, no spike
confound_remover = fMRIPrepConfoundRemover( confound_remover = fMRIPrepConfoundRemover(
masks={"compute_brain_mask": {"threshold": 0.2}} masks={"compute_brain_mask": {"threshold": 0.2}}
) )
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg: with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
input = dg["sub-01"] element_data = DefaultDataReader().fit_transform(dg["sub-01"])
input = reader.fit_transform(input) orig_bold = element_data["BOLD"]["data"].get_fdata().copy()
orig_bold = input["BOLD"]["data"].get_fdata().copy() output = confound_remover.fit_transform(element_data)
output = confound_remover.fit_transform(input)
trans_bold = output["BOLD"]["data"].get_fdata() trans_bold = output["BOLD"]["data"].get_fdata()
# Transformation is in place # 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 # Data should have the same shape
assert orig_bold.shape == trans_bold.shape assert orig_bold.shape == trans_bold.shape

View file

@ -4,7 +4,7 @@
# License: AGPL # License: AGPL
import socket import socket
from typing import TYPE_CHECKING, List, Tuple from typing import TYPE_CHECKING, Tuple
import pytest import pytest
from numpy.testing import assert_array_equal, assert_raises 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"] assert bold_warper.get_valid_inputs() == ["BOLD"]
@pytest.mark.parametrize( def test_BOLDWarper_get_output_type() -> None:
"input_", """Test BOLDWarper get_output_type."""
[
["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.
"""
bold_warper = BOLDWarper(reference="T1w") bold_warper = BOLDWarper(reference="T1w")
assert bold_warper.get_output_type(input_) == input_ assert bold_warper.get_output_type("BOLD") == "BOLD"
@pytest.mark.parametrize( @pytest.mark.parametrize(
@ -101,11 +86,10 @@ def test_BOLDWarper_preprocess_to_native(
# Read data # Read data
element_data = DefaultDataReader().fit_transform(dg[element]) element_data = DefaultDataReader().fit_transform(dg[element])
# Preprocess data # Preprocess data
data_type, data = BOLDWarper(reference="T1w").preprocess( data, _ = BOLDWarper(reference="T1w").preprocess(
input=element_data["BOLD"], input=element_data["BOLD"],
extra_input=element_data, extra_input=element_data,
) )
assert data_type == "BOLD"
assert isinstance(data, dict) assert isinstance(data, dict)
@ -165,11 +149,10 @@ def test_BOLDWarper_preprocess_to_multi_mni(
element_data = DefaultDataReader().fit_transform(dg[element]) element_data = DefaultDataReader().fit_transform(dg[element])
pre_xfm_data = element_data["BOLD"]["data"].get_fdata().copy() pre_xfm_data = element_data["BOLD"]["data"].get_fdata().copy()
# Preprocess data # Preprocess data
data_type, data = BOLDWarper(reference=space).preprocess( data, _ = BOLDWarper(reference=space).preprocess(
input=element_data["BOLD"], input=element_data["BOLD"],
extra_input=element_data, extra_input=element_data,
) )
assert data_type == "BOLD"
assert isinstance(data, dict) assert isinstance(data, dict)
assert data["space"] == space assert data["space"] == space
with assert_raises(AssertionError): with assert_raises(AssertionError):

View file

@ -27,12 +27,12 @@ def test_base_preprocessor_subclassing() -> None:
def get_valid_inputs(self): def get_valid_inputs(self):
return ["BOLD", "T1w"] return ["BOLD", "T1w"]
def get_output_type(self, input): def get_output_type(self, input_type):
return ["timeseries"] return input_type
def preprocess(self, input, extra_input): def preprocess(self, input, extra_input=None):
input["data"] = f"mofidied_{input['data']}" input["data"] = f"modified_{input['data']}"
return "BOLD", input return input, extra_input
with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"): with pytest.raises(ValueError, match=r"cannot be computed on \['T2w'\]"):
MyBasePreprocessor(on=["BOLD", "T2w"]) MyBasePreprocessor(on=["BOLD", "T2w"])
@ -70,7 +70,7 @@ def test_base_preprocessor_subclassing() -> None:
# Check output # Check output
assert "BOLD" in output assert "BOLD" in output
assert "data" in output["BOLD"] assert "data" in output["BOLD"]
assert output["BOLD"]["data"] == "mofidied_data" assert output["BOLD"]["data"] == "modified_data"
assert "path" in output["BOLD"] assert "path" in output["BOLD"]
assert "meta" in output["BOLD"] assert "meta" in output["BOLD"]