[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
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.

View file

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

View file

@ -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()