[ENH]: Add TemporalFilter #432

Merged
fraimondo merged 7 commits from enh/temporal_filter into main 2025-07-21 14:28:41 +00:00
6 changed files with 347 additions and 0 deletions

View file

@ -167,6 +167,10 @@ Available
- Slice ``BOLD`` data temporally
- | Done
- :gh:`443`
* - ``TemporalFilter``
- Filter (clean) ``BOLD`` data temporally
- | Done
- :gh:`432`
..

View file

@ -0,0 +1 @@
Introduce :class:`.TemporalFilter` preprocessor for temporally filtering BOLD data by `Fede Raimondo`_

View file

@ -76,6 +76,7 @@ class PipelineComponentRegistry(metaclass=Singleton):
"SpaceWarper": "SpaceWarper",
"fMRIPrepConfoundRemover": "fMRIPrepConfoundRemover",
"TemporalSlicer": "TemporalSlicer",
"TemporalFilter": "TemporalFilter",
},
"marker": {
"ALFFParcels": "ALFFParcels",

View file

@ -4,6 +4,7 @@ __all__ = [
"SpaceWarper",
"Smoothing",
"TemporalSlicer",
"TemporalFilter",
]
from .base import BasePreprocessor
@ -11,3 +12,4 @@ from .confounds import fMRIPrepConfoundRemover
from .warping import SpaceWarper
from .smoothing import Smoothing
from ._temporal_slicer import TemporalSlicer
from ._temporal_filter import TemporalFilter

View file

@ -0,0 +1,240 @@
"""Provide class for temporal filtering."""
fraimondo commented 2025-07-21 12:02:05 +00:00 (Migrated from github.com)

Copy/paste error. It should say filtered data

Copy/paste error. It should say `filtered data`
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
# Leonard Sasse <l.sasse@fz-juelich.de>
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import (
Any,
ClassVar,
Optional,
Union,
)
import nibabel as nib
from nilearn import image as nimg
from nilearn._utils.niimg_conversions import check_niimg_4d
from ..api.decorators import register_preprocessor
from ..data import get_data
from ..pipeline import WorkDirManager
from ..typing import Dependencies
from ..utils import logger
from .base import BasePreprocessor
__all__ = ["TemporalFilter"]
@register_preprocessor
class TemporalFilter(BasePreprocessor):
"""Class for temporal filtering.
Temporal filtering is based on :func:`nilearn.image.clean_img`.
Parameters
----------
detrend : bool, optional
If True, detrending will be applied on timeseries (default True).
standardize : bool, optional
If True, returned signals are set to unit variance (default True).
low_pass : float, optional
Low cutoff frequencies, in Hertz. If None, no filtering is applied
(default None).
high_pass : float, optional
High cutoff frequencies, in Hertz. If None, no filtering is
applied (default None).
t_r : float, optional
Repetition time, in second (sampling period).
If None, it will use t_r from nifti header (default None).
masks : str, dict or list of dict or str, optional
The specification of the masks to apply to regions before extracting
signals. Check :ref:`Using Masks <using_masks>` for more details.
If None, will not apply any mask (default None).
"""
_DEPENDENCIES: ClassVar[Dependencies] = {"numpy", "nilearn"}
def __init__(
self,
detrend: bool = True,
standardize: bool = True,
low_pass: Optional[float] = None,
high_pass: Optional[float] = None,
t_r: Optional[float] = None,
masks: Union[str, dict, list[Union[dict, str]], None] = None,
) -> None:
"""Initialize the class."""
self.detrend = detrend
self.standardize = standardize
self.low_pass = low_pass
self.high_pass = high_pass
self.t_r = t_r
self.masks = masks
super().__init__(on="BOLD", required_data_types=["BOLD"])
def get_valid_inputs(self) -> list[str]:
"""Get valid data types for input.
Returns
-------
list of str
The list of data types that can be used as input for this
preprocessor.
"""
return ["BOLD"]
def get_output_type(self, input_type: str) -> str:
"""Get output type.
Parameters
----------
input_type : str
The input to the preprocessor.
Returns
-------
str
The data type output by the preprocessor.
"""
# Does not add any new keys
return input_type
def _validate_data(
self,
input: dict[str, Any],
) -> None:
"""Validate input data.
Parameters
----------
input : dict
Dictionary containing the ``BOLD`` data from the
Junifer Data object.
Raises
------
ValueError
If ``"data"`` is not 4D
"""
# BOLD must be 4D niimg
check_niimg_4d(input["data"])
def preprocess(
self,
input: dict[str, Any],
extra_input: Optional[dict[str, Any]] = None,
) -> tuple[dict[str, Any], Optional[dict[str, dict[str, Any]]]]:
"""Preprocess.
Parameters
----------
input : dict
A single input from the Junifer Data object to preprocess.
extra_input : dict, optional
The other fields in the Junifer Data object.
Returns
-------
dict
The computed result as dictionary. If `self.masks` is not None,
then the target data computed mask is updated for further steps.
None
Extra "helper" data types as dictionary to add to the Junifer Data
object.
"""
# Validate data
self._validate_data(input)
# 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}",
)
# Create element-specific tempdir for storing generated data
# and / or mask
element_tempdir = WorkDirManager().get_element_tempdir(
prefix="temporal_filter"
)
# Set mask data
mask_img = None
if self.masks is not None:
# Generate mask
logger.debug(f"Masking with {self.masks}")
mask_img = get_data(
kind="mask",
names=self.masks,
target_data=input,
extra_input=extra_input,
)
# Save generated mask for use later
generated_mask_img_path = element_tempdir / "generated_mask.nii.gz"
nib.save(mask_img, generated_mask_img_path)
# Save BOLD mask and link it to the BOLD data type dict;
# this allows to use "inherit" down the pipeline
logger.debug("Setting `BOLD.mask`")
input.update(
{
"mask": {
# Update path to sync with "data"
"path": generated_mask_img_path,
# Update data
"data": mask_img,
# Should be in the same space as target data
"space": input["space"],
}
}
)
signal_clean_kwargs = {}
# Clean image
logger.info("Temporal filter 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}")
logger.debug(f"\tt_r: {self.t_r}")
cleaned_img = nimg.clean_img(
imgs=bold_img,
detrend=self.detrend,
standardize=self.standardize,
low_pass=self.low_pass,
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
# Save filtered data
filtered_data_path = element_tempdir / "filtered_data.nii.gz"
nib.save(cleaned_img, filtered_data_path)
logger.debug("Updating `BOLD`")
input.update(
{
# Update path to sync with "data"
"path": filtered_data_path,
# Update data
"data": cleaned_img,
}
)
return input, None

View file

@ -0,0 +1,99 @@
"""Provide tests for TemporalFilter."""
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from typing import Optional
import pytest
from junifer.datareader import DefaultDataReader
from junifer.preprocess import TemporalFilter
from junifer.testing.datagrabbers import PartlyCloudyTestingDataGrabber
@pytest.mark.parametrize(
"detrend, standardize, low_pass, high_pass, t_r, masks",
(
[
True,
True,
None,
None,
None,
None,
],
[
False,
True,
0.1,
None,
None,
"compute_brain_mask",
],
[
True,
False,
None,
0.08,
None,
"compute_background_mask",
],
[
False,
False,
None,
None,
2,
None,
],
[
True,
True,
0.1,
0.08,
2,
"compute_brain_mask",
],
),
)
def test_TemporalFilter(
detrend: bool,
standardize: bool,
low_pass: Optional[float],
high_pass: Optional[float],
t_r: Optional[float],
masks: Optional[str],
) -> None:
"""Test TemporalFilter.
Parameters
----------
detrend : bool
The parametrized detrending flag.
standardize : bool
The parametrized standardization flag.
low_pass : float or None
The parametrized low pass value.
high_pass : float or None
The parametrized high pass value.
t_r : float or None
The parametrized repetition time.
masks : str or None
The parametrized mask.
"""
with PartlyCloudyTestingDataGrabber() as dg:
# Read data
element_data = DefaultDataReader().fit_transform(dg["sub-01"])
# Preprocess data
output = TemporalFilter(
detrend=detrend,
standardize=standardize,
low_pass=low_pass,
high_pass=high_pass,
t_r=t_r,
masks=masks,
).fit_transform(element_data)
assert isinstance(output, dict)