[ENH]: Add scrubbing support in fMRIPrepConfoundRemover #421
5 changed files with 278 additions and 47 deletions
1
docs/changes/newsfragments/421.change
Normal file
1
docs/changes/newsfragments/421.change
Normal file
|
|
@ -0,0 +1 @@
|
|||
Add ``scrub``, ``fd_threshold`` and ``std_vars_threshold`` parameters to :class:`.fMRIPrepConfoundRemover` and allow ``"scrubbing"`` key of type bool for ``strategy`` by `Synchon Mandal`_
|
||||
1
docs/changes/newsfragments/421.enh
Normal file
1
docs/changes/newsfragments/421.enh
Normal file
|
|
@ -0,0 +1 @@
|
|||
Add scrubbing support to :class:`.fMRIPrepConfoundRemover` by `Synchon Mandal`_
|
||||
|
|
@ -185,7 +185,7 @@ def signals() -> list[np.ndarray]:
|
|||
|
||||
@pytest.fixture
|
||||
def signals_and_covariances(
|
||||
cov_estimator: Union[LedoitWolf, EmpiricalCovariance]
|
||||
cov_estimator: Union[LedoitWolf, EmpiricalCovariance],
|
||||
) -> tuple[list[np.ndarray], list[float]]:
|
||||
"""Return signals and covariances for a covariance estimator.
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,8 @@ 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.interfaces.fmriprep.load_confounds_components import _load_scrub
|
||||
from nilearn.interfaces.fmriprep.load_confounds_utils import prepare_output
|
||||
|
||||
from ...api.decorators import register_preprocessor
|
||||
from ...data import get_data
|
||||
|
|
@ -110,13 +112,15 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
----------
|
||||
strategy : dict, optional
|
||||
The strategy to use for each component. If None, will use the *full*
|
||||
strategy for all components (default None).
|
||||
strategy for all components except ``"scrubbing"`` which will be set
|
||||
to False (default None).
|
||||
The keys of the dictionary should correspond to names of noise
|
||||
components to include:
|
||||
|
||||
* ``motion``
|
||||
* ``wm_csf``
|
||||
* ``global_signal``
|
||||
* ``scrubbing``
|
||||
|
||||
The values of dictionary should correspond to types of confounds
|
||||
extracted from each signal:
|
||||
|
|
@ -126,10 +130,29 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
* ``derivatives`` : signal + derivatives
|
||||
* ``full`` : signal + deriv. + quadratic terms + power2 deriv.
|
||||
|
||||
except ``scrubbing`` which needs to be bool.
|
||||
spike : float, optional
|
||||
If None, no spike regressor is added. If spike is a float, it will
|
||||
add a spike regressor for every point at which framewise displacement
|
||||
exceeds the specified float (default None).
|
||||
scrub : int, optional
|
||||
After accounting for time frames with excessive motion, further remove
|
||||
segments shorter than the given number. When the value is 0, remove
|
||||
time frames based on excessive framewise displacement and DVARS only.
|
||||
If None and no ``"scrubbing"`` in ``strategy``, no scrubbing is
|
||||
performed, else the default value is 0. The default value is referred
|
||||
as full scrubbing (default None).
|
||||
fd_threshold : float, optional
|
||||
Framewise displacement threshold for scrub in mm. If None no
|
||||
``"scrubbing"`` in ``strategy``, no scrubbing is performed, else the
|
||||
default value is 0.5 (default None).
|
||||
std_dvars_threshold : float, optional
|
||||
Standardized DVARS threshold for scrub. DVARs is defined as root mean
|
||||
squared intensity difference of volume N to volume N+1. D referring to
|
||||
temporal derivative of timecourses, VARS referring to root mean squared
|
||||
variance over voxels. If None and no ``"scrubbing"`` in ``strategy``,
|
||||
no scrubbing is performed, else the default value is 1.5
|
||||
(default None).
|
||||
detrend : bool, optional
|
||||
If True, detrending will be applied on timeseries, before confound
|
||||
removal (default True).
|
||||
|
|
@ -155,8 +178,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
|
||||
def __init__(
|
||||
self,
|
||||
strategy: Optional[dict[str, str]] = None,
|
||||
strategy: Optional[dict[str, Union[str, bool]]] = None,
|
||||
spike: Optional[float] = None,
|
||||
scrub: Optional[int] = None,
|
||||
fd_threshold: Optional[float] = None,
|
||||
std_dvars_threshold: Optional[float] = None,
|
||||
detrend: bool = True,
|
||||
standardize: bool = True,
|
||||
low_pass: Optional[float] = None,
|
||||
|
|
@ -170,9 +196,13 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
"motion": "full",
|
||||
"wm_csf": "full",
|
||||
"global_signal": "full",
|
||||
"scrubbing": False,
|
||||
}
|
||||
self.strategy = strategy
|
||||
self.spike = spike
|
||||
self.scrub = scrub
|
||||
self.fd_threshold = fd_threshold
|
||||
self.std_dvars_threshold = std_dvars_threshold
|
||||
self.detrend = detrend
|
||||
self.standardize = standardize
|
||||
self.low_pass = low_pass
|
||||
|
|
@ -180,13 +210,22 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
self.t_r = t_r
|
||||
self.masks = masks
|
||||
|
||||
self._valid_components = ["motion", "wm_csf", "global_signal"]
|
||||
self._valid_components = [
|
||||
"motion",
|
||||
"wm_csf",
|
||||
"global_signal",
|
||||
"scrubbing",
|
||||
]
|
||||
self._valid_confounds = ["basic", "power2", "derivatives", "full"]
|
||||
|
||||
if any(not isinstance(k, str) for k in strategy.keys()):
|
||||
raise_error("Strategy keys must be strings", ValueError)
|
||||
|
||||
if any(not isinstance(v, str) for v in strategy.values()):
|
||||
if any(
|
||||
not isinstance(v, str)
|
||||
for k, v in strategy.items()
|
||||
if k != "scrubbing"
|
||||
):
|
||||
raise_error("Strategy values must be strings", ValueError)
|
||||
|
||||
if any(x not in self._valid_components for x in strategy.keys()):
|
||||
|
|
@ -199,7 +238,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
klass=ValueError,
|
||||
)
|
||||
|
||||
if any(x not in self._valid_confounds for x in strategy.values()):
|
||||
if any(
|
||||
v not in self._valid_confounds
|
||||
for k, v in strategy.items()
|
||||
if k != "scrubbing"
|
||||
):
|
||||
raise_error(
|
||||
msg=f"Invalid confound types {list(strategy.values())}. "
|
||||
f"Valid confound types are {self._valid_confounds}.\n"
|
||||
|
|
@ -302,7 +345,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
|
||||
Raises
|
||||
------
|
||||
ValueError
|
||||
RuntimeError
|
||||
If invalid confounds file is found.
|
||||
|
||||
"""
|
||||
|
|
@ -315,15 +358,20 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
derivatives_to_compute = {} # the dictionary of missing derivatives
|
||||
|
||||
for t_kind, t_strategy in self.strategy.items():
|
||||
if t_kind != "scrubbing":
|
||||
t_basics = FMRIPREP_BASICS[t_kind]
|
||||
|
||||
if any(x not in available_vars for x in t_basics):
|
||||
missing = [x for x in t_basics if x not in available_vars]
|
||||
raise_error(
|
||||
msg=(
|
||||
"Invalid confounds file. Missing basic confounds: "
|
||||
f"{missing}. "
|
||||
"Check if this file is really an fmriprep confounds file. "
|
||||
"You can also modify the confound removal strategy."
|
||||
"Check if this file is really an fmriprep "
|
||||
"confounds file. You can also modify the confound "
|
||||
"removal strategy."
|
||||
),
|
||||
klass=RuntimeError,
|
||||
)
|
||||
|
||||
to_select.extend(t_basics)
|
||||
|
|
@ -345,15 +393,22 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
x_derivative_2 = f"{x_derivative}_power2"
|
||||
to_select.append(x_derivative_2)
|
||||
if x_derivative_2 not in available_vars:
|
||||
squares_to_compute[x_derivative_2] = x_derivative
|
||||
squares_to_compute[x_derivative_2] = (
|
||||
x_derivative
|
||||
)
|
||||
# Add spike
|
||||
spike_name = "framewise_displacement"
|
||||
if self.spike is not None:
|
||||
if spike_name not in available_vars:
|
||||
raise_error(
|
||||
"Invalid confounds file. Missing framewise_displacement "
|
||||
"(spike) confound. "
|
||||
"Check if this file is really an fmriprep confounds file. "
|
||||
"You can also deactivate spike (set spike = None)."
|
||||
msg=(
|
||||
"Invalid confounds file. Missing "
|
||||
"framewise_displacement (spike) confound. "
|
||||
"Check if this file is really an fmriprep confounds "
|
||||
"file. You can also deactivate spike "
|
||||
"(set spike = None)."
|
||||
),
|
||||
klass=RuntimeError,
|
||||
)
|
||||
out = to_select, squares_to_compute, derivatives_to_compute, spike_name
|
||||
return out
|
||||
|
|
@ -377,14 +432,12 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
if confounds_format == "adhoc":
|
||||
self._map_adhoc_to_fmriprep(input)
|
||||
|
||||
processed_spec = self._process_fmriprep_spec(input)
|
||||
|
||||
(
|
||||
to_select,
|
||||
squares_to_compute,
|
||||
derivatives_to_compute,
|
||||
spike_name,
|
||||
) = processed_spec
|
||||
) = self._process_fmriprep_spec(input)
|
||||
# Copy the confounds
|
||||
out_df = input["data"].copy()
|
||||
|
||||
|
|
@ -415,6 +468,67 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
|
||||
return out_df
|
||||
|
||||
def _get_scrub_regressors(
|
||||
self, confounds_df: pd.DataFrame
|
||||
) -> pd.DataFrame:
|
||||
"""Get motion outline regressors.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
confounds_df : pandas.DataFrame
|
||||
pandas.DataFrame with all confounds.
|
||||
|
||||
Returns
|
||||
-------
|
||||
pandas.DataFrame
|
||||
pandas.DataFrame of motion outline regressors.
|
||||
|
||||
Raises
|
||||
------
|
||||
RuntimeError
|
||||
If ``std_dvars`` and / or ``framewise_displacement`` is not found
|
||||
in ``confounds_df``.
|
||||
|
||||
"""
|
||||
# Check columns
|
||||
if (
|
||||
self.std_dvars_threshold is not None
|
||||
and "std_dvars" not in confounds_df.columns
|
||||
):
|
||||
raise_error(
|
||||
msg=(
|
||||
"Invalid confounds file. Missing std_dvars "
|
||||
"(standardized DVARS) confound. "
|
||||
"Check if this file is really an fMRIPrep confounds file. "
|
||||
),
|
||||
klass=RuntimeError,
|
||||
)
|
||||
if (
|
||||
self.fd_threshold is not None
|
||||
and "framewise_displacement" not in confounds_df.columns
|
||||
):
|
||||
raise_error(
|
||||
msg=(
|
||||
"Invalid confounds file. Missing framewise_displacement "
|
||||
"confound. "
|
||||
"Check if this file is really an fMRIPrep confounds file. "
|
||||
),
|
||||
klass=RuntimeError,
|
||||
)
|
||||
# Use function from nilearn to not reinvent the wheel
|
||||
return _load_scrub(
|
||||
confounds_raw=confounds_df,
|
||||
scrub=self.scrub if self.scrub is not None else 0,
|
||||
fd_threshold=(
|
||||
self.fd_threshold if self.fd_threshold is not None else 0.5
|
||||
),
|
||||
std_dvars_threshold=(
|
||||
self.std_dvars_threshold
|
||||
if self.std_dvars_threshold is not None
|
||||
else 1.5
|
||||
),
|
||||
)
|
||||
|
||||
def _validate_data(
|
||||
self,
|
||||
input: dict[str, Any],
|
||||
|
|
@ -578,12 +692,38 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
}
|
||||
}
|
||||
)
|
||||
|
||||
signal_clean_kwargs = {}
|
||||
# Set up scrubbing mask if needed
|
||||
if self.strategy.get("scrubbing", False):
|
||||
motion_outline_regressors = self._get_scrub_regressors(
|
||||
input["confounds"]["data"]
|
||||
)
|
||||
# Add regressors to confounds
|
||||
confounds_df = pd.concat(
|
||||
[confounds_df, motion_outline_regressors], axis=1
|
||||
)
|
||||
# Get sample mask
|
||||
sample_mask, confounds_df = prepare_output(
|
||||
confounds=confounds_df, demean=False
|
||||
)
|
||||
signal_clean_kwargs.update(
|
||||
{
|
||||
"clean__sample_mask": sample_mask,
|
||||
}
|
||||
)
|
||||
# Clean image
|
||||
logger.info("Cleaning image using nilearn")
|
||||
logger.debug(f"\tstrategy: {self.strategy}")
|
||||
logger.debug(f"\tspike: {self.spike}")
|
||||
logger.debug(f"\tscrub: {self.scrub}")
|
||||
logger.debug(f"\tfd_threshold: {self.fd_threshold}")
|
||||
logger.debug(f"\tstd_dvars_threshold: {self.std_dvars_threshold}")
|
||||
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}")
|
||||
logger.debug(f"\tt_r: {self.t_r}")
|
||||
|
||||
# Deconfound data
|
||||
cleaned_img = nimg.clean_img(
|
||||
|
|
@ -595,6 +735,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
|
|||
high_pass=self.high_pass,
|
||||
t_r=t_r,
|
||||
mask_img=mask_img,
|
||||
**signal_clean_kwargs,
|
||||
)
|
||||
# Fix t_r as nilearn messes it up
|
||||
cleaned_img.header["pixdim"][4] = t_r
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import numpy as np
|
|||
import pandas as pd
|
||||
import pytest
|
||||
from nilearn._utils.exceptions import DimensionError
|
||||
from nilearn.interfaces.fmriprep.load_confounds_utils import prepare_output
|
||||
from numpy.testing import assert_array_equal, assert_raises
|
||||
from pandas.testing import assert_frame_equal
|
||||
|
||||
|
|
@ -211,7 +212,7 @@ def test_fMRIPrepConfoundRemover__process_fmriprep_spec() -> None:
|
|||
)
|
||||
|
||||
msg = r"Missing basic confounds: \['white_matter'\]"
|
||||
with pytest.raises(ValueError, match=msg):
|
||||
with pytest.raises(RuntimeError, match=msg):
|
||||
confound_remover._process_fmriprep_spec({"data": confounds_df})
|
||||
|
||||
var_names = ["csf", "white_matter"]
|
||||
|
|
@ -220,7 +221,7 @@ def test_fMRIPrepConfoundRemover__process_fmriprep_spec() -> None:
|
|||
)
|
||||
|
||||
msg = r"Missing framewise_displacement"
|
||||
with pytest.raises(ValueError, match=msg):
|
||||
with pytest.raises(RuntimeError, match=msg):
|
||||
confound_remover._process_fmriprep_spec({"data": confounds_df})
|
||||
|
||||
|
||||
|
|
@ -320,6 +321,32 @@ def test_fMRIPRepConfoundRemover__pick_confounds_fmriprep_compute() -> None:
|
|||
assert_frame_equal(out_junifer, out_fmriprep)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"preprocessor",
|
||||
[
|
||||
fMRIPrepConfoundRemover(
|
||||
std_dvars_threshold=1.5,
|
||||
),
|
||||
fMRIPrepConfoundRemover(
|
||||
fd_threshold=0.5,
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_fMRIPrepConfoundRemover__get_scrub_regressors_errors(
|
||||
preprocessor: type,
|
||||
) -> None:
|
||||
"""Test fMRIPrepConfoundRemover scrub regressors errors.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
preprocessor : object
|
||||
The parametrized preprocessor.
|
||||
|
||||
"""
|
||||
with pytest.raises(RuntimeError, match="Invalid confounds file."):
|
||||
preprocessor._get_scrub_regressors(pd.DataFrame({"a": [1, 2]}))
|
||||
|
||||
|
||||
def test_fMRIPrepConfoundRemover__validate_data() -> None:
|
||||
"""Test fMRIPrepConfoundRemover validate data."""
|
||||
confound_remover = fMRIPrepConfoundRemover(strategy={"wm_csf": "full"})
|
||||
|
|
@ -567,3 +594,64 @@ def test_fMRIPrepConfoundRemover_fit_transform_masks() -> None:
|
|||
assert "dependencies" in output["BOLD"]["meta"]
|
||||
dependencies = output["BOLD"]["meta"]["dependencies"]
|
||||
assert dependencies == {"numpy", "nilearn"}
|
||||
|
||||
|
||||
def test_fMRIPrepConfoundRemover_scrubbing() -> None:
|
||||
"""Test fMRIPrepConfoundRemover with scrubbing."""
|
||||
confound_remover = fMRIPrepConfoundRemover(
|
||||
strategy={
|
||||
"motion": "full",
|
||||
"wm_csf": "full",
|
||||
"global_signal": "full",
|
||||
"scrubbing": True,
|
||||
},
|
||||
)
|
||||
|
||||
with PartlyCloudyTestingDataGrabber(reduce_confounds=False) as dg:
|
||||
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
|
||||
orig_bold = element_data["BOLD"]["data"].get_fdata().copy()
|
||||
orig_confounds = element_data["BOLD"]["confounds"].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, element_data["BOLD"]["data"].get_fdata()
|
||||
)
|
||||
|
||||
# Data should have different shape
|
||||
assert_raises(
|
||||
AssertionError,
|
||||
assert_array_equal,
|
||||
orig_bold.shape,
|
||||
trans_bold.shape,
|
||||
)
|
||||
# and be different
|
||||
assert_raises(
|
||||
AssertionError, assert_array_equal, orig_bold, trans_bold
|
||||
)
|
||||
|
||||
# Check scrubbing process
|
||||
# Should be at the start
|
||||
confounds_df = confound_remover._pick_confounds(orig_confounds)
|
||||
assert confounds_df.shape == (168, 36)
|
||||
# Should have 4 motion outliers based on threshold
|
||||
motion_outline_regressors = confound_remover._get_scrub_regressors(
|
||||
orig_confounds["data"]
|
||||
)
|
||||
assert motion_outline_regressors.shape == (168, 4)
|
||||
# Add regressors to confounds
|
||||
concat_confounds_df = pd.concat(
|
||||
[confounds_df, motion_outline_regressors], axis=1
|
||||
)
|
||||
assert concat_confounds_df.shape == (168, 40)
|
||||
# Get sample mask and correct confounds
|
||||
sample_mask, final_confounds_df = prepare_output(
|
||||
confounds=concat_confounds_df, demean=False
|
||||
)
|
||||
assert not confounds_df.equals(final_confounds_df)
|
||||
assert sample_mask.shape == (164,)
|
||||
assert not (np.isin([91, 92, 93, 113], sample_mask)).all()
|
||||
|
|
|
|||
Loading…
Reference in a new issue