[ENH]: Add scrubbing support in fMRIPrepConfoundRemover #421

Merged
synchon merged 14 commits from feat/fmriprepconfoundremover-dvars-scrub into main 2025-02-13 17:05:38 +00:00
5 changed files with 278 additions and 47 deletions

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

View file

@ -0,0 +1 @@
Add scrubbing support to :class:`.fMRIPrepConfoundRemover` by `Synchon Mandal`_

View file

@ -185,7 +185,7 @@ def signals() -> list[np.ndarray]:
@pytest.fixture @pytest.fixture
def signals_and_covariances( def signals_and_covariances(
cov_estimator: Union[LedoitWolf, EmpiricalCovariance] cov_estimator: Union[LedoitWolf, EmpiricalCovariance],
) -> tuple[list[np.ndarray], list[float]]: ) -> tuple[list[np.ndarray], list[float]]:
"""Return signals and covariances for a covariance estimator. """Return signals and covariances for a covariance estimator.

View file

@ -17,6 +17,8 @@ import numpy as np
import pandas as pd import pandas as pd
from nilearn import image as nimg 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.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 ...api.decorators import register_preprocessor
from ...data import get_data from ...data import get_data
@ -110,13 +112,15 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
---------- ----------
strategy : dict, optional strategy : dict, optional
The strategy to use for each component. If None, will use the *full* 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 The keys of the dictionary should correspond to names of noise
components to include: components to include:
* ``motion`` * ``motion``
* ``wm_csf`` * ``wm_csf``
* ``global_signal`` * ``global_signal``
* ``scrubbing``
The values of dictionary should correspond to types of confounds The values of dictionary should correspond to types of confounds
extracted from each signal: extracted from each signal:
@ -126,10 +130,29 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
* ``derivatives`` : signal + derivatives * ``derivatives`` : signal + derivatives
* ``full`` : signal + deriv. + quadratic terms + power2 deriv. * ``full`` : signal + deriv. + quadratic terms + power2 deriv.
except ``scrubbing`` which needs to be bool.
spike : float, optional spike : float, optional
If None, no spike regressor is added. If spike is a float, it will 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 add a spike regressor for every point at which framewise displacement
exceeds the specified float (default None). 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 detrend : bool, optional
If True, detrending will be applied on timeseries, before confound If True, detrending will be applied on timeseries, before confound
removal (default True). removal (default True).
@ -155,8 +178,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
def __init__( def __init__(
self, self,
strategy: Optional[dict[str, str]] = None, strategy: Optional[dict[str, Union[str, bool]]] = None,
spike: Optional[float] = None, spike: Optional[float] = None,
scrub: Optional[int] = None,
fd_threshold: Optional[float] = None,
std_dvars_threshold: Optional[float] = None,
detrend: bool = True, detrend: bool = True,
standardize: bool = True, standardize: bool = True,
low_pass: Optional[float] = None, low_pass: Optional[float] = None,
@ -170,9 +196,13 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
"motion": "full", "motion": "full",
"wm_csf": "full", "wm_csf": "full",
"global_signal": "full", "global_signal": "full",
"scrubbing": False,
} }
self.strategy = strategy self.strategy = strategy
self.spike = spike self.spike = spike
self.scrub = scrub
self.fd_threshold = fd_threshold
self.std_dvars_threshold = std_dvars_threshold
self.detrend = detrend self.detrend = detrend
self.standardize = standardize self.standardize = standardize
self.low_pass = low_pass self.low_pass = low_pass
@ -180,13 +210,22 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
self.t_r = t_r self.t_r = t_r
self.masks = masks 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"] self._valid_confounds = ["basic", "power2", "derivatives", "full"]
if any(not isinstance(k, str) for k in strategy.keys()): if any(not isinstance(k, str) for k in strategy.keys()):
raise_error("Strategy keys must be strings", ValueError) 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) raise_error("Strategy values must be strings", ValueError)
if any(x not in self._valid_components for x in strategy.keys()): if any(x not in self._valid_components for x in strategy.keys()):
@ -199,7 +238,11 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
klass=ValueError, 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( raise_error(
msg=f"Invalid confound types {list(strategy.values())}. " msg=f"Invalid confound types {list(strategy.values())}. "
f"Valid confound types are {self._valid_confounds}.\n" f"Valid confound types are {self._valid_confounds}.\n"
@ -302,7 +345,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
Raises Raises
------ ------
ValueError RuntimeError
If invalid confounds file is found. If invalid confounds file is found.
""" """
@ -315,15 +358,20 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
derivatives_to_compute = {} # the dictionary of missing derivatives derivatives_to_compute = {} # the dictionary of missing derivatives
for t_kind, t_strategy in self.strategy.items(): for t_kind, t_strategy in self.strategy.items():
if t_kind != "scrubbing":
t_basics = FMRIPREP_BASICS[t_kind] t_basics = FMRIPREP_BASICS[t_kind]
if any(x not in available_vars for x in t_basics): 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] missing = [x for x in t_basics if x not in available_vars]
raise_error( raise_error(
msg=(
"Invalid confounds file. Missing basic confounds: " "Invalid confounds file. Missing basic confounds: "
f"{missing}. " f"{missing}. "
"Check if this file is really an fmriprep confounds file. " "Check if this file is really an fmriprep "
"You can also modify the confound removal strategy." "confounds file. You can also modify the confound "
"removal strategy."
),
klass=RuntimeError,
) )
to_select.extend(t_basics) to_select.extend(t_basics)
@ -345,15 +393,22 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
x_derivative_2 = f"{x_derivative}_power2" x_derivative_2 = f"{x_derivative}_power2"
to_select.append(x_derivative_2) to_select.append(x_derivative_2)
if x_derivative_2 not in available_vars: 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" spike_name = "framewise_displacement"
if self.spike is not None: if self.spike is not None:
if spike_name not in available_vars: if spike_name not in available_vars:
raise_error( raise_error(
"Invalid confounds file. Missing framewise_displacement " msg=(
"(spike) confound. " "Invalid confounds file. Missing "
"Check if this file is really an fmriprep confounds file. " "framewise_displacement (spike) confound. "
"You can also deactivate spike (set spike = None)." "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 out = to_select, squares_to_compute, derivatives_to_compute, spike_name
return out return out
@ -377,14 +432,12 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
if confounds_format == "adhoc": if confounds_format == "adhoc":
self._map_adhoc_to_fmriprep(input) self._map_adhoc_to_fmriprep(input)
processed_spec = self._process_fmriprep_spec(input)
( (
to_select, to_select,
squares_to_compute, squares_to_compute,
derivatives_to_compute, derivatives_to_compute,
spike_name, spike_name,
) = processed_spec ) = self._process_fmriprep_spec(input)
# Copy the confounds # Copy the confounds
out_df = input["data"].copy() out_df = input["data"].copy()
@ -415,6 +468,67 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
return out_df 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( def _validate_data(
self, self,
input: dict[str, Any], 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 # Clean image
logger.info("Cleaning image using nilearn") 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"\tdetrend: {self.detrend}")
logger.debug(f"\tstandardize: {self.standardize}") logger.debug(f"\tstandardize: {self.standardize}")
logger.debug(f"\tlow_pass: {self.low_pass}") logger.debug(f"\tlow_pass: {self.low_pass}")
logger.debug(f"\thigh_pass: {self.high_pass}") logger.debug(f"\thigh_pass: {self.high_pass}")
logger.debug(f"\tt_r: {self.t_r}")
# Deconfound data # Deconfound data
cleaned_img = nimg.clean_img( cleaned_img = nimg.clean_img(
@ -595,6 +735,7 @@ class fMRIPrepConfoundRemover(BasePreprocessor):
high_pass=self.high_pass, high_pass=self.high_pass,
t_r=t_r, t_r=t_r,
mask_img=mask_img, mask_img=mask_img,
**signal_clean_kwargs,
) )
# Fix t_r as nilearn messes it up # Fix t_r as nilearn messes it up
cleaned_img.header["pixdim"][4] = t_r cleaned_img.header["pixdim"][4] = t_r

View file

@ -10,6 +10,7 @@ import numpy as np
import pandas as pd import pandas as pd
import pytest import pytest
from nilearn._utils.exceptions import DimensionError 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 numpy.testing import assert_array_equal, assert_raises
from pandas.testing import assert_frame_equal 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'\]" 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}) confound_remover._process_fmriprep_spec({"data": confounds_df})
var_names = ["csf", "white_matter"] var_names = ["csf", "white_matter"]
@ -220,7 +221,7 @@ def test_fMRIPrepConfoundRemover__process_fmriprep_spec() -> None:
) )
msg = r"Missing framewise_displacement" 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}) 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) 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: 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"})
@ -567,3 +594,64 @@ def test_fMRIPrepConfoundRemover_fit_transform_masks() -> None:
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"}
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()